mirror of
https://github.com/mudler/LocalAI.git
synced 2026-09-08 12:14:06 -04:00
Compare commits
205
Commits
No files matched your search
@@ -49,6 +49,40 @@ AI agents MUST NOT add `Co-Authored-By` trailers for themselves either.
|
||||
A human reviewer owns the contribution; the AI's involvement is recorded
|
||||
via `Assisted-by` (see below).
|
||||
|
||||
### Exception: automation operated by a maintainer
|
||||
|
||||
The rule above addresses the common case, an AI assistant helping a human
|
||||
contributor who then signs off. It does not fit automation that a
|
||||
maintainer runs themselves, which opens pull requests with no human
|
||||
submitter to sign. Applied literally there, nothing ever signs and the
|
||||
DCO check blocks the pull request permanently.
|
||||
|
||||
A maintainer-operated bot MUST therefore add a `Signed-off-by` trailer
|
||||
naming **the maintainer who operates it**, not the bot and not the model:
|
||||
|
||||
```
|
||||
Assisted-by: Codex:gpt-5
|
||||
Signed-off-by: Ettore Di Giacinto <mudler@localai.io>
|
||||
```
|
||||
|
||||
This is not the AI certifying the DCO. The maintainer is, exactly as they
|
||||
do for a commit they typed by hand: they configured the automation, they
|
||||
own its output, and they take responsibility for it when they merge it.
|
||||
The `Assisted-by` trailer still records that a model produced the code, so
|
||||
the provenance trail is unchanged.
|
||||
|
||||
The exception is narrow and does not widen the rule for anyone else:
|
||||
|
||||
- It applies only to automation a LocalAI maintainer operates and whose
|
||||
output that maintainer reviews before merge.
|
||||
- The sign-off names a real person who accepts DCO responsibility.
|
||||
- An AI assistant helping an outside contributor still MUST NOT sign off.
|
||||
That contributor adds their own trailer.
|
||||
- A bot MUST NOT sign off on behalf of anyone other than its operator, and
|
||||
MUST NOT add a trailer for a contributor whose branch it pushes to. If
|
||||
automation contributes to someone else's branch, it leaves the sign-off
|
||||
to that contributor.
|
||||
|
||||
## Attribution
|
||||
|
||||
When AI tools contribute to LocalAI development, proper attribution helps
|
||||
|
||||
@@ -236,6 +236,58 @@ Use these HTTP status codes:
|
||||
|
||||
If your endpoint should be tracked for usage (token counts, request counts), add the `usageMiddleware` to its middleware chain. See `core/http/middleware/usage.go` and how it's applied in `routes/openai.go`.
|
||||
|
||||
## Control-plane database health metrics
|
||||
|
||||
In distributed mode the frontend registers three OpenTelemetry gauges over the
|
||||
PostgreSQL control-plane database (`core/services/monitoring/control_plane_db.go`,
|
||||
wired in `core/application/distributed.go`). They reach `/metrics` through the
|
||||
same Prometheus exporter as the rest of the API metrics.
|
||||
|
||||
| Metric | Meaning | Page when |
|
||||
|--------|---------|-----------|
|
||||
| `localai_control_plane_oldest_xmin_age` | Transactions elapsed since the oldest snapshot any backend still holds | above a few million, and rising |
|
||||
| `localai_control_plane_longest_transaction_seconds` | Age of the longest open transaction | above 3600 |
|
||||
| `localai_control_plane_dead_tuple_ratio` | Dead tuples per live tuple, labelled by `table`, on `backend_nodes`, `node_models` and `gallery_operations` | sustained above ~10 on a small table |
|
||||
|
||||
A sustained high `localai_control_plane_oldest_xmin_age` is the one to page on.
|
||||
While it grows, autovacuum can reclaim nothing anywhere in the database no
|
||||
matter how often it runs, so the dead tuple ratio keeps climbing and a six-row
|
||||
registry table can reach hundreds of megabytes. Tuning autovacuum does not help.
|
||||
The fix is to find the transaction holding the horizon open and clear it:
|
||||
|
||||
```sql
|
||||
SELECT pid, state, age(backend_xmin) AS xmin_age, now() - xact_start AS xact_age, query
|
||||
FROM pg_stat_activity
|
||||
WHERE backend_xmin IS NOT NULL
|
||||
ORDER BY age(backend_xmin) DESC;
|
||||
```
|
||||
|
||||
Then `pg_terminate_backend(pid)` on the offenders, and `VACUUM (VERBOSE)` the
|
||||
bloated tables once the horizon has moved.
|
||||
|
||||
**A healthy-looking xmin age does not on its own prove the horizon is free.**
|
||||
The gauge reads `pg_stat_activity`, which only sees live backends. Two other
|
||||
things pin the very same horizon and are invisible there, so either one can hold
|
||||
vacuum back while the gauge reads 0:
|
||||
|
||||
```sql
|
||||
SELECT gid, prepared, database, transaction FROM pg_prepared_xacts;
|
||||
SELECT slot_name, active, xmin, catalog_xmin FROM pg_replication_slots;
|
||||
```
|
||||
|
||||
An orphaned prepared transaction is cleared with `ROLLBACK PREPARED '<gid>'`,
|
||||
and a stale slot with `pg_drop_replication_slot('<slot_name>')`. Check both
|
||||
before concluding that a bloated table has some other cause.
|
||||
|
||||
Sampling is scrape-driven behind a 30 second cache, so scrape frequency does not
|
||||
translate into database load. Failed and timed-out samples cost the same interval
|
||||
as successful ones, so a database that is already struggling is not retried on
|
||||
every scrape. A failed sample reports the last good values rather than failing the
|
||||
scrape, because these gauges matter most when the database is struggling. Before
|
||||
the first successful sample the gauges are absent rather than zero, since a zero
|
||||
xmin age would read as a healthy horizon: alert on `absent()` too if you need to
|
||||
distinguish "healthy" from "never sampled".
|
||||
|
||||
## Advertising surfaces — where to register a new capability
|
||||
|
||||
Beyond routing and auth, LocalAI publishes its capability surface in **four independent places**. When you add an endpoint — especially one introducing a net-new capability like a new media type or a new auth-gated feature — you must update every relevant surface. These aren't optional: missing them means the endpoint works but is invisible to clients, admins, and the UI.
|
||||
|
||||
@@ -77,6 +77,56 @@ spectrum. **Metal (Darwin) only** - it is a no-op on CUDA/CPU. Enable with
|
||||
budget). Gallery entries built on this: `deepseek-v4-flash-q4-ssd` (153 GB Flash
|
||||
on a 128 GB Mac) and `deepseek-v4-pro-q2-ssd` (433 GB Pro, experimental).
|
||||
|
||||
## CUDA architecture (do not build without one)
|
||||
|
||||
`backend/cpp/ds4/Makefile` drives upstream's **object targets** directly
|
||||
(`$(MAKE) -C ds4 ds4.o ds4_cuda.o ...`), which bypasses upstream's own guard:
|
||||
its `cuda` target refuses to build unless `CUDA_ARCH` is set, and offers
|
||||
`cuda-spark` (sm_121, DGX Spark / GB10) and `cuda-generic` (native) instead.
|
||||
Built with no `-arch`, nvcc targets its default architecture and the kernels run
|
||||
as JIT'd PTX. On GB10 that silently corrupted every prefill batch of >=128
|
||||
tokens - the model emitted text unrelated to the prompt and never closed its
|
||||
thinking block, so `content` came back empty - and cost close to two orders of
|
||||
magnitude of prefill throughput (4.21 t/s vs 325.70 t/s, same box, same model).
|
||||
Short prompts stayed correct, which is why it went unnoticed.
|
||||
|
||||
The Makefile therefore picks a gencode list from `CUDA_MAJOR_VERSION` (a build
|
||||
arg the backend matrix already declares, forwarded by `Dockerfile.ds4`) and
|
||||
`uname -m`, and passes it as `NVCC_ARCH_FLAGS` to the sub-make. Upstream's
|
||||
`CUDA_ARCH` accepts a single value, so it cannot express the fat binary the
|
||||
shipped images need; a command-line assignment beats its `:=`. An empty
|
||||
`CUDA_MAJOR_VERSION` falls back to upstream's `native` for local developer
|
||||
builds, and an unrecognised one is a hard error - no CI runner has a GPU, so a
|
||||
silent `native` there is exactly the failure mode this guards against.
|
||||
|
||||
`DS4_CUDA_HAVE_MXF4` is deliberately unset: upstream defines it only for
|
||||
single-arch sm_120/sm_121 builds and guards it with a plain `#ifdef` rather than
|
||||
`__CUDA_ARCH__`, so it cannot be combined with older archs. It gates an optional
|
||||
MXFP4 indexer fast path whose `#ifndef` branch returns 0, so omitting it costs
|
||||
speed, not correctness.
|
||||
|
||||
### Verifying a build
|
||||
|
||||
Check which flags a configuration resolves to, without compiling anything:
|
||||
|
||||
```
|
||||
make -C backend/cpp/ds4 BUILD_TYPE=cublas CUDA_MAJOR_VERSION=13 NATIVE=false \
|
||||
--eval='show: ; @echo [$(DS4_ARCH_MAKEVARS)]' show
|
||||
```
|
||||
|
||||
Do not use `make -n` for this: the recipe is `+$(MAKE) ...`, and the `+` prefix
|
||||
makes it run even under `-n`.
|
||||
|
||||
Then exercise the failure mode itself against a built backend. It only appears
|
||||
above one prefill batch, so the ordinary `predict` spec cannot catch it:
|
||||
|
||||
```
|
||||
BACKEND_BINARY=$(pwd)/backend/cpp/ds4/package/run.sh \
|
||||
BACKEND_TEST_MODEL_FILE=/path/to/ds4flash.gguf \
|
||||
BACKEND_TEST_CAPS=health,load,predict,long_prefill \
|
||||
go test -count=1 -timeout=30m -v ./tests/e2e-backends/...
|
||||
```
|
||||
|
||||
## Build matrix
|
||||
|
||||
| Build | Where | Notes |
|
||||
|
||||
@@ -5,7 +5,7 @@ This PR fixes #
|
||||
**Notes for Reviewers**
|
||||
|
||||
|
||||
**[Signed commits](../CONTRIBUTING.md#signing-off-on-commits-developer-certificate-of-origin)**
|
||||
**[Signed commits](../CONTRIBUTING.md#commit-messages)**
|
||||
- [ ] Yes, I signed my commits.
|
||||
- [ ] Documentation updated (docs/content/) for user-facing changes, or not applicable
|
||||
|
||||
|
||||
@@ -3754,6 +3754,19 @@ include:
|
||||
dockerfile: "./backend/Dockerfile.golang"
|
||||
context: "./"
|
||||
ubuntu-version: '2404'
|
||||
- build-type: 'hipblas'
|
||||
cuda-major-version: ""
|
||||
cuda-minor-version: ""
|
||||
platforms: 'linux/amd64'
|
||||
tag-latest: 'auto'
|
||||
tag-suffix: '-gpu-rocm-hipblas-stablediffusion-ggml'
|
||||
runs-on: 'ubuntu-latest'
|
||||
base-image: "rocm/dev-ubuntu-24.04:7.2.1"
|
||||
skip-drivers: 'false'
|
||||
backend: "stablediffusion-ggml"
|
||||
dockerfile: "./backend/Dockerfile.golang"
|
||||
context: "./"
|
||||
ubuntu-version: '2404'
|
||||
- build-type: 'sycl_f16'
|
||||
cuda-major-version: ""
|
||||
cuda-minor-version: ""
|
||||
|
||||
@@ -3,9 +3,9 @@
|
||||
# darwin (Apple Silicon) install path. The macOS/Metal build
|
||||
# (backend/python/vllm/install.sh, Darwin branch) installs vllm-metal, which is
|
||||
# version-locked to a specific vLLM source release. install.sh derives that vLLM
|
||||
# version at build time from vllm-metal's own installer at the pinned
|
||||
# tag, so there is only ONE value to bump here -- mirroring bump_vllm_wheel.sh,
|
||||
# which bumps the Linux cu130 wheel pin.
|
||||
# version, and the wheel asset name, at build time from the pinned tag, so there
|
||||
# is only ONE value to bump here -- mirroring bump_vllm_wheel.sh, which bumps the
|
||||
# Linux cu130 wheel pin.
|
||||
#
|
||||
# This deliberately tracks vllm-project/vllm-metal, NOT vllm-project/vllm: the
|
||||
# darwin build can only use the exact vLLM version vllm-metal supports, so it may
|
||||
@@ -23,15 +23,20 @@ if [ -z "$FILE" ] || [ -z "$REPO" ] || [ -z "$VAR" ]; then
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# vllm-metal ships frequent dev releases, all flagged as non-prerelease, so
|
||||
# /releases/latest returns the newest one (with its cp312 wheel asset).
|
||||
# vllm-metal ships frequent .dev releases, flagged as prereleases, alongside the
|
||||
# stable ones. /releases/latest skips the prereleases and returns the newest
|
||||
# stable tag, which is what darwin should pin: upstream deletes and re-cuts .dev
|
||||
# tags, and a pin to a deleted tag 404s the whole build.
|
||||
LATEST_TAG=$(gh_curl -H "Accept: application/vnd.github+json" \
|
||||
"https://api.github.com/repos/$REPO/releases/latest" \
|
||||
| python3 -c "import json,sys; print(json.load(sys.stdin)['tag_name'])")
|
||||
|
||||
# The coupled vLLM source version lives in vllm-metal's installer at that tag.
|
||||
NEW_VLLM_VERSION=$(gh_curl \
|
||||
"https://raw.githubusercontent.com/$REPO/$LATEST_TAG/install.sh" \
|
||||
# The coupled vLLM release lives in .github/vllm-release-tag.commit at that tag
|
||||
# (since vllm-metal 0.28); releases predating that file pinned it inline in their
|
||||
# own install.sh. The extractor reads both forms.
|
||||
NEW_VLLM_VERSION=$( { gh_curl \
|
||||
"https://raw.githubusercontent.com/$REPO/$LATEST_TAG/.github/vllm-release-tag.commit" \
|
||||
|| gh_curl "https://raw.githubusercontent.com/$REPO/$LATEST_TAG/install.sh"; } \
|
||||
| "$(dirname "${BASH_SOURCE[0]}")/../scripts/lib/extract-vllm-metal-version.sh")
|
||||
|
||||
if [ -z "$LATEST_TAG" ] || [ -z "$NEW_VLLM_VERSION" ]; then
|
||||
|
||||
+1
-65
@@ -29,10 +29,6 @@ updates:
|
||||
schedule:
|
||||
# Check for updates to GitHub Actions every weekday
|
||||
interval: "weekly"
|
||||
- package-ecosystem: "pip"
|
||||
directory: "/backend/python/bark"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
- package-ecosystem: "pip"
|
||||
directory: "/backend/python/common/template"
|
||||
schedule:
|
||||
@@ -55,30 +51,10 @@ updates:
|
||||
ignore:
|
||||
- dependency-name: "torch"
|
||||
- dependency-name: "transformers"
|
||||
- package-ecosystem: "pip"
|
||||
directory: "/backend/python/exllama"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
- package-ecosystem: "pip"
|
||||
directory: "/backend/python/exllama2"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
- package-ecosystem: "pip"
|
||||
directory: "/backend/python/mamba"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
- package-ecosystem: "pip"
|
||||
directory: "/backend/python/openvoice"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
- package-ecosystem: "pip"
|
||||
directory: "/backend/python/rerankers"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
- package-ecosystem: "pip"
|
||||
directory: "/backend/python/sentencetransformers"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
- package-ecosystem: "pip"
|
||||
directory: "/backend/python/transformers"
|
||||
schedule:
|
||||
@@ -86,44 +62,4 @@ updates:
|
||||
- package-ecosystem: "pip"
|
||||
directory: "/backend/python/vllm"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
- package-ecosystem: "pip"
|
||||
directory: "/examples/chainlit"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
- package-ecosystem: "pip"
|
||||
directory: "/examples/functions"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
- package-ecosystem: "pip"
|
||||
directory: "/examples/langchain/langchainpy-localai-example"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
- package-ecosystem: "pip"
|
||||
directory: "/examples/langchain-chroma"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
- package-ecosystem: "pip"
|
||||
directory: "/examples/streamlit-bot"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
- package-ecosystem: "docker"
|
||||
directory: "/examples/k8sgpt"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
- package-ecosystem: "docker"
|
||||
directory: "/examples/kubernetes"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
- package-ecosystem: "docker"
|
||||
directory: "/examples/langchain"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
- package-ecosystem: "gomod"
|
||||
directory: "/examples/semantic-todo"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
- package-ecosystem: "docker"
|
||||
directory: "/examples/telegram-bot"
|
||||
schedule:
|
||||
interval: "weekly"
|
||||
interval: "weekly"
|
||||
@@ -166,7 +166,7 @@ jobs:
|
||||
push-to-fork: ci-forks/LocalAI
|
||||
commit-message: ':arrow_up: Update ${{ matrix.repository }}'
|
||||
title: 'chore: :arrow_up: Update ${{ matrix.repository }} to `${{ steps.bump.outputs.commit }}`'
|
||||
branch: "update/${{ matrix.variable }}"
|
||||
branch: "bump/${{ matrix.variable }}"
|
||||
body: ${{ steps.bump.outputs.message }}
|
||||
signoff: true
|
||||
|
||||
@@ -203,7 +203,7 @@ jobs:
|
||||
push-to-fork: ci-forks/LocalAI
|
||||
commit-message: ':arrow_up: Update vllm-project/vllm cu130 wheel'
|
||||
title: 'chore: :arrow_up: Update vllm-project/vllm cu130 wheel to `${{ steps.bump.outputs.commit }}`'
|
||||
branch: "update/VLLM_VERSION"
|
||||
branch: "bump/VLLM_VERSION"
|
||||
body: ${{ steps.bump.outputs.message }}
|
||||
signoff: true
|
||||
|
||||
@@ -241,6 +241,6 @@ jobs:
|
||||
push-to-fork: ci-forks/LocalAI
|
||||
commit-message: ':arrow_up: Update vllm-project/vllm-metal (darwin)'
|
||||
title: 'chore: :arrow_up: Update vllm-metal (darwin) to `${{ steps.bump.outputs.commit }}`'
|
||||
branch: "update/VLLM_METAL_VERSION"
|
||||
branch: "bump/VLLM_METAL_VERSION"
|
||||
body: ${{ steps.bump.outputs.message }}
|
||||
signoff: true
|
||||
@@ -31,13 +31,14 @@ jobs:
|
||||
messages: [
|
||||
{
|
||||
role: "system",
|
||||
content: "Write a discord message with a bullet point summary of the release notes."
|
||||
content: "Write a Discord message with a bullet point summary of the release notes. Keep the complete message under 1800 characters."
|
||||
},
|
||||
{
|
||||
role: "user",
|
||||
content: $input
|
||||
}
|
||||
]
|
||||
],
|
||||
max_tokens: 450
|
||||
}')
|
||||
|
||||
# Send the request to LocalAI API
|
||||
@@ -46,7 +47,7 @@ jobs:
|
||||
-d "$json_payload")
|
||||
|
||||
# Extract the summary from the response
|
||||
summary=$(echo $response | jq -r '.choices[0].message.content')
|
||||
summary=$(printf '%s' "$response" | jq -er '.choices[0].message.content | strings | .[0:1800]')
|
||||
|
||||
# Print the summary
|
||||
# -H "Authorization: Bearer $API_KEY" \
|
||||
|
||||
@@ -80,8 +80,13 @@ jobs:
|
||||
coverage/coverage.out
|
||||
coverage/coverage.html
|
||||
if-no-files-found: ignore
|
||||
# tmate keeps the runner busy until the 6 hour job limit, so a single
|
||||
# failure costs a whole runner slot. Only open a session when someone
|
||||
# asked for one by labelling the pull request `ci-debug`, and cap the
|
||||
# session so a forgotten label cannot idle a runner either.
|
||||
- name: Setup tmate session if tests fail
|
||||
if: ${{ failure() }}
|
||||
if: ${{ failure() && contains(github.event.pull_request.labels.*.name, 'ci-debug') }}
|
||||
timeout-minutes: 30
|
||||
uses: mxschmitt/action-tmate@v3.23
|
||||
with:
|
||||
detached: true
|
||||
@@ -125,8 +130,13 @@ jobs:
|
||||
export PATH="/opt/homebrew/opt/make/libexec/gnubin:$PATH"
|
||||
PATH="$PATH:$HOME/go/bin" make protogen-go
|
||||
PATH="$PATH:$HOME/go/bin" BUILD_TYPE="GITHUB_CI_HAS_BROKEN_METAL" CMAKE_ARGS="-DGGML_F16C=OFF -DGGML_AVX512=OFF -DGGML_AVX2=OFF -DGGML_FMA=OFF" make --jobs 4 --output-sync=target test
|
||||
# tmate keeps the runner busy until the 6 hour job limit, so a single
|
||||
# failure costs a whole runner slot. Only open a session when someone
|
||||
# asked for one by labelling the pull request `ci-debug`, and cap the
|
||||
# session so a forgotten label cannot idle a runner either.
|
||||
- name: Setup tmate session if tests fail
|
||||
if: ${{ failure() }}
|
||||
if: ${{ failure() && contains(github.event.pull_request.labels.*.name, 'ci-debug') }}
|
||||
timeout-minutes: 30
|
||||
uses: mxschmitt/action-tmate@v3.23
|
||||
with:
|
||||
detached: true
|
||||
|
||||
@@ -77,8 +77,13 @@ jobs:
|
||||
- name: Test
|
||||
run: |
|
||||
PATH="$PATH:$HOME/go/bin" make backends/local-store backends/silero-vad backends/llama-cpp backends/whisper backends/piper backends/stablediffusion-ggml docker-build-e2e e2e-aio
|
||||
# tmate keeps the runner busy until the 6 hour job limit, so a single
|
||||
# failure costs a whole runner slot. Only open a session when someone
|
||||
# asked for one by labelling the pull request `ci-debug`, and cap the
|
||||
# session so a forgotten label cannot idle a runner either.
|
||||
- name: Setup tmate session if tests fail
|
||||
if: ${{ failure() }}
|
||||
if: ${{ failure() && contains(github.event.pull_request.labels.*.name, 'ci-debug') }}
|
||||
timeout-minutes: 30
|
||||
uses: mxschmitt/action-tmate@v3.23
|
||||
with:
|
||||
detached: true
|
||||
|
||||
@@ -63,8 +63,13 @@ jobs:
|
||||
- name: Test Backend E2E
|
||||
run: |
|
||||
PATH="$PATH:$HOME/go/bin" make build-mock-backend test-e2e
|
||||
# tmate keeps the runner busy until the 6 hour job limit, so a single
|
||||
# failure costs a whole runner slot. Only open a session when someone
|
||||
# asked for one by labelling the pull request `ci-debug`, and cap the
|
||||
# session so a forgotten label cannot idle a runner either.
|
||||
- name: Setup tmate session if tests fail
|
||||
if: ${{ failure() }}
|
||||
if: ${{ failure() && contains(github.event.pull_request.labels.*.name, 'ci-debug') }}
|
||||
timeout-minutes: 30
|
||||
uses: mxschmitt/action-tmate@v3.23
|
||||
with:
|
||||
detached: true
|
||||
|
||||
@@ -88,8 +88,13 @@ jobs:
|
||||
# CPU and runs the token_classify capability spec (byte-offset contract).
|
||||
- name: Run live PII NER backend E2E
|
||||
run: PATH="$PATH:$HOME/go/bin" make test-extra-backend-privacy-filter
|
||||
# tmate keeps the runner busy until the 6 hour job limit, so a single
|
||||
# failure costs a whole runner slot. Only open a session when someone
|
||||
# asked for one by labelling the pull request `ci-debug`, and cap the
|
||||
# session so a forgotten label cannot idle a runner either.
|
||||
- name: Setup tmate session if tests fail
|
||||
if: ${{ failure() }}
|
||||
if: ${{ failure() && contains(github.event.pull_request.labels.*.name, 'ci-debug') }}
|
||||
timeout-minutes: 30
|
||||
uses: mxschmitt/action-tmate@v3.23
|
||||
with:
|
||||
detached: true
|
||||
|
||||
@@ -75,8 +75,13 @@ jobs:
|
||||
path: core/http/react-ui/coverage/
|
||||
if-no-files-found: ignore
|
||||
retention-days: 7
|
||||
# tmate keeps the runner busy until the 6 hour job limit, so a single
|
||||
# failure costs a whole runner slot. Only open a session when someone
|
||||
# asked for one by labelling the pull request `ci-debug`, and cap the
|
||||
# session so a forgotten label cannot idle a runner either.
|
||||
- name: Setup tmate session if tests fail
|
||||
if: ${{ failure() }}
|
||||
if: ${{ failure() && contains(github.event.pull_request.labels.*.name, 'ci-debug') }}
|
||||
timeout-minutes: 30
|
||||
uses: mxschmitt/action-tmate@v3.23
|
||||
with:
|
||||
detached: true
|
||||
|
||||
@@ -8,7 +8,7 @@ Human contributors: see [CONTRIBUTING.md](CONTRIBUTING.md) for the development w
|
||||
|
||||
LocalAI follows the Linux kernel project's [guidelines for AI coding assistants](https://docs.kernel.org/process/coding-assistants.html). Before submitting AI-assisted code, read [.agents/ai-coding-assistants.md](.agents/ai-coding-assistants.md). Key rules:
|
||||
|
||||
- **No `Signed-off-by` from AI.** Only the human submitter may sign off on the Developer Certificate of Origin.
|
||||
- **No `Signed-off-by` from AI.** Only the human submitter may sign off on the Developer Certificate of Origin. One exception: automation a maintainer operates signs off with *that maintainer's* identity, since no other human submitter exists to certify it. See [.agents/ai-coding-assistants.md](.agents/ai-coding-assistants.md).
|
||||
- **No `Co-Authored-By: <AI>` trailers.** The human contributor owns the change.
|
||||
- **Use an `Assisted-by:` trailer** to attribute AI involvement. Format: `Assisted-by: AGENT_NAME:MODEL_VERSION [TOOL1] [TOOL2]`.
|
||||
- **The human submitter is responsible** for reviewing, testing, and understanding every line of generated code.
|
||||
|
||||
+1
-1
@@ -218,7 +218,7 @@ LocalAI follows the **same guidelines as the Linux kernel project** for AI-assis
|
||||
|
||||
The full policy for this repository lives in [`.agents/ai-coding-assistants.md`](.agents/ai-coding-assistants.md). Summary:
|
||||
|
||||
- **AI agents MUST NOT add `Signed-off-by` tags.** Only humans can certify the Developer Certificate of Origin.
|
||||
- **AI agents MUST NOT add `Signed-off-by` tags.** Only humans can certify the Developer Certificate of Origin. Automation operated by a maintainer is the one exception: it signs off with that maintainer's identity, because there is no other human submitter to certify it.
|
||||
- **AI agents MUST NOT add `Co-Authored-By` trailers** attributing themselves as co-authors.
|
||||
- **Attribute AI involvement with an `Assisted-by` trailer** in the commit message:
|
||||
|
||||
|
||||
@@ -34,6 +34,11 @@ TEST_FLAKES?=5
|
||||
RANDOM := $(shell bash -c 'echo $$RANDOM')
|
||||
|
||||
VERSION?=$(shell git describe --always --tags || echo "dev" )
|
||||
# fyne package only accepts numeric x[.y[.z]] app versions, so reduce git
|
||||
# describe output (v4.9.0, v4.9.0-14-gabc1234, or a bare sha on untagged
|
||||
# checkouts) to its numeric core; anything non-numeric falls back to 0.0.0.
|
||||
# Without this the packaged launcher reports itself as version 0.0.0 (#11673).
|
||||
LAUNCHER_APP_VERSION?=$(shell v=$$(echo "$(VERSION)" | sed -E 's/^v//; s/[+-].*$$//'); echo "$$v" | grep -qE '^[0-9]+(\.[0-9]+){0,2}$$' && echo "$$v" || echo "0.0.0")
|
||||
# go tool nm ./local-ai | grep Commit
|
||||
LD_FLAGS?=-s -w
|
||||
override LD_FLAGS += -X "github.com/mudler/LocalAI/internal.Version=$(VERSION)"
|
||||
@@ -388,9 +393,17 @@ test-e2e: build-mock-backend build-cloud-proxy-backend prepare-e2e run-e2e-image
|
||||
$(MAKE) teardown-e2e
|
||||
docker rmi localai-tests
|
||||
|
||||
# `docker stop` returns as soon as the container exits, but Docker reaps a
|
||||
# `--rm` container asynchronously after that. The `docker rmi localai-tests` in
|
||||
# test-e2e then loses the race against the reaper and fails on a still
|
||||
# referenced image, turning a green suite red. Removing the container ourselves
|
||||
# is synchronous, so the image reference is gone before we return. It also
|
||||
# covers the case where nothing is running, which `docker stop` could not
|
||||
# because it rejects an empty argument list.
|
||||
teardown-e2e:
|
||||
rm -rf $(TEST_DIR) || true
|
||||
docker stop $$(docker ps -q --filter ancestor=localai-tests)
|
||||
@CONTAINERS=$$(docker ps -aq --filter ancestor=localai-tests 2>/dev/null); \
|
||||
if [ -n "$$CONTAINERS" ]; then docker rm -f $$CONTAINERS || true; fi
|
||||
|
||||
########################################################
|
||||
## Integration and unit tests
|
||||
@@ -1622,7 +1635,7 @@ site-serve: site
|
||||
build-launcher-darwin:
|
||||
rm -rf dist/LocalAI.app cmd/launcher/LocalAI.app
|
||||
mkdir -p dist
|
||||
cd cmd/launcher && go run fyne.io/tools/cmd/fyne@latest package -os darwin -icon ../../core/http/static/logo.png --executable $(LAUNCHER_BINARY_NAME)
|
||||
cd cmd/launcher && go run fyne.io/tools/cmd/fyne@latest package -os darwin -icon ../../core/http/static/logo.png --executable $(LAUNCHER_BINARY_NAME) --app-version $(LAUNCHER_APP_VERSION)
|
||||
mv cmd/launcher/LocalAI.app dist/LocalAI.app
|
||||
bash contrib/macos/sign-and-notarize.sh sign dist/LocalAI.app
|
||||
|
||||
@@ -1649,4 +1662,4 @@ release-launcher-darwin: notarize-launcher-darwin
|
||||
@echo "dist/LocalAI.dmg is ready"
|
||||
|
||||
build-launcher-linux:
|
||||
cd cmd/launcher && go run fyne.io/tools/cmd/fyne@latest package -os linux -icon ../../core/http/static/logo.png --executable $(LAUNCHER_BINARY_NAME)-linux && mv LocalAI.tar.xz ../../$(LAUNCHER_BINARY_NAME)-linux.tar.xz
|
||||
cd cmd/launcher && go run fyne.io/tools/cmd/fyne@latest package -os linux -icon ../../core/http/static/logo.png --executable $(LAUNCHER_BINARY_NAME)-linux --app-version $(LAUNCHER_APP_VERSION) && mv LocalAI.tar.xz ../../$(LAUNCHER_BINARY_NAME)-linux.tar.xz
|
||||
@@ -10,6 +10,7 @@ FROM ${BASE_IMAGE} AS builder
|
||||
ARG BUILD_TYPE
|
||||
ARG TARGETARCH
|
||||
ARG TARGETVARIANT
|
||||
ARG CUDA_MAJOR_VERSION
|
||||
|
||||
ENV BUILD_TYPE=${BUILD_TYPE} \
|
||||
DEBIAN_FRONTEND=noninteractive \
|
||||
@@ -35,7 +36,8 @@ RUN apt-get update && \
|
||||
COPY . /LocalAI
|
||||
|
||||
RUN --mount=type=cache,target=/root/.ccache,id=ds4-ccache-${TARGETARCH}-${BUILD_TYPE},sharing=locked \
|
||||
make -C /LocalAI/backend/cpp/ds4 BUILD_TYPE=${BUILD_TYPE} NATIVE=false grpc-server package
|
||||
make -C /LocalAI/backend/cpp/ds4 BUILD_TYPE=${BUILD_TYPE} \
|
||||
CUDA_MAJOR_VERSION=${CUDA_MAJOR_VERSION} NATIVE=false grpc-server package
|
||||
|
||||
FROM scratch
|
||||
COPY --from=builder /LocalAI/backend/cpp/ds4/package/. ./
|
||||
@@ -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?=43001a7e0f452d80f4588e613f13332940dd4d3a
|
||||
AUDIO_CPP_VERSION?=9c6a282337cc83f227cc10428867a478947706ad
|
||||
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?=312bb2a93ea2bf798333fa859614fbf913ecb9e2
|
||||
LLAMA_REPO?=https://github.com/PrismML-Eng/llama.cpp
|
||||
|
||||
CMAKE_ARGS?=
|
||||
@@ -41,6 +41,7 @@ define bonsai-build
|
||||
# and are applied by apply-patches.sh below.
|
||||
rm -rf $(CURRENT_MAKEFILE_DIR)/../bonsai-$(1)-build/patches
|
||||
$(MAKE) -C $(CURRENT_MAKEFILE_DIR)/../bonsai-$(1)-build purge
|
||||
bash $(CURRENT_MAKEFILE_DIR)/patch-grpc-server.sh $(CURRENT_MAKEFILE_DIR)/../bonsai-$(1)-build/grpc-server.cpp
|
||||
bash $(LLAMA_CPP_DIR)/disable-score-task.sh $(CURRENT_MAKEFILE_DIR)/../bonsai-$(1)-build/grpc-server.cpp
|
||||
bash $(LLAMA_CPP_DIR)/disable-tts-task.sh $(CURRENT_MAKEFILE_DIR)/../bonsai-$(1)-build/grpc-server.cpp
|
||||
$(info $(GREEN)I bonsai build info:$(1)$(RESET))
|
||||
@@ -79,6 +80,7 @@ bonsai-cpu-all:
|
||||
# and are applied by apply-patches.sh below.
|
||||
rm -rf $(CURRENT_MAKEFILE_DIR)/../bonsai-cpu-all-build/patches
|
||||
$(MAKE) -C $(CURRENT_MAKEFILE_DIR)/../bonsai-cpu-all-build purge
|
||||
bash $(CURRENT_MAKEFILE_DIR)/patch-grpc-server.sh $(CURRENT_MAKEFILE_DIR)/../bonsai-cpu-all-build/grpc-server.cpp
|
||||
bash $(LLAMA_CPP_DIR)/disable-score-task.sh $(CURRENT_MAKEFILE_DIR)/../bonsai-cpu-all-build/grpc-server.cpp
|
||||
bash $(LLAMA_CPP_DIR)/disable-tts-task.sh $(CURRENT_MAKEFILE_DIR)/../bonsai-cpu-all-build/grpc-server.cpp
|
||||
$(info $(GREEN)I bonsai build info:cpu-all-variants$(RESET))
|
||||
|
||||
@@ -0,0 +1,24 @@
|
||||
#!/bin/bash
|
||||
# Adapt the shared llama.cpp gRPC source to the older JSON API in Bonsai.
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
if [[ $# -ne 1 ]]; then
|
||||
echo "usage: $0 <grpc-server.cpp>" >&2
|
||||
exit 2
|
||||
fi
|
||||
|
||||
SRC=$1
|
||||
if [[ ! -f "$SRC" ]]; then
|
||||
echo "grpc-server.cpp not found at $SRC" >&2
|
||||
exit 2
|
||||
fi
|
||||
|
||||
if grep -q 'common_json_error' "$SRC"; then
|
||||
echo "==> patching $SRC to use the Bonsai JSON exception type"
|
||||
awk '{ gsub(/common_json_error/, "json::parse_error"); print }' "$SRC" > "$SRC.tmp"
|
||||
mv "$SRC.tmp" "$SRC"
|
||||
echo "==> Bonsai JSON exception patch OK"
|
||||
else
|
||||
echo "==> $SRC already uses a Bonsai-compatible JSON exception type, skipping"
|
||||
fi
|
||||
@@ -84,9 +84,10 @@ elseif(DS4_GPU STREQUAL "cpu")
|
||||
set(DS4_OBJS "${DS4_DIR}/ds4_cpu.o")
|
||||
endif()
|
||||
|
||||
# Upstream splits distributed inference, tensor-parallel transport, the SSD
|
||||
# expert cache, and layer placement into GPU-agnostic translation units. Link
|
||||
# them regardless of DS4_GPU.
|
||||
# Upstream splits image preprocessing, distributed inference, tensor-parallel
|
||||
# transport, the SSD expert cache, and layer placement into GPU-agnostic
|
||||
# translation units. Link them regardless of DS4_GPU.
|
||||
list(APPEND DS4_OBJS "${DS4_DIR}/ds4_image.o")
|
||||
list(APPEND DS4_OBJS "${DS4_DIR}/ds4_distributed.o")
|
||||
list(APPEND DS4_OBJS "${DS4_DIR}/ds4_tp.o")
|
||||
list(APPEND DS4_OBJS "${DS4_DIR}/ds4_ssd.o")
|
||||
|
||||
+73
-11
@@ -1,10 +1,10 @@
|
||||
# ds4 backend Makefile.
|
||||
#
|
||||
# Upstream pin lives below as DS4_VERSION?=84cc882352757baf628a1776badf7cc54d584e28
|
||||
# Upstream pin lives below as DS4_VERSION?=f62ca29a308724cde5bc99134ede19104b2a3260
|
||||
# (.github/bump_deps.sh) can find and update it - matches the
|
||||
# llama-cpp / ik-llama-cpp / turboquant convention.
|
||||
|
||||
DS4_VERSION?=84cc882352757baf628a1776badf7cc54d584e28
|
||||
DS4_VERSION?=f62ca29a308724cde5bc99134ede19104b2a3260
|
||||
DS4_REPO?=https://github.com/antirez/ds4
|
||||
|
||||
CURRENT_MAKEFILE_DIR := $(dir $(abspath $(lastword $(MAKEFILE_LIST))))
|
||||
@@ -18,21 +18,83 @@ UNAME_S := $(shell uname -s)
|
||||
|
||||
CMAKE_ARGS ?= -DCMAKE_BUILD_TYPE=Release
|
||||
|
||||
# Upstream splits distributed inference, tensor-parallel transport, the SSD
|
||||
# expert cache, and layer placement into GPU-agnostic translation units. They
|
||||
# are shared by every GPU mode, so append them unconditionally below.
|
||||
# nvcc must be told the target architecture explicitly for a cublas build, and
|
||||
# this is not a tuning knob. Upstream's Makefile leaves CUDA_ARCH empty and its
|
||||
# `cuda` target REFUSES to build without one, offering `cuda-spark`
|
||||
# (CUDA_ARCH=sm_121) and `cuda-generic` (CUDA_ARCH=native) instead. We drive its
|
||||
# object targets directly, which bypasses that guard: nvcc then compiles with no
|
||||
# -arch at all, and the kernels run as JIT'd PTX for its default architecture.
|
||||
# On GB10 (sm_121) that silently produced corrupt inference output above a
|
||||
# ~128-token prefill batch and ~77x slower prefill (4.21 t/s vs 325.70 t/s,
|
||||
# measured on the same box with the same model). No CI runner has a GPU, so
|
||||
# `native` has nothing to enumerate there.
|
||||
#
|
||||
# Upstream's CUDA_ARCH takes a SINGLE value (see its sm_120/sm_121 special cases
|
||||
# and the `-arch=$(CUDA_ARCH)` fallback), so it cannot express the fat binary
|
||||
# these images need. NVCC_ARCH_FLAGS is overridden instead: a command-line
|
||||
# assignment wins over the `:=` in upstream's Makefile, and its NVCCFLAGS
|
||||
# expands whatever we pass.
|
||||
#
|
||||
# The architecture lists are copied from backend/go/vllm-cpp/Makefile rather
|
||||
# than invented, so the two CUDA images cover the same GPUs: amd64 datacenter +
|
||||
# consumer, and l4t/arm64 covering Orin (87), Thor (110) and GB10 (121a).
|
||||
#
|
||||
# -DDS4_CUDA_HAVE_MXF4=1 is deliberately NOT set. Upstream only defines it for
|
||||
# single-arch sm_120/sm_121 builds and guards the code with a plain #ifdef
|
||||
# rather than __CUDA_ARCH__, so it cannot be combined with older archs in one
|
||||
# fat binary. It gates an optional MXFP4 indexer fast path whose #ifndef branch
|
||||
# returns 0 and falls back to the generic path, so omitting it costs some speed
|
||||
# on GB10, not correctness. Revisit if upstream adds __CUDA_ARCH__ guards.
|
||||
#
|
||||
# An EMPTY CUDA_MAJOR_VERSION means a local developer build, not CI: fall back
|
||||
# to upstream's own `native` handling, which needs a GPU present but is what a
|
||||
# developer building on their own machine wants. Both variables are `?=` so an
|
||||
# explicit value on the command line always wins.
|
||||
UNAME_M := $(shell uname -m)
|
||||
CUDA_MAJOR_VERSION ?=
|
||||
ifeq ($(BUILD_TYPE),cublas)
|
||||
ifeq ($(CUDA_MAJOR_VERSION),13)
|
||||
ifeq ($(UNAME_M),aarch64)
|
||||
DS4_NVCC_ARCH_FLAGS ?= -gencode arch=compute_87,code=sm_87 \
|
||||
-gencode arch=compute_90a,code=sm_90a \
|
||||
-gencode arch=compute_100a,code=sm_100a \
|
||||
-gencode arch=compute_110,code=sm_110 \
|
||||
-gencode arch=compute_121a,code=sm_121a
|
||||
else
|
||||
DS4_NVCC_ARCH_FLAGS ?= -gencode arch=compute_80,code=sm_80 \
|
||||
-gencode arch=compute_86,code=sm_86 \
|
||||
-gencode arch=compute_89,code=sm_89 \
|
||||
-gencode arch=compute_90a,code=sm_90a \
|
||||
-gencode arch=compute_100a,code=sm_100a \
|
||||
-gencode arch=compute_103a,code=sm_103a \
|
||||
-gencode arch=compute_120a,code=sm_120a \
|
||||
-gencode arch=compute_121a,code=sm_121a
|
||||
endif
|
||||
DS4_ARCH_MAKEVARS := NVCC_ARCH_FLAGS="$(DS4_NVCC_ARCH_FLAGS)"
|
||||
else ifeq ($(CUDA_MAJOR_VERSION),)
|
||||
# Local build: let upstream resolve the host GPU.
|
||||
DS4_ARCH_MAKEVARS := CUDA_ARCH=native
|
||||
else
|
||||
$(error CUDA_MAJOR_VERSION=$(CUDA_MAJOR_VERSION) has no architecture list here (13 does). Leave it empty for a native build, or pass DS4_NVCC_ARCH_FLAGS explicitly.)
|
||||
endif
|
||||
endif
|
||||
|
||||
# Upstream splits image preprocessing, distributed inference, tensor-parallel
|
||||
# transport, the SSD expert cache, and layer placement into GPU-agnostic
|
||||
# translation units. They are shared by every GPU mode, so append them
|
||||
# unconditionally below.
|
||||
ifeq ($(BUILD_TYPE),cublas)
|
||||
CMAKE_ARGS += -DDS4_GPU=cuda
|
||||
DS4_OBJ_TARGET := ds4.o ds4_cuda.o ds4_distributed.o ds4_tp.o ds4_ssd.o ds4_layer_pack.o \
|
||||
DS4_OBJ_TARGET := ds4.o ds4_image.o ds4_cuda.o ds4_distributed.o ds4_tp.o ds4_ssd.o ds4_layer_pack.o \
|
||||
cuda/mmq/ds4_ggml_stubs.o cuda/mmq/ds4_mmq.o cuda/mmq/ds4_mmq_d2r.o \
|
||||
cuda/mmq/quantize.o cuda/mmq/mmid.o cuda/mmq/mmvq.o cuda/mmq/ds4_repack.o
|
||||
else ifeq ($(UNAME_S),Darwin)
|
||||
CMAKE_ARGS += -DDS4_GPU=metal
|
||||
DS4_OBJ_TARGET := ds4.o ds4_metal.o ds4_distributed.o ds4_tp.o ds4_ssd.o ds4_layer_pack.o
|
||||
DS4_OBJ_TARGET := ds4.o ds4_image.o ds4_metal.o ds4_distributed.o ds4_tp.o ds4_ssd.o ds4_layer_pack.o
|
||||
else
|
||||
# CPU reference path (Linux only - macOS CPU path is broken by VM bug per ds4 README).
|
||||
CMAKE_ARGS += -DDS4_GPU=cpu
|
||||
DS4_OBJ_TARGET := ds4_cpu.o ds4_distributed.o ds4_tp.o ds4_ssd.o ds4_layer_pack.o
|
||||
DS4_OBJ_TARGET := ds4_cpu.o ds4_image.o ds4_distributed.o ds4_tp.o ds4_ssd.o ds4_layer_pack.o
|
||||
endif
|
||||
|
||||
ifneq ($(NATIVE),true)
|
||||
@@ -57,11 +119,11 @@ ds4:
|
||||
# the right per-platform compile flags (Objective-C/Metal on Darwin, nvcc on Linux+CUDA).
|
||||
ds4/ds4.o: ds4
|
||||
ifeq ($(BUILD_TYPE),cublas)
|
||||
+$(MAKE) -C ds4 $(DS4_OBJ_TARGET)
|
||||
+$(MAKE) -C ds4 $(DS4_ARCH_MAKEVARS) $(DS4_OBJ_TARGET)
|
||||
else ifeq ($(UNAME_S),Darwin)
|
||||
+$(MAKE) -C ds4 ds4.o ds4_metal.o ds4_distributed.o ds4_tp.o ds4_ssd.o ds4_layer_pack.o
|
||||
+$(MAKE) -C ds4 ds4.o ds4_image.o ds4_metal.o ds4_distributed.o ds4_tp.o ds4_ssd.o ds4_layer_pack.o
|
||||
else
|
||||
+$(MAKE) -C ds4 ds4_cpu.o ds4_distributed.o ds4_tp.o ds4_ssd.o ds4_layer_pack.o
|
||||
+$(MAKE) -C ds4 ds4_cpu.o ds4_image.o ds4_distributed.o ds4_tp.o ds4_ssd.o ds4_layer_pack.o
|
||||
endif
|
||||
|
||||
grpc-server: ds4/ds4.o
|
||||
|
||||
@@ -92,7 +92,8 @@ std::string json_escape(const std::string &in) {
|
||||
|
||||
} // namespace
|
||||
|
||||
DsmlParser::DsmlParser() = default;
|
||||
DsmlParser::DsmlParser(bool starts_in_thinking)
|
||||
: state_(starts_in_thinking ? State::THINK : State::TEXT) {}
|
||||
|
||||
bool DsmlParser::IsInDsmlStructural() const {
|
||||
switch (state_) {
|
||||
|
||||
@@ -17,7 +17,9 @@ struct ParserEvent {
|
||||
// Streaming parser. Stateless across instances; one per Predict call.
|
||||
class DsmlParser {
|
||||
public:
|
||||
DsmlParser();
|
||||
// The chat prompt may already contain the opening thinking marker, so the
|
||||
// generated text can begin directly with reasoning bytes.
|
||||
explicit DsmlParser(bool starts_in_thinking = false);
|
||||
|
||||
// Feed a chunk of raw model-emitted text. Appends classified events to
|
||||
// `out`. May buffer the tail of `chunk` internally if it looks like a
|
||||
@@ -43,7 +45,7 @@ public:
|
||||
|
||||
private:
|
||||
enum class State { TEXT, THINK, TOOL_CALLS, INVOKE, PARAM_VALUE };
|
||||
State state_ = State::TEXT;
|
||||
State state_;
|
||||
std::string buf_;
|
||||
std::string current_tool_name_;
|
||||
int tool_index_ = -1;
|
||||
|
||||
@@ -0,0 +1,133 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
// Standalone regression tests for the DSML streaming parser.
|
||||
//
|
||||
// The repository's backend/cpp/run-unit-tests.sh harness compiles each
|
||||
// *_test.cpp as a single translation unit, so include the implementation here.
|
||||
|
||||
#include "dsml_parser.cpp"
|
||||
|
||||
#include <cstdio>
|
||||
#include <string>
|
||||
#include <type_traits>
|
||||
#include <vector>
|
||||
|
||||
namespace {
|
||||
|
||||
struct ParsedText {
|
||||
std::string content;
|
||||
std::string reasoning;
|
||||
};
|
||||
|
||||
int failures = 0;
|
||||
|
||||
void check_equal(const std::string &got, const std::string &want,
|
||||
const char *name) {
|
||||
if (got == want) return;
|
||||
std::fprintf(stderr, "FAIL %s: got \"%s\", want \"%s\"\n",
|
||||
name, got.c_str(), want.c_str());
|
||||
failures++;
|
||||
}
|
||||
|
||||
void collect_text(const std::vector<ds4cpp::ParserEvent> &events,
|
||||
ParsedText *parsed) {
|
||||
for (const auto &event : events) {
|
||||
if (event.type == ds4cpp::ParserEvent::CONTENT) {
|
||||
parsed->content += event.text;
|
||||
} else if (event.type == ds4cpp::ParserEvent::REASONING) {
|
||||
parsed->reasoning += event.text;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
ParsedText parse_chunks(ds4cpp::DsmlParser *parser,
|
||||
const std::vector<std::string> &chunks) {
|
||||
ParsedText parsed;
|
||||
for (const auto &chunk : chunks) {
|
||||
std::vector<ds4cpp::ParserEvent> events;
|
||||
parser->Feed(chunk, events);
|
||||
collect_text(events, &parsed);
|
||||
}
|
||||
std::vector<ds4cpp::ParserEvent> events;
|
||||
parser->Flush(events);
|
||||
collect_text(events, &parsed);
|
||||
return parsed;
|
||||
}
|
||||
|
||||
template <typename Parser>
|
||||
void test_reasoning_opened_by_prompt() {
|
||||
if constexpr (!std::is_constructible_v<Parser, bool>) {
|
||||
std::fprintf(stderr,
|
||||
"FAIL reasoning_opened_by_prompt: parser cannot start in thinking state\n");
|
||||
failures++;
|
||||
} else {
|
||||
Parser parser(true);
|
||||
ParsedText parsed = parse_chunks(
|
||||
&parser,
|
||||
{"We need to calculate factorial recursively.</think>Here is the answer."});
|
||||
check_equal(parsed.reasoning,
|
||||
"We need to calculate factorial recursively.",
|
||||
"reasoning_opened_by_prompt:reasoning");
|
||||
check_equal(parsed.content, "Here is the answer.",
|
||||
"reasoning_opened_by_prompt:content");
|
||||
}
|
||||
}
|
||||
|
||||
template <typename Parser>
|
||||
Parser text_parser() {
|
||||
if constexpr (std::is_constructible_v<Parser, bool>) {
|
||||
return Parser(false);
|
||||
} else {
|
||||
return Parser();
|
||||
}
|
||||
}
|
||||
|
||||
void test_reasoning_disabled() {
|
||||
auto parser = text_parser<ds4cpp::DsmlParser>();
|
||||
ParsedText parsed = parse_chunks(&parser, {"Here is the answer."});
|
||||
check_equal(parsed.reasoning, "", "reasoning_disabled:reasoning");
|
||||
check_equal(parsed.content, "Here is the answer.",
|
||||
"reasoning_disabled:content");
|
||||
}
|
||||
|
||||
void test_explicit_think_tag() {
|
||||
auto parser = text_parser<ds4cpp::DsmlParser>();
|
||||
ParsedText parsed = parse_chunks(
|
||||
&parser, {"<think>reasoning</think>answer"});
|
||||
check_equal(parsed.reasoning, "reasoning", "explicit_think_tag:reasoning");
|
||||
check_equal(parsed.content, "answer", "explicit_think_tag:content");
|
||||
}
|
||||
|
||||
template <typename Parser>
|
||||
void test_split_think_close_marker() {
|
||||
if constexpr (!std::is_constructible_v<Parser, bool>) {
|
||||
std::fprintf(stderr,
|
||||
"FAIL split_think_close_marker: parser cannot start in thinking state\n");
|
||||
failures++;
|
||||
} else {
|
||||
Parser parser(true);
|
||||
ParsedText parsed = parse_chunks(
|
||||
&parser,
|
||||
{"We need ", "to calculate ", "factorial", "</thi", "nk>",
|
||||
"Here is ", "the answer."});
|
||||
check_equal(parsed.reasoning, "We need to calculate factorial",
|
||||
"split_think_close_marker:reasoning");
|
||||
check_equal(parsed.content, "Here is the answer.",
|
||||
"split_think_close_marker:content");
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
int main() {
|
||||
test_reasoning_opened_by_prompt<ds4cpp::DsmlParser>();
|
||||
test_reasoning_disabled();
|
||||
test_explicit_think_tag();
|
||||
test_split_think_close_marker<ds4cpp::DsmlParser>();
|
||||
|
||||
if (failures == 0) {
|
||||
std::fprintf(stderr, "all dsml_parser checks passed\n");
|
||||
return 0;
|
||||
}
|
||||
std::fprintf(stderr, "%d check(s) failed\n", failures);
|
||||
return 1;
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
#pragma once
|
||||
|
||||
#include <algorithm>
|
||||
|
||||
namespace ds4cpp {
|
||||
|
||||
inline int EffectiveGenerationLimit(int requested, int context_size,
|
||||
int session_position) {
|
||||
const int limit = requested > 0 ? requested : 256;
|
||||
const int room = context_size - session_position;
|
||||
if (room <= 1) return 0;
|
||||
return std::min(limit, room - 1);
|
||||
}
|
||||
|
||||
inline int RemainingGenerationBudget(int effective_limit, int produced) {
|
||||
if (effective_limit <= produced) return 0;
|
||||
return effective_limit - produced;
|
||||
}
|
||||
|
||||
inline int SpeculativeAcceptedCapacity(int remaining, int draft_allowance,
|
||||
int buffer_capacity) {
|
||||
if (remaining <= 0 || draft_allowance < 0 || buffer_capacity <= 0) return 0;
|
||||
return std::min({remaining, draft_allowance + 1, buffer_capacity});
|
||||
}
|
||||
|
||||
} // namespace ds4cpp
|
||||
@@ -0,0 +1,92 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
#include "generation_limits.h"
|
||||
|
||||
#include <cstdio>
|
||||
|
||||
namespace {
|
||||
|
||||
int failures = 0;
|
||||
|
||||
void check_equal(int got, int want, const char *name) {
|
||||
if (got == want) return;
|
||||
std::fprintf(stderr, "FAIL %s: got %d, want %d\n", name, got, want);
|
||||
failures++;
|
||||
}
|
||||
|
||||
// Mutation caught: treating omitted or negative max_tokens as unlimited instead
|
||||
// of preserving DS4's legacy 256-token default.
|
||||
void test_nonpositive_uses_legacy_default_when_space_permits() {
|
||||
check_equal(ds4cpp::EffectiveGenerationLimit(0, 4096, 100), 256,
|
||||
"zero max_tokens uses legacy default");
|
||||
check_equal(ds4cpp::EffectiveGenerationLimit(-1, 4096, 100), 256,
|
||||
"negative max_tokens uses legacy default");
|
||||
}
|
||||
|
||||
// Mutation caught: applying the legacy default without clamping it to the
|
||||
// post-prefill context room and reserved slot.
|
||||
void test_legacy_default_is_clamped_by_context() {
|
||||
check_equal(ds4cpp::EffectiveGenerationLimit(0, 300, 100), 199,
|
||||
"legacy default is context-clamped");
|
||||
}
|
||||
|
||||
// Mutation caught: allowing an explicitly large request to overrun the
|
||||
// post-prefill context boundary.
|
||||
void test_large_positive_limit_is_clamped_to_context() {
|
||||
check_equal(ds4cpp::EffectiveGenerationLimit(32768, 32768, 100), 32667,
|
||||
"large positive is context-clamped");
|
||||
}
|
||||
|
||||
// Mutation caught: replacing every positive request with the legacy default
|
||||
// rather than preserving a smaller configured limit.
|
||||
void test_smaller_positive_limit_is_preserved() {
|
||||
check_equal(ds4cpp::EffectiveGenerationLimit(64, 4096, 100), 64,
|
||||
"smaller positive is preserved");
|
||||
}
|
||||
|
||||
// Mutation caught: consuming the final context slot instead of reserving it as
|
||||
// required by DS4's generation loop.
|
||||
void test_no_usable_room_returns_zero() {
|
||||
check_equal(ds4cpp::EffectiveGenerationLimit(32, 100, 99), 0,
|
||||
"one remaining context slot is not usable");
|
||||
}
|
||||
|
||||
// Mutation caught: sending the original generation limit to a later
|
||||
// speculative cycle instead of subtracting tokens already produced.
|
||||
void test_remaining_budget_accounts_for_produced_tokens() {
|
||||
check_equal(ds4cpp::RemainingGenerationBudget(10, 4), 6,
|
||||
"remaining budget subtracts produced tokens");
|
||||
check_equal(ds4cpp::RemainingGenerationBudget(10, 12), 0,
|
||||
"remaining budget never becomes negative");
|
||||
}
|
||||
|
||||
// Mutation caught: giving speculative evaluation capacity beyond either the
|
||||
// output budget, the draft allowance plus its first target token, or the fixed
|
||||
// accepted-token buffer.
|
||||
void test_speculative_capacity_obeys_all_bounds() {
|
||||
check_equal(ds4cpp::SpeculativeAcceptedCapacity(3, 8, 8), 3,
|
||||
"capacity respects remaining output budget");
|
||||
check_equal(ds4cpp::SpeculativeAcceptedCapacity(20, 4, 8), 5,
|
||||
"capacity includes one target token beyond draft allowance");
|
||||
check_equal(ds4cpp::SpeculativeAcceptedCapacity(20, 8, 6), 6,
|
||||
"capacity respects fixed buffer");
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
int main() {
|
||||
test_nonpositive_uses_legacy_default_when_space_permits();
|
||||
test_legacy_default_is_clamped_by_context();
|
||||
test_large_positive_limit_is_clamped_to_context();
|
||||
test_smaller_positive_limit_is_preserved();
|
||||
test_no_usable_room_returns_zero();
|
||||
test_remaining_budget_accounts_for_produced_tokens();
|
||||
test_speculative_capacity_obeys_all_bounds();
|
||||
|
||||
if (failures == 0) {
|
||||
std::fprintf(stderr, "all generation limit checks passed\n");
|
||||
return 0;
|
||||
}
|
||||
std::fprintf(stderr, "%d check(s) failed\n", failures);
|
||||
return 1;
|
||||
}
|
||||
+184
-55
@@ -10,7 +10,9 @@
|
||||
|
||||
#include "dsml_parser.h" // populated in Task 12
|
||||
#include "dsml_renderer.h" // populated in Task 16
|
||||
#include "generation_limits.h"
|
||||
#include "kv_cache.h" // populated in Task 17
|
||||
#include "request_lifecycle.h"
|
||||
|
||||
extern "C" {
|
||||
#include "ds4.h"
|
||||
@@ -35,6 +37,7 @@ extern "C" {
|
||||
#include <mutex>
|
||||
#include <string>
|
||||
#include <thread>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
using grpc::Server;
|
||||
@@ -69,6 +72,21 @@ int g_route_timeout_sec = 60;
|
||||
|
||||
std::atomic<Server *> g_server{nullptr};
|
||||
|
||||
static bool server_context_cancelled(void *ud) {
|
||||
return static_cast<ServerContext *>(ud)->IsCancelled();
|
||||
}
|
||||
|
||||
static void set_session_cancel(void *target, ds4cpp::CancelCallback callback,
|
||||
void *userdata) noexcept {
|
||||
ds4_session_set_cancel(static_cast<ds4_session *>(target), callback, userdata);
|
||||
}
|
||||
|
||||
static bool request_should_continue(ds4cpp::RequestLifecycle *request,
|
||||
ServerContext *context) {
|
||||
request->ObserveContextCancellation(context->IsCancelled());
|
||||
return request->ShouldContinue();
|
||||
}
|
||||
|
||||
// Parse a "key:value" option string. Returns empty when no colon.
|
||||
static std::pair<std::string, std::string> split_option(const std::string &opt) {
|
||||
auto colon = opt.find(':');
|
||||
@@ -238,37 +256,58 @@ static bool apply_engine_option(ds4_engine_options *opt, const std::string &key,
|
||||
|
||||
// When acting as a distributed coordinator, block until the worker route
|
||||
// covers all layers (ds4_session_distributed_route_ready == 1) or the timeout
|
||||
// elapses. Returns an empty string on success, or an error message to return
|
||||
// to the client. No-op when not distributed.
|
||||
// elapses. No-op when not distributed.
|
||||
//
|
||||
// Takes the g_engine_mu lock by reference and RELEASES it during each poll
|
||||
// sleep. The wait can span up to g_route_timeout_sec seconds while workers
|
||||
// connect; holding g_engine_mu the whole time would block the Status/Health
|
||||
// readiness probes (they also lock g_engine_mu), making LocalAI's loader treat
|
||||
// a still-starting worker as hung.
|
||||
static std::string wait_route_ready(std::unique_lock<std::mutex> &lock) {
|
||||
if (!g_distributed) return "";
|
||||
struct RouteWaitResult {
|
||||
ds4cpp::RouteWaitDecision decision;
|
||||
std::string error;
|
||||
};
|
||||
|
||||
static RouteWaitResult wait_route_ready(std::unique_lock<std::mutex> &lock,
|
||||
ServerContext *context) {
|
||||
if (!g_distributed) return {ds4cpp::RouteWaitDecision::Ready, ""};
|
||||
char err[256] = {0};
|
||||
const int deadline_polls = g_route_timeout_sec * 10; // 100ms per poll
|
||||
for (int i = 0; i <= deadline_polls; ++i) {
|
||||
int ready = ds4_session_distributed_route_ready(g_session, err, sizeof(err));
|
||||
if (ready == 1) return "";
|
||||
if (ready < 0) {
|
||||
return std::string("ds4 distributed route error: ") +
|
||||
(err[0] ? err : "unknown");
|
||||
switch (ds4cpp::DecideRouteWait(ready, context->IsCancelled())) {
|
||||
case ds4cpp::RouteWaitDecision::Ready:
|
||||
return {ds4cpp::RouteWaitDecision::Ready, ""};
|
||||
case ds4cpp::RouteWaitDecision::Error:
|
||||
return {ds4cpp::RouteWaitDecision::Error,
|
||||
std::string("ds4 distributed route error: ") +
|
||||
(err[0] ? err : "unknown")};
|
||||
case ds4cpp::RouteWaitDecision::Cancelled:
|
||||
return {ds4cpp::RouteWaitDecision::Cancelled, ""};
|
||||
case ds4cpp::RouteWaitDecision::Pending:
|
||||
break;
|
||||
}
|
||||
if (i == deadline_polls) break;
|
||||
// Release the lock while sleeping so Status/Health and other RPCs can
|
||||
// interleave during worker startup.
|
||||
lock.unlock();
|
||||
struct timespec ts = {0, 100L * 1000L * 1000L}; // 100ms
|
||||
nanosleep(&ts, nullptr);
|
||||
lock.lock();
|
||||
if (context->IsCancelled()) {
|
||||
return {ds4cpp::RouteWaitDecision::Cancelled, ""};
|
||||
}
|
||||
// A concurrent Free() may have torn down the engine while we slept.
|
||||
if (!g_engine || !g_session) {
|
||||
return "ds4: model unloaded while waiting for distributed route";
|
||||
return {ds4cpp::RouteWaitDecision::Error,
|
||||
"ds4: model unloaded while waiting for distributed route"};
|
||||
}
|
||||
}
|
||||
return "ds4 distributed route incomplete: workers not connected (layers uncovered)";
|
||||
if (context->IsCancelled()) {
|
||||
return {ds4cpp::RouteWaitDecision::Cancelled, ""};
|
||||
}
|
||||
return {ds4cpp::RouteWaitDecision::Error,
|
||||
"ds4 distributed route incomplete: workers not connected (layers uncovered)"};
|
||||
}
|
||||
|
||||
static void append_token_text(ds4_engine *engine, int token, std::string &out) {
|
||||
@@ -341,9 +380,9 @@ static void collect_done(void *) {}
|
||||
struct StreamCtx {
|
||||
ds4_engine *engine;
|
||||
ServerWriter<backend::Reply> *writer;
|
||||
ds4cpp::RequestLifecycle *request;
|
||||
ds4cpp::DsmlParser parser;
|
||||
int tokens;
|
||||
bool aborted;
|
||||
// Track which tool indices we've seen TOOL_START for, so subsequent
|
||||
// ARGS deltas can elide the redundant id/name fields.
|
||||
std::vector<bool> tool_started;
|
||||
@@ -351,7 +390,7 @@ struct StreamCtx {
|
||||
|
||||
static void stream_emit(void *ud, int token) {
|
||||
auto *s = static_cast<StreamCtx *>(ud);
|
||||
if (s->aborted) return;
|
||||
if (!s->request->ShouldContinue()) return;
|
||||
if (token == ds4_token_eos(s->engine)) return;
|
||||
size_t len = 0;
|
||||
const char *text = ds4_token_text(s->engine, token, &len);
|
||||
@@ -401,7 +440,7 @@ static void stream_emit(void *ud, int token) {
|
||||
reply.set_message(chunk);
|
||||
reply.set_tokens(1);
|
||||
if (any_field) {
|
||||
if (!s->writer->Write(reply)) s->aborted = true;
|
||||
s->request->ObserveStreamWrite(s->writer->Write(reply));
|
||||
}
|
||||
s->tokens++;
|
||||
}
|
||||
@@ -757,21 +796,30 @@ public:
|
||||
return GStatus::OK;
|
||||
}
|
||||
|
||||
GStatus Predict(ServerContext *, const backend::PredictOptions *request,
|
||||
GStatus Predict(ServerContext *context, const backend::PredictOptions *request,
|
||||
backend::Reply *reply) override {
|
||||
std::unique_lock<std::mutex> lock(g_engine_mu);
|
||||
if (!g_engine || !g_session) {
|
||||
return GStatus(StatusCode::FAILED_PRECONDITION, "ds4: model not loaded");
|
||||
}
|
||||
if (GStatus id = check_model_identity(request); !id.ok()) return id;
|
||||
if (std::string route_err = wait_route_ready(lock); !route_err.empty()) {
|
||||
return GStatus(StatusCode::UNAVAILABLE, route_err);
|
||||
RouteWaitResult route = wait_route_ready(lock, context);
|
||||
if (route.decision == ds4cpp::RouteWaitDecision::Cancelled) {
|
||||
return GStatus(StatusCode::CANCELLED, "ds4 request cancelled");
|
||||
}
|
||||
if (route.decision == ds4cpp::RouteWaitDecision::Error) {
|
||||
return GStatus(StatusCode::UNAVAILABLE, route.error);
|
||||
}
|
||||
ds4_tokens prompt = {};
|
||||
build_prompt(g_engine, request, &prompt);
|
||||
int n_predict = request->tokens() > 0 ? request->tokens() : 256;
|
||||
|
||||
CollectCtx collect = {g_engine, "", {}, reply, 0, {}, "", ""};
|
||||
const bool think_enabled = ds4_think_mode_enabled(parse_think_mode(request));
|
||||
const bool starts_in_thinking = think_enabled &&
|
||||
request->usetokenizertemplate() && request->messages_size() > 0;
|
||||
CollectCtx collect = {
|
||||
g_engine, "", ds4cpp::DsmlParser(starts_in_thinking),
|
||||
reply, 0, {}, "", ""};
|
||||
ds4cpp::RequestLifecycle lifecycle;
|
||||
std::string cache_key = render_prompt_text(request);
|
||||
size_t cache_hit = maybe_load_cache(cache_key);
|
||||
(void)cache_hit; // future: skip prompt prefix if hit covers full prompt
|
||||
@@ -783,15 +831,27 @@ public:
|
||||
// Either way g_session advances so the disk KV cache picks up a
|
||||
// real checkpoint after the call (see maybe_save_cache below).
|
||||
char err[256] = {0};
|
||||
int rc = ds4_session_sync(g_session, &prompt, err, sizeof(err));
|
||||
int rc;
|
||||
{
|
||||
ds4cpp::CancelCallbackScope cancel_scope(
|
||||
g_session, set_session_cancel, server_context_cancelled, context);
|
||||
rc = ds4_session_sync(g_session, &prompt, err, sizeof(err));
|
||||
}
|
||||
int prompt_len = prompt.len;
|
||||
ds4_tokens_free(&prompt);
|
||||
if (rc == 0) {
|
||||
if (rc == DS4_SESSION_SYNC_INTERRUPTED) {
|
||||
lifecycle.ObserveContextCancellation(true);
|
||||
}
|
||||
const bool generation_started = rc == 0;
|
||||
if (generation_started) {
|
||||
const int n_predict = ds4cpp::EffectiveGenerationLimit(
|
||||
request->tokens(), ds4_session_ctx(g_session),
|
||||
ds4_session_pos(g_session));
|
||||
const int eos = ds4_token_eos(g_engine);
|
||||
const int draft_max = ds4_engine_mtp_draft_tokens(g_engine);
|
||||
const bool think_enabled = ds4_think_mode_enabled(parse_think_mode(request));
|
||||
int produced = 0;
|
||||
while (produced < n_predict) {
|
||||
if (!request_should_continue(&lifecycle, context)) break;
|
||||
SampleParams sp = compute_sample_params(request, collect.parser, think_enabled);
|
||||
int first;
|
||||
if (sp.temperature <= 0.0f) {
|
||||
@@ -806,13 +866,20 @@ public:
|
||||
if (draft_max > 0 && sp.temperature <= 0.0f) {
|
||||
constexpr int kAcceptedMax = 8;
|
||||
int accepted[kAcceptedMax];
|
||||
int cap = std::min(kAcceptedMax, draft_max + 1);
|
||||
const int remaining = ds4cpp::RemainingGenerationBudget(
|
||||
n_predict, produced);
|
||||
const int cap = ds4cpp::SpeculativeAcceptedCapacity(
|
||||
remaining, draft_max, kAcceptedMax);
|
||||
int n = ds4_session_eval_speculative_argmax(
|
||||
g_session, first, draft_max, eos,
|
||||
g_session, first, remaining, eos,
|
||||
accepted, cap, err, sizeof(err));
|
||||
if (n < 0) { rc = -1; break; }
|
||||
bool stop = false;
|
||||
for (int j = 0; j < n; ++j) {
|
||||
if (!request_should_continue(&lifecycle, context)) {
|
||||
stop = true;
|
||||
break;
|
||||
}
|
||||
if (accepted[j] == eos) { stop = true; break; }
|
||||
collect_emit(&collect, accepted[j]);
|
||||
if (++produced >= n_predict) { stop = true; break; }
|
||||
@@ -821,12 +888,26 @@ public:
|
||||
} else {
|
||||
collect_emit(&collect, first);
|
||||
if (++produced >= n_predict) break;
|
||||
if (!request_should_continue(&lifecycle, context)) break;
|
||||
rc = ds4_session_eval(g_session, first, err, sizeof(err));
|
||||
if (rc != 0) break;
|
||||
}
|
||||
}
|
||||
collect_done(&collect);
|
||||
}
|
||||
|
||||
request_should_continue(&lifecycle, context);
|
||||
ds4cpp::TerminalDecision terminal = ds4cpp::ResolveTerminalDecision(
|
||||
rc == DS4_SESSION_SYNC_INTERRUPTED, rc != 0,
|
||||
!lifecycle.ShouldFinalize());
|
||||
if (!terminal.should_finalize) {
|
||||
if (terminal.cause == ds4cpp::TerminalCause::EngineError) {
|
||||
return GStatus(StatusCode::INTERNAL,
|
||||
std::string("ds4 generation failed: ") + err);
|
||||
}
|
||||
return GStatus(StatusCode::CANCELLED,
|
||||
"ds4 request cancelled");
|
||||
}
|
||||
if (generation_started) collect_done(&collect);
|
||||
maybe_save_cache(cache_key);
|
||||
|
||||
// Flush any buffered parser state.
|
||||
@@ -834,7 +915,7 @@ public:
|
||||
collect.parser.Flush(events);
|
||||
apply_events(&collect, events);
|
||||
|
||||
if (rc != 0) {
|
||||
if (terminal.cause == ds4cpp::TerminalCause::EngineError) {
|
||||
return GStatus(StatusCode::INTERNAL,
|
||||
std::string("ds4 generation failed: ") + err);
|
||||
}
|
||||
@@ -857,21 +938,30 @@ public:
|
||||
return GStatus::OK;
|
||||
}
|
||||
|
||||
GStatus PredictStream(ServerContext *, const backend::PredictOptions *request,
|
||||
GStatus PredictStream(ServerContext *context, const backend::PredictOptions *request,
|
||||
ServerWriter<backend::Reply> *writer) override {
|
||||
std::unique_lock<std::mutex> lock(g_engine_mu);
|
||||
if (!g_engine || !g_session) {
|
||||
return GStatus(StatusCode::FAILED_PRECONDITION, "ds4: model not loaded");
|
||||
}
|
||||
if (GStatus id = check_model_identity(request); !id.ok()) return id;
|
||||
if (std::string route_err = wait_route_ready(lock); !route_err.empty()) {
|
||||
return GStatus(StatusCode::UNAVAILABLE, route_err);
|
||||
RouteWaitResult route = wait_route_ready(lock, context);
|
||||
if (route.decision == ds4cpp::RouteWaitDecision::Cancelled) {
|
||||
return GStatus(StatusCode::CANCELLED, "ds4 request cancelled");
|
||||
}
|
||||
if (route.decision == ds4cpp::RouteWaitDecision::Error) {
|
||||
return GStatus(StatusCode::UNAVAILABLE, route.error);
|
||||
}
|
||||
ds4_tokens prompt = {};
|
||||
build_prompt(g_engine, request, &prompt);
|
||||
int n_predict = request->tokens() > 0 ? request->tokens() : 256;
|
||||
|
||||
StreamCtx s = {g_engine, writer, {}, 0, false, {}};
|
||||
const bool think_enabled = ds4_think_mode_enabled(parse_think_mode(request));
|
||||
const bool starts_in_thinking = think_enabled &&
|
||||
request->usetokenizertemplate() && request->messages_size() > 0;
|
||||
ds4cpp::RequestLifecycle lifecycle;
|
||||
StreamCtx s = {
|
||||
g_engine, writer, &lifecycle,
|
||||
ds4cpp::DsmlParser(starts_in_thinking), 0, {}};
|
||||
std::string cache_key = render_prompt_text(request);
|
||||
size_t cache_hit = maybe_load_cache(cache_key);
|
||||
(void)cache_hit;
|
||||
@@ -879,14 +969,26 @@ public:
|
||||
// Manual loop on g_session - see Predict() above for the rationale.
|
||||
// MTP speculative path used when ds4_engine_mtp_draft_tokens > 0.
|
||||
char err[256] = {0};
|
||||
int rc = ds4_session_sync(g_session, &prompt, err, sizeof(err));
|
||||
int rc;
|
||||
{
|
||||
ds4cpp::CancelCallbackScope cancel_scope(
|
||||
g_session, set_session_cancel, server_context_cancelled, context);
|
||||
rc = ds4_session_sync(g_session, &prompt, err, sizeof(err));
|
||||
}
|
||||
ds4_tokens_free(&prompt);
|
||||
if (rc == 0) {
|
||||
if (rc == DS4_SESSION_SYNC_INTERRUPTED) {
|
||||
lifecycle.ObserveContextCancellation(true);
|
||||
}
|
||||
const bool generation_started = rc == 0;
|
||||
if (generation_started) {
|
||||
const int n_predict = ds4cpp::EffectiveGenerationLimit(
|
||||
request->tokens(), ds4_session_ctx(g_session),
|
||||
ds4_session_pos(g_session));
|
||||
const int eos = ds4_token_eos(g_engine);
|
||||
const int draft_max = ds4_engine_mtp_draft_tokens(g_engine);
|
||||
const bool think_enabled = ds4_think_mode_enabled(parse_think_mode(request));
|
||||
int produced = 0;
|
||||
while (produced < n_predict && !s.aborted) {
|
||||
while (produced < n_predict) {
|
||||
if (!request_should_continue(&lifecycle, context)) break;
|
||||
SampleParams sp = compute_sample_params(request, s.parser, think_enabled);
|
||||
int first;
|
||||
if (sp.temperature <= 0.0f) {
|
||||
@@ -900,50 +1002,77 @@ public:
|
||||
if (draft_max > 0 && sp.temperature <= 0.0f) {
|
||||
constexpr int kAcceptedMax = 8;
|
||||
int accepted[kAcceptedMax];
|
||||
int cap = std::min(kAcceptedMax, draft_max + 1);
|
||||
const int remaining = ds4cpp::RemainingGenerationBudget(
|
||||
n_predict, produced);
|
||||
const int cap = ds4cpp::SpeculativeAcceptedCapacity(
|
||||
remaining, draft_max, kAcceptedMax);
|
||||
int n = ds4_session_eval_speculative_argmax(
|
||||
g_session, first, draft_max, eos,
|
||||
g_session, first, remaining, eos,
|
||||
accepted, cap, err, sizeof(err));
|
||||
if (n < 0) { rc = -1; break; }
|
||||
bool stop = false;
|
||||
for (int j = 0; j < n; ++j) {
|
||||
if (!request_should_continue(&lifecycle, context)) {
|
||||
stop = true;
|
||||
break;
|
||||
}
|
||||
if (accepted[j] == eos) { stop = true; break; }
|
||||
stream_emit(&s, accepted[j]);
|
||||
if (s.aborted) { stop = true; break; }
|
||||
if (!lifecycle.ShouldContinue()) { stop = true; break; }
|
||||
if (++produced >= n_predict) { stop = true; break; }
|
||||
}
|
||||
if (stop) break;
|
||||
} else {
|
||||
stream_emit(&s, first);
|
||||
if (s.aborted || ++produced >= n_predict) break;
|
||||
if (!lifecycle.ShouldContinue() || ++produced >= n_predict) break;
|
||||
if (!request_should_continue(&lifecycle, context)) break;
|
||||
rc = ds4_session_eval(g_session, first, err, sizeof(err));
|
||||
if (rc != 0) break;
|
||||
}
|
||||
}
|
||||
stream_done(&s);
|
||||
}
|
||||
maybe_save_cache(cache_key);
|
||||
|
||||
// Flush parser state.
|
||||
std::vector<ds4cpp::ParserEvent> events;
|
||||
s.parser.Flush(events);
|
||||
if (!events.empty() && !s.aborted) {
|
||||
backend::Reply reply;
|
||||
auto *delta = reply.add_chat_deltas();
|
||||
for (const auto &e : events) {
|
||||
if (e.type == ds4cpp::ParserEvent::CONTENT) {
|
||||
delta->set_content(delta->content() + e.text);
|
||||
} else if (e.type == ds4cpp::ParserEvent::REASONING) {
|
||||
delta->set_reasoning_content(delta->reasoning_content() + e.text);
|
||||
request_should_continue(&lifecycle, context);
|
||||
ds4cpp::TerminalDecision terminal = ds4cpp::ResolveTerminalDecision(
|
||||
rc == DS4_SESSION_SYNC_INTERRUPTED, rc != 0,
|
||||
!lifecycle.ShouldFinalize());
|
||||
terminal = ds4cpp::RunPostlude(
|
||||
terminal,
|
||||
[&]() {
|
||||
ds4cpp::DsmlParser staged_parser = s.parser;
|
||||
std::vector<ds4cpp::ParserEvent> events;
|
||||
staged_parser.Flush(events);
|
||||
bool write_succeeded = true;
|
||||
if (!events.empty()) {
|
||||
backend::Reply reply;
|
||||
auto *delta = reply.add_chat_deltas();
|
||||
for (const auto &e : events) {
|
||||
if (e.type == ds4cpp::ParserEvent::CONTENT) {
|
||||
delta->set_content(delta->content() + e.text);
|
||||
} else if (e.type == ds4cpp::ParserEvent::REASONING) {
|
||||
delta->set_reasoning_content(
|
||||
delta->reasoning_content() + e.text);
|
||||
}
|
||||
}
|
||||
write_succeeded = s.writer->Write(reply);
|
||||
}
|
||||
}
|
||||
s.writer->Write(reply);
|
||||
}
|
||||
lifecycle.ObserveStreamWrite(write_succeeded);
|
||||
request_should_continue(&lifecycle, context);
|
||||
if (!lifecycle.ShouldFinalize()) return false;
|
||||
s.parser = std::move(staged_parser);
|
||||
if (generation_started) stream_done(&s);
|
||||
return true;
|
||||
},
|
||||
[&]() { maybe_save_cache(cache_key); });
|
||||
|
||||
if (rc != 0 && !s.aborted) {
|
||||
if (terminal.cause == ds4cpp::TerminalCause::EngineError) {
|
||||
return GStatus(StatusCode::INTERNAL,
|
||||
std::string("ds4 generation failed: ") + err);
|
||||
}
|
||||
if (terminal.cause == ds4cpp::TerminalCause::Cancelled) {
|
||||
return GStatus(StatusCode::CANCELLED,
|
||||
"ds4 request cancelled");
|
||||
}
|
||||
return GStatus::OK;
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,111 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
#pragma once
|
||||
|
||||
namespace ds4cpp {
|
||||
|
||||
using CancelCallback = bool (*)(void *);
|
||||
using CancelSetter = void (*)(void *, CancelCallback, void *) noexcept;
|
||||
|
||||
class CancelCallbackScope {
|
||||
public:
|
||||
CancelCallbackScope(void *target, CancelSetter setter,
|
||||
CancelCallback callback, void *userdata) noexcept
|
||||
: target_(target), setter_(setter) {
|
||||
setter_(target_, callback, userdata);
|
||||
}
|
||||
|
||||
~CancelCallbackScope() noexcept {
|
||||
setter_(target_, nullptr, nullptr);
|
||||
}
|
||||
|
||||
CancelCallbackScope(const CancelCallbackScope &) = delete;
|
||||
CancelCallbackScope &operator=(const CancelCallbackScope &) = delete;
|
||||
|
||||
private:
|
||||
void *target_;
|
||||
CancelSetter setter_;
|
||||
};
|
||||
|
||||
enum class RouteWaitDecision {
|
||||
Pending,
|
||||
Ready,
|
||||
Error,
|
||||
Cancelled,
|
||||
};
|
||||
|
||||
inline RouteWaitDecision DecideRouteWait(int route_status, bool cancelled) {
|
||||
if (cancelled) return RouteWaitDecision::Cancelled;
|
||||
if (route_status > 0) return RouteWaitDecision::Ready;
|
||||
if (route_status < 0) return RouteWaitDecision::Error;
|
||||
return RouteWaitDecision::Pending;
|
||||
}
|
||||
|
||||
enum class TerminalCause {
|
||||
Success,
|
||||
Cancelled,
|
||||
EngineError,
|
||||
};
|
||||
|
||||
inline TerminalCause DecideTerminalCause(bool sync_interrupted,
|
||||
bool engine_error,
|
||||
bool abandoned) {
|
||||
if (sync_interrupted) return TerminalCause::Cancelled;
|
||||
if (engine_error) return TerminalCause::EngineError;
|
||||
if (abandoned) return TerminalCause::Cancelled;
|
||||
return TerminalCause::Success;
|
||||
}
|
||||
|
||||
struct TerminalDecision {
|
||||
TerminalCause cause;
|
||||
bool should_finalize;
|
||||
};
|
||||
|
||||
inline TerminalDecision ResolveTerminalDecision(bool sync_interrupted,
|
||||
bool engine_error,
|
||||
bool abandoned) {
|
||||
return {
|
||||
DecideTerminalCause(sync_interrupted, engine_error, abandoned),
|
||||
!sync_interrupted && !abandoned,
|
||||
};
|
||||
}
|
||||
|
||||
template <typename Finalize, typename Persist>
|
||||
TerminalDecision RunPostlude(TerminalDecision terminal,
|
||||
Finalize transactional_finalize,
|
||||
Persist persist) {
|
||||
if (!terminal.should_finalize) return terminal;
|
||||
if (!transactional_finalize()) {
|
||||
terminal.should_finalize = false;
|
||||
if (terminal.cause != TerminalCause::EngineError) {
|
||||
terminal.cause = TerminalCause::Cancelled;
|
||||
}
|
||||
return terminal;
|
||||
}
|
||||
persist();
|
||||
return terminal;
|
||||
}
|
||||
|
||||
class RequestLifecycle {
|
||||
public:
|
||||
void ObserveContextCancellation(bool cancelled) {
|
||||
context_cancelled_ = context_cancelled_ || cancelled;
|
||||
}
|
||||
|
||||
void ObserveStreamWrite(bool succeeded) {
|
||||
stream_write_aborted_ = stream_write_aborted_ || !succeeded;
|
||||
}
|
||||
|
||||
bool ShouldContinue() const {
|
||||
return !context_cancelled_ && !stream_write_aborted_;
|
||||
}
|
||||
|
||||
bool ShouldFinalize() const {
|
||||
return ShouldContinue();
|
||||
}
|
||||
|
||||
private:
|
||||
bool context_cancelled_ = false;
|
||||
bool stream_write_aborted_ = false;
|
||||
};
|
||||
|
||||
} // namespace ds4cpp
|
||||
@@ -0,0 +1,414 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
// Standalone regression tests for DS4 request cancellation policy.
|
||||
|
||||
#include "request_lifecycle.h"
|
||||
|
||||
#include <cstdio>
|
||||
|
||||
namespace {
|
||||
|
||||
int failures = 0;
|
||||
|
||||
struct FakeCancelTarget {
|
||||
ds4cpp::CancelCallback callback = nullptr;
|
||||
void *userdata = nullptr;
|
||||
int installs = 0;
|
||||
int clears = 0;
|
||||
};
|
||||
|
||||
struct PostludeCounts {
|
||||
int finalize_attempts = 0;
|
||||
int finalize_commits = 0;
|
||||
int cache_persists = 0;
|
||||
bool cache_followed_commit = true;
|
||||
};
|
||||
|
||||
ds4cpp::TerminalDecision run_fake_postlude(
|
||||
ds4cpp::TerminalDecision terminal, bool finalize_succeeds,
|
||||
PostludeCounts *counts) {
|
||||
return ds4cpp::RunPostlude(
|
||||
terminal,
|
||||
[=]() {
|
||||
counts->finalize_attempts++;
|
||||
if (!finalize_succeeds) return false;
|
||||
counts->finalize_commits++;
|
||||
return true;
|
||||
},
|
||||
[=]() {
|
||||
counts->cache_followed_commit = counts->finalize_commits == 1;
|
||||
counts->cache_persists++;
|
||||
});
|
||||
}
|
||||
|
||||
bool fake_cancel(void *) {
|
||||
return false;
|
||||
}
|
||||
|
||||
void fake_set_cancel(void *target, ds4cpp::CancelCallback callback,
|
||||
void *userdata) noexcept {
|
||||
auto *fake = static_cast<FakeCancelTarget *>(target);
|
||||
fake->callback = callback;
|
||||
fake->userdata = userdata;
|
||||
if (callback) {
|
||||
fake->installs++;
|
||||
} else {
|
||||
fake->clears++;
|
||||
}
|
||||
}
|
||||
|
||||
void check(bool condition, const char *name) {
|
||||
if (condition) return;
|
||||
std::fprintf(stderr, "FAIL %s\n", name);
|
||||
failures++;
|
||||
}
|
||||
|
||||
// Production mutation caught: treating an active request as abandoned would
|
||||
// skip its parser finalization and cache save.
|
||||
void test_active_request_continues_and_finalizes() {
|
||||
ds4cpp::RequestLifecycle request;
|
||||
|
||||
check(request.ShouldContinue(), "active:continue");
|
||||
check(request.ShouldFinalize(), "active:finalize");
|
||||
}
|
||||
|
||||
// Production mutation caught: omitting the ServerContext cancellation branch
|
||||
// would continue decoding and finalize a partial response.
|
||||
void test_context_cancellation_stops_without_finalizing() {
|
||||
ds4cpp::RequestLifecycle request;
|
||||
|
||||
request.ObserveContextCancellation(true);
|
||||
|
||||
check(!request.ShouldContinue(), "context_cancelled:stop");
|
||||
check(!request.ShouldFinalize(), "context_cancelled:no_finalize");
|
||||
}
|
||||
|
||||
// Production mutation caught: ignoring ServerWriter::Write failure would keep
|
||||
// streaming and finalize a response whose client has gone away.
|
||||
void test_stream_write_abort_stops_without_finalizing() {
|
||||
ds4cpp::RequestLifecycle request;
|
||||
|
||||
request.ObserveStreamWrite(false);
|
||||
|
||||
check(!request.ShouldContinue(), "write_abort:stop");
|
||||
check(!request.ShouldFinalize(), "write_abort:no_finalize");
|
||||
}
|
||||
|
||||
// Production mutation caught: combining cancellation and write failure with
|
||||
// AND would fail to stop when either signal occurs on its own.
|
||||
void test_cancellation_and_write_abort_are_independent_or_conditions() {
|
||||
ds4cpp::RequestLifecycle cancelled;
|
||||
cancelled.ObserveContextCancellation(true);
|
||||
cancelled.ObserveStreamWrite(true);
|
||||
|
||||
ds4cpp::RequestLifecycle write_aborted;
|
||||
write_aborted.ObserveContextCancellation(false);
|
||||
write_aborted.ObserveStreamWrite(false);
|
||||
|
||||
check(!cancelled.ShouldContinue(), "or:context_only");
|
||||
check(!write_aborted.ShouldContinue(), "or:write_only");
|
||||
}
|
||||
|
||||
// Production mutation caught: treating an incomplete distributed route as an
|
||||
// error would return before workers have time to connect.
|
||||
void test_route_wait_pending() {
|
||||
check(ds4cpp::DecideRouteWait(0, false) ==
|
||||
ds4cpp::RouteWaitDecision::Pending,
|
||||
"route_wait:pending");
|
||||
}
|
||||
|
||||
// Production mutation caught: failing to recognize a complete route would
|
||||
// keep a ready inference request in the polling loop.
|
||||
void test_route_wait_ready() {
|
||||
check(ds4cpp::DecideRouteWait(1, false) ==
|
||||
ds4cpp::RouteWaitDecision::Ready,
|
||||
"route_wait:ready");
|
||||
}
|
||||
|
||||
// Production mutation caught: ignoring a route probe error would poll until a
|
||||
// misleading timeout instead of returning UNAVAILABLE promptly.
|
||||
void test_route_wait_error() {
|
||||
check(ds4cpp::DecideRouteWait(-1, false) ==
|
||||
ds4cpp::RouteWaitDecision::Error,
|
||||
"route_wait:error");
|
||||
}
|
||||
|
||||
// Production mutation caught: omitting cancellation from route waiting would
|
||||
// leave an abandoned request blocked until the distributed timeout.
|
||||
void test_route_wait_cancellation() {
|
||||
check(ds4cpp::DecideRouteWait(0, true) ==
|
||||
ds4cpp::RouteWaitDecision::Cancelled,
|
||||
"route_wait:cancelled");
|
||||
}
|
||||
|
||||
// Production mutation caught: checking route errors before cancellation would
|
||||
// report UNAVAILABLE for a request the client already abandoned.
|
||||
void test_route_wait_cancellation_precedes_error() {
|
||||
check(ds4cpp::DecideRouteWait(-1, true) ==
|
||||
ds4cpp::RouteWaitDecision::Cancelled,
|
||||
"route_wait:cancellation_precedence");
|
||||
}
|
||||
|
||||
// Production mutation caught: classifying a successful active request as a
|
||||
// terminal failure would suppress its normal response finalization.
|
||||
void test_terminal_success() {
|
||||
check(ds4cpp::DecideTerminalCause(false, false, false) ==
|
||||
ds4cpp::TerminalCause::Success,
|
||||
"terminal:success");
|
||||
}
|
||||
|
||||
// Production mutation caught: treating DS4's cooperative sync interruption
|
||||
// as an ordinary engine error would return INTERNAL instead of CANCELLED.
|
||||
void test_terminal_sync_interruption_is_cancelled() {
|
||||
check(ds4cpp::DecideTerminalCause(true, true, true) ==
|
||||
ds4cpp::TerminalCause::Cancelled,
|
||||
"terminal:sync_interrupted");
|
||||
}
|
||||
|
||||
// Production mutation caught: treating every nonzero engine result as client
|
||||
// abandonment would hide genuine DS4 failures behind CANCELLED.
|
||||
void test_terminal_engine_error() {
|
||||
check(ds4cpp::DecideTerminalCause(false, true, false) ==
|
||||
ds4cpp::TerminalCause::EngineError,
|
||||
"terminal:engine_error");
|
||||
}
|
||||
|
||||
// Production mutation caught: ignoring an rc==0 context cancellation would
|
||||
// finalize and cache an abandoned request.
|
||||
void test_terminal_context_abandonment() {
|
||||
ds4cpp::RequestLifecycle request;
|
||||
request.ObserveContextCancellation(true);
|
||||
|
||||
check(ds4cpp::DecideTerminalCause(
|
||||
false, false, !request.ShouldFinalize()) ==
|
||||
ds4cpp::TerminalCause::Cancelled,
|
||||
"terminal:context_abandonment");
|
||||
}
|
||||
|
||||
// Production mutation caught: ignoring an rc==0 stream write failure would
|
||||
// finalize and cache an abandoned streaming request.
|
||||
void test_terminal_write_abandonment() {
|
||||
ds4cpp::RequestLifecycle request;
|
||||
request.ObserveStreamWrite(false);
|
||||
|
||||
check(ds4cpp::DecideTerminalCause(
|
||||
false, false, !request.ShouldFinalize()) ==
|
||||
ds4cpp::TerminalCause::Cancelled,
|
||||
"terminal:write_abandonment");
|
||||
}
|
||||
|
||||
// Production mutation caught: checking late cancellation or write failure
|
||||
// before a determined ordinary DS4 error would replace INTERNAL with CANCELLED.
|
||||
void test_terminal_engine_error_precedes_late_abandonment() {
|
||||
ds4cpp::RequestLifecycle cancelled;
|
||||
cancelled.ObserveContextCancellation(true);
|
||||
ds4cpp::RequestLifecycle write_aborted;
|
||||
write_aborted.ObserveStreamWrite(false);
|
||||
|
||||
check(ds4cpp::DecideTerminalCause(
|
||||
false, true, !cancelled.ShouldFinalize()) ==
|
||||
ds4cpp::TerminalCause::EngineError,
|
||||
"terminal:engine_error_precedes_cancellation");
|
||||
check(ds4cpp::DecideTerminalCause(
|
||||
false, true, !write_aborted.ShouldFinalize()) ==
|
||||
ds4cpp::TerminalCause::EngineError,
|
||||
"terminal:engine_error_precedes_write_abort");
|
||||
}
|
||||
|
||||
// Production mutation caught: using status precedence alone to gate side
|
||||
// effects would finalize and persist an engine-error request abandoned later.
|
||||
void test_abandoned_engine_error_keeps_internal_without_finalizing() {
|
||||
ds4cpp::RequestLifecycle request;
|
||||
request.ObserveContextCancellation(true);
|
||||
|
||||
ds4cpp::TerminalDecision terminal = ds4cpp::ResolveTerminalDecision(
|
||||
false, true, !request.ShouldFinalize());
|
||||
|
||||
check(terminal.cause == ds4cpp::TerminalCause::EngineError,
|
||||
"terminal_decision:abandoned_engine_error_status");
|
||||
check(!terminal.should_finalize,
|
||||
"terminal_decision:abandoned_engine_error_no_finalize");
|
||||
}
|
||||
|
||||
// Production mutation caught: suppressing side effects for every engine error
|
||||
// would change the existing finalization and cache behavior of active failures.
|
||||
void test_active_engine_error_still_finalizes() {
|
||||
ds4cpp::RequestLifecycle request;
|
||||
|
||||
ds4cpp::TerminalDecision terminal = ds4cpp::ResolveTerminalDecision(
|
||||
false, true, !request.ShouldFinalize());
|
||||
|
||||
check(terminal.cause == ds4cpp::TerminalCause::EngineError,
|
||||
"terminal_decision:active_engine_error_status");
|
||||
check(terminal.should_finalize,
|
||||
"terminal_decision:active_engine_error_finalize");
|
||||
}
|
||||
|
||||
// Production mutation caught: persisting before committed finalization would
|
||||
// cache a state whose final buffered stream reply was never completed.
|
||||
void test_postlude_active_success_commits_then_persists() {
|
||||
PostludeCounts counts;
|
||||
|
||||
ds4cpp::TerminalDecision terminal = run_fake_postlude(
|
||||
{ds4cpp::TerminalCause::Success, true}, true, &counts);
|
||||
|
||||
check(terminal.cause == ds4cpp::TerminalCause::Success,
|
||||
"postlude:success_outcome");
|
||||
check(terminal.should_finalize, "postlude:success_committed");
|
||||
check(counts.finalize_attempts == 1, "postlude:success_attempts");
|
||||
check(counts.finalize_commits == 1, "postlude:success_commits");
|
||||
check(counts.cache_persists == 1, "postlude:success_cache");
|
||||
check(counts.cache_followed_commit, "postlude:success_cache_order");
|
||||
}
|
||||
|
||||
// Production mutation caught: starting the postlude for an already-cancelled
|
||||
// request would flush buffered parser state or persist an abandoned session.
|
||||
void test_postlude_cancellation_skips_all_side_effects() {
|
||||
PostludeCounts counts;
|
||||
|
||||
ds4cpp::TerminalDecision terminal = run_fake_postlude(
|
||||
{ds4cpp::TerminalCause::Cancelled, false}, true, &counts);
|
||||
|
||||
check(terminal.cause == ds4cpp::TerminalCause::Cancelled,
|
||||
"postlude:cancelled_outcome");
|
||||
check(counts.finalize_attempts == 0, "postlude:cancelled_attempts");
|
||||
check(counts.finalize_commits == 0, "postlude:cancelled_commits");
|
||||
check(counts.cache_persists == 0, "postlude:cancelled_cache");
|
||||
}
|
||||
|
||||
// Production mutation caught: committing the live parser or cache after a
|
||||
// failed final Write would publish an abandoned streaming postlude.
|
||||
void test_postlude_finalize_failure_cancels_without_commit_or_cache() {
|
||||
PostludeCounts counts;
|
||||
|
||||
ds4cpp::TerminalDecision terminal = run_fake_postlude(
|
||||
{ds4cpp::TerminalCause::Success, true}, false, &counts);
|
||||
|
||||
check(terminal.cause == ds4cpp::TerminalCause::Cancelled,
|
||||
"postlude:write_failure_outcome");
|
||||
check(!terminal.should_finalize, "postlude:write_failure_not_committed");
|
||||
check(counts.finalize_attempts == 1, "postlude:write_failure_attempts");
|
||||
check(counts.finalize_commits == 0, "postlude:write_failure_commits");
|
||||
check(counts.cache_persists == 0, "postlude:write_failure_cache");
|
||||
}
|
||||
|
||||
// Production mutation caught: skipping the postlude for every engine error
|
||||
// would change active internal-error finalization and cache behavior.
|
||||
void test_postlude_active_engine_error_finalizes_and_persists() {
|
||||
PostludeCounts counts;
|
||||
|
||||
ds4cpp::TerminalDecision terminal = run_fake_postlude(
|
||||
{ds4cpp::TerminalCause::EngineError, true}, true, &counts);
|
||||
|
||||
check(terminal.cause == ds4cpp::TerminalCause::EngineError,
|
||||
"postlude:engine_error_outcome");
|
||||
check(counts.finalize_attempts == 1, "postlude:engine_error_attempts");
|
||||
check(counts.finalize_commits == 1, "postlude:engine_error_commits");
|
||||
check(counts.cache_persists == 1, "postlude:engine_error_cache");
|
||||
check(counts.cache_followed_commit, "postlude:engine_error_cache_order");
|
||||
}
|
||||
|
||||
// Production mutation caught: replacing every failed transactional finalize
|
||||
// with cancellation would hide an already-determined engine error.
|
||||
void test_postlude_engine_error_finalize_failure_preserves_internal() {
|
||||
PostludeCounts counts;
|
||||
|
||||
ds4cpp::TerminalDecision terminal = run_fake_postlude(
|
||||
{ds4cpp::TerminalCause::EngineError, true}, false, &counts);
|
||||
|
||||
check(terminal.cause == ds4cpp::TerminalCause::EngineError,
|
||||
"postlude:engine_error_write_failure_outcome");
|
||||
check(!terminal.should_finalize,
|
||||
"postlude:engine_error_write_failure_not_committed");
|
||||
check(counts.finalize_attempts == 1,
|
||||
"postlude:engine_error_write_failure_attempts");
|
||||
check(counts.finalize_commits == 0,
|
||||
"postlude:engine_error_write_failure_commits");
|
||||
check(counts.cache_persists == 0,
|
||||
"postlude:engine_error_write_failure_cache");
|
||||
}
|
||||
|
||||
// Production mutation caught: status precedence must not grant side-effect
|
||||
// permission to an engine-error request that was also abandoned.
|
||||
void test_postlude_abandoned_engine_error_skips_all_side_effects() {
|
||||
PostludeCounts counts;
|
||||
|
||||
ds4cpp::TerminalDecision terminal = run_fake_postlude(
|
||||
{ds4cpp::TerminalCause::EngineError, false}, true, &counts);
|
||||
|
||||
check(terminal.cause == ds4cpp::TerminalCause::EngineError,
|
||||
"postlude:abandoned_engine_error_outcome");
|
||||
check(counts.finalize_attempts == 0,
|
||||
"postlude:abandoned_engine_error_attempts");
|
||||
check(counts.finalize_commits == 0,
|
||||
"postlude:abandoned_engine_error_commits");
|
||||
check(counts.cache_persists == 0,
|
||||
"postlude:abandoned_engine_error_cache");
|
||||
}
|
||||
|
||||
// Production mutation caught: failing to install the request callback would
|
||||
// make DS4 prompt synchronization unable to observe client cancellation.
|
||||
void test_cancel_callback_scope_installs_callback() {
|
||||
FakeCancelTarget target;
|
||||
int request_context = 42;
|
||||
|
||||
{
|
||||
ds4cpp::CancelCallbackScope scope(
|
||||
&target, fake_set_cancel, fake_cancel, &request_context);
|
||||
check(target.callback == fake_cancel, "cancel_scope:callback_installed");
|
||||
check(target.userdata == &request_context, "cancel_scope:userdata_installed");
|
||||
check(target.installs == 1, "cancel_scope:installed_once");
|
||||
}
|
||||
}
|
||||
|
||||
// Production mutation caught: failing to clear the callback at every scope
|
||||
// exit would leave DS4 pointing at a destroyed stack-owned ServerContext.
|
||||
void test_cancel_callback_scope_clears_callback() {
|
||||
FakeCancelTarget target;
|
||||
int request_context = 42;
|
||||
|
||||
{
|
||||
ds4cpp::CancelCallbackScope scope(
|
||||
&target, fake_set_cancel, fake_cancel, &request_context);
|
||||
}
|
||||
|
||||
check(target.callback == nullptr, "cancel_scope:callback_cleared");
|
||||
check(target.userdata == nullptr, "cancel_scope:userdata_cleared");
|
||||
check(target.clears == 1, "cancel_scope:cleared_once");
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
int main() {
|
||||
test_active_request_continues_and_finalizes();
|
||||
test_context_cancellation_stops_without_finalizing();
|
||||
test_stream_write_abort_stops_without_finalizing();
|
||||
test_cancellation_and_write_abort_are_independent_or_conditions();
|
||||
test_route_wait_pending();
|
||||
test_route_wait_ready();
|
||||
test_route_wait_error();
|
||||
test_route_wait_cancellation();
|
||||
test_route_wait_cancellation_precedes_error();
|
||||
test_terminal_success();
|
||||
test_terminal_sync_interruption_is_cancelled();
|
||||
test_terminal_engine_error();
|
||||
test_terminal_context_abandonment();
|
||||
test_terminal_write_abandonment();
|
||||
test_terminal_engine_error_precedes_late_abandonment();
|
||||
test_abandoned_engine_error_keeps_internal_without_finalizing();
|
||||
test_active_engine_error_still_finalizes();
|
||||
test_postlude_active_success_commits_then_persists();
|
||||
test_postlude_cancellation_skips_all_side_effects();
|
||||
test_postlude_finalize_failure_cancels_without_commit_or_cache();
|
||||
test_postlude_active_engine_error_finalizes_and_persists();
|
||||
test_postlude_engine_error_finalize_failure_preserves_internal();
|
||||
test_postlude_abandoned_engine_error_skips_all_side_effects();
|
||||
test_cancel_callback_scope_installs_callback();
|
||||
test_cancel_callback_scope_clears_callback();
|
||||
|
||||
if (failures == 0) {
|
||||
std::fprintf(stderr, "all request_lifecycle checks passed\n");
|
||||
return 0;
|
||||
}
|
||||
std::fprintf(stderr, "%d check(s) failed\n", failures);
|
||||
return 1;
|
||||
}
|
||||
@@ -1,5 +1,5 @@
|
||||
|
||||
IK_LLAMA_VERSION?=8337e4cd3861406fc04e0854b1409cd1b027fbc9
|
||||
IK_LLAMA_VERSION?=fe215a8ccdce6b844d2a3a3bbde08ae76a6284bf
|
||||
LLAMA_REPO?=https://github.com/ikawrakow/ik_llama.cpp
|
||||
|
||||
CMAKE_ARGS?=
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
|
||||
LLAMA_VERSION?=d59d455fd8ea09e5a2e87ce2a9d668267ffb5ccd
|
||||
LLAMA_VERSION?=67672dc5b76f8bc17785a19d3dc6d1463fc2902c
|
||||
LLAMA_REPO?=https://github.com/ggerganov/llama.cpp
|
||||
|
||||
CMAKE_ARGS?=
|
||||
|
||||
@@ -88,6 +88,12 @@ using grpc::ServerBuilder;
|
||||
using grpc::ServerContext;
|
||||
using grpc::Status;
|
||||
|
||||
#if LOCALAI_HAS_MTMD_INIT_OPT
|
||||
#define LOCALAI_MTMD_INIT_OPT_ARG(value) , value
|
||||
#else
|
||||
#define LOCALAI_MTMD_INIT_OPT_ARG(value)
|
||||
#endif
|
||||
|
||||
// gRPC bearer token auth for distributed mode.
|
||||
// Reads LOCALAI_GRPC_AUTH_TOKEN from the environment. When set, rejects
|
||||
// requests without a matching "authorization: Bearer <token>" metadata header.
|
||||
@@ -294,7 +300,7 @@ json parse_options(bool streaming, const backend::PredictOptions* predict, const
|
||||
} else {
|
||||
SRV_WRN("[TOOLS DEBUG] parse_options: Parsed tools JSON is not an array: %s\n", tools_json.dump().c_str());
|
||||
}
|
||||
} catch (const json::parse_error& e) {
|
||||
} catch (const common_json_error& e) {
|
||||
SRV_WRN("Failed to parse tools JSON from proto: %s\n", e.what());
|
||||
SRV_WRN("[TOOLS DEBUG] parse_options: Tools string that failed to parse: %s\n", predict->tools().c_str());
|
||||
}
|
||||
@@ -324,7 +330,7 @@ json parse_options(bool streaming, const backend::PredictOptions* predict, const
|
||||
SRV_DBG("[TOOLS DEBUG] Received tool_choice object from Go layer: %s\n", tool_choice_json.dump().c_str());
|
||||
}
|
||||
SRV_INF("Extracted tool_choice from proto: %s\n", predict->toolchoice().c_str());
|
||||
} catch (const json::parse_error& e) {
|
||||
} catch (const common_json_error& e) {
|
||||
// If parsing fails, treat as string
|
||||
data["tool_choice"] = predict->toolchoice();
|
||||
SRV_INF("Extracted tool_choice as string: %s\n", predict->toolchoice().c_str());
|
||||
@@ -353,7 +359,7 @@ json parse_options(bool streaming, const backend::PredictOptions* predict, const
|
||||
// Add to data - llama.cpp server expects it as an object (map)
|
||||
data["logit_bias"] = logit_bias_json;
|
||||
SRV_INF("Using logit_bias: %s\n", predict->logitbias().c_str());
|
||||
} catch (const json::parse_error& e) {
|
||||
} catch (const common_json_error& e) {
|
||||
SRV_ERR("Failed to parse logit_bias JSON from proto: %s\n", e.what());
|
||||
}
|
||||
}
|
||||
@@ -398,7 +404,10 @@ json parse_options(bool streaming, const backend::PredictOptions* predict, const
|
||||
});
|
||||
}
|
||||
|
||||
data["stop"] = predict->stopprompts();
|
||||
data["stop"] = json::array();
|
||||
for (const auto & stop : predict->stopprompts()) {
|
||||
data["stop"].push_back(stop);
|
||||
}
|
||||
// data["n_probs"] = predict->nprobs();
|
||||
//TODO: images,
|
||||
|
||||
@@ -1116,14 +1125,16 @@ static void params_parse(server_context& /*ctx_server*/, const backend::ModelOpt
|
||||
try {
|
||||
int n = std::stoi(optval_str);
|
||||
if (n < 0) n = 0;
|
||||
// Keep override-name storage alive for the lifetime of the params struct
|
||||
// (mirrors upstream arg.cpp behavior with a function-local static).
|
||||
#if LOCALAI_HAS_N_CPU_FFN_HELPER
|
||||
llm_add_n_cpu_ffn_overrides(n, LLM_FFN_EXPS_REGEX, params.speculative.draft.tensor_buft_overrides);
|
||||
#else
|
||||
static std::list<std::string> buft_overrides_draft;
|
||||
for (int i = 0; i < n; ++i) {
|
||||
buft_overrides_draft.push_back(llm_ffn_exps_block_regex(i));
|
||||
params.speculative.draft.tensor_buft_overrides.push_back(
|
||||
{buft_overrides_draft.back().c_str(), ggml_backend_cpu_buffer_type()});
|
||||
}
|
||||
#endif
|
||||
} catch (...) {}
|
||||
}
|
||||
|
||||
@@ -1141,14 +1152,16 @@ static void params_parse(server_context& /*ctx_server*/, const backend::ModelOpt
|
||||
try {
|
||||
int n = std::stoi(optval_str);
|
||||
if (n < 0) n = 0;
|
||||
// Keep override-name storage alive for the lifetime of the
|
||||
// params struct (mirrors upstream arg.cpp's function-local static).
|
||||
#if LOCALAI_HAS_N_CPU_FFN_HELPER
|
||||
llm_add_n_cpu_ffn_overrides(n, LLM_FFN_EXPS_REGEX, params.tensor_buft_overrides);
|
||||
#else
|
||||
static std::list<std::string> buft_overrides_main;
|
||||
for (int i = 0; i < n; ++i) {
|
||||
buft_overrides_main.push_back(llm_ffn_exps_block_regex(i));
|
||||
params.tensor_buft_overrides.push_back(
|
||||
{buft_overrides_main.back().c_str(), ggml_backend_cpu_buffer_type()});
|
||||
}
|
||||
#endif
|
||||
} catch (...) {}
|
||||
}
|
||||
|
||||
@@ -1795,7 +1808,7 @@ public:
|
||||
for (int j = 0; j < request->audios_size(); j++) rin.audios.push_back(request->audios(j));
|
||||
for (int j = 0; j < request->videos_size(); j++) rin.videos.push_back(request->videos(j));
|
||||
}
|
||||
messages_json.push_back(llama_grpc::build_reconstructed_message(rin));
|
||||
messages_json.push_back(json::parse(llama_grpc::build_reconstructed_message(rin).dump()));
|
||||
}
|
||||
|
||||
// Final safety check: Ensure no message has null content (Jinja templates require strings)
|
||||
@@ -1988,7 +2001,7 @@ public:
|
||||
if (!body_json.contains("chat_template_kwargs")) {
|
||||
body_json["chat_template_kwargs"] = json::object();
|
||||
}
|
||||
for (auto& el : ctk.items()) {
|
||||
for (auto el : ctk.items()) {
|
||||
body_json["chat_template_kwargs"][el.key()] = el.value();
|
||||
}
|
||||
}
|
||||
@@ -2074,30 +2087,27 @@ public:
|
||||
// If not using chat templates, extract files from image_data/audio_data fields
|
||||
// (If using chat templates, files were already extracted by oaicompat_chat_params_parse)
|
||||
if (!request->usetokenizertemplate() || request->messages_size() == 0 || ctx_server.impl->chat_params.tmpls == nullptr) {
|
||||
const auto &images_data = data.find("image_data");
|
||||
if (images_data != data.end() && images_data->is_array())
|
||||
if (data.contains("image_data") && data.at("image_data").is_array())
|
||||
{
|
||||
for (const auto &img : *images_data)
|
||||
for (const auto &img : data.at("image_data"))
|
||||
{
|
||||
auto decoded_data = base64_decode(img["data"].get<std::string>());
|
||||
files.push_back(decoded_data);
|
||||
}
|
||||
}
|
||||
|
||||
const auto &audio_data = data.find("audio_data");
|
||||
if (audio_data != data.end() && audio_data->is_array())
|
||||
if (data.contains("audio_data") && data.at("audio_data").is_array())
|
||||
{
|
||||
for (const auto &audio : *audio_data)
|
||||
for (const auto &audio : data.at("audio_data"))
|
||||
{
|
||||
auto decoded_data = base64_decode(audio["data"].get<std::string>());
|
||||
files.push_back(decoded_data);
|
||||
}
|
||||
}
|
||||
|
||||
const auto &video_data = data.find("video_data");
|
||||
if (video_data != data.end() && video_data->is_array())
|
||||
if (data.contains("video_data") && data.at("video_data").is_array())
|
||||
{
|
||||
for (const auto &video : *video_data)
|
||||
for (const auto &video : data.at("video_data"))
|
||||
{
|
||||
auto decoded_data = base64_decode(video["data"].get<std::string>());
|
||||
files.push_back(decoded_data);
|
||||
@@ -2111,10 +2121,10 @@ public:
|
||||
std::vector<server_tokens> inputs;
|
||||
if (has_mtmd) {
|
||||
// multimodal
|
||||
inputs.push_back(process_mtmd_prompt(ctx_server.impl->mctx, prompt_str, files));
|
||||
inputs.push_back(process_mtmd_prompt(ctx_server.impl->mctx, prompt_str, files LOCALAI_MTMD_INIT_OPT_ARG(ctx_server.impl->init_opt)));
|
||||
} else {
|
||||
// Everything else, including multimodal completions.
|
||||
inputs = tokenize_input_prompts(ctx_server.impl->vocab, ctx_server.impl->mctx, prompt_str, true, true);
|
||||
inputs = tokenize_input_prompts(ctx_server.impl->vocab, ctx_server.impl->mctx, prompt_str, true, true LOCALAI_MTMD_INIT_OPT_ARG(ctx_server.impl->init_opt));
|
||||
}
|
||||
|
||||
tasks.reserve(inputs.size());
|
||||
@@ -2370,7 +2380,7 @@ public:
|
||||
for (int j = 0; j < request->audios_size(); j++) rin.audios.push_back(request->audios(j));
|
||||
for (int j = 0; j < request->videos_size(); j++) rin.videos.push_back(request->videos(j));
|
||||
}
|
||||
messages_json.push_back(llama_grpc::build_reconstructed_message(rin));
|
||||
messages_json.push_back(json::parse(llama_grpc::build_reconstructed_message(rin).dump()));
|
||||
}
|
||||
|
||||
// Final safety check: Ensure no message has null content (Jinja templates require strings)
|
||||
@@ -2563,7 +2573,7 @@ public:
|
||||
if (!body_json.contains("chat_template_kwargs")) {
|
||||
body_json["chat_template_kwargs"] = json::object();
|
||||
}
|
||||
for (auto& el : ctk.items()) {
|
||||
for (auto el : ctk.items()) {
|
||||
body_json["chat_template_kwargs"][el.key()] = el.value();
|
||||
}
|
||||
}
|
||||
@@ -2649,11 +2659,10 @@ public:
|
||||
// If not using chat templates, extract files from image_data/audio_data fields
|
||||
// (If using chat templates, files were already extracted by oaicompat_chat_params_parse)
|
||||
if (!request->usetokenizertemplate() || request->messages_size() == 0 || ctx_server.impl->chat_params.tmpls == nullptr) {
|
||||
const auto &images_data = data.find("image_data");
|
||||
if (images_data != data.end() && images_data->is_array())
|
||||
if (data.contains("image_data") && data.at("image_data").is_array())
|
||||
{
|
||||
std::cout << "[PREDICT] Processing " << images_data->size() << " images" << std::endl;
|
||||
for (const auto &img : *images_data)
|
||||
std::cout << "[PREDICT] Processing " << data.at("image_data").size() << " images" << std::endl;
|
||||
for (const auto &img : data.at("image_data"))
|
||||
{
|
||||
std::cout << "[PREDICT] Processing image" << std::endl;
|
||||
auto decoded_data = base64_decode(img["data"].get<std::string>());
|
||||
@@ -2661,20 +2670,18 @@ public:
|
||||
}
|
||||
}
|
||||
|
||||
const auto &audio_data = data.find("audio_data");
|
||||
if (audio_data != data.end() && audio_data->is_array())
|
||||
if (data.contains("audio_data") && data.at("audio_data").is_array())
|
||||
{
|
||||
for (const auto &audio : *audio_data)
|
||||
for (const auto &audio : data.at("audio_data"))
|
||||
{
|
||||
auto decoded_data = base64_decode(audio["data"].get<std::string>());
|
||||
files.push_back(decoded_data);
|
||||
}
|
||||
}
|
||||
|
||||
const auto &video_data = data.find("video_data");
|
||||
if (video_data != data.end() && video_data->is_array())
|
||||
if (data.contains("video_data") && data.at("video_data").is_array())
|
||||
{
|
||||
for (const auto &video : *video_data)
|
||||
for (const auto &video : data.at("video_data"))
|
||||
{
|
||||
auto decoded_data = base64_decode(video["data"].get<std::string>());
|
||||
files.push_back(decoded_data);
|
||||
@@ -2689,10 +2696,10 @@ public:
|
||||
std::vector<server_tokens> inputs;
|
||||
if (has_mtmd) {
|
||||
// multimodal
|
||||
inputs.push_back(process_mtmd_prompt(ctx_server.impl->mctx, prompt_str, files));
|
||||
inputs.push_back(process_mtmd_prompt(ctx_server.impl->mctx, prompt_str, files LOCALAI_MTMD_INIT_OPT_ARG(ctx_server.impl->init_opt)));
|
||||
} else {
|
||||
// Everything else, including multimodal completions.
|
||||
inputs = tokenize_input_prompts(ctx_server.impl->vocab, ctx_server.impl->mctx, prompt_str, true, true);
|
||||
inputs = tokenize_input_prompts(ctx_server.impl->vocab, ctx_server.impl->mctx, prompt_str, true, true LOCALAI_MTMD_INIT_OPT_ARG(ctx_server.impl->init_opt));
|
||||
}
|
||||
|
||||
tasks.reserve(inputs.size());
|
||||
@@ -2879,7 +2886,7 @@ public:
|
||||
json prompt = body.at("embeddings");
|
||||
|
||||
|
||||
auto tokenized_prompts = tokenize_input_prompts(ctx_server.impl->vocab, ctx_server.impl->mctx, prompt, true, true);
|
||||
auto tokenized_prompts = tokenize_input_prompts(ctx_server.impl->vocab, ctx_server.impl->mctx, prompt, true, true LOCALAI_MTMD_INIT_OPT_ARG(ctx_server.impl->init_opt));
|
||||
for (const auto & tokens : tokenized_prompts) {
|
||||
// this check is necessary for models that do not add BOS token to the input
|
||||
if (tokens.empty()) {
|
||||
@@ -2984,7 +2991,7 @@ public:
|
||||
|
||||
tasks.reserve(documents.size());
|
||||
for (size_t i = 0; i < documents.size(); i++) {
|
||||
auto tmp = format_prompt_rerank(ctx_server.impl->model_tgt, ctx_server.impl->vocab, ctx_server.impl->mctx, request->query(), documents[i]);
|
||||
auto tmp = format_prompt_rerank(ctx_server.impl->model_tgt, ctx_server.impl->vocab, ctx_server.impl->mctx, request->query(), documents[i] LOCALAI_MTMD_INIT_OPT_ARG(ctx_server.impl->init_opt));
|
||||
server_task task = server_task(SERVER_TASK_TYPE_RERANK);
|
||||
task.id = rd.queue_tasks.get_new_id();
|
||||
task.index = i;
|
||||
@@ -3005,7 +3012,7 @@ public:
|
||||
}
|
||||
|
||||
// Collect responses
|
||||
json responses = json::array();
|
||||
std::vector<json> responses;
|
||||
for (auto & res : all_results.results) {
|
||||
GGML_ASSERT(dynamic_cast<server_task_result_rerank*>(res.get()) != nullptr);
|
||||
responses.push_back(res->to_json());
|
||||
@@ -3018,7 +3025,7 @@ public:
|
||||
// Crop results by request.top_n if specified
|
||||
int top_n = request->top_n();
|
||||
if (top_n > 0 && top_n < static_cast<int>(responses.size())) {
|
||||
responses = json(responses.begin(), responses.begin() + top_n);
|
||||
responses.resize(top_n);
|
||||
}
|
||||
// Set usage information
|
||||
backend::Usage* usage = rerankResult->mutable_usage();
|
||||
@@ -3065,7 +3072,7 @@ public:
|
||||
return grpc::Status(grpc::StatusCode::INVALID_ARGUMENT, opts.error);
|
||||
}
|
||||
|
||||
auto wrapper = mtmd_helper_bitmap_init_from_file(ctx_server.impl->mctx, opts.voice_path.c_str(), false);
|
||||
auto wrapper = mtmd_helper_bitmap_init_from_file(ctx_server.impl->mctx, opts.voice_path.c_str(), false LOCALAI_MTMD_INIT_OPT_ARG(ctx_server.impl->init_opt));
|
||||
if (!wrapper.bitmap) {
|
||||
return grpc::Status(grpc::StatusCode::INVALID_ARGUMENT,
|
||||
"failed to read speaker reference audio: " + opts.voice_path);
|
||||
|
||||
@@ -52,14 +52,15 @@ inline nlohmann::ordered_json normalize_message_content(const std::string& role,
|
||||
// (#7528). A multimodal user message legitimately carries a typed-part array
|
||||
// ({type:text}, {type:image_url}, ...), which must be left intact. Shared by the
|
||||
// streaming and non-streaming paths so this invariant cannot drift between them.
|
||||
inline void normalize_template_message(nlohmann::ordered_json& msg) {
|
||||
template <typename Json>
|
||||
inline void normalize_template_message(Json& msg) {
|
||||
if (!msg.contains("content")) {
|
||||
msg["content"] = ""; // templates expect the field to exist
|
||||
return;
|
||||
}
|
||||
nlohmann::ordered_json& content = msg["content"];
|
||||
auto& content = msg["content"];
|
||||
const std::string role = (msg.contains("role") && msg["role"].is_string())
|
||||
? msg["role"].get<std::string>()
|
||||
? msg["role"].template get<std::string>()
|
||||
: std::string();
|
||||
if (content.is_null()) {
|
||||
content = ""; // #7324: null would crash content[:N] slicing
|
||||
|
||||
@@ -6,10 +6,9 @@ Subject: [PATCH 1/2] score-patch
|
||||
---
|
||||
common/common.cpp | 6 +-
|
||||
common/common.h | 3 +
|
||||
tools/CMakeLists.txt | 1 +
|
||||
tools/server/server-context.cpp | 358 +++++++++++++++++++++++++++++++-
|
||||
tools/server/server-task.h | 47 +++++
|
||||
5 files changed, 406 insertions(+), 9 deletions(-)
|
||||
4 files changed, 405 insertions(+), 9 deletions(-)
|
||||
|
||||
diff --git a/common/common.cpp b/common/common.cpp
|
||||
index 2e3f14c..0cec0dc 100644
|
||||
@@ -42,15 +41,6 @@ index 878534d..4001df2 100644
|
||||
int32_t n_sequences = 1; // number of sequences to decode
|
||||
int32_t n_outputs_max = 0; // max outputs in a batch (0 = n_batch)
|
||||
int32_t n_outputs_max_per_seq = 1; // max outputs per sequence
|
||||
diff --git a/tools/CMakeLists.txt b/tools/CMakeLists.txt
|
||||
index 780df32..1d2fe8f 100644
|
||||
--- a/tools/CMakeLists.txt
|
||||
+++ b/tools/CMakeLists.txt
|
||||
@@ -41,3 +41,4 @@ else()
|
||||
add_subdirectory(fit-params)
|
||||
add_subdirectory(results)
|
||||
endif()
|
||||
+add_subdirectory(grpc-server)
|
||||
diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp
|
||||
index 3b5f6a1..d0e18e6 100644
|
||||
--- a/tools/server/server-context.cpp
|
||||
|
||||
@@ -659,7 +659,7 @@ index 9069463fe..b7fa1e534 100644
|
||||
+ }
|
||||
+
|
||||
+ if (speaker_ref_len > 0) {
|
||||
+ auto wrapper = mtmd_helper_bitmap_init_from_buf(ctx_server.mctx, speaker_ref_data, speaker_ref_len, false);
|
||||
+ auto wrapper = mtmd_helper_bitmap_init_from_buf(ctx_server.mctx, speaker_ref_data, speaker_ref_len, false, ctx_server.init_opt);
|
||||
+ if (!wrapper.bitmap) {
|
||||
+ res->error(format_error_response("failed to decode \"speaker_ref\"", ERROR_TYPE_INVALID_REQUEST));
|
||||
+ return res;
|
||||
|
||||
@@ -15,6 +15,30 @@ if [ -d "patches" ]; then
|
||||
done
|
||||
fi
|
||||
|
||||
## Apple RDMA link fixup.
|
||||
|
||||
## ggml-rpc hands Apple's librdma to the linker with
|
||||
## target_link_options(ggml-rpc PRIVATE "LINKER:-weak_library,..."). Link options are not
|
||||
## a usage requirement of a static library, so in our BUILD_SHARED_LIBS=OFF build the flag
|
||||
## dies with libggml-rpc.a and every ibv_* symbol transport-apple.cpp reaches for comes out
|
||||
## undefined when grpc-server and ggml-rpc-server link. Re-declare the same weak link as
|
||||
## INTERFACE so it travels to whoever links the static library.
|
||||
##
|
||||
## Guarded on the marker so a second prepare.sh over the same checkout is a no-op, and on
|
||||
## GGML_RPC_RDMA_APPLE so forks that branched before the Apple RDMA transport (turboquant,
|
||||
## bonsai) are left alone.
|
||||
RPC_CMAKE=llama.cpp/ggml/src/ggml-rpc/CMakeLists.txt
|
||||
if [ -f "$RPC_CMAKE" ] && grep -q "GGML_RPC_RDMA_APPLE" "$RPC_CMAKE" && ! grep -q "LOCALAI_RDMA_IFACE" "$RPC_CMAKE"; then
|
||||
echo "==> ggml-rpc carries the Apple RDMA transport, re-declaring its weak librdma link as INTERFACE"
|
||||
cat >> "$RPC_CMAKE" <<'EOF'
|
||||
|
||||
# LOCALAI_RDMA_IFACE: added by backend/cpp/llama-cpp/prepare.sh
|
||||
if (GGML_RPC_RDMA AND APPLE AND NOT BUILD_SHARED_LIBS)
|
||||
target_link_options(ggml-rpc INTERFACE "LINKER:-weak_library,${RDMA_LIB}")
|
||||
endif()
|
||||
EOF
|
||||
fi
|
||||
|
||||
for file in $(ls llama.cpp/tools/server/); do
|
||||
cp -rfv llama.cpp/tools/server/$file llama.cpp/tools/grpc-server/
|
||||
done
|
||||
@@ -61,11 +85,23 @@ if grep -q "server_metrics metrics;" llama.cpp/tools/server/server-task.h; then
|
||||
else
|
||||
HAS_SERVER_METRICS=0
|
||||
fi
|
||||
if grep -q "mtmd_helper_init_opt" llama.cpp/tools/mtmd/mtmd-helper.h; then
|
||||
HAS_MTMD_INIT_OPT=1
|
||||
else
|
||||
HAS_MTMD_INIT_OPT=0
|
||||
fi
|
||||
if grep -q "llm_add_n_cpu_ffn_overrides" llama.cpp/common/common.h; then
|
||||
HAS_N_CPU_FFN_HELPER=1
|
||||
else
|
||||
HAS_N_CPU_FFN_HELPER=0
|
||||
fi
|
||||
cat > llama.cpp/tools/grpc-server/llama_compat.h <<EOF
|
||||
// Generated by backend/cpp/llama-cpp/prepare.sh. Do not edit.
|
||||
#pragma once
|
||||
#define LOCALAI_LEGACY_LOAD_MODE ${LEGACY_LOAD_MODE}
|
||||
#define LOCALAI_HAS_SERVER_METRICS ${HAS_SERVER_METRICS}
|
||||
#define LOCALAI_HAS_MTMD_INIT_OPT ${HAS_MTMD_INIT_OPT}
|
||||
#define LOCALAI_HAS_N_CPU_FFN_HELPER ${HAS_N_CPU_FFN_HELPER}
|
||||
EOF
|
||||
|
||||
set +e
|
||||
|
||||
@@ -8,6 +8,8 @@
|
||||
# so the grpc-server option parser skips the two references to
|
||||
# common_params::checkpoint_min_step (the default and the option handler).
|
||||
# That field does not exist in the fork yet; drop this once it does.
|
||||
# 3. Use nlohmann's parse_error type in JSON catch clauses because the fork
|
||||
# predates upstream's common_json_error wrapper.
|
||||
#
|
||||
# The fork used to lag upstream on the whole common_params_speculative refactor
|
||||
# (ggml-org/llama.cpp#22397/#22838/#22964), the model_tgt rename (#22838) and
|
||||
@@ -100,4 +102,16 @@ else
|
||||
echo "==> LOCALAI_TURBOQUANT_NO_CHECKPOINT_MIN_STEP define OK"
|
||||
fi
|
||||
|
||||
# 3. The shared source follows current upstream and catches common_json_error.
|
||||
# TurboQuant still exposes nlohmann::json directly, so its equivalent parse
|
||||
# failures use json::parse_error instead.
|
||||
if grep -q 'common_json_error' "$SRC"; then
|
||||
echo "==> patching $SRC to use the TurboQuant JSON exception type"
|
||||
awk '{ gsub(/common_json_error/, "json::parse_error"); print }' "$SRC" > "$SRC.tmp"
|
||||
mv "$SRC.tmp" "$SRC"
|
||||
echo "==> TurboQuant JSON exception patch OK"
|
||||
else
|
||||
echo "==> $SRC already uses a TurboQuant-compatible JSON exception type, skipping"
|
||||
fi
|
||||
|
||||
echo "==> all patches applied"
|
||||
@@ -8,7 +8,7 @@ JOBS?=$(shell nproc --ignore=1)
|
||||
|
||||
# CrispASR version (release tag)
|
||||
CRISPASR_REPO?=https://github.com/CrispStrobe/CrispASR
|
||||
CRISPASR_VERSION?=a153b09b37c90cd55cd9336fccbdf3ba7a289596
|
||||
CRISPASR_VERSION?=301acd87b036764973b8bfba71e0a21818036d33
|
||||
SO_TARGET?=libgocrispasr.so
|
||||
|
||||
CMAKE_ARGS+=-DBUILD_SHARED_LIBS=OFF
|
||||
|
||||
@@ -14,7 +14,7 @@ JOBS?=$(shell nproc --ignore=1)
|
||||
# It is kept alive by the upstream tag da2-support (survives a squash-merge);
|
||||
# repoint to the master merge commit once mudler/depth-anything.cpp PR #1 lands.
|
||||
DEPTHANYTHING_REPO?=https://github.com/mudler/depth-anything.cpp.git
|
||||
DEPTHANYTHING_VERSION?=54abd5c0abfd1f394e01cb3c38f2e3af4daedf85
|
||||
DEPTHANYTHING_VERSION?=14f7461d1f704761a038ac9f50dbde8fdb7275e2
|
||||
|
||||
ifeq ($(NATIVE),false)
|
||||
CMAKE_ARGS+=-DGGML_NATIVE=OFF
|
||||
|
||||
@@ -12,7 +12,7 @@
|
||||
# runs 'make -C backend/go/$(BACKEND) build' and then copies package/), so it
|
||||
# has to produce the binary and the package, not just the shared libraries.
|
||||
|
||||
NEMO_SPEECH_VERSION?=4f9676226f667d14608487df744f375db87127f8
|
||||
NEMO_SPEECH_VERSION?=ffa38cb2408f1e832a36d46fef5e3e1e80d07e6c
|
||||
NEMO_SPEECH_REPO?=https://github.com/NVIDIA/NeMo-Speech.cpp
|
||||
|
||||
GOCMD?=go
|
||||
|
||||
@@ -8,7 +8,7 @@ JOBS?=$(shell nproc --ignore=1)
|
||||
|
||||
# omnivoice.cpp version
|
||||
OMNIVOICE_REPO?=https://github.com/ServeurpersoCom/omnivoice.cpp
|
||||
OMNIVOICE_VERSION?=4f33af825d66e6ef1cb185e87b4589cacf747291
|
||||
OMNIVOICE_VERSION?=040c8b344d8c670ce1475194751d119b5ef82c78
|
||||
SO_TARGET?=libgomnivoicecpp.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?=97d2990807fe6d558e395f8764198d7c7e7b411c
|
||||
STABLEDIFFUSION_GGML_VERSION?=d04e8950c1ec8d30248cbe996682b3182fb1adf6
|
||||
|
||||
CMAKE_ARGS+=-DGGML_MAX_NAME=128
|
||||
|
||||
@@ -38,8 +38,11 @@ else ifeq ($(BUILD_TYPE),hipblas)
|
||||
ROCM_PATH ?= /opt/rocm
|
||||
export CXX=$(ROCM_HOME)/llvm/bin/clang++
|
||||
export CC=$(ROCM_HOME)/llvm/bin/clang
|
||||
AMDGPU_TARGETS?=gfx908,gfx90a,gfx942,gfx950,gfx1030,gfx1100,gfx1101,gfx1102,gfx1200,gfx1201
|
||||
CMAKE_ARGS+=-DSD_HIPBLAS=ON -DGGML_HIPBLAS=ON -DAMDGPU_TARGETS=$(AMDGPU_TARGETS)
|
||||
AMDGPU_TARGETS?=gfx908,gfx90a,gfx942,gfx950,gfx1030,gfx1100,gfx1101,gfx1102,gfx1151,gfx1200,gfx1201
|
||||
# SD_HIPBLAS turns on ggml's HIP backend itself; GGML_HIPBLAS is the name ggml
|
||||
# used before it was renamed to GGML_HIP, so passing it here only produced an
|
||||
# unused-variable warning.
|
||||
CMAKE_ARGS+=-DSD_HIPBLAS=ON -DAMDGPU_TARGETS=$(AMDGPU_TARGETS)
|
||||
else ifeq ($(BUILD_TYPE),vulkan)
|
||||
CMAKE_ARGS+=-DSD_VULKAN=ON -DGGML_VULKAN=ON
|
||||
else ifeq ($(BUILD_TYPE),metal)
|
||||
|
||||
@@ -401,7 +401,6 @@ int load_model(const char *model, char *model_path, char* options[], int threads
|
||||
const char *params_backend_arg = "";
|
||||
const char *rpc_servers_arg = "";
|
||||
const char *max_vram_arg = "";
|
||||
bool stream_layers = false;
|
||||
|
||||
int n_threads = threads;
|
||||
enum sd_type_t wtype = SD_TYPE_COUNT;
|
||||
@@ -510,7 +509,10 @@ int load_model(const char *model, char *model_path, char* options[], int threads
|
||||
if (!strcmp(optname, "params_backend")) params_backend_arg = strdup(optval);
|
||||
if (!strcmp(optname, "rpc_servers")) rpc_servers_arg = strdup(optval);
|
||||
if (!strcmp(optname, "max_vram")) max_vram_arg = strdup(optval);
|
||||
if (!strcmp(optname, "stream_layers")) stream_layers = (strcmp(optval, "true") == 0 || strcmp(optval, "1") == 0);
|
||||
if (!strcmp(optname, "stream_layers")) {
|
||||
// Retained as a no-op for existing configurations. Upstream now
|
||||
// selects segmented weight streaming automatically.
|
||||
}
|
||||
|
||||
// vae_decode_only is still accepted for backwards compatibility with
|
||||
// existing gallery configs, but upstream dropped the option (the model
|
||||
@@ -650,11 +652,9 @@ int load_model(const char *model, char *model_path, char* options[], int threads
|
||||
ctx_params.rpc_servers = env_rpc_servers;
|
||||
}
|
||||
}
|
||||
// max_vram: GiB budget or per-backend spec for graph-cut segmented param
|
||||
// offload ("0" = disabled, "-1" = auto). stream_layers only has effect when
|
||||
// max_vram is set.
|
||||
// max_vram is an optional GiB budget or per-backend spec for automatic
|
||||
// graph-cut execution. A zero value uses the live free-VRAM budget.
|
||||
if (strlen(max_vram_arg) > 0) ctx_params.max_vram = max_vram_arg;
|
||||
ctx_params.stream_layers = stream_layers;
|
||||
ctx_params.diffusion_flash_attn = diffusion_flash_attn;
|
||||
ctx_params.tae_preview_only = tae_preview_only;
|
||||
ctx_params.diffusion_conv_direct = diffusion_conv_direct;
|
||||
@@ -1438,4 +1438,3 @@ int unload() {
|
||||
free_sd_ctx(sd_c);
|
||||
return 0;
|
||||
}
|
||||
|
||||
@@ -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?=438305e1577768ec0f75729456a4c8b9f425e2ee
|
||||
VLLM_CPP_VERSION?=6bf3abb580982f4fd2e4525ef37802ee0ce28981
|
||||
|
||||
# MLX GEMM provider (darwin/metal only; see the metal branch below for why).
|
||||
# Consumed as the prebuilt pip wheel: building MLX from source needs `xcrun
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
package main
|
||||
|
||||
// purego bindings for the vllm.cpp stable C ABI (include/vllm.h, ABI v21).
|
||||
// purego bindings for the vllm.cpp stable C ABI (include/vllm.h, ABI v23).
|
||||
//
|
||||
// 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
|
||||
@@ -21,7 +21,7 @@ import (
|
||||
// the header of the VLLM_CPP_VERSION pinned in the Makefile: the build checks
|
||||
// the two against each other, because a mismatch is only caught at runtime by
|
||||
// registerLib, where it takes the backend down on every load (issue #11379).
|
||||
const abiVersion = 21
|
||||
const abiVersion = 23
|
||||
|
||||
// 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
|
||||
@@ -83,6 +83,7 @@ type cModelParams struct {
|
||||
LanguageModelOnly int32 // 0 = multimodal inputs enabled (ABI v19)
|
||||
_ [4]byte
|
||||
LimitMMPerPrompt uintptr // const char* JSON; NULL = default limits (ABI v19)
|
||||
MMProjPath uintptr // const char*; NULL = no GGUF projector (ABI v22)
|
||||
}
|
||||
|
||||
// cSamplingParams mirrors vllm_sampling_params (structured fields included).
|
||||
|
||||
@@ -128,9 +128,40 @@ func parseOptions(opts *pb.ModelOptions) loadOptions {
|
||||
lo := loadOptions{}
|
||||
applyOptionsList(&lo, opts.GetOptions())
|
||||
applyEngineArgs(&lo, opts.GetEngineArgs())
|
||||
applyDraftModelOption(&lo, opts.GetOptions())
|
||||
return lo
|
||||
}
|
||||
|
||||
// applyDraftModelOption binds a managed companion snapshot after engine_args
|
||||
// has supplied the speculative document. Companion paths do not exist until
|
||||
// LocalAI materializes the artifact, so they must replace the gallery's static
|
||||
// repository reference without disturbing the method or token budget.
|
||||
func applyDraftModelOption(lo *loadOptions, options []string) {
|
||||
if strings.TrimSpace(lo.speculativeConfig) == "" {
|
||||
return
|
||||
}
|
||||
var draftModel string
|
||||
for _, option := range options {
|
||||
key, value, found := strings.Cut(option, ":")
|
||||
if found && strings.TrimSpace(key) == "draft_model" {
|
||||
draftModel = strings.TrimSpace(value)
|
||||
}
|
||||
}
|
||||
if draftModel == "" {
|
||||
return
|
||||
}
|
||||
|
||||
var spec map[string]any
|
||||
if err := json.Unmarshal([]byte(lo.speculativeConfig), &spec); err != nil {
|
||||
return
|
||||
}
|
||||
spec["model"] = draftModel
|
||||
encoded, err := json.Marshal(spec)
|
||||
if err == nil {
|
||||
lo.speculativeConfig = string(encoded)
|
||||
}
|
||||
}
|
||||
|
||||
// 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) {
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
)
|
||||
|
||||
var _ = Describe("managed DFlash companion options", func() {
|
||||
It("replaces only the draft model in an existing speculative configuration", func() {
|
||||
managedPath := ".artifacts/huggingface/0123456789abcdef/snapshot"
|
||||
lo := parseOptions(&pb.ModelOptions{
|
||||
Options: []string{"draft_model:" + managedPath},
|
||||
EngineArgs: `{
|
||||
"speculative_config": {
|
||||
"method": "dflash",
|
||||
"model": "Mia-AiLab/Qwen3.8-27B-DFlash2-EXL3-5.0bpw",
|
||||
"num_speculative_tokens": 7
|
||||
}
|
||||
}`,
|
||||
})
|
||||
|
||||
Expect(lo.speculativeConfig).To(MatchJSON(`{
|
||||
"method": "dflash",
|
||||
"model": ".artifacts/huggingface/0123456789abcdef/snapshot",
|
||||
"num_speculative_tokens": 7
|
||||
}`))
|
||||
})
|
||||
|
||||
It("ignores a draft companion when speculative decoding is not configured", func() {
|
||||
lo := parseOptions(&pb.ModelOptions{
|
||||
Options: []string{"draft_model:.artifacts/huggingface/0123456789abcdef/snapshot"},
|
||||
})
|
||||
|
||||
Expect(lo.speculativeConfig).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
@@ -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 v21)
|
||||
// The Go POD mirrors must match the C struct layout of vllm.h (ABI v23)
|
||||
// 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() {
|
||||
@@ -24,7 +24,7 @@ var _ = Describe("C ABI struct mirrors", func() {
|
||||
// VLLM_ABI_VERSION in the vllm.h of VLLM_CPP_VERSION (Makefile).
|
||||
// Moving the pin past this without growing the mirrors below ships a
|
||||
// backend that refuses every load at startup (issue #11379).
|
||||
Expect(abiVersion).To(Equal(21))
|
||||
Expect(abiVersion).To(Equal(23))
|
||||
})
|
||||
|
||||
It("cModelParams matches vllm_model_params", func() {
|
||||
@@ -51,7 +51,8 @@ var _ = Describe("C ABI struct mirrors", func() {
|
||||
Expect(unsafe.Offsetof(p.KVCacheMemoryBytes)).To(Equal(uintptr(104)))
|
||||
Expect(unsafe.Offsetof(p.LanguageModelOnly)).To(Equal(uintptr(112)))
|
||||
Expect(unsafe.Offsetof(p.LimitMMPerPrompt)).To(Equal(uintptr(120)))
|
||||
Expect(unsafe.Sizeof(p)).To(Equal(uintptr(128)))
|
||||
Expect(unsafe.Offsetof(p.MMProjPath)).To(Equal(uintptr(128)))
|
||||
Expect(unsafe.Sizeof(p)).To(Equal(uintptr(136)))
|
||||
})
|
||||
|
||||
It("cSamplingParams matches vllm_sampling_params (ABI v8)", func() {
|
||||
|
||||
@@ -8,7 +8,7 @@ JOBS?=$(shell nproc --ignore=1)
|
||||
|
||||
# whisper.cpp version
|
||||
WHISPER_REPO?=https://github.com/ggml-org/whisper.cpp
|
||||
WHISPER_CPP_VERSION?=4834a2327d008ace3ec5a9ed00f51454bcabbc1c
|
||||
WHISPER_CPP_VERSION?=52a939a2a762224e255d366c1182b2af4dd1a032
|
||||
SO_TARGET?=libgowhisper.so
|
||||
|
||||
CMAKE_ARGS+=-DBUILD_SHARED_LIBS=OFF
|
||||
|
||||
+12
-2
@@ -510,7 +510,7 @@
|
||||
default: "cpu-stablediffusion-ggml"
|
||||
nvidia: "cuda12-stablediffusion-ggml"
|
||||
intel: "intel-sycl-f16-stablediffusion-ggml"
|
||||
# amd: "rocm-stablediffusion-ggml"
|
||||
amd: "rocm-stablediffusion-ggml"
|
||||
vulkan: "vulkan-stablediffusion-ggml"
|
||||
nvidia-l4t: "nvidia-l4t-arm64-stablediffusion-ggml"
|
||||
metal: "metal-stablediffusion-ggml"
|
||||
@@ -2109,7 +2109,7 @@
|
||||
default: "cpu-stablediffusion-ggml-development"
|
||||
nvidia: "cuda12-stablediffusion-ggml-development"
|
||||
intel: "intel-sycl-f16-stablediffusion-ggml-development"
|
||||
# amd: "rocm-stablediffusion-ggml-development"
|
||||
amd: "rocm-stablediffusion-ggml-development"
|
||||
vulkan: "vulkan-stablediffusion-ggml-development"
|
||||
nvidia-l4t: "nvidia-l4t-arm64-stablediffusion-ggml-development"
|
||||
metal: "metal-stablediffusion-ggml-development"
|
||||
@@ -3904,6 +3904,11 @@
|
||||
uri: "quay.io/go-skynet/local-ai-backends:latest-gpu-nvidia-cuda-12-stablediffusion-ggml"
|
||||
mirrors:
|
||||
- localai/localai-backends:latest-gpu-nvidia-cuda-12-stablediffusion-ggml
|
||||
- !!merge <<: *stablediffusionggml
|
||||
name: "rocm-stablediffusion-ggml"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:latest-gpu-rocm-hipblas-stablediffusion-ggml"
|
||||
mirrors:
|
||||
- localai/localai-backends:latest-gpu-rocm-hipblas-stablediffusion-ggml
|
||||
- !!merge <<: *stablediffusionggml
|
||||
name: "intel-sycl-f32-stablediffusion-ggml"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:latest-gpu-intel-sycl-f32-stablediffusion-ggml"
|
||||
@@ -3917,6 +3922,11 @@
|
||||
uri: "quay.io/go-skynet/local-ai-backends:master-gpu-nvidia-cuda-12-stablediffusion-ggml"
|
||||
mirrors:
|
||||
- localai/localai-backends:master-gpu-nvidia-cuda-12-stablediffusion-ggml
|
||||
- !!merge <<: *stablediffusionggml
|
||||
name: "rocm-stablediffusion-ggml-development"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:master-gpu-rocm-hipblas-stablediffusion-ggml"
|
||||
mirrors:
|
||||
- localai/localai-backends:master-gpu-rocm-hipblas-stablediffusion-ggml
|
||||
- !!merge <<: *stablediffusionggml
|
||||
name: "intel-sycl-f32-stablediffusion-ggml-development"
|
||||
uri: "quay.io/go-skynet/local-ai-backends:master-gpu-intel-sycl-f32-stablediffusion-ggml"
|
||||
|
||||
@@ -37,6 +37,46 @@ def parse_options(options_list):
|
||||
return opts
|
||||
|
||||
|
||||
def attach_media_parts(messages_dicts, n_images=0, n_videos=0):
|
||||
"""Rebuild the last user message as content *parts* carrying media markers.
|
||||
|
||||
Backends that let the tokenizer do the templating hand plain string content
|
||||
to ``apply_chat_template``, but a chat template only emits the model's own
|
||||
media tokens (``<|vision_start|><|image_pad|><|vision_end|>`` for the
|
||||
Qwen-VL family, and the equivalents elsewhere) when the content is a list
|
||||
of parts. Without those markers the engine's multimodal processor finds
|
||||
nothing to substitute and silently discards the pixels, even though they
|
||||
were forwarded correctly out of band.
|
||||
|
||||
Returns a new list whose last user message has
|
||||
``[{"type": "image"} * n_images, {"type": "video"} * n_videos, text]`` as
|
||||
its content, or ``None`` when there is nothing to attach - no media, no
|
||||
user turn, or content that is already a list of parts - so the caller can
|
||||
keep using the original string-content list.
|
||||
"""
|
||||
if not n_images and not n_videos:
|
||||
return None
|
||||
idx = next(
|
||||
(
|
||||
i
|
||||
for i in reversed(range(len(messages_dicts)))
|
||||
if messages_dicts[i].get("role") == "user"
|
||||
),
|
||||
None,
|
||||
)
|
||||
if idx is None:
|
||||
return None
|
||||
text = messages_dicts[idx].get("content") or ""
|
||||
if not isinstance(text, str):
|
||||
return None
|
||||
parts = [{"type": "image"}] * n_images + [{"type": "video"}] * n_videos
|
||||
if text:
|
||||
parts.append({"type": "text", "text": text})
|
||||
patched = list(messages_dicts)
|
||||
patched[idx] = dict(patched[idx], content=parts)
|
||||
return patched
|
||||
|
||||
|
||||
def messages_to_dicts(proto_messages):
|
||||
"""Convert proto ``Message`` objects to dicts suitable for ``apply_chat_template``.
|
||||
|
||||
|
||||
@@ -14,7 +14,7 @@ import json
|
||||
import types
|
||||
import unittest
|
||||
|
||||
from python_utils import messages_to_dicts, parse_options
|
||||
from python_utils import attach_media_parts, messages_to_dicts, parse_options
|
||||
|
||||
|
||||
def _msg(**fields):
|
||||
@@ -118,5 +118,63 @@ class TestMessagesToDicts(unittest.TestCase):
|
||||
self.assertNotIn("tool_calls", out[0])
|
||||
|
||||
|
||||
class TestAttachMediaParts(unittest.TestCase):
|
||||
def test_image_marker_added_to_last_user_turn(self):
|
||||
messages = [
|
||||
{"role": "system", "content": "be brief"},
|
||||
{"role": "user", "content": "first"},
|
||||
{"role": "assistant", "content": "ok"},
|
||||
{"role": "user", "content": "how high is the water?"},
|
||||
]
|
||||
out = attach_media_parts(messages, n_images=1)
|
||||
self.assertEqual(
|
||||
out[3]["content"],
|
||||
[{"type": "image"}, {"type": "text", "text": "how high is the water?"}],
|
||||
)
|
||||
# Earlier turns and the input list itself are untouched.
|
||||
self.assertEqual(out[:3], messages[:3])
|
||||
self.assertEqual(messages[3]["content"], "how high is the water?")
|
||||
|
||||
def test_counts_and_order_images_then_videos(self):
|
||||
out = attach_media_parts(
|
||||
[{"role": "user", "content": "describe"}], n_images=2, n_videos=1
|
||||
)
|
||||
self.assertEqual(
|
||||
out[0]["content"],
|
||||
[
|
||||
{"type": "image"},
|
||||
{"type": "image"},
|
||||
{"type": "video"},
|
||||
{"type": "text", "text": "describe"},
|
||||
],
|
||||
)
|
||||
|
||||
def test_empty_text_yields_media_only_parts(self):
|
||||
out = attach_media_parts([{"role": "user", "content": ""}], n_images=1)
|
||||
self.assertEqual(out[0]["content"], [{"type": "image"}])
|
||||
|
||||
def test_other_message_keys_are_preserved(self):
|
||||
out = attach_media_parts(
|
||||
[{"role": "user", "content": "hi", "name": "bob"}], n_images=1
|
||||
)
|
||||
self.assertEqual(out[0]["name"], "bob")
|
||||
|
||||
def test_no_media_is_a_no_op(self):
|
||||
self.assertIsNone(attach_media_parts([{"role": "user", "content": "hi"}]))
|
||||
|
||||
def test_no_user_turn_is_a_no_op(self):
|
||||
self.assertIsNone(
|
||||
attach_media_parts([{"role": "system", "content": "hi"}], n_images=1)
|
||||
)
|
||||
|
||||
def test_content_already_parts_is_a_no_op(self):
|
||||
self.assertIsNone(
|
||||
attach_media_parts(
|
||||
[{"role": "user", "content": [{"type": "text", "text": "hi"}]}],
|
||||
n_images=1,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,4 +1,4 @@
|
||||
--extra-index-url https://download.pytorch.org/whl/xpu
|
||||
torch==2.13.0+xpu
|
||||
torch==2.14.0+xpu
|
||||
oneccl_bind_pt==2.8.0+xpu
|
||||
optimum[openvino]
|
||||
@@ -1,3 +1,3 @@
|
||||
grpcio==1.82.1
|
||||
grpcio==1.83.1
|
||||
protobuf
|
||||
grpcio-tools
|
||||
@@ -1,4 +1,4 @@
|
||||
grpcio==1.83.0
|
||||
grpcio==1.83.1
|
||||
protobuf
|
||||
certifi
|
||||
packaging==26.3
|
||||
@@ -122,6 +122,21 @@ from diffusers.schedulers import (
|
||||
UniPCMultistepScheduler,
|
||||
)
|
||||
|
||||
def select_device(request_cuda, device_option, cuda_available, xpu, mps_available):
|
||||
"""Pick the pipeline device. An explicit `device:` model option wins;
|
||||
otherwise CUDA is used whenever torch reports it available (ROCm
|
||||
builds included) or the model config forces it with `cuda: true`,
|
||||
keeping the pre-existing XPU/MPS overrides. CPU is the fallback, not
|
||||
the default."""
|
||||
if device_option:
|
||||
return device_option
|
||||
device = "cuda" if (request_cuda or cuda_available) else "cpu"
|
||||
if xpu:
|
||||
device = "xpu"
|
||||
if mps_available:
|
||||
device = "mps"
|
||||
return device
|
||||
|
||||
def is_float(s):
|
||||
"""Check if a string can be converted to float."""
|
||||
try:
|
||||
@@ -627,12 +642,13 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
# modify LoraAdapter to be relative to modelFileBase
|
||||
request.LoraAdapter = os.path.join(request.ModelPath, request.LoraAdapter)
|
||||
|
||||
device = "cpu" if not request.CUDA else "cuda"
|
||||
if XPU:
|
||||
device = "xpu"
|
||||
mps_available = hasattr(torch.backends, "mps") and torch.backends.mps.is_available()
|
||||
if mps_available:
|
||||
device = "mps"
|
||||
device = select_device(
|
||||
request.CUDA,
|
||||
self.options.pop("device", None),
|
||||
torch.cuda.is_available(),
|
||||
XPU,
|
||||
hasattr(torch.backends, "mps") and torch.backends.mps.is_available(),
|
||||
)
|
||||
self.device = device
|
||||
if request.LoraAdapter:
|
||||
# Check if its a local file and not a directory ( we load lora differently for a safetensor file )
|
||||
@@ -800,12 +816,12 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
image = image.resize((1024, 576))
|
||||
|
||||
generator = torch.manual_seed(request.seed)
|
||||
frames = self.pipe(image, guidance_scale=self.cfg_scale, decode_chunk_size=CHUNK_SIZE, generator=generator).frames[0]
|
||||
frames = self.pipe(image=image, guidance_scale=self.cfg_scale, decode_chunk_size=CHUNK_SIZE, generator=generator).frames[0]
|
||||
export_to_video(frames, request.dst, fps=FPS)
|
||||
return backend_pb2.Result(message="Media generated successfully", success=True)
|
||||
|
||||
if self.txt2vid:
|
||||
video_frames = self.pipe(prompt, guidance_scale=self.cfg_scale, num_inference_steps=steps, num_frames=int(FRAMES)).frames
|
||||
video_frames = self.pipe(prompt=prompt, guidance_scale=self.cfg_scale, num_inference_steps=steps, num_frames=int(FRAMES)).frames
|
||||
export_to_video(video_frames, request.dst)
|
||||
return backend_pb2.Result(message="Media generated successfully", success=True)
|
||||
|
||||
@@ -868,7 +884,7 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
else:
|
||||
# pass the kwargs dictionary to the self.pipe method
|
||||
image = self.pipe(
|
||||
prompt,
|
||||
prompt=prompt,
|
||||
guidance_scale=self.cfg_scale,
|
||||
**kwargs
|
||||
).images[0]
|
||||
|
||||
@@ -7,6 +7,7 @@ import time
|
||||
from unittest.mock import patch, MagicMock
|
||||
|
||||
# Import dynamic loader for testing (these don't need gRPC)
|
||||
import backend
|
||||
import diffusers_dynamic_loader as loader
|
||||
from diffusers import DiffusionPipeline, StableDiffusionPipeline
|
||||
|
||||
@@ -373,3 +374,74 @@ class TestGenerateImageOptionsKwargsMerge(unittest.TestCase):
|
||||
finally:
|
||||
os.unlink(src_file.name)
|
||||
os.unlink(dst_file.name)
|
||||
|
||||
def test_text_to_image_prompt_is_passed_by_keyword(self):
|
||||
"""Test compatibility with pipelines that take image before prompt."""
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
from PIL import Image
|
||||
|
||||
from backend import BackendServicer
|
||||
|
||||
class Flux2CompatiblePipeline:
|
||||
"""Model the FLUX.2 call signature: image is before prompt."""
|
||||
|
||||
def __call__(self, image=None, prompt=None, **kwargs):
|
||||
if prompt is None:
|
||||
raise ValueError("prompt was not passed by keyword")
|
||||
self.prompt = prompt
|
||||
self.kwargs = kwargs
|
||||
return MagicMock(images=[Image.new("RGB", (4, 4))])
|
||||
|
||||
pipeline = Flux2CompatiblePipeline()
|
||||
svc = BackendServicer.__new__(BackendServicer)
|
||||
svc.pipe = pipeline
|
||||
svc.cfg_scale = 7.5
|
||||
svc.controlnet = None
|
||||
svc.img2vid = False
|
||||
svc.txt2vid = False
|
||||
svc.clip_skip = 0
|
||||
svc.PipelineType = "Flux2KleinPipeline"
|
||||
svc.options = {}
|
||||
|
||||
with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as dst_file:
|
||||
dst_path = dst_file.name
|
||||
|
||||
try:
|
||||
request = MagicMock()
|
||||
request.positive_prompt = "a red apple on a wooden table"
|
||||
request.negative_prompt = ""
|
||||
request.step = 4
|
||||
request.seed = 0
|
||||
request.width = 0
|
||||
request.height = 0
|
||||
request.src = ""
|
||||
request.ref_images = []
|
||||
request.dst = dst_path
|
||||
|
||||
svc.GenerateImage(request, context=None)
|
||||
|
||||
self.assertEqual(pipeline.prompt, request.positive_prompt)
|
||||
self.assertEqual(pipeline.kwargs["num_inference_steps"], 4)
|
||||
finally:
|
||||
os.unlink(dst_path)
|
||||
|
||||
|
||||
class TestDeviceSelection(unittest.TestCase):
|
||||
"""Unit tests for backend.select_device (no GPU required)."""
|
||||
|
||||
def test_autodetect_cuda(self):
|
||||
self.assertEqual(backend.select_device(False, None, True, False, False), "cuda")
|
||||
|
||||
def test_cpu_fallback(self):
|
||||
self.assertEqual(backend.select_device(False, None, False, False, False), "cpu")
|
||||
|
||||
def test_forced_cuda(self):
|
||||
self.assertEqual(backend.select_device(True, None, False, False, False), "cuda")
|
||||
|
||||
def test_device_option_wins(self):
|
||||
self.assertEqual(backend.select_device(True, "cpu", True, True, True), "cpu")
|
||||
|
||||
def test_mps_overrides(self):
|
||||
self.assertEqual(backend.select_device(False, None, True, False, True), "mps")
|
||||
@@ -11,7 +11,7 @@ RPC. It supports:
|
||||
systems such as NVIDIA DGX Spark.
|
||||
|
||||
Install the `longcat-video` or `longcat-video-avatar-1.5` recipe from the
|
||||
LocalAI Model Gallery. See the [LongCat user guide](../../../docs/content/features/longcat-video.md)
|
||||
LocalAI Model Gallery. LongCat video backend
|
||||
for Studio and API examples, hardware requirements, and manual configuration.
|
||||
|
||||
The upstream source is pinned in `Makefile` and patched at build time. The
|
||||
|
||||
@@ -18,6 +18,7 @@ sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..', 'common'))
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), 'common'))
|
||||
from grpc_auth import get_auth_interceptors
|
||||
from model_utils import resolve_model_reference
|
||||
from device_utils import device_map_for, select_device
|
||||
|
||||
|
||||
|
||||
@@ -95,13 +96,7 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
return backend_pb2.Reply(message=bytes("OK", 'utf-8'))
|
||||
|
||||
def LoadModel(self, request, context):
|
||||
if torch.cuda.is_available():
|
||||
device = "cuda"
|
||||
else:
|
||||
device = "cpu"
|
||||
mps_available = hasattr(torch.backends, "mps") and torch.backends.mps.is_available()
|
||||
if mps_available:
|
||||
device = "mps"
|
||||
device = select_device(torch)
|
||||
if not torch.cuda.is_available() and request.CUDA:
|
||||
return backend_pb2.Result(success=False, message="CUDA is not available")
|
||||
|
||||
@@ -123,7 +118,7 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
model_path, local_only = resolve_model_reference(
|
||||
request, "Qwen/Qwen3-ASR-1.7B"
|
||||
)
|
||||
default_dtype = torch.bfloat16 if self.device == "cuda" else torch.float32
|
||||
default_dtype = torch.bfloat16 if self.device in ("cuda", "xpu") else torch.float32
|
||||
load_dtype = default_dtype
|
||||
if "torch_dtype" in self.options:
|
||||
d = str(self.options["torch_dtype"]).lower()
|
||||
@@ -145,12 +140,7 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
if attn_implementation is not None and isinstance(attn_implementation, str):
|
||||
attn_implementation = attn_implementation.strip() or None
|
||||
|
||||
if self.device == "mps":
|
||||
device_map = None
|
||||
elif self.device == "cuda":
|
||||
device_map = "cuda:0"
|
||||
else:
|
||||
device_map = "cpu"
|
||||
device_map = device_map_for(self.device)
|
||||
|
||||
load_kwargs = dict(
|
||||
dtype=load_dtype,
|
||||
@@ -423,4 +413,4 @@ if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="Run the gRPC server.")
|
||||
parser.add_argument("--addr", default="localhost:50051", help="The address to bind the server to.")
|
||||
args = parser.parse_args()
|
||||
serve(args.addr)
|
||||
serve(args.addr)
|
||||
@@ -0,0 +1,18 @@
|
||||
def select_device(torch_module):
|
||||
mps = getattr(getattr(torch_module, "backends", None), "mps", None)
|
||||
if mps is not None and mps.is_available():
|
||||
return "mps"
|
||||
if torch_module.cuda.is_available():
|
||||
return "cuda"
|
||||
xpu = getattr(torch_module, "xpu", None)
|
||||
if xpu is not None and xpu.is_available():
|
||||
return "xpu"
|
||||
return "cpu"
|
||||
|
||||
|
||||
def device_map_for(device):
|
||||
if device == "mps":
|
||||
return None
|
||||
if device in ("cuda", "xpu"):
|
||||
return f"{device}:0"
|
||||
return "cpu"
|
||||
@@ -0,0 +1,58 @@
|
||||
import unittest
|
||||
|
||||
from device_utils import device_map_for, select_device
|
||||
|
||||
|
||||
class Availability:
|
||||
def __init__(self, available):
|
||||
self._available = available
|
||||
|
||||
def is_available(self):
|
||||
return self._available
|
||||
|
||||
|
||||
class TorchStub:
|
||||
def __init__(self, *, cuda=False, mps=False, xpu=False):
|
||||
self.cuda = Availability(cuda)
|
||||
self.backends = type("Backends", (), {"mps": Availability(mps)})()
|
||||
self.xpu = Availability(xpu)
|
||||
|
||||
|
||||
class SelectDeviceTest(unittest.TestCase):
|
||||
def test_preserves_cuda_selection(self):
|
||||
torch_module = TorchStub(cuda=True)
|
||||
|
||||
self.assertEqual(select_device(torch_module), "cuda")
|
||||
|
||||
def test_preserves_mps_selection(self):
|
||||
torch_module = TorchStub(mps=True)
|
||||
|
||||
self.assertEqual(select_device(torch_module), "mps")
|
||||
|
||||
def test_selects_xpu_when_intel_gpu_is_available(self):
|
||||
torch_module = TorchStub(xpu=True)
|
||||
|
||||
self.assertEqual(select_device(torch_module), "xpu")
|
||||
|
||||
def test_falls_back_to_cpu(self):
|
||||
torch_module = TorchStub()
|
||||
|
||||
self.assertEqual(select_device(torch_module), "cpu")
|
||||
|
||||
|
||||
class DeviceMapTest(unittest.TestCase):
|
||||
def test_preserves_cuda_model_placement(self):
|
||||
self.assertEqual(device_map_for("cuda"), "cuda:0")
|
||||
|
||||
def test_preserves_mps_model_placement(self):
|
||||
self.assertIsNone(device_map_for("mps"))
|
||||
|
||||
def test_places_the_model_on_the_first_xpu(self):
|
||||
self.assertEqual(device_map_for("xpu"), "xpu:0")
|
||||
|
||||
def test_preserves_cpu_model_placement(self):
|
||||
self.assertEqual(device_map_for("cpu"), "cpu")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,3 +1,3 @@
|
||||
grpcio==1.82.1
|
||||
grpcio==1.83.1
|
||||
protobuf
|
||||
certifi
|
||||
@@ -40,6 +40,7 @@ import grpc
|
||||
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..', 'common'))
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), 'common'))
|
||||
from python_utils import attach_media_parts
|
||||
from grpc_auth import get_auth_interceptors
|
||||
from model_utils import resolve_model_reference
|
||||
|
||||
@@ -90,6 +91,14 @@ except Exception:
|
||||
|
||||
|
||||
_ONE_DAY_IN_SECONDS = 60 * 60 * 24
|
||||
|
||||
# proto3 has no field presence, so an explicit 0 is indistinguishable from
|
||||
# "unset" and the zero-filter below would drop it. These two fields have a
|
||||
# meaningful zero a caller can actually intend: temperature 0 is greedy
|
||||
# decoding, and 0 is a valid seed. Silently substituting a default for either
|
||||
# turns a reproducible request into a random one.
|
||||
_EXPLICIT_ZERO_FIELDS = ("Temperature", "Seed")
|
||||
|
||||
MAX_WORKERS = int(os.environ.get('PYTHON_GRPC_MAX_WORKERS', '1'))
|
||||
|
||||
|
||||
@@ -323,7 +332,7 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
if not hasattr(request, proto_field):
|
||||
continue
|
||||
value = getattr(request, proto_field)
|
||||
if value in (None, 0, 0.0, [], False, ""):
|
||||
if proto_field not in _EXPLICIT_ZERO_FIELDS and value in (None, 0, 0.0, [], False, ""):
|
||||
continue
|
||||
# repeated fields come back as RepeatedScalarContainer — convert
|
||||
if hasattr(value, "__iter__") and not isinstance(value, (str, bytes)):
|
||||
@@ -363,8 +372,27 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
template_kwargs["tools"] = json.loads(request.Tools)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
if request.Metadata.get("enable_thinking", "").lower() == "true":
|
||||
template_kwargs["enable_thinking"] = True
|
||||
_thinking = request.Metadata.get("enable_thinking", "").lower()
|
||||
if _thinking in ("true", "false"):
|
||||
template_kwargs["enable_thinking"] = (_thinking == "true")
|
||||
|
||||
# sglang locates the attached images/videos by scanning the rendered
|
||||
# prompt for the model's own media token, so the template has to be
|
||||
# given content *parts* - string content renders a prompt with no
|
||||
# placeholder and the media are dropped without a word (#11621).
|
||||
media_dicts = attach_media_parts(
|
||||
messages_dicts, len(request.Images), len(request.Videos)
|
||||
)
|
||||
if media_dicts is not None:
|
||||
try:
|
||||
return self.tokenizer.apply_chat_template(media_dicts, **template_kwargs)
|
||||
except Exception as e:
|
||||
# A text-only template cannot iterate content parts; fall
|
||||
# through to the text-only prompt instead of failing.
|
||||
print(
|
||||
f"chat template rejected multimodal content parts: {e!r}",
|
||||
file=sys.stderr,
|
||||
)
|
||||
|
||||
try:
|
||||
return self.tokenizer.apply_chat_template(messages_dicts, **template_kwargs)
|
||||
@@ -373,10 +401,67 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
messages_dicts, tokenize=False, add_generation_prompt=True,
|
||||
)
|
||||
|
||||
def _make_parsers(self, request):
|
||||
def _new_reasoning_parser(self, stream_reasoning: bool, prompt: str = "",
|
||||
grammar_constrained: bool = False):
|
||||
"""Build a ReasoningParser for one request, or None.
|
||||
|
||||
Reasoning templates come in two flavours. Some let the model emit the
|
||||
opening tag, others put it into the *prompt* — Qwen3's template appends
|
||||
``<think>`` when thinking is on, so the completion starts straight in
|
||||
the reasoning block and only the closing ``</think>`` ever shows up.
|
||||
sglang's detector keys off the opening tag, so in that second case it
|
||||
classifies the whole completion as normal content and
|
||||
``reasoning_content`` stays empty.
|
||||
|
||||
sglang's own OpenAI server covers this with
|
||||
``template_manager.force_reasoning``; this backend has no template
|
||||
manager, so it derives the same signal from the rendered prompt.
|
||||
``force_reasoning`` is only passed when we mean True, leaving detector
|
||||
defaults (e.g. DeepSeek-R1's built-in True) untouched.
|
||||
|
||||
``grammar_constrained`` suppresses the prefill heuristic. A structured
|
||||
decoding constraint applies from the first token, so the model cannot
|
||||
emit the closing tag even though the template opened the block: the
|
||||
whole completion is schema output and belongs in ``content``. Forcing
|
||||
there files the answer as reasoning and leaves content empty. sglang's
|
||||
own server keeps the two apart for the same reason — its grammar
|
||||
backend owns the reasoning prefix when a reasoning parser is set.
|
||||
"""
|
||||
if grammar_constrained:
|
||||
prompt = ""
|
||||
|
||||
if not (HAS_REASONING_PARSERS and self.reasoning_parser_name):
|
||||
return None
|
||||
|
||||
kwargs = {
|
||||
"model_type": self.reasoning_parser_name,
|
||||
"stream_reasoning": stream_reasoning,
|
||||
}
|
||||
try:
|
||||
parser = ReasoningParser(**kwargs)
|
||||
except Exception as e:
|
||||
print(f"ReasoningParser init failed: {e!r}", file=sys.stderr)
|
||||
return None
|
||||
|
||||
start = getattr(getattr(parser, "detector", None), "think_start_token", None)
|
||||
if start and prompt and prompt.rstrip().endswith(start):
|
||||
try:
|
||||
parser = ReasoningParser(force_reasoning=True, **kwargs)
|
||||
except TypeError:
|
||||
# sglang without the force_reasoning kwarg: keep the default
|
||||
# parser rather than failing the request.
|
||||
pass
|
||||
except Exception as e:
|
||||
print(
|
||||
f"ReasoningParser(force_reasoning=True) failed: {e!r}",
|
||||
file=sys.stderr,
|
||||
)
|
||||
|
||||
return parser
|
||||
|
||||
def _make_parsers(self, request, prompt: str = ""):
|
||||
"""Construct fresh per-request parser instances (stateful)."""
|
||||
tool_parser = None
|
||||
reasoning_parser = None
|
||||
|
||||
if HAS_TOOL_PARSERS and self.tool_parser_name and request.Tools:
|
||||
try:
|
||||
@@ -388,14 +473,9 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
except Exception as e:
|
||||
print(f"FunctionCallParser init failed: {e!r}", file=sys.stderr)
|
||||
|
||||
if HAS_REASONING_PARSERS and self.reasoning_parser_name:
|
||||
try:
|
||||
reasoning_parser = ReasoningParser(
|
||||
model_type=self.reasoning_parser_name,
|
||||
stream_reasoning=True,
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"ReasoningParser init failed: {e!r}", file=sys.stderr)
|
||||
reasoning_parser = self._new_reasoning_parser(
|
||||
True, prompt, bool(getattr(request, "Grammar", "")),
|
||||
)
|
||||
|
||||
return tool_parser, reasoning_parser
|
||||
|
||||
@@ -403,7 +483,7 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
sampling_params = self._build_sampling_params(request)
|
||||
prompt = self._build_prompt(request)
|
||||
|
||||
tool_parser, reasoning_parser = self._make_parsers(request)
|
||||
tool_parser, reasoning_parser = self._make_parsers(request, prompt)
|
||||
|
||||
image_data = list(request.Images) if request.Images else None
|
||||
video_data = list(request.Videos) if request.Videos else None
|
||||
@@ -499,15 +579,9 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
final_tool_calls: List[backend_pb2.ToolCallDelta] = []
|
||||
|
||||
if not streaming:
|
||||
final_reasoning_parser = None
|
||||
if HAS_REASONING_PARSERS and self.reasoning_parser_name:
|
||||
try:
|
||||
final_reasoning_parser = ReasoningParser(
|
||||
model_type=self.reasoning_parser_name,
|
||||
stream_reasoning=False,
|
||||
)
|
||||
except Exception:
|
||||
final_reasoning_parser = None
|
||||
final_reasoning_parser = self._new_reasoning_parser(
|
||||
False, prompt, bool(getattr(request, "Grammar", "")),
|
||||
)
|
||||
|
||||
if final_reasoning_parser is not None:
|
||||
try:
|
||||
|
||||
@@ -96,6 +96,126 @@ class TestSglangHelpers(unittest.TestCase):
|
||||
servicer._apply_engine_args({}, "[1,2,3]")
|
||||
self.assertIn("must be a JSON object", str(ctx.exception))
|
||||
|
||||
def test_build_prompt_forwards_enable_thinking(self):
|
||||
from types import SimpleNamespace
|
||||
|
||||
class Tok:
|
||||
def __init__(self):
|
||||
self.kwargs = None
|
||||
|
||||
def apply_chat_template(self, messages, **kwargs):
|
||||
self.kwargs = kwargs
|
||||
return "PROMPT"
|
||||
|
||||
def kwargs_for(metadata):
|
||||
servicer = self._servicer()
|
||||
tok = Tok()
|
||||
servicer.tokenizer = tok
|
||||
msg = SimpleNamespace(
|
||||
role="user", content="hi", name="",
|
||||
tool_call_id="", reasoning_content="", tool_calls="",
|
||||
)
|
||||
req = SimpleNamespace(
|
||||
Prompt="", UseTokenizerTemplate=True,
|
||||
Messages=[msg], Tools="", Metadata=metadata,
|
||||
)
|
||||
self.assertEqual(servicer._build_prompt(req), "PROMPT")
|
||||
return tok.kwargs
|
||||
|
||||
self.assertIs(kwargs_for({"enable_thinking": "true"})["enable_thinking"], True)
|
||||
# "false" used to be dropped, so Qwen3 kept thinking on
|
||||
self.assertIs(kwargs_for({"enable_thinking": "false"})["enable_thinking"], False)
|
||||
self.assertNotIn("enable_thinking", kwargs_for({}))
|
||||
self.assertIs(kwargs_for({"enable_thinking": "FALSE"})["enable_thinking"], False)
|
||||
|
||||
def test_reasoning_parser_forced_when_template_prefills_think_tag(self):
|
||||
"""Qwen3's template puts ``<think>`` in the prompt, so the completion
|
||||
never contains it. Without force_reasoning the detector treats the whole
|
||||
completion as normal text and reasoning_content stays empty."""
|
||||
servicer = self._servicer()
|
||||
servicer.reasoning_parser_name = "qwen3"
|
||||
|
||||
# What the model actually emits when the prompt ends in "<think>".
|
||||
completion = "adding two and two</think>4"
|
||||
|
||||
forced = servicer._new_reasoning_parser(False, prompt="user: hi\n<think>\n")
|
||||
reasoning, content = forced.parse_non_stream(completion)
|
||||
self.assertEqual(reasoning, "adding two and two")
|
||||
self.assertEqual(content, "4")
|
||||
|
||||
# No prefilled tag in the prompt: detector default, unchanged behaviour.
|
||||
unforced = servicer._new_reasoning_parser(False, prompt="user: hi\n")
|
||||
reasoning, content = unforced.parse_non_stream(completion)
|
||||
self.assertFalse(reasoning)
|
||||
self.assertEqual(content, completion)
|
||||
|
||||
def test_reasoning_parser_not_forced_when_thinking_is_off(self):
|
||||
"""Thinking off means no ``<think>`` in the prompt either, so the answer
|
||||
must not be swallowed into reasoning_content."""
|
||||
servicer = self._servicer()
|
||||
servicer.reasoning_parser_name = "qwen3"
|
||||
|
||||
parser = servicer._new_reasoning_parser(False, prompt="user: primes?\n")
|
||||
reasoning, content = parser.parse_non_stream("2,3,5,7,11")
|
||||
self.assertFalse(reasoning)
|
||||
self.assertEqual(content, "2,3,5,7,11")
|
||||
|
||||
def test_grammar_constrained_output_is_not_forced_into_reasoning(self):
|
||||
"""Structured decoding applies from the first token, so the model cannot
|
||||
emit the closing tag even though the template opened the block. The whole
|
||||
completion is schema output and must stay in content."""
|
||||
servicer = self._servicer()
|
||||
servicer.reasoning_parser_name = "qwen3"
|
||||
|
||||
schema_out = '{"findings": [{"line": 42, "issue": "off-by-one"}]}'
|
||||
parser = servicer._new_reasoning_parser(
|
||||
False, prompt="audit this\n<think>\n", grammar_constrained=True,
|
||||
)
|
||||
reasoning, content = parser.parse_non_stream(schema_out)
|
||||
self.assertFalse(reasoning)
|
||||
self.assertEqual(content, schema_out)
|
||||
|
||||
def test_reasoning_parser_absent_without_configured_parser(self):
|
||||
servicer = self._servicer()
|
||||
servicer.reasoning_parser_name = None
|
||||
self.assertIsNone(servicer._new_reasoning_parser(False, prompt="<think>"))
|
||||
|
||||
def test_explicit_zero_temperature_and_seed_are_preserved(self):
|
||||
"""Temperature=0 is greedy decoding and 0 is a valid seed — neither is
|
||||
an unset value. A dropped seed turns a reproducible request random."""
|
||||
from types import SimpleNamespace
|
||||
|
||||
servicer = self._servicer()
|
||||
import sys as _sys
|
||||
_SEED_KEY_FOR_TEST = _sys.modules["backend"]._SEED_KEY
|
||||
request = SimpleNamespace(
|
||||
Temperature=0,
|
||||
N=0,
|
||||
PresencePenalty=0,
|
||||
FrequencyPenalty=0,
|
||||
RepetitionPenalty=0,
|
||||
TopP=0,
|
||||
TopK=0,
|
||||
MinP=0,
|
||||
Seed=0,
|
||||
StopPrompts=[],
|
||||
StopTokenIds=[],
|
||||
IgnoreEOS=False,
|
||||
Tokens=0,
|
||||
MinTokens=0,
|
||||
SkipSpecialTokens=False,
|
||||
Grammar="",
|
||||
)
|
||||
|
||||
params = servicer._build_sampling_params(request)
|
||||
self.assertEqual(params["temperature"], 0)
|
||||
self.assertEqual(params[_SEED_KEY_FOR_TEST], 0)
|
||||
# Other protobuf-default scalar fields must remain filtered. top_k=0 in
|
||||
# particular is not a value sglang accepts (-1 disables it), so it must
|
||||
# keep falling through to the engine default.
|
||||
self.assertNotIn("top_p", params)
|
||||
self.assertNotIn("top_k", params)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
+122
-18
@@ -20,6 +20,7 @@ import backend_pb2_grpc
|
||||
import grpc
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..', 'common'))
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), 'common'))
|
||||
from python_utils import attach_media_parts
|
||||
from grpc_auth import get_auth_interceptors
|
||||
from model_utils import resolve_model_reference
|
||||
from vllm_utils import apply_options_to_engine_args, normalize_option_key
|
||||
@@ -60,6 +61,12 @@ except ImportError:
|
||||
|
||||
_ONE_DAY_IN_SECONDS = 60 * 60 * 24
|
||||
|
||||
# proto3 has no field presence, so an explicit 0 is indistinguishable from
|
||||
# "unset". These two fields have a meaningful zero a caller can intend:
|
||||
# temperature 0 is greedy decoding, and 0 is a valid seed.
|
||||
_EXPLICIT_ZERO_FIELDS = ("Temperature", "Seed")
|
||||
|
||||
|
||||
# If MAX_WORKERS are specified in the environment use it, otherwise default to 1
|
||||
MAX_WORKERS = int(os.environ.get('PYTHON_GRPC_MAX_WORKERS', '1'))
|
||||
|
||||
@@ -523,9 +530,7 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
context.set_details(str(e))
|
||||
return backend_pb2.ScoreResponse()
|
||||
|
||||
async def _predict(self, request, context, streaming=False):
|
||||
# Build the sampling parameters
|
||||
# NOTE: this must stay in sync with the vllm backend
|
||||
def _build_sampling_params(self, request):
|
||||
request_to_sampling_params = {
|
||||
"N": "n",
|
||||
"PresencePenalty": "presence_penalty",
|
||||
@@ -555,9 +560,84 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
for request_field, param_field in request_to_sampling_params.items():
|
||||
if hasattr(request, request_field):
|
||||
value = getattr(request, request_field)
|
||||
if value not in (None, 0, [], False):
|
||||
# See _EXPLICIT_ZERO_FIELDS: temperature 0 is greedy decoding
|
||||
# and 0 is a valid seed, so neither may be filtered out.
|
||||
if request_field in _EXPLICIT_ZERO_FIELDS or value not in (None, 0, [], False):
|
||||
setattr(sampling_params, param_field, value)
|
||||
|
||||
return sampling_params
|
||||
|
||||
def _new_reasoning_parser(self, chat_template_kwargs):
|
||||
"""Build the reasoning parser, telling it whether thinking is on.
|
||||
|
||||
vLLM's newer parser engines decide their *initial state* from
|
||||
``chat_template_kwargs``: ``Qwen3Parser`` reads
|
||||
``chat_template_kwargs["enable_thinking"]`` and defaults to ``True``,
|
||||
starting in the REASONING state. Constructed without it, a completion
|
||||
produced with thinking disabled is classified as reasoning end to end,
|
||||
and the answer is reported in both ``reasoning_content`` and
|
||||
``content``.
|
||||
|
||||
vLLM's own OpenAI server forwards the request's chat template kwargs
|
||||
here; this backend renders the template itself, so it forwards the
|
||||
same dict. Older parsers do not accept the argument — fall back to the
|
||||
plain constructor for those.
|
||||
"""
|
||||
try:
|
||||
return self.reasoning_parser_cls(
|
||||
self.tokenizer, chat_template_kwargs=chat_template_kwargs or {},
|
||||
)
|
||||
except TypeError:
|
||||
return self.reasoning_parser_cls(self.tokenizer)
|
||||
|
||||
@staticmethod
|
||||
def _split_reasoning(rp, generated_text, prompt, reasoning, content):
|
||||
"""Decide what the reasoning parser's output actually means.
|
||||
|
||||
Covers the *older* parser shape, which has no initial state to set:
|
||||
``BaseThinkingReasoningParser.extract_reasoning`` documents its own
|
||||
fallback — "For models that may not generate start token, assume the
|
||||
reasoning content is always at the start." When no end token is
|
||||
present it returns *everything* as reasoning and ``None`` as content,
|
||||
which is right for a truncated reasoning run and wrong for a
|
||||
completion that never contained reasoning at all.
|
||||
|
||||
Taking ``None`` content to mean "keep the raw text" then duplicates
|
||||
the answer into both fields.
|
||||
|
||||
The prompt says which case it is. A template with thinking on leaves
|
||||
the reasoning block open (the prompt ends with the start token); with
|
||||
thinking off it closes the block in the prompt, so the completion is
|
||||
plain content. Parsers that expose no token pair (the engine-based
|
||||
adapters, which take the ``chat_template_kwargs`` route above) keep
|
||||
the parser's verdict unchanged.
|
||||
"""
|
||||
start = getattr(rp, "start_token", None)
|
||||
end = getattr(rp, "end_token", None)
|
||||
|
||||
if end and end in generated_text:
|
||||
# The parser split on the end token. Empty content here means the
|
||||
# model stopped right after it, not that parsing failed.
|
||||
return reasoning or "", content or ""
|
||||
|
||||
if not start:
|
||||
# Unknown token layout — keep the previous behaviour rather than
|
||||
# guess.
|
||||
return reasoning or "", content if content is not None else generated_text
|
||||
|
||||
if not (start in generated_text or (prompt or "").rstrip().endswith(start)):
|
||||
# No end token and the block was never open: the "reasoning starts
|
||||
# at the beginning" fallback does not apply to this completion.
|
||||
return "", generated_text
|
||||
|
||||
# Block was open and the end token never arrived — reasoning ran out of
|
||||
# budget. It is all reasoning, and there is no answer to report.
|
||||
return reasoning or "", content or ""
|
||||
|
||||
async def _predict(self, request, context, streaming=False):
|
||||
# Build the sampling parameters
|
||||
sampling_params = self._build_sampling_params(request)
|
||||
|
||||
# Structured-output decoding: use Grammar field to pass JSON schema or BNF
|
||||
if HAS_GUIDED_DECODING and request.Grammar:
|
||||
try:
|
||||
@@ -568,6 +648,9 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
|
||||
# Extract image paths and process images
|
||||
prompt = request.Prompt
|
||||
# Kept in scope: the reasoning parser needs to know which chat
|
||||
# template kwargs produced this prompt.
|
||||
template_kwargs = {}
|
||||
|
||||
image_paths = request.Images
|
||||
image_data = [self.load_image(img_path) for img_path in image_paths]
|
||||
@@ -578,7 +661,7 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
# If tokenizer template is enabled and messages are provided instead of prompt, apply the tokenizer template
|
||||
if not request.Prompt and request.UseTokenizerTemplate and request.Messages:
|
||||
messages_dicts = self._messages_to_dicts(request.Messages)
|
||||
template_kwargs = {"tokenize": False, "add_generation_prompt": True}
|
||||
template_kwargs.update({"tokenize": False, "add_generation_prompt": True})
|
||||
|
||||
# Pass tools for tool calling
|
||||
if request.Tools:
|
||||
@@ -587,17 +670,37 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
# Enable thinking mode if requested
|
||||
if request.Metadata.get("enable_thinking", "").lower() == "true":
|
||||
template_kwargs["enable_thinking"] = True
|
||||
_thinking = request.Metadata.get("enable_thinking", "").lower()
|
||||
if _thinking in ("true", "false"):
|
||||
template_kwargs["enable_thinking"] = (_thinking == "true")
|
||||
|
||||
try:
|
||||
prompt = self.tokenizer.apply_chat_template(messages_dicts, **template_kwargs)
|
||||
except TypeError:
|
||||
# Some tokenizers don't support tools/enable_thinking kwargs — retry without them
|
||||
prompt = self.tokenizer.apply_chat_template(
|
||||
messages_dicts, tokenize=False, add_generation_prompt=True
|
||||
)
|
||||
# vLLM substitutes multi_modal_data into the model's own media
|
||||
# token, so the template has to be given content *parts* - string
|
||||
# content renders a prompt with no placeholder and the media are
|
||||
# dropped without a word (#11621).
|
||||
prompt = None
|
||||
media_dicts = attach_media_parts(
|
||||
messages_dicts, len(image_data), len(video_data)
|
||||
)
|
||||
if media_dicts is not None:
|
||||
try:
|
||||
prompt = self.tokenizer.apply_chat_template(media_dicts, **template_kwargs)
|
||||
except Exception as e:
|
||||
# A text-only template cannot iterate content parts; fall
|
||||
# through to the text-only prompt instead of failing.
|
||||
print(
|
||||
f"chat template rejected multimodal content parts: {e!r}",
|
||||
file=sys.stderr,
|
||||
)
|
||||
|
||||
if prompt is None:
|
||||
try:
|
||||
prompt = self.tokenizer.apply_chat_template(messages_dicts, **template_kwargs)
|
||||
except TypeError:
|
||||
# Some tokenizers don't support tools/enable_thinking kwargs — retry without them
|
||||
prompt = self.tokenizer.apply_chat_template(
|
||||
messages_dicts, tokenize=False, add_generation_prompt=True
|
||||
)
|
||||
|
||||
# Generate text using the LLM engine
|
||||
request_id = random_uuid()
|
||||
@@ -753,10 +856,11 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
|
||||
if self.reasoning_parser_cls:
|
||||
try:
|
||||
rp = self.reasoning_parser_cls(self.tokenizer)
|
||||
rp = self._new_reasoning_parser(template_kwargs)
|
||||
r, c = rp.extract_reasoning(generated_text, request=None)
|
||||
reasoning_content = r or ""
|
||||
content = c if c is not None else generated_text
|
||||
reasoning_content, content = self._split_reasoning(
|
||||
rp, generated_text, prompt, r, c,
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"Reasoning parser error: {e}", file=sys.stderr)
|
||||
|
||||
|
||||
@@ -119,14 +119,18 @@ if [ "$(uname -s)" = "Darwin" ]; then
|
||||
# can rewrite it. Darwin therefore follows vllm-metal and can lag the Linux
|
||||
# vllm pin (requirements-cublas13-after.txt, bumped independently against
|
||||
# vllm/vllm) until vllm-metal supports a newer vLLM.
|
||||
VLLM_METAL_VERSION="v0.3.0.dev20260818075955"
|
||||
VLLM_METAL_VERSION="v0.28.0"
|
||||
|
||||
# The coupled vLLM source version is whatever this vllm-metal release builds
|
||||
# against. Derive it from
|
||||
# the PINNED tag rather than hardcoding a second value that could drift. The
|
||||
# tag is immutable, so this stays reproducible across rebuilds.
|
||||
VLLM_VERSION=$(curl -fsSL "https://raw.githubusercontent.com/vllm-project/vllm-metal/${VLLM_METAL_VERSION}/install.sh" \
|
||||
| "$backend_dir/../../../scripts/lib/extract-vllm-metal-version.sh")
|
||||
# against. Derive it from the PINNED tag rather than hardcoding a second value
|
||||
# that could drift. The tag is immutable, so this stays reproducible across
|
||||
# rebuilds. Since vllm-metal 0.28 the coupling is declared in
|
||||
# .github/vllm-release-tag.commit; older releases pinned it inline in their
|
||||
# own install.sh, so fall back to that. The extractor reads both forms.
|
||||
_vllm_metal_raw="https://raw.githubusercontent.com/vllm-project/vllm-metal/${VLLM_METAL_VERSION}"
|
||||
VLLM_VERSION=$( { curl -fsSL "${_vllm_metal_raw}/.github/vllm-release-tag.commit" \
|
||||
|| curl -fsSL "${_vllm_metal_raw}/install.sh"; } \
|
||||
| "$backend_dir/../../../scripts/lib/extract-vllm-metal-version.sh" || true)
|
||||
if [ -z "${VLLM_VERSION}" ]; then
|
||||
echo "ERROR: could not derive the vLLM version from vllm-metal ${VLLM_METAL_VERSION}" >&2
|
||||
exit 1
|
||||
@@ -153,10 +157,18 @@ if [ "$(uname -s)" = "Darwin" ]; then
|
||||
# 2) Install the prebuilt vllm-metal wheel for the PINNED release. It pulls
|
||||
# mlx / mlx-metal as deps and registers the `metal` platform plugin that
|
||||
# backend.py resolves to at engine-init time. Build the release-asset URL
|
||||
# deterministically (tag + the cp312/arm64 wheel name) rather than querying
|
||||
# api.github.com, whose unauthenticated rate limit (60/hr per IP) 403s on
|
||||
# shared CI runners. The wheel version is the tag without its leading 'v'.
|
||||
_metal_wheel="vllm_metal-${VLLM_METAL_VERSION#v}-cp312-cp312-macosx_11_0_arm64.whl"
|
||||
# from the release's OWN asset listing rather than composing it from a
|
||||
# hardcoded platform tag: upstream raised its macOS deployment target
|
||||
# (macosx_11_0 -> macosx_15_0) and every composed URL started to 404.
|
||||
# expanded_assets is the plain release page, not api.github.com, whose
|
||||
# unauthenticated rate limit (60/hr per IP) 403s on shared CI runners.
|
||||
# The wheel version is the tag without its leading 'v'.
|
||||
_metal_wheel=$(curl -fsSL "https://github.com/vllm-project/vllm-metal/releases/expanded_assets/${VLLM_METAL_VERSION}" \
|
||||
| grep -oE "vllm_metal-${VLLM_METAL_VERSION#v}-cp312-cp312-[A-Za-z0-9_]+\.whl" | head -1 || true)
|
||||
if [ -z "${_metal_wheel}" ]; then
|
||||
echo "ERROR: no cp312 wheel asset on vllm-metal release ${VLLM_METAL_VERSION}" >&2
|
||||
exit 1
|
||||
fi
|
||||
_metal_wheel_url="https://github.com/vllm-project/vllm-metal/releases/download/${VLLM_METAL_VERSION}/${_metal_wheel}"
|
||||
echo "Installing vllm-metal wheel: ${_metal_wheel_url}"
|
||||
uv pip install "${_metal_wheel_url}"
|
||||
|
||||
@@ -3,8 +3,8 @@
|
||||
# on a cu130 host. Pull the cu130-flavoured wheel from vLLM's per-tag index
|
||||
# instead — the cublas13 case in install.sh adds --index-strategy=unsafe-best-match
|
||||
# so uv consults this index alongside PyPI.
|
||||
--extra-index-url https://wheels.vllm.ai/0.27.1/cu130
|
||||
--extra-index-url https://wheels.vllm.ai/0.28.0/cu130
|
||||
# VERSION COUPLING: darwin/Apple-Silicon builds use vllm-metal (see install.sh),
|
||||
# which pins this exact vLLM version. Bumping vllm here means coordinating with a
|
||||
# vllm-metal release that supports the new version, or macOS/Metal builds break.
|
||||
vllm==0.27.1
|
||||
vllm==0.28.0
|
||||
@@ -9,4 +9,4 @@
|
||||
# memory architecture crash deterministically with an empty "Engine core init
|
||||
# failed" set (mudler/LocalAI#10722). Leaving this unpinned let the L4T image
|
||||
# drift onto whatever wheel was latest at build time.
|
||||
vllm==0.26.0
|
||||
vllm==0.28.0
|
||||
@@ -1,4 +1,4 @@
|
||||
grpcio==1.83.0
|
||||
grpcio==1.83.1
|
||||
protobuf
|
||||
certifi
|
||||
setuptools
|
||||
|
||||
@@ -121,6 +121,21 @@ class TestBackendServicer(unittest.TestCase):
|
||||
finally:
|
||||
self.tearDown()
|
||||
|
||||
def test_explicit_zero_temperature_and_seed_are_preserved(self):
|
||||
"""Temperature=0 is greedy decoding and 0 is a valid seed — neither is
|
||||
an unset value. A dropped seed turns a reproducible request random."""
|
||||
import sys, os
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
from backend import BackendServicer
|
||||
|
||||
servicer = BackendServicer()
|
||||
request = backend_pb2.PredictOptions(Prompt="hello", Temperature=0, Seed=0)
|
||||
sampling_params = servicer._build_sampling_params(request)
|
||||
self.assertEqual(sampling_params.temperature, 0)
|
||||
self.assertEqual(sampling_params.seed, 0)
|
||||
# Other protobuf-default scalar fields must remain filtered.
|
||||
self.assertEqual(sampling_params.top_p, 0.9)
|
||||
|
||||
|
||||
def test_messages_to_dicts(self):
|
||||
"""
|
||||
@@ -536,3 +551,116 @@ class TestStreamingToolParser(unittest.TestCase):
|
||||
intermediate, ["Hello ", "world", "!"],
|
||||
f"plain streaming changed; got {intermediate!r}",
|
||||
)
|
||||
|
||||
|
||||
class TestReasoningSplit(unittest.TestCase):
|
||||
"""Server-less tests for BackendServicer._split_reasoning.
|
||||
|
||||
vLLM's BaseThinkingReasoningParser returns the whole completion as
|
||||
reasoning and None as content whenever the end token is missing. Taken
|
||||
literally that duplicates a thinking-disabled answer into both fields.
|
||||
"""
|
||||
|
||||
class _Parser:
|
||||
start_token = "<think>"
|
||||
end_token = "</think>"
|
||||
|
||||
def _split(self, generated, prompt, reasoning, content):
|
||||
import sys, os
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
from backend import BackendServicer
|
||||
return BackendServicer._split_reasoning(
|
||||
self._Parser(), generated, prompt, reasoning, content,
|
||||
)
|
||||
|
||||
def test_thinking_off_is_not_duplicated_into_reasoning(self):
|
||||
"""No tags anywhere: the answer is content, and only content."""
|
||||
r, c = self._split(
|
||||
"391", "user: 17*23?\n<think>\n\n</think>\n\n",
|
||||
reasoning="391", content=None,
|
||||
)
|
||||
self.assertEqual(r, "")
|
||||
self.assertEqual(c, "391")
|
||||
|
||||
def test_prefilled_start_tag_keeps_truncated_reasoning(self):
|
||||
"""Prompt left the block open and the end token never arrived
|
||||
(budget exhausted): that really is all reasoning."""
|
||||
r, c = self._split(
|
||||
"thinking and thinking", "user: hi\n<think>\n",
|
||||
reasoning="thinking and thinking", content=None,
|
||||
)
|
||||
self.assertEqual(r, "thinking and thinking")
|
||||
self.assertEqual(c, "")
|
||||
|
||||
def test_end_token_present_keeps_parser_split(self):
|
||||
r, c = self._split(
|
||||
"adding two and two</think>4", "user: hi\n<think>\n",
|
||||
reasoning="adding two and two", content="4",
|
||||
)
|
||||
self.assertEqual(r, "adding two and two")
|
||||
self.assertEqual(c, "4")
|
||||
|
||||
def test_stop_right_after_end_token_yields_empty_content(self):
|
||||
"""Content must not fall back to the raw text — that would put the
|
||||
reasoning into the answer."""
|
||||
r, c = self._split(
|
||||
"reasoned</think>", "user: hi\n<think>\n",
|
||||
reasoning="reasoned", content=None,
|
||||
)
|
||||
self.assertEqual(r, "reasoned")
|
||||
self.assertEqual(c, "")
|
||||
|
||||
def test_unknown_token_layout_keeps_previous_behaviour(self):
|
||||
class _Bare:
|
||||
pass
|
||||
import sys, os
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
from backend import BackendServicer
|
||||
r, c = BackendServicer._split_reasoning(
|
||||
_Bare(), "raw", "prompt", "raw", None,
|
||||
)
|
||||
self.assertEqual(r, "raw")
|
||||
self.assertEqual(c, "raw")
|
||||
|
||||
|
||||
class TestReasoningParserConstruction(unittest.TestCase):
|
||||
"""The parser must learn whether thinking was on for this request.
|
||||
|
||||
vLLM's engine-based parsers (Qwen3Parser and friends) read
|
||||
chat_template_kwargs["enable_thinking"] and default to True, so a parser
|
||||
built without it treats a thinking-disabled completion as pure reasoning.
|
||||
"""
|
||||
|
||||
def _servicer(self):
|
||||
import sys, os
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
from backend import BackendServicer
|
||||
s = BackendServicer()
|
||||
s.tokenizer = object()
|
||||
return s
|
||||
|
||||
def test_chat_template_kwargs_are_forwarded(self):
|
||||
seen = {}
|
||||
|
||||
class _Parser:
|
||||
def __init__(self, tokenizer, **kwargs):
|
||||
seen.update(kwargs)
|
||||
|
||||
s = self._servicer()
|
||||
s.reasoning_parser_cls = _Parser
|
||||
s._new_reasoning_parser({"enable_thinking": False})
|
||||
self.assertEqual(
|
||||
seen.get("chat_template_kwargs"), {"enable_thinking": False},
|
||||
)
|
||||
|
||||
def test_parser_without_the_kwarg_still_builds(self):
|
||||
"""Older parsers take only the tokenizer — must not break them."""
|
||||
class _Old:
|
||||
def __init__(self, tokenizer):
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
s = self._servicer()
|
||||
s.reasoning_parser_cls = _Old
|
||||
self.assertIsInstance(
|
||||
s._new_reasoning_parser({"enable_thinking": False}), _Old,
|
||||
)
|
||||
@@ -16,6 +16,7 @@ import grpc
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..', 'common'))
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), 'common'))
|
||||
from grpc_auth import get_auth_interceptors
|
||||
from transcript_utils import require_diarization_token, seconds_to_nanoseconds
|
||||
|
||||
|
||||
|
||||
@@ -81,6 +82,11 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
import whisperx
|
||||
from whisperx.diarize import DiarizationPipeline
|
||||
|
||||
try:
|
||||
require_diarization_token(request.diarize, self.hf_token)
|
||||
except ValueError as err:
|
||||
context.abort(grpc.StatusCode.FAILED_PRECONDITION, str(err))
|
||||
|
||||
resultSegments = []
|
||||
text = ""
|
||||
try:
|
||||
@@ -117,8 +123,8 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
# Build result segments
|
||||
for idx, seg in enumerate(transcript["segments"]):
|
||||
seg_text = seg.get("text", "")
|
||||
start = int(seg.get("start", 0))
|
||||
end = int(seg.get("end", 0))
|
||||
start = seconds_to_nanoseconds(seg.get("start", 0))
|
||||
end = seconds_to_nanoseconds(seg.get("end", 0))
|
||||
speaker = seg.get("speaker", "")
|
||||
|
||||
resultSegments.append(backend_pb2.TranscriptSegment(
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
import unittest
|
||||
|
||||
import transcript_utils
|
||||
|
||||
|
||||
class TestTranscriptUtils(unittest.TestCase):
|
||||
def test_diarization_requires_hugging_face_token(self):
|
||||
with self.assertRaisesRegex(
|
||||
ValueError,
|
||||
"HF_TOKEN is required for WhisperX diarization",
|
||||
):
|
||||
transcript_utils.require_diarization_token(True, None)
|
||||
|
||||
def test_diarization_does_not_require_token_when_disabled(self):
|
||||
transcript_utils.require_diarization_token(False, None)
|
||||
|
||||
def test_seconds_are_serialized_as_nanoseconds(self):
|
||||
self.assertEqual(
|
||||
transcript_utils.seconds_to_nanoseconds(3.25),
|
||||
3_250_000_000,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,12 @@
|
||||
"""Helpers for WhisperX transcript responses."""
|
||||
|
||||
|
||||
def require_diarization_token(diarize, token):
|
||||
"""Reject diarization when WhisperX cannot load its gated pipeline."""
|
||||
if diarize and not token:
|
||||
raise ValueError("HF_TOKEN is required for WhisperX diarization")
|
||||
|
||||
|
||||
def seconds_to_nanoseconds(seconds):
|
||||
"""Convert WhisperX timestamps to the duration unit used by LocalAI."""
|
||||
return int(seconds * 1_000_000_000)
|
||||
@@ -24,10 +24,17 @@ import (
|
||||
|
||||
// Config represents the launcher configuration
|
||||
type Config struct {
|
||||
ModelsPath string `json:"models_path"`
|
||||
BackendsPath string `json:"backends_path"`
|
||||
Address string `json:"address"`
|
||||
AutoStart bool `json:"auto_start"`
|
||||
ModelsPath string `json:"models_path"`
|
||||
BackendsPath string `json:"backends_path"`
|
||||
Address string `json:"address"`
|
||||
// AutoStart controls whether the launcher starts the LocalAI server as
|
||||
// soon as the launcher itself opens (and right after a fresh install).
|
||||
// Unset means enabled: launching the app must yield a serving endpoint,
|
||||
// which is what the quickstart docs promise. The JSON key is deliberately
|
||||
// not the legacy "auto_start": that field was never honored nor exposed
|
||||
// in any UI, so every existing launcher.json carries an unintentional
|
||||
// false that would keep auto-start permanently off (#11673).
|
||||
AutoStart *bool `json:"auto_start_server"`
|
||||
StartOnBoot bool `json:"start_on_boot"`
|
||||
LogLevel string `json:"log_level"`
|
||||
EnvironmentVars map[string]string `json:"environment_vars"`
|
||||
@@ -122,9 +129,6 @@ func (l *Launcher) Initialize() error {
|
||||
log.Printf("Warning: failed to cleanup partial downloads: %v", err)
|
||||
}
|
||||
|
||||
if l.config.StartOnBoot {
|
||||
l.StartLocalAI()
|
||||
}
|
||||
// Set default paths if not configured (only if not already loaded from config)
|
||||
if l.config.ModelsPath == "" {
|
||||
homeDir, _ := os.UserHomeDir()
|
||||
@@ -156,6 +160,12 @@ func (l *Launcher) Initialize() error {
|
||||
log.Printf("Setting default ShowWelcome: true")
|
||||
}
|
||||
|
||||
if l.config.AutoStart == nil {
|
||||
enabled := true
|
||||
l.config.AutoStart = &enabled
|
||||
log.Printf("Setting default AutoStart: true")
|
||||
}
|
||||
|
||||
// Create directories
|
||||
os.MkdirAll(l.config.ModelsPath, 0755)
|
||||
os.MkdirAll(l.config.BackendsPath, 0755)
|
||||
@@ -177,6 +187,11 @@ func (l *Launcher) Initialize() error {
|
||||
l.showDownloadLocalAIDialog()
|
||||
}
|
||||
})
|
||||
} else if l.ShouldAutoStartServer() {
|
||||
// The launcher is a tray-only app: without this the user launches it,
|
||||
// sees no window and no server, and concludes it does nothing (#11673).
|
||||
log.Printf("Auto-starting LocalAI server")
|
||||
l.autoStartServer()
|
||||
}
|
||||
|
||||
// Check for updates periodically
|
||||
@@ -185,6 +200,35 @@ func (l *Launcher) Initialize() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// ShouldAutoStartServer reports whether the launcher should start the server
|
||||
// without user interaction: at launcher startup and right after a fresh
|
||||
// install. Defaults to enabled; StartOnBoot forces a start even when
|
||||
// auto-start was explicitly disabled, preserving its historical behavior.
|
||||
func (l *Launcher) ShouldAutoStartServer() bool {
|
||||
if l.config == nil {
|
||||
return false
|
||||
}
|
||||
if l.config.StartOnBoot {
|
||||
return true
|
||||
}
|
||||
return l.config.AutoStart == nil || *l.config.AutoStart
|
||||
}
|
||||
|
||||
// autoStartServer starts LocalAI in the background and surfaces failures
|
||||
// through the systray error dialog: during an auto-start there is no visible
|
||||
// window for a regular error dialog to attach to.
|
||||
func (l *Launcher) autoStartServer() {
|
||||
go func() {
|
||||
if err := l.StartLocalAI(); err != nil {
|
||||
log.Printf("Failed to auto-start LocalAI: %v", err)
|
||||
l.updateStatus(fmt.Sprintf("Failed to start LocalAI: %v", err))
|
||||
if l.systray != nil {
|
||||
l.systray.showStartupErrorDialog(err)
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
// StartLocalAI starts the LocalAI server
|
||||
func (l *Launcher) StartLocalAI() error {
|
||||
if l.isRunning {
|
||||
@@ -644,14 +688,22 @@ func (l *Launcher) showDownloadError(title, message string) {
|
||||
// after a fresh install (no LocalAI binary present yet).
|
||||
func (l *Launcher) showDownloadProgress(version, title string) {
|
||||
l.showDownloadProgressWindow(version, title, func(win fyne.Window) {
|
||||
dialog.ShowConfirm("Installation Complete",
|
||||
"LocalAI has been downloaded and installed successfully. You can now start LocalAI from the launcher.",
|
||||
message := "LocalAI has been downloaded and installed successfully. You can now start LocalAI from the launcher."
|
||||
if l.ShouldAutoStartServer() {
|
||||
message = "LocalAI has been downloaded and installed successfully. It will start now: manage it and open the WebUI from the system tray icon."
|
||||
}
|
||||
dialog.ShowConfirm("Installation Complete", message,
|
||||
func(bool) {
|
||||
win.Close()
|
||||
l.updateStatus("LocalAI installed successfully")
|
||||
if l.systray != nil {
|
||||
l.systray.recreateMenu()
|
||||
}
|
||||
// A fresh install should end with a running server, not with
|
||||
// the user hunting for a start button in the tray (#11673).
|
||||
if l.ShouldAutoStartServer() && !l.isRunning {
|
||||
l.autoStartServer()
|
||||
}
|
||||
}, win)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package launcher_test
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
@@ -55,7 +56,8 @@ var _ = Describe("Launcher", func() {
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
config := launcherInstance.GetConfig()
|
||||
Expect(config.ShowWelcome).To(BeTrue())
|
||||
Expect(config.ShowWelcome).ToNot(BeNil())
|
||||
Expect(*config.ShowWelcome).To(BeTrue())
|
||||
Expect(config.Address).To(Equal("127.0.0.1:8080"))
|
||||
Expect(config.LogLevel).To(Equal("info"))
|
||||
})
|
||||
@@ -177,13 +179,53 @@ var _ = Describe("Launcher", func() {
|
||||
|
||||
assertFlagValue("--generated-content-path", filepath.Join(dataPath, "generated"))
|
||||
assertFlagValue("--upload-path", filepath.Join(dataPath, "uploads"))
|
||||
// The bug was the server resolving these to shared /tmp paths.
|
||||
// The bug was the server resolving these to its shared /tmp
|
||||
// defaults. Only reject those specific paths: on Linux the test's
|
||||
// own temp directory legitimately lives under /tmp.
|
||||
for _, a := range args {
|
||||
Expect(a).ToNot(HavePrefix("/tmp/"), "run args must not reference shared /tmp paths, got %s", a)
|
||||
Expect(a).ToNot(HavePrefix("/tmp/generated"), "run args must not reference the shared /tmp generated-content default, got %s", a)
|
||||
Expect(a).ToNot(HavePrefix("/tmp/upload"), "run args must not reference the shared /tmp upload default, got %s", a)
|
||||
}
|
||||
})
|
||||
})
|
||||
|
||||
// Regression for "Mac dmg launcher launches nothing" (issue #11673): the
|
||||
// launcher created empty log files and served nothing because nothing ever
|
||||
// started the server unless the unrelated "start on system boot" option was
|
||||
// enabled. Launching the app must yield a serving endpoint by default.
|
||||
Describe("ShouldAutoStartServer", func() {
|
||||
It("should auto-start by default when nothing is configured", func() {
|
||||
Expect(launcherInstance.ShouldAutoStartServer()).To(BeTrue())
|
||||
})
|
||||
|
||||
It("should respect an explicit opt-out", func() {
|
||||
config := launcherInstance.GetConfig()
|
||||
err := json.Unmarshal([]byte(`{"auto_start_server": false}`), config)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
Expect(launcherInstance.ShouldAutoStartServer()).To(BeFalse())
|
||||
})
|
||||
|
||||
It("should still auto-start when StartOnBoot is set even if auto-start is off", func() {
|
||||
config := launcherInstance.GetConfig()
|
||||
err := json.Unmarshal([]byte(`{"auto_start_server": false, "start_on_boot": true}`), config)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
Expect(launcherInstance.ShouldAutoStartServer()).To(BeTrue())
|
||||
})
|
||||
|
||||
It("should ignore the legacy auto_start key older launchers persisted as false", func() {
|
||||
// Old launchers marshaled the never-honored AutoStart field as
|
||||
// "auto_start": false into every launcher.json. That stale value
|
||||
// carries no user intent and must not disable auto-start.
|
||||
config := launcherInstance.GetConfig()
|
||||
err := json.Unmarshal([]byte(`{"auto_start": false}`), config)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
Expect(launcherInstance.ShouldAutoStartServer()).To(BeTrue())
|
||||
})
|
||||
})
|
||||
|
||||
Describe("Logs", func() {
|
||||
It("should return empty logs initially", func() {
|
||||
logs := launcherInstance.GetLogs()
|
||||
@@ -210,13 +252,38 @@ var _ = Describe("Launcher", func() {
|
||||
})
|
||||
})
|
||||
|
||||
// Regression for the welcome window suppressing itself (part of issue
|
||||
// #11673): the "don't show this welcome window again" checkbox was
|
||||
// initialized with the ShowWelcome value itself, so on the very first
|
||||
// showing it came up checked AND its change callback persisted
|
||||
// ShowWelcome=false, hiding the welcome window forever.
|
||||
var _ = Describe("WelcomeDontShowAgainChecked", func() {
|
||||
It("should be unchecked when the welcome window is enabled", func() {
|
||||
show := true
|
||||
config := &launcher.Config{ShowWelcome: &show}
|
||||
Expect(launcher.WelcomeDontShowAgainChecked(config)).To(BeFalse())
|
||||
})
|
||||
|
||||
It("should be checked when the user opted out", func() {
|
||||
show := false
|
||||
config := &launcher.Config{ShowWelcome: &show}
|
||||
Expect(launcher.WelcomeDontShowAgainChecked(config)).To(BeTrue())
|
||||
})
|
||||
|
||||
It("should be unchecked when the preference is unset", func() {
|
||||
Expect(launcher.WelcomeDontShowAgainChecked(&launcher.Config{})).To(BeFalse())
|
||||
Expect(launcher.WelcomeDontShowAgainChecked(nil)).To(BeFalse())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("Config", func() {
|
||||
It("should have proper JSON tags", func() {
|
||||
autoStart := true
|
||||
config := &launcher.Config{
|
||||
ModelsPath: "/test/models",
|
||||
BackendsPath: "/test/backends",
|
||||
Address: ":8080",
|
||||
AutoStart: true,
|
||||
AutoStart: &autoStart,
|
||||
LogLevel: "info",
|
||||
EnvironmentVars: map[string]string{"TEST": "value"},
|
||||
}
|
||||
@@ -224,7 +291,7 @@ var _ = Describe("Config", func() {
|
||||
Expect(config.ModelsPath).To(Equal("/test/models"))
|
||||
Expect(config.BackendsPath).To(Equal("/test/backends"))
|
||||
Expect(config.Address).To(Equal(":8080"))
|
||||
Expect(config.AutoStart).To(BeTrue())
|
||||
Expect(*config.AutoStart).To(BeTrue())
|
||||
Expect(config.LogLevel).To(Equal("info"))
|
||||
Expect(config.EnvironmentVars).To(HaveKeyWithValue("TEST", "value"))
|
||||
})
|
||||
|
||||
@@ -34,6 +34,7 @@ type LauncherUI struct {
|
||||
backendsPathEntry *widget.Entry
|
||||
addressEntry *widget.Entry
|
||||
logLevelSelect *widget.Select
|
||||
autoStartCheck *widget.Check
|
||||
startOnBootCheck *widget.Check
|
||||
|
||||
// Environment Variables
|
||||
@@ -75,6 +76,7 @@ func NewLauncherUI() *LauncherUI {
|
||||
backendsPathEntry: widget.NewEntry(),
|
||||
addressEntry: widget.NewEntry(),
|
||||
logLevelSelect: widget.NewSelect([]string{"error", "warn", "info", "debug", "trace"}, nil),
|
||||
autoStartCheck: widget.NewCheck("Start LocalAI when the launcher opens", nil),
|
||||
startOnBootCheck: widget.NewCheck("Start LocalAI on system boot", nil),
|
||||
logText: widget.NewMultiLineEntry(),
|
||||
progressBar: widget.NewProgressBar(),
|
||||
@@ -117,6 +119,7 @@ func (ui *LauncherUI) createConfigTab() *fyne.Container {
|
||||
widget.NewLabel("Log Level:"),
|
||||
ui.logLevelSelect,
|
||||
),
|
||||
ui.autoStartCheck,
|
||||
ui.startOnBootCheck,
|
||||
))
|
||||
|
||||
@@ -401,6 +404,8 @@ func (ui *LauncherUI) saveConfiguration() {
|
||||
config.BackendsPath = ui.backendsPathEntry.Text
|
||||
config.Address = ui.addressEntry.Text
|
||||
config.LogLevel = ui.logLevelSelect.Selected
|
||||
autoStart := ui.autoStartCheck.Checked
|
||||
config.AutoStart = &autoStart
|
||||
config.StartOnBoot = ui.startOnBootCheck.Checked
|
||||
|
||||
// Ensure environment variables are included in the configuration
|
||||
@@ -583,6 +588,7 @@ func (ui *LauncherUI) LoadConfiguration() {
|
||||
ui.backendsPathEntry.SetText(config.BackendsPath)
|
||||
ui.addressEntry.SetText(config.Address)
|
||||
ui.logLevelSelect.SetSelected(config.LogLevel)
|
||||
ui.autoStartCheck.SetChecked(config.AutoStart == nil || *config.AutoStart)
|
||||
ui.startOnBootCheck.SetChecked(config.StartOnBoot)
|
||||
|
||||
// Load environment variables
|
||||
@@ -616,6 +622,14 @@ func (ui *LauncherUI) UpdateRunningState(isRunning bool) {
|
||||
})
|
||||
}
|
||||
|
||||
// WelcomeDontShowAgainChecked reports the initial state of the welcome
|
||||
// window's "don't show this welcome window again" checkbox for the given
|
||||
// config: checked only when the user has already opted out of the welcome
|
||||
// window.
|
||||
func WelcomeDontShowAgainChecked(config *Config) bool {
|
||||
return config != nil && config.ShowWelcome != nil && !*config.ShowWelcome
|
||||
}
|
||||
|
||||
// ShowWelcomeWindow displays the welcome window with helpful information
|
||||
func (ui *LauncherUI) ShowWelcomeWindow() {
|
||||
if ui.launcher == nil || ui.launcher.window == nil {
|
||||
@@ -677,19 +691,20 @@ Getting Started:
|
||||
ui.openURL("https://discord.gg/XgwjKptP7Z")
|
||||
})
|
||||
|
||||
// Checkbox to disable welcome window
|
||||
dontShowAgainCheck := widget.NewCheck("Don't show this welcome window again", func(checked bool) {
|
||||
// Checkbox to disable welcome window. The initial state is applied
|
||||
// BEFORE the change callback is attached: SetChecked fires OnChanged,
|
||||
// and letting the initialization itself persist a ShowWelcome flip is
|
||||
// exactly the bug that suppressed this window forever after its first
|
||||
// showing (#11673).
|
||||
dontShowAgainCheck := widget.NewCheck("Don't show this welcome window again", nil)
|
||||
dontShowAgainCheck.SetChecked(WelcomeDontShowAgainChecked(ui.launcher.GetConfig()))
|
||||
dontShowAgainCheck.OnChanged = func(checked bool) {
|
||||
if ui.launcher != nil {
|
||||
config := ui.launcher.GetConfig()
|
||||
v := !checked
|
||||
config.ShowWelcome = &v
|
||||
ui.launcher.SetConfig(config)
|
||||
}
|
||||
})
|
||||
|
||||
config := ui.launcher.GetConfig()
|
||||
if config.ShowWelcome != nil {
|
||||
dontShowAgainCheck.SetChecked(*config.ShowWelcome)
|
||||
}
|
||||
|
||||
// Close button
|
||||
|
||||
@@ -15,6 +15,7 @@ import (
|
||||
"github.com/mudler/LocalAI/core/services/distributed"
|
||||
"github.com/mudler/LocalAI/core/services/jobs"
|
||||
"github.com/mudler/LocalAI/core/services/messaging"
|
||||
"github.com/mudler/LocalAI/core/services/monitoring"
|
||||
"github.com/mudler/LocalAI/core/services/nodes"
|
||||
"github.com/mudler/LocalAI/core/services/nodes/prefixcache"
|
||||
"github.com/mudler/LocalAI/core/services/storage"
|
||||
@@ -162,6 +163,26 @@ func initDistributed(cfg *config.ApplicationConfig, authDB *gorm.DB, configLoade
|
||||
}
|
||||
xlog.Info("Node registry initialized")
|
||||
|
||||
// Bound durable heartbeat writes: a beat that only carries a fresher
|
||||
// timestamp is what turned backend_nodes into a 460 MB six-row table.
|
||||
registry.SetHeartbeatCheckpoint(cfg.Distributed.NodeHeartbeatCheckpointOrDefault())
|
||||
|
||||
// Measure the vacuum horizon. The 42 days it stayed open went unnoticed
|
||||
// because no gauge reported it until models started failing to load.
|
||||
if err := monitoring.RegisterControlPlaneDBMetrics(authDB, 30*time.Second); err != nil {
|
||||
// Metrics are diagnostic; a failure here must not stop the frontend.
|
||||
xlog.Warn("Control-plane database metrics unavailable", "error", err)
|
||||
}
|
||||
|
||||
// Let scheduling rules be keyed by a model alias. The registry resolves a
|
||||
// rule's name through the config loader to find the model it governs, so an
|
||||
// operator can pin placement to a stable name like "production" and have it
|
||||
// follow the alias when the alias is repointed. Wired before the seed below
|
||||
// and before the reconciler starts, so the first tick already resolves.
|
||||
if configLoader != nil {
|
||||
registry.SetAliasResolver(configLoader)
|
||||
}
|
||||
|
||||
// Seed declarative per-model scheduling config (LOCALAI_MODEL_SCHEDULING /
|
||||
// LOCALAI_MODEL_SCHEDULING_CONFIG). Authoritative: overwrites matching models
|
||||
// on every boot. Runs before the reconciler starts so the first tick already
|
||||
|
||||
@@ -267,6 +267,10 @@ func New(opts ...config.AppOption) (*Application, error) {
|
||||
}
|
||||
|
||||
// Initialize distributed mode services (NATS, object storage, node registry)
|
||||
// revisionStore is built inside the distributed block below but used after
|
||||
// the model configs are loaded, so it is declared out here.
|
||||
var revisionStore modeladmin.RevisionStore
|
||||
|
||||
distSvc, err := initDistributed(options, application.authDB, application.ModelConfigLoader())
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("distributed mode initialization failed: %w", err)
|
||||
@@ -373,6 +377,9 @@ func New(opts ...config.AppOption) (*Application, error) {
|
||||
cfgLoaderOpts := options.ToConfigLoaderOptions()
|
||||
modelRevisionLifecycle := modeladmin.NewDistributedModelRevisionLifecycle(distSvc.Registry, distSvc.ModelCleanup)
|
||||
gs.SetModelRevisionLifecycle(modelRevisionLifecycle)
|
||||
// Captured here, used after the model configs are loaded below: the
|
||||
// resync reads the loader, which is still empty at this point.
|
||||
revisionStore = modeladmin.NewRevisionStore(distSvc.Registry, modelRevisionLifecycle)
|
||||
gs.OnModelsChanged = func(evt messaging.CacheInvalidateEvent) {
|
||||
// ApplyRemoteChange honors the op: a "delete" prunes the element
|
||||
// (a reload-from-path is additive and cannot drop it), anything
|
||||
@@ -419,6 +426,18 @@ func New(opts ...config.AppOption) (*Application, error) {
|
||||
xlog.Error("error loading config files", "error", err)
|
||||
}
|
||||
|
||||
// Bring the controller's stored revisions back in line with the
|
||||
// configuration just loaded. An inference request may only establish a
|
||||
// revision, never replace one, so a model whose stored value has drifted
|
||||
// stays unroutable until something republishes it. This has to run after
|
||||
// the load above: the loader is empty until then, and a resync against an
|
||||
// empty loader silently reconciles nothing.
|
||||
if revisionStore != nil {
|
||||
if err := modeladmin.ResyncModelConfigRevisions(options.Context, application.ModelConfigLoader(), options, revisionStore); err != nil {
|
||||
xlog.Warn("Failed to resync model config revisions", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := gallery.RegisterBackends(options.SystemState, application.ModelLoader()); err != nil {
|
||||
xlog.Error("error registering external backends", "error", err)
|
||||
}
|
||||
@@ -446,13 +465,20 @@ func New(opts ...config.AppOption) (*Application, error) {
|
||||
// Wire gallery generation counter into VRAM caches so they invalidate
|
||||
// when gallery data refreshes instead of using a fixed TTL.
|
||||
vram.SetGalleryGenerationFunc(gallery.GalleryGeneration)
|
||||
if options.AutoloadGalleries {
|
||||
if options.VRAMPersistentCache {
|
||||
// Remote GGUF probes can transfer substantial metadata. Keep successful
|
||||
// results across restarts so the startup warmer does not repeat that work.
|
||||
vram.ConfigurePersistentCache(filepath.Join(options.SystemState.Model.ModelsPath, "..", "cache", "vram"), 24*time.Hour)
|
||||
}
|
||||
|
||||
// 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())
|
||||
// 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 {
|
||||
|
||||
+10
-2
@@ -202,10 +202,18 @@ func ModelOptions(c config.ModelConfig, so *config.ApplicationConfig, opts ...mo
|
||||
model.WithContext(so.Context),
|
||||
model.WithModelID(c.ModelID()),
|
||||
}
|
||||
if revision, err := config.ModelConfigRevision(&c); err == nil {
|
||||
// Use the revision stamped when the configuration was parsed, and only
|
||||
// that. By this point c has been merged with the request's prediction
|
||||
// parameters and had SetDefaults applied, so hashing it here would produce
|
||||
// a revision that depends on the request body and on whether the model file
|
||||
// parsed, which the controller reads as a config change and rejects. Every
|
||||
// config the loader hands out is stamped; an unstamped one was synthesized
|
||||
// elsewhere and is routed without a revision rather than with a wrong one.
|
||||
if revision := c.PersistedConfigRevision(); revision != "" {
|
||||
defOpts = append(defOpts, model.WithConfigRevision(revision))
|
||||
} else {
|
||||
xlog.Warn("Failed to compute model configuration revision", "model", c.ModelID(), "error", err)
|
||||
xlog.Warn("Model configuration carries no revision stamp; routing without one",
|
||||
"model", c.ModelID())
|
||||
}
|
||||
managedPrimary := len(c.Artifacts) > 0 && c.Artifacts[0].Resolved != nil
|
||||
if managedPrimary {
|
||||
|
||||
@@ -53,6 +53,7 @@ type RunCMD struct {
|
||||
BackendGalleries string `env:"LOCALAI_BACKEND_GALLERIES,BACKEND_GALLERIES" help:"JSON list of backend galleries" group:"backends" default:"${backends}"`
|
||||
Galleries string `env:"LOCALAI_GALLERIES,GALLERIES" help:"JSON list of galleries" group:"models" default:"${galleries}"`
|
||||
AutoloadGalleries bool `env:"LOCALAI_AUTOLOAD_GALLERIES,AUTOLOAD_GALLERIES" group:"models" default:"true"`
|
||||
VRAMPersistentCache bool `env:"LOCALAI_VRAM_PERSISTENT_CACHE,VRAM_PERSISTENT_CACHE" group:"models" default:"true" help:"Persist successful remote VRAM metadata probes across restarts"`
|
||||
AutoloadBackendGalleries bool `env:"LOCALAI_AUTOLOAD_BACKEND_GALLERIES,AUTOLOAD_BACKEND_GALLERIES" group:"backends" default:"true"`
|
||||
BackendImagesReleaseTag string `env:"LOCALAI_BACKEND_IMAGES_RELEASE_TAG,BACKEND_IMAGES_RELEASE_TAG" help:"Fallback release tag for backend images" group:"backends" default:"latest"`
|
||||
BackendImagesBranchTag string `env:"LOCALAI_BACKEND_IMAGES_BRANCH_TAG,BACKEND_IMAGES_BRANCH_TAG" help:"Fallback branch tag for backend images" group:"backends" default:"master"`
|
||||
@@ -181,6 +182,8 @@ type RunCMD struct {
|
||||
BackendUpgradeTimeout string `env:"LOCALAI_NATS_BACKEND_UPGRADE_TIMEOUT" help:"NATS round-trip timeout for backend.upgrade requests (default 15m)." group:"distributed"`
|
||||
ModelLoadTimeout string `env:"LOCALAI_NATS_MODEL_LOAD_TIMEOUT" help:"Fixed gRPC deadline for the remote LoadModel call sent to a worker node once its backend is installed and model files are staged. Unset (the default), the deadline is derived from the checkpoint size instead: 5m plus 20s per GiB, capped at 6h, so multi-tens-of-GB diffusion/video checkpoints get the minutes they need without a fixed cliff. Set this only to pin a specific budget; the value is used verbatim, including when it is shorter than the derived one." group:"distributed"`
|
||||
ModelLoadWait string `env:"LOCALAI_MODEL_LOAD_WAIT" help:"How long an inference request waits for a model that is still cold-loading onto a worker before it is answered with 503, a Retry-After header and live staging progress (default 60s). The request is served the moment the model becomes ready, so a model already most of the way staged needs no client retry. Set to 0 to wait as long as the load takes — only safe when no ingress or load balancer with an idle timeout sits in front." group:"distributed"`
|
||||
StaleNodeThreshold string `env:"LOCALAI_STALE_NODE_THRESHOLD" help:"How long a worker node may go without a durable heartbeat before the health monitor marks it offline (default 5m). Because a beat that only carries a fresher timestamp is held back by --node-heartbeat-checkpoint, this must stay comfortably wider than that interval; raise both together. Dead-node detection through the per-model gRPC health check and through request-time failure is unaffected by this knob." group:"distributed"`
|
||||
NodeHeartbeatCheckpoint string `env:"LOCALAI_NODE_HEARTBEAT_CHECKPOINT" help:"Minimum gap between durable heartbeat writes for a worker node (default 60s). A beat that only carries a fresher timestamp is dropped until this interval elapses; every field is compared against the value last written, so a node's first beat, a changed total VRAM/total disk/GPU vendor, and a free VRAM/RAM/disk reading that has moved more than 256 MiB from the written value all still write immediately, and a node that is not active is never suppressed. Set below the worker heartbeat interval to write on every beat." group:"distributed"`
|
||||
NatsAccountSeed string `env:"LOCALAI_NATS_ACCOUNT_SEED" help:"NATS account signing seed (SU...) used to mint per-node worker JWTs at registration" group:"distributed"`
|
||||
NatsServiceJWT string `env:"LOCALAI_NATS_SERVICE_JWT" help:"NATS user JWT for the frontend (and agent workers) to publish control-plane messages" group:"distributed"`
|
||||
NatsServiceSeed string `env:"LOCALAI_NATS_SERVICE_SEED" help:"NATS user signing seed (SU...) paired with LOCALAI_NATS_SERVICE_JWT" group:"distributed"`
|
||||
@@ -302,6 +305,7 @@ func (r *RunCMD) Run(ctx *cliContext.Context) error {
|
||||
config.WithF16(r.F16),
|
||||
config.WithStringGalleries(r.Galleries),
|
||||
config.WithBackendGalleries(r.BackendGalleries),
|
||||
config.WithVRAMPersistentCache(r.VRAMPersistentCache),
|
||||
config.WithCors(r.CORS),
|
||||
config.WithCorsAllowOrigins(r.CORSAllowOrigins),
|
||||
config.WithDisableCSRF(r.DisableCSRF),
|
||||
@@ -395,6 +399,20 @@ func (r *RunCMD) Run(ctx *cliContext.Context) error {
|
||||
}
|
||||
opts = append(opts, config.WithModelLoadWait(d))
|
||||
}
|
||||
if r.StaleNodeThreshold != "" {
|
||||
d, err := parseDistributedDuration("LOCALAI_STALE_NODE_THRESHOLD", r.StaleNodeThreshold)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
opts = append(opts, config.WithStaleNodeThreshold(d))
|
||||
}
|
||||
if r.NodeHeartbeatCheckpoint != "" {
|
||||
d, err := parseDistributedDuration("LOCALAI_NODE_HEARTBEAT_CHECKPOINT", r.NodeHeartbeatCheckpoint)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
opts = append(opts, config.WithNodeHeartbeatCheckpoint(d))
|
||||
}
|
||||
if r.RegistrationToken != "" {
|
||||
opts = append(opts, config.WithRegistrationToken(r.RegistrationToken))
|
||||
}
|
||||
|
||||
@@ -10,7 +10,7 @@ func ParseNodeLabels(input string) map[string]string {
|
||||
if input == "" {
|
||||
return labels
|
||||
}
|
||||
for _, pair := range strings.Split(input, ",") {
|
||||
for pair := range strings.SplitSeq(input, ",") {
|
||||
pair = strings.TrimSpace(pair)
|
||||
if k, v, ok := strings.Cut(pair, "="); ok {
|
||||
labels[strings.TrimSpace(k)] = strings.TrimSpace(v)
|
||||
|
||||
@@ -125,6 +125,7 @@ type ApplicationConfig struct {
|
||||
ExternalGRPCBackends map[string]string
|
||||
|
||||
AutoloadGalleries, AutoloadBackendGalleries bool
|
||||
VRAMPersistentCache bool
|
||||
AutoUpgradeBackends bool
|
||||
PreferDevelopmentBackends bool
|
||||
|
||||
@@ -284,6 +285,7 @@ func NewApplicationConfig(o ...AppOption) *ApplicationConfig {
|
||||
// toggle can still turn it off (a persisted false wins - see
|
||||
// loadRuntimeSettingsFromFile).
|
||||
EnableBackendLogging: true,
|
||||
VRAMPersistentCache: true,
|
||||
ArtifactDownloadConcurrency: modelartifacts.DefaultDownloadConcurrency,
|
||||
AgentJobRetentionDays: 30, // Default: 30 days
|
||||
LRUEvictionMaxRetries: 30, // Default: 30 retries
|
||||
@@ -596,6 +598,10 @@ func WithAutoUpgradeBackends(v bool) AppOption {
|
||||
return func(o *ApplicationConfig) { o.AutoUpgradeBackends = v }
|
||||
}
|
||||
|
||||
func WithVRAMPersistentCache(v bool) AppOption {
|
||||
return func(o *ApplicationConfig) { o.VRAMPersistentCache = v }
|
||||
}
|
||||
|
||||
func WithRequireBackendIntegrity(v bool) AppOption {
|
||||
return func(o *ApplicationConfig) { o.RequireBackendIntegrity = v }
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"time"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
@@ -9,6 +10,15 @@ import (
|
||||
|
||||
var _ = Describe("ApplicationConfig RuntimeSettings Conversion", func() {
|
||||
Describe("ToRuntimeSettings", func() {
|
||||
It("includes the persistent VRAM cache toggle", func() {
|
||||
encoded, err := json.Marshal(NewApplicationConfig().ToRuntimeSettings())
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
var settings map[string]any
|
||||
Expect(json.Unmarshal(encoded, &settings)).To(Succeed())
|
||||
Expect(settings).To(HaveKeyWithValue("vram_persistent_cache", true))
|
||||
})
|
||||
|
||||
It("should convert all fields correctly", func() {
|
||||
appConfig := &ApplicationConfig{
|
||||
WatchDog: true,
|
||||
|
||||
@@ -57,12 +57,13 @@ type DistributedConfig struct {
|
||||
StorageSecretKey string // --storage-secret-key / LOCALAI_STORAGE_SECRET_KEY
|
||||
|
||||
// Timeout configuration (all have sensible defaults — zero means use default)
|
||||
MCPToolTimeout time.Duration // MCP tool execution timeout (default 360s)
|
||||
MCPDiscoveryTimeout time.Duration // MCP discovery timeout (default 60s)
|
||||
WorkerWaitTimeout time.Duration // Max wait for healthy worker at startup (default 5m)
|
||||
DrainTimeout time.Duration // Time to wait for in-flight requests during drain (default 30s)
|
||||
HealthCheckInterval time.Duration // Health monitor check interval (default 15s)
|
||||
StaleNodeThreshold time.Duration // Time before a node is considered stale (default 60s)
|
||||
MCPToolTimeout time.Duration // MCP tool execution timeout (default 360s)
|
||||
MCPDiscoveryTimeout time.Duration // MCP discovery timeout (default 60s)
|
||||
WorkerWaitTimeout time.Duration // Max wait for healthy worker at startup (default 5m)
|
||||
DrainTimeout time.Duration // Time to wait for in-flight requests during drain (default 30s)
|
||||
HealthCheckInterval time.Duration // Health monitor check interval (default 15s)
|
||||
StaleNodeThreshold time.Duration // Time before a node is considered stale (default 5m)
|
||||
NodeHeartbeatCheckpoint time.Duration // Minimum gap between durable heartbeat writes (default 60s, 0 = every beat)
|
||||
// DisablePerModelHealthCheck turns off the health monitor's per-model
|
||||
// gRPC probe. When enabled (the default), the monitor pings each model's
|
||||
// gRPC address and removes stale node_models rows whose backend has
|
||||
@@ -165,16 +166,17 @@ func (c DistributedConfig) Validate() error {
|
||||
c.NatsAuthConfig().WarnIfInsecure(true)
|
||||
// Check for negative durations
|
||||
for name, d := range map[string]time.Duration{
|
||||
FlagMCPToolTimeout: c.MCPToolTimeout,
|
||||
FlagMCPDiscoveryTimeout: c.MCPDiscoveryTimeout,
|
||||
FlagWorkerWaitTimeout: c.WorkerWaitTimeout,
|
||||
FlagDrainTimeout: c.DrainTimeout,
|
||||
FlagHealthCheckInterval: c.HealthCheckInterval,
|
||||
FlagStaleNodeThreshold: c.StaleNodeThreshold,
|
||||
FlagMCPCIJobTimeout: c.MCPCIJobTimeout,
|
||||
FlagBackendInstallTimeout: c.BackendInstallTimeout,
|
||||
FlagBackendUpgradeTimeout: c.BackendUpgradeTimeout,
|
||||
FlagModelLoadTimeout: c.ModelLoadTimeout,
|
||||
FlagMCPToolTimeout: c.MCPToolTimeout,
|
||||
FlagMCPDiscoveryTimeout: c.MCPDiscoveryTimeout,
|
||||
FlagWorkerWaitTimeout: c.WorkerWaitTimeout,
|
||||
FlagDrainTimeout: c.DrainTimeout,
|
||||
FlagHealthCheckInterval: c.HealthCheckInterval,
|
||||
FlagStaleNodeThreshold: c.StaleNodeThreshold,
|
||||
FlagNodeHeartbeatCheckpoint: c.NodeHeartbeatCheckpoint,
|
||||
FlagMCPCIJobTimeout: c.MCPCIJobTimeout,
|
||||
FlagBackendInstallTimeout: c.BackendInstallTimeout,
|
||||
FlagBackendUpgradeTimeout: c.BackendUpgradeTimeout,
|
||||
FlagModelLoadTimeout: c.ModelLoadTimeout,
|
||||
} {
|
||||
if d < 0 {
|
||||
return fmt.Errorf("%s must not be negative", name)
|
||||
@@ -337,6 +339,27 @@ func WithModelLoadWait(d time.Duration) AppOption {
|
||||
}
|
||||
}
|
||||
|
||||
// WithStaleNodeThreshold sets how long a node may go without a durable
|
||||
// heartbeat before the health monitor marks it offline. It has to be raised
|
||||
// alongside WithNodeHeartbeatCheckpoint: a checkpoint interval wider than this
|
||||
// threshold makes every healthy node look dead the moment its beats start
|
||||
// being suppressed.
|
||||
func WithStaleNodeThreshold(d time.Duration) AppOption {
|
||||
return func(o *ApplicationConfig) {
|
||||
o.Distributed.StaleNodeThreshold = d
|
||||
}
|
||||
}
|
||||
|
||||
// WithNodeHeartbeatCheckpoint bounds durable heartbeat writes. A zero d is
|
||||
// deliberately not special-cased into "unbounded": NodeHeartbeatCheckpointOrDefault
|
||||
// reads zero as unset, and an operator who wants a write per beat sets a value
|
||||
// below the worker's heartbeat interval instead.
|
||||
func WithNodeHeartbeatCheckpoint(d time.Duration) AppOption {
|
||||
return func(o *ApplicationConfig) {
|
||||
o.Distributed.NodeHeartbeatCheckpoint = d
|
||||
}
|
||||
}
|
||||
|
||||
var EnableAutoApproveNodes = func(o *ApplicationConfig) {
|
||||
o.Distributed.AutoApproveNodes = true
|
||||
}
|
||||
@@ -391,17 +414,18 @@ func WithModelSchedulingConfigPath(path string) AppOption {
|
||||
// them as constants prevents the string from drifting from the actual
|
||||
// flag a future rename would produce.
|
||||
const (
|
||||
FlagMCPToolTimeout = "mcp-tool-timeout"
|
||||
FlagMCPDiscoveryTimeout = "mcp-discovery-timeout"
|
||||
FlagWorkerWaitTimeout = "worker-wait-timeout"
|
||||
FlagDrainTimeout = "drain-timeout"
|
||||
FlagHealthCheckInterval = "health-check-interval"
|
||||
FlagStaleNodeThreshold = "stale-node-threshold"
|
||||
FlagMCPCIJobTimeout = "mcp-ci-job-timeout"
|
||||
FlagBackendInstallTimeout = "backend-install-timeout"
|
||||
FlagBackendUpgradeTimeout = "backend-upgrade-timeout"
|
||||
FlagModelLoadTimeout = "model-load-timeout"
|
||||
FlagModelLoadWait = "model-load-wait"
|
||||
FlagMCPToolTimeout = "mcp-tool-timeout"
|
||||
FlagMCPDiscoveryTimeout = "mcp-discovery-timeout"
|
||||
FlagWorkerWaitTimeout = "worker-wait-timeout"
|
||||
FlagDrainTimeout = "drain-timeout"
|
||||
FlagHealthCheckInterval = "health-check-interval"
|
||||
FlagStaleNodeThreshold = "stale-node-threshold"
|
||||
FlagNodeHeartbeatCheckpoint = "node-heartbeat-checkpoint"
|
||||
FlagMCPCIJobTimeout = "mcp-ci-job-timeout"
|
||||
FlagBackendInstallTimeout = "backend-install-timeout"
|
||||
FlagBackendUpgradeTimeout = "backend-upgrade-timeout"
|
||||
FlagModelLoadTimeout = "model-load-timeout"
|
||||
FlagModelLoadWait = "model-load-wait"
|
||||
// FlagDiskHeadroomCheck names the disk-headroom toggle. It is quoted in
|
||||
// the warning the check emits while disabled, so the operator reading a
|
||||
// log line knows exactly which knob produced it.
|
||||
@@ -410,16 +434,22 @@ const (
|
||||
|
||||
// Defaults for distributed timeouts.
|
||||
const (
|
||||
DefaultMCPToolTimeout = 360 * time.Second
|
||||
DefaultMCPDiscoveryTimeout = 60 * time.Second
|
||||
DefaultWorkerWaitTimeout = 5 * time.Minute
|
||||
DefaultDrainTimeout = 30 * time.Second
|
||||
DefaultHealthCheckInterval = 15 * time.Second
|
||||
DefaultStaleNodeThreshold = 60 * time.Second
|
||||
DefaultMCPCIJobTimeout = 10 * time.Minute
|
||||
DefaultBackendInstallTimeout = 15 * time.Minute
|
||||
DefaultBackendUpgradeTimeout = 15 * time.Minute
|
||||
DefaultModelLoadTimeout = 5 * time.Minute
|
||||
DefaultMCPToolTimeout = 360 * time.Second
|
||||
DefaultMCPDiscoveryTimeout = 60 * time.Second
|
||||
DefaultWorkerWaitTimeout = 5 * time.Minute
|
||||
DefaultDrainTimeout = 30 * time.Second
|
||||
DefaultHealthCheckInterval = 15 * time.Second
|
||||
// A beat that only refreshes the timestamp is now dropped until the
|
||||
// checkpoint interval elapses, so the persisted column is up to one
|
||||
// interval stale by design. The threshold covers that plus jitter.
|
||||
// A genuinely dead node is still caught sooner by the per-model gRPC
|
||||
// health check and by request-time failure, neither of which reads this.
|
||||
DefaultStaleNodeThreshold = 5 * time.Minute
|
||||
DefaultNodeHeartbeatCheckpoint = 60 * time.Second
|
||||
DefaultMCPCIJobTimeout = 10 * time.Minute
|
||||
DefaultBackendInstallTimeout = 15 * time.Minute
|
||||
DefaultBackendUpgradeTimeout = 15 * time.Minute
|
||||
DefaultModelLoadTimeout = 5 * time.Minute
|
||||
// DefaultModelLoadWait is how long a request waits for a cold-loading model
|
||||
// before it is answered with 503 and live progress. Chosen to sit under the
|
||||
// idle timeout of typical ingress/LB defaults, so the answer comes from
|
||||
@@ -519,6 +549,14 @@ func (c DistributedConfig) StaleNodeThresholdOrDefault() time.Duration {
|
||||
return cmp.Or(c.StaleNodeThreshold, DefaultStaleNodeThreshold)
|
||||
}
|
||||
|
||||
// NodeHeartbeatCheckpointOrDefault returns the configured interval or the
|
||||
// default. A configured zero is indistinguishable from unset here, which is
|
||||
// intentional: cmp.Or falls back to the default, and an operator who wants a
|
||||
// write per beat sets a value below the heartbeat interval instead.
|
||||
func (c DistributedConfig) NodeHeartbeatCheckpointOrDefault() time.Duration {
|
||||
return cmp.Or(c.NodeHeartbeatCheckpoint, DefaultNodeHeartbeatCheckpoint)
|
||||
}
|
||||
|
||||
// MCPCIJobTimeoutOrDefault returns the configured MCP CI job timeout or the default.
|
||||
func (c DistributedConfig) MCPCIJobTimeoutOrDefault() time.Duration {
|
||||
return cmp.Or(c.MCPCIJobTimeout, DefaultMCPCIJobTimeout)
|
||||
|
||||
@@ -47,6 +47,27 @@ var _ = Describe("DistributedConfig backend NATS timeouts", func() {
|
||||
})
|
||||
})
|
||||
|
||||
// Heartbeat checkpointing makes last_heartbeat up to one checkpoint interval
|
||||
// stale by design, which is why the threshold defaults to 5 minutes. An
|
||||
// operator who widens the checkpoint has to widen this to match, so it has to
|
||||
// be reachable from the CLI rather than being a compile-time constant.
|
||||
var _ = Describe("DistributedConfig stale node threshold", func() {
|
||||
It("defaults to 5 minutes, wide enough to cover a suppressed beat", func() {
|
||||
Expect(config.DistributedConfig{}.StaleNodeThresholdOrDefault()).
|
||||
To(Equal(5 * time.Minute))
|
||||
Expect(config.DefaultStaleNodeThreshold).
|
||||
To(BeNumerically(">", config.DefaultNodeHeartbeatCheckpoint),
|
||||
"a threshold at or below the checkpoint interval marks healthy, "+
|
||||
"beating nodes offline every cycle")
|
||||
})
|
||||
|
||||
It("is configurable, so a widened checkpoint can be matched", func() {
|
||||
o := config.NewApplicationConfig(config.WithStaleNodeThreshold(20 * time.Minute))
|
||||
Expect(o.Distributed.StaleNodeThreshold).To(Equal(20 * time.Minute))
|
||||
Expect(o.Distributed.StaleNodeThresholdOrDefault()).To(Equal(20 * time.Minute))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("DistributedConfig flag-name constants", func() {
|
||||
// Pin the kebab-case strings so a rename of the Go field name (or a
|
||||
// CLI flag naming convention change) forces the constant to update,
|
||||
@@ -62,6 +83,7 @@ var _ = Describe("DistributedConfig flag-name constants", func() {
|
||||
Entry("drain timeout", config.FlagDrainTimeout, "drain-timeout"),
|
||||
Entry("health check interval", config.FlagHealthCheckInterval, "health-check-interval"),
|
||||
Entry("stale node threshold", config.FlagStaleNodeThreshold, "stale-node-threshold"),
|
||||
Entry("node heartbeat checkpoint", config.FlagNodeHeartbeatCheckpoint, "node-heartbeat-checkpoint"),
|
||||
Entry("MCP CI job timeout", config.FlagMCPCIJobTimeout, "mcp-ci-job-timeout"),
|
||||
Entry("backend install timeout", config.FlagBackendInstallTimeout, "backend-install-timeout"),
|
||||
Entry("backend upgrade timeout", config.FlagBackendUpgradeTimeout, "backend-upgrade-timeout"),
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
{
|
||||
"_comment": "Auto-generated from unsloth inference_defaults.json. DO NOT EDIT. Run go generate ./core/config/ to update.",
|
||||
"families": {
|
||||
"moss-tts-local-transformer-v1.5": {"min_p":0,"repeat_penalty":1,"temperature":1.7,"top_k":25,"top_p":0.8},
|
||||
"moss-tts-nano": {"min_p":0,"repeat_penalty":1,"temperature":1.7,"top_k":25,"top_p":0.8},
|
||||
"qwen3.8": {"min_p":0,"presence_penalty":1.5,"repeat_penalty":1,"temperature":0.7,"top_k":20,"top_p":0.8},
|
||||
"qwen3.6": {"min_p":0,"presence_penalty":1.5,"repeat_penalty":1,"temperature":0.7,"top_k":20,"top_p":0.8},
|
||||
"qwen3.5": {"min_p":0,"presence_penalty":1.5,"repeat_penalty":1,"temperature":0.7,"top_k":20,"top_p":0.8},
|
||||
@@ -60,5 +62,5 @@
|
||||
"grok": {"min_p":0.01,"repeat_penalty":1,"temperature":1,"top_k":-1,"top_p":0.95},
|
||||
"mimo": {"min_p":0.01,"repeat_penalty":1,"temperature":0.7,"top_k":-1,"top_p":0.95}
|
||||
},
|
||||
"patterns": ["qwen3.8","qwen3.6","qwen3.5","qwen3-coder","qwen3-next","qwen3-vl","qwen3","qwen2.5-coder","qwen2.5-vl","qwen2.5-omni","qwen2.5-math","qwen2.5","qwen2-vl","qwen2","qwq","gemma-4","gemma-3n","gemma-3","medgemma","gemma-2","muse-glimmer","llama-4","llama-3.3","llama-3.2","llama-3.1","llama-3","phi-4","phi-3","mistral-nemo","mistral-small","mistral-large","magistral","ministral","devstral","pixtral","deepseek-v4","deepseek-r1","deepseek-v3","deepseek-ocr","glm-5","glm-4","nemotron","minimax-m2.7","minimax-m2.5","minimax","gpt-oss","granite-4","kimi-k3","kimi-k2","kimi","lfm2","smollm","olmo","falcon","ernie","seed","grok","mimo"]
|
||||
"patterns": ["moss-tts-local-transformer-v1.5","moss-tts-nano","qwen3.8","qwen3.6","qwen3.5","qwen3-coder","qwen3-next","qwen3-vl","qwen3","qwen2.5-coder","qwen2.5-vl","qwen2.5-omni","qwen2.5-math","qwen2.5","qwen2-vl","qwen2","qwq","gemma-4","gemma-3n","gemma-3","medgemma","gemma-2","muse-glimmer","llama-4","llama-3.3","llama-3.2","llama-3.1","llama-3","phi-4","phi-3","mistral-nemo","mistral-small","mistral-large","magistral","ministral","devstral","pixtral","deepseek-v4","deepseek-r1","deepseek-v3","deepseek-ocr","glm-5","glm-4","nemotron","minimax-m2.7","minimax-m2.5","minimax","gpt-oss","granite-4","kimi-k3","kimi-k2","kimi","lfm2","smollm","olmo","falcon","ernie","seed","grok","mimo"]
|
||||
}
|
||||
@@ -43,8 +43,17 @@ type TTSConfig struct {
|
||||
|
||||
// @Description ModelConfig represents a model configuration
|
||||
type ModelConfig struct {
|
||||
modelConfigFile string `yaml:"-" json:"-"`
|
||||
modelTemplate string `yaml:"-" json:"-"`
|
||||
modelConfigFile string `yaml:"-" json:"-"`
|
||||
modelTemplate string `yaml:"-" json:"-"`
|
||||
// persistedConfigRevision is the revision of this model's persisted
|
||||
// configuration, stamped when the loader materializes it and therefore
|
||||
// before any per-request override is merged in. The request pipeline
|
||||
// mutates its copy of a ModelConfig with the caller's sampling parameters
|
||||
// (temperature, top_p, stop, ...), so hashing the config at load time is
|
||||
// the only way the controller sees one revision per configuration rather
|
||||
// than one per request body. Unexported, so it never enters the hash it
|
||||
// describes and never reaches YAML or JSON.
|
||||
persistedConfigRevision string `yaml:"-" json:"-"`
|
||||
schema.PredictionOptions `yaml:"parameters,omitempty" json:"parameters,omitempty"`
|
||||
Name string `yaml:"name,omitempty" json:"name,omitempty"`
|
||||
Artifacts []modelartifacts.Spec `yaml:"artifacts,omitempty" json:"artifacts,omitempty"`
|
||||
@@ -1360,6 +1369,12 @@ func (c *ModelConfig) syncKnownUsecasesFromString() {
|
||||
c.KnownUsecaseStrings = append(c.KnownUsecaseStrings, k)
|
||||
}
|
||||
}
|
||||
// GetAllModelConfigUsecases returns a map, and ranging one yields a random
|
||||
// order per call. KnownUsecaseStrings is part of the serialized config, so
|
||||
// an unsorted list gives the same file a different config revision on every
|
||||
// load. In distributed mode that reads as a config change and the router
|
||||
// rejects the request with ErrStaleModelConfigRevision.
|
||||
slices.Sort(c.KnownUsecaseStrings)
|
||||
}
|
||||
|
||||
func (c *ModelConfig) UnmarshalYAML(value *yaml.Node) error {
|
||||
@@ -1836,6 +1851,27 @@ func (c *ModelConfig) GetModelConfigFile() string {
|
||||
return c.modelConfigFile
|
||||
}
|
||||
|
||||
// PersistedConfigRevision returns the revision stamped when this configuration
|
||||
// was loaded, or "" when it was never stamped (a config synthesized outside the
|
||||
// loader). Callers that need a revision for a request must prefer this over
|
||||
// recomputing one from the config they hold: by then the request pipeline has
|
||||
// merged the caller's prediction parameters into it.
|
||||
func (c *ModelConfig) PersistedConfigRevision() string {
|
||||
return c.persistedConfigRevision
|
||||
}
|
||||
|
||||
// StampPersistedConfigRevision records the revision of this configuration as
|
||||
// persisted. It is computed from the receiver as-is, so callers must invoke it
|
||||
// only on a configuration that has not been merged with request overrides.
|
||||
func (c *ModelConfig) StampPersistedConfigRevision() error {
|
||||
revision, err := modelConfigRevision(c)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
c.persistedConfigRevision = revision
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetModelTemplate returns the model's chat template if available
|
||||
func (c *ModelConfig) GetModelTemplate() string {
|
||||
return c.modelTemplate
|
||||
|
||||
@@ -168,6 +168,14 @@ func readModelConfigsFromFile(file string, opts ...ConfigLoaderOption) ([]*Model
|
||||
if err := yaml.Unmarshal(f, &configs); err == nil && len(configs) > 0 {
|
||||
for _, cc := range configs {
|
||||
cc.modelConfigFile = file
|
||||
// Stamp before SetDefaults: the revision describes what is on disk.
|
||||
// SetDefaults folds in the GGUF guess, hardware defaults and
|
||||
// app-level options, none of which are persisted configuration, and
|
||||
// the GGUF guess in particular depends on whether the model file
|
||||
// parses at that moment.
|
||||
if err := cc.StampPersistedConfigRevision(); err != nil {
|
||||
return nil, fmt.Errorf("stamping config revision for %q: %w", cc.Name, err)
|
||||
}
|
||||
cc.SetDefaults(opts...)
|
||||
cc.syncKnownUsecasesFromString()
|
||||
}
|
||||
@@ -182,6 +190,9 @@ func readModelConfigsFromFile(file string, opts ...ConfigLoaderOption) ([]*Model
|
||||
|
||||
c.modelConfigFile = file
|
||||
c.syncKnownUsecasesFromString()
|
||||
if err := c.StampPersistedConfigRevision(); err != nil {
|
||||
return nil, fmt.Errorf("stamping config revision for %q: %w", c.Name, err)
|
||||
}
|
||||
c.SetDefaults(opts...)
|
||||
|
||||
return []*ModelConfig{c}, nil
|
||||
@@ -218,6 +229,16 @@ func (bcl *ModelConfigLoader) LoadModelConfigFileByName(modelName, modelPath str
|
||||
}
|
||||
}
|
||||
|
||||
// Stamp before SetDefaults, and only when this config did not come from
|
||||
// disk already carrying one (a name with no config file on disk is
|
||||
// synthesized above). Re-stamping a loaded config here would hash it after
|
||||
// SetDefaults and reintroduce the dependency on the GGUF guess.
|
||||
if cfg.PersistedConfigRevision() == "" {
|
||||
if err := cfg.StampPersistedConfigRevision(); err != nil {
|
||||
return nil, fmt.Errorf("stamping config revision for %q: %w", modelName, err)
|
||||
}
|
||||
}
|
||||
|
||||
cfg.SetDefaults(append(opts, ModelPath(modelPath))...)
|
||||
|
||||
return cfg, nil
|
||||
@@ -420,6 +441,26 @@ func (bcl *ModelConfigLoader) ResolveAlias(cfg *ModelConfig) (*ModelConfig, bool
|
||||
return &target, true, nil
|
||||
}
|
||||
|
||||
// ResolveAliasName maps a model name to the name of the model that actually
|
||||
// serves it: an alias resolves to its target, anything else resolves to
|
||||
// itself. The second return reports whether name was an alias.
|
||||
//
|
||||
// Unlike ResolveAlias this never errors. A name with no config (a rule may be
|
||||
// authored before the model is installed), a dangling alias, and a chained
|
||||
// alias all resolve to themselves, so callers keep a usable name that simply
|
||||
// has no model behind it rather than silently governing a different model.
|
||||
func (bcl *ModelConfigLoader) ResolveAliasName(name string) (string, bool) {
|
||||
cfg, exists := bcl.GetModelConfig(name)
|
||||
if !exists || !cfg.IsAlias() {
|
||||
return name, false
|
||||
}
|
||||
target, exists := bcl.GetModelConfig(cfg.Alias)
|
||||
if !exists || target.IsAlias() {
|
||||
return name, true
|
||||
}
|
||||
return target.Name, true
|
||||
}
|
||||
|
||||
// ValidateAliasTarget checks an alias config's target at create/swap time:
|
||||
// the target must exist, must not be an alias, and must not be disabled.
|
||||
// Returns nil for non-alias configs.
|
||||
@@ -944,3 +985,38 @@ func hasAnyMappingKey(mapping *yaml.Node, keys ...string) bool {
|
||||
func nonemptyScalar(node *yaml.Node) bool {
|
||||
return node != nil && node.Kind == yaml.ScalarNode && node.Tag == "!!str" && strings.TrimSpace(node.Value) != ""
|
||||
}
|
||||
|
||||
// RevisionFor returns the config revision for modelName: the one an inference
|
||||
// request for that model will carry.
|
||||
//
|
||||
// This is the only way to obtain a revision outside this package. Every
|
||||
// publisher must use it, so that what is published and what is checked are
|
||||
// the same value by construction rather than by two implementations happening
|
||||
// to agree. Hashing a ModelConfig directly is not available to callers, because
|
||||
// a config that has been through SetDefaults or the request middleware hashes
|
||||
// to something no request will ever present.
|
||||
func (bcl *ModelConfigLoader) RevisionFor(modelName string, appConfig *ApplicationConfig) (string, error) {
|
||||
cfg, err := bcl.LoadModelConfigFileByNameDefaultOptions(modelName, appConfig)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("resolving config revision for %q: %w", modelName, err)
|
||||
}
|
||||
return stampedRevision(cfg, modelName)
|
||||
}
|
||||
|
||||
// RevisionForPath is RevisionFor for callers that hold loader options and a
|
||||
// models path rather than an ApplicationConfig.
|
||||
func (bcl *ModelConfigLoader) RevisionForPath(modelName, modelPath string, opts ...ConfigLoaderOption) (string, error) {
|
||||
cfg, err := bcl.LoadModelConfigFileByName(modelName, modelPath, opts...)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("resolving config revision for %q: %w", modelName, err)
|
||||
}
|
||||
return stampedRevision(cfg, modelName)
|
||||
}
|
||||
|
||||
func stampedRevision(cfg *ModelConfig, modelName string) (string, error) {
|
||||
revision := cfg.PersistedConfigRevision()
|
||||
if revision == "" {
|
||||
return "", fmt.Errorf("no config revision stamped for %q", modelName)
|
||||
}
|
||||
return revision, nil
|
||||
}
|
||||
@@ -314,3 +314,57 @@ var _ = Describe("ModelConfigLoader alias resolution", func() {
|
||||
Expect(loader.ValidateAliasTarget(&bad)).To(MatchError(ContainSubstring("itself an alias")))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("ModelConfigLoader ResolveAliasName", func() {
|
||||
var loader *ModelConfigLoader
|
||||
|
||||
BeforeEach(func() {
|
||||
loader = NewModelConfigLoader("")
|
||||
loader.configs["real"] = ModelConfig{Name: "real", Backend: "llama-cpp"}
|
||||
loader.configs["production"] = ModelConfig{Name: "production", Alias: "real"}
|
||||
loader.configs["chain"] = ModelConfig{Name: "chain", Alias: "production"}
|
||||
loader.configs["dangling"] = ModelConfig{Name: "dangling", Alias: "nope"}
|
||||
})
|
||||
|
||||
It("maps an alias name to the model that actually serves it", func() {
|
||||
target, isAlias := loader.ResolveAliasName("production")
|
||||
Expect(isAlias).To(BeTrue())
|
||||
Expect(target).To(Equal("real"))
|
||||
})
|
||||
|
||||
It("maps a real model name to itself", func() {
|
||||
target, isAlias := loader.ResolveAliasName("real")
|
||||
Expect(isAlias).To(BeFalse())
|
||||
Expect(target).To(Equal("real"))
|
||||
})
|
||||
|
||||
// A rule may be authored for a model that is not installed yet (pre-staging
|
||||
// placement before standing up a node), so an unknown name must resolve to
|
||||
// itself rather than to the empty string.
|
||||
It("maps an unknown name to itself", func() {
|
||||
target, isAlias := loader.ResolveAliasName("not-installed-yet")
|
||||
Expect(isAlias).To(BeFalse())
|
||||
Expect(target).To(Equal("not-installed-yet"))
|
||||
})
|
||||
|
||||
// A broken alias has no model behind it. Resolving to itself keeps the
|
||||
// caller on a name that simply has no replicas, instead of silently
|
||||
// governing some other model.
|
||||
It("maps a dangling alias to itself", func() {
|
||||
target, isAlias := loader.ResolveAliasName("dangling")
|
||||
Expect(isAlias).To(BeTrue())
|
||||
Expect(target).To(Equal("dangling"))
|
||||
})
|
||||
|
||||
It("maps a chained alias to itself rather than following the chain", func() {
|
||||
target, isAlias := loader.ResolveAliasName("chain")
|
||||
Expect(isAlias).To(BeTrue())
|
||||
Expect(target).To(Equal("chain"))
|
||||
})
|
||||
|
||||
It("maps the empty name to itself", func() {
|
||||
target, isAlias := loader.ResolveAliasName("")
|
||||
Expect(isAlias).To(BeFalse())
|
||||
Expect(target).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
@@ -10,10 +10,17 @@ import (
|
||||
"google.golang.org/protobuf/proto"
|
||||
)
|
||||
|
||||
// ModelConfigRevision returns a stable revision of the persisted semantic
|
||||
// modelConfigRevision returns a stable revision of the persisted semantic
|
||||
// configuration. ModelConfig's JSON tags exclude runtime-derived state and
|
||||
// source bookkeeping, while encoding/json orders map keys deterministically.
|
||||
func ModelConfigRevision(cfg *ModelConfig) (string, error) {
|
||||
//
|
||||
// Deliberately unexported. It must only ever be called on a configuration as
|
||||
// parsed from disk, before SetDefaults folds in the GGUF guess, the hardware
|
||||
// defaults and app-level options. Callers outside this package cannot tell
|
||||
// which they hold, and every time one hashed a defaulted or request-merged
|
||||
// config it published a revision no inference request would carry, which makes
|
||||
// the model unroutable. Use ModelConfigLoader.RevisionFor instead.
|
||||
func modelConfigRevision(cfg *ModelConfig) (string, error) {
|
||||
if cfg == nil {
|
||||
return "", errors.New("model config is nil")
|
||||
}
|
||||
|
||||
@@ -0,0 +1,146 @@
|
||||
package config_test
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/pkg/system"
|
||||
)
|
||||
|
||||
// The distributed controller pins a model's replicas to its config revision and
|
||||
// rejects any request carrying a different one. A revision that is not stable
|
||||
// for one unchanged file on disk therefore wedges the model.
|
||||
var _ = Describe("Model config revision stability", func() {
|
||||
// A chat model with an mmproj derives two usecase flags, FLAG_CHAT and
|
||||
// FLAG_VISION. syncKnownUsecasesFromString builds that list by ranging a
|
||||
// map, so an unstable order shows up with two or more flags and stays
|
||||
// hidden with one.
|
||||
const multiUsecaseModel = `backend: llama-cpp
|
||||
context_size: 50000
|
||||
known_usecases:
|
||||
- chat
|
||||
mmproj: llama-cpp/mmproj/example/mmproj.gguf
|
||||
name: example
|
||||
options:
|
||||
- use_jinja:true
|
||||
- parallel:2
|
||||
parameters:
|
||||
model: llama-cpp/models/example/example.gguf
|
||||
template:
|
||||
use_tokenizer_template: true
|
||||
`
|
||||
|
||||
var (
|
||||
dir string
|
||||
appConfig *config.ApplicationConfig
|
||||
)
|
||||
|
||||
BeforeEach(func() {
|
||||
dir = GinkgoT().TempDir()
|
||||
Expect(os.WriteFile(filepath.Join(dir, "example.yaml"), []byte(multiUsecaseModel), 0o600)).To(Succeed())
|
||||
appConfig = config.NewApplicationConfig()
|
||||
appConfig.SystemState = &system.SystemState{Model: system.Model{ModelsPath: dir}}
|
||||
})
|
||||
|
||||
loadRevision := func() string {
|
||||
loader := config.NewModelConfigLoader(dir)
|
||||
Expect(loader.LoadModelConfigsFromPath(dir, appConfig.ToConfigLoaderOptions()...)).To(Succeed())
|
||||
cfg, ok := loader.GetModelConfig("example")
|
||||
Expect(ok).To(BeTrue())
|
||||
Expect(cfg.PersistedConfigRevision()).ToNot(BeEmpty())
|
||||
return cfg.PersistedConfigRevision()
|
||||
}
|
||||
|
||||
It("does not change when the same file is loaded repeatedly", func() {
|
||||
baseline := loadRevision()
|
||||
for i := 0; i < 20; i++ {
|
||||
Expect(loadRevision()).To(Equal(baseline), "revision changed between two loads of one unchanged file")
|
||||
}
|
||||
})
|
||||
|
||||
It("orders the derived usecases deterministically", func() {
|
||||
loader := config.NewModelConfigLoader(dir)
|
||||
Expect(loader.LoadModelConfigsFromPath(dir, appConfig.ToConfigLoaderOptions()...)).To(Succeed())
|
||||
cfg, ok := loader.GetModelConfig("example")
|
||||
Expect(ok).To(BeTrue())
|
||||
Expect(len(cfg.KnownUsecaseStrings)).To(BeNumerically(">=", 2), "fixture must derive several usecases to expose ordering")
|
||||
Expect(cfg.KnownUsecaseStrings).To(Equal([]string{"FLAG_CHAT", "FLAG_VISION"}))
|
||||
})
|
||||
|
||||
// The request pipeline reloads the config through LoadModelConfigFileByName,
|
||||
// which applies SetDefaults a second time. The stamp is taken before those
|
||||
// defaults, so both the stored config and the one a request resolves carry
|
||||
// the same revision.
|
||||
It("survives the extra SetDefaults the request path applies", func() {
|
||||
loader := config.NewModelConfigLoader(dir)
|
||||
Expect(loader.LoadModelConfigsFromPath(dir, appConfig.ToConfigLoaderOptions()...)).To(Succeed())
|
||||
stored, ok := loader.GetModelConfig("example")
|
||||
Expect(ok).To(BeTrue())
|
||||
Expect(stored.PersistedConfigRevision()).ToNot(BeEmpty())
|
||||
|
||||
requestCfg, err := loader.LoadModelConfigFileByNameDefaultOptions("example", appConfig)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(requestCfg.PersistedConfigRevision()).To(Equal(stored.PersistedConfigRevision()))
|
||||
})
|
||||
})
|
||||
|
||||
// The revision must describe the configuration as persisted, and nothing else.
|
||||
// SetDefaults folds in values that are not persisted config: the GGUF guess
|
||||
// (which reads the model file and can fail on slow or remote storage), the
|
||||
// hardware defaults, and app-level options like threads. Hashing after that
|
||||
// made the revision a function of whether a multi-gigabyte file happened to
|
||||
// parse, so one unchanged YAML produced two different revisions depending on
|
||||
// the moment, and the controller rejected every request carrying the other one.
|
||||
var _ = Describe("Model config revision independence from runtime defaults", func() {
|
||||
It("does not change when SetDefaults is applied", func() {
|
||||
dir := GinkgoT().TempDir()
|
||||
body := "backend: llama-cpp\ncontext_size: 50000\nknown_usecases:\n - chat\n" +
|
||||
"mmproj: llama-cpp/mmproj/example/mmproj.gguf\nname: example\n" +
|
||||
"parameters:\n model: llama-cpp/models/example/example.gguf\n"
|
||||
Expect(os.WriteFile(filepath.Join(dir, "example.yaml"), []byte(body), 0o600)).To(Succeed())
|
||||
|
||||
appConfig := config.NewApplicationConfig()
|
||||
appConfig.SystemState = &system.SystemState{Model: system.Model{ModelsPath: dir}}
|
||||
loader := config.NewModelConfigLoader(dir)
|
||||
Expect(loader.LoadModelConfigsFromPath(dir, appConfig.ToConfigLoaderOptions()...)).To(Succeed())
|
||||
|
||||
stored, ok := loader.GetModelConfig("example")
|
||||
Expect(ok).To(BeTrue())
|
||||
before := stored.PersistedConfigRevision()
|
||||
Expect(before).ToNot(BeEmpty())
|
||||
|
||||
// Applying defaults again is what the request path does.
|
||||
stored.SetDefaults(appConfig.ToConfigLoaderOptions()...)
|
||||
Expect(stored.PersistedConfigRevision()).To(Equal(before))
|
||||
|
||||
resolved, err := loader.LoadModelConfigFileByNameDefaultOptions("example", appConfig)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(resolved.PersistedConfigRevision()).To(Equal(before),
|
||||
"the request path must carry the same revision as the stored config")
|
||||
})
|
||||
|
||||
It("does not change when app-level defaults differ", func() {
|
||||
dir := GinkgoT().TempDir()
|
||||
Expect(os.WriteFile(filepath.Join(dir, "example.yaml"),
|
||||
[]byte("name: example\nbackend: llama-cpp\nparameters:\n model: m.gguf\n"), 0o600)).To(Succeed())
|
||||
|
||||
revWith := func(threads int, f16 bool) string {
|
||||
appConfig := config.NewApplicationConfig()
|
||||
appConfig.SystemState = &system.SystemState{Model: system.Model{ModelsPath: dir}}
|
||||
appConfig.Threads = threads
|
||||
appConfig.F16 = f16
|
||||
loader := config.NewModelConfigLoader(dir)
|
||||
Expect(loader.LoadModelConfigsFromPath(dir, appConfig.ToConfigLoaderOptions()...)).To(Succeed())
|
||||
cfg, err := loader.LoadModelConfigFileByNameDefaultOptions("example", appConfig)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
return cfg.PersistedConfigRevision()
|
||||
}
|
||||
|
||||
Expect(revWith(8, false)).To(Equal(revWith(1, true)),
|
||||
"an operator changing threads must not make every model unroutable")
|
||||
})
|
||||
})
|
||||
@@ -19,10 +19,11 @@ var _ = Describe("Model configuration revisions", func() {
|
||||
return cfg
|
||||
}
|
||||
|
||||
// The raw hash is unexported on purpose, so these specs exercise it the way
|
||||
// every caller now must: by stamping the parsed config.
|
||||
revision := func(cfg *config.ModelConfig) string {
|
||||
value, err := config.ModelConfigRevision(cfg)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
return value
|
||||
Expect(cfg.StampPersistedConfigRevision()).To(Succeed())
|
||||
return cfg.PersistedConfigRevision()
|
||||
}
|
||||
|
||||
It("is stable across equivalent YAML formatting and map order", func() {
|
||||
|
||||
@@ -59,6 +59,7 @@ type RuntimeSettings struct {
|
||||
BackendGalleries *[]Gallery `json:"backend_galleries,omitempty"`
|
||||
AutoloadGalleries *bool `json:"autoload_galleries,omitempty"`
|
||||
AutoloadBackendGalleries *bool `json:"autoload_backend_galleries,omitempty"`
|
||||
VRAMPersistentCache *bool `json:"vram_persistent_cache,omitempty"`
|
||||
|
||||
// API keys - No omitempty as we need to save empty arrays to clear keys
|
||||
ApiKeys *[]string `json:"api_keys"`
|
||||
|
||||
@@ -328,6 +328,10 @@ var runtimeSettingsFields = []fieldSpec{
|
||||
func(s *RuntimeSettings) **bool { return &s.AutoloadBackendGalleries },
|
||||
func(o *ApplicationConfig) bool { return o.AutoloadBackendGalleries },
|
||||
func(o *ApplicationConfig, v bool) { o.AutoloadBackendGalleries = v }),
|
||||
field("vram_persistent_cache",
|
||||
func(s *RuntimeSettings) **bool { return &s.VRAMPersistentCache },
|
||||
func(o *ApplicationConfig) bool { return o.VRAMPersistentCache },
|
||||
func(o *ApplicationConfig, v bool) { o.VRAMPersistentCache = v }),
|
||||
|
||||
// API keys: echoed for the UI, but the apply loops never touch them.
|
||||
// The settings endpoint and the file watcher own the env+runtime merge
|
||||
|
||||
@@ -45,6 +45,7 @@ func DefaultRuntimeBaseline() *ApplicationConfig {
|
||||
o.BackendGalleries = mustGalleries(DefaultBackendGalleriesJSON)
|
||||
o.AutoloadGalleries = true
|
||||
o.AutoloadBackendGalleries = true
|
||||
o.VRAMPersistentCache = true
|
||||
// core/cli/run.go injects WithMemoryReclaimer(enabled, threshold)
|
||||
// unconditionally, so the kong threshold default (0.95) reaches the
|
||||
// config even when the reclaimer flag is off - this overlay must match
|
||||
|
||||
@@ -0,0 +1,94 @@
|
||||
package gallery_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
|
||||
"github.com/mudler/LocalAI/core/gallery"
|
||||
"github.com/mudler/LocalAI/pkg/system"
|
||||
)
|
||||
|
||||
// On a distributed controller the GPUs live on the workers, so a variant
|
||||
// picker sized against the controller tells admins a cluster of A100s can only
|
||||
// run the smallest CPU build.
|
||||
var _ = Describe("ClusterResolveEnv", func() {
|
||||
gib := func(n uint64) uint64 { return n * 1024 * 1024 * 1024 }
|
||||
|
||||
// The controller as Argus actually runs it: no GPU at all.
|
||||
var controller *system.SystemState
|
||||
|
||||
BeforeEach(func() {
|
||||
controller = system.NewCapabilityState("default")
|
||||
})
|
||||
|
||||
It("sizes models against the cluster reading rather than the controller", func() {
|
||||
env := gallery.ClusterResolveEnv(context.Background(), controller, gib(80), []string{"nvidia-cuda-13"})
|
||||
|
||||
Expect(env.AvailableMemory).To(Equal(gib(80)))
|
||||
})
|
||||
|
||||
It("accepts a CUDA backend that only the workers can run", func() {
|
||||
env := gallery.ClusterResolveEnv(context.Background(), controller, gib(80), []string{"nvidia-cuda-13"})
|
||||
|
||||
Expect(env.BackendCompatible).ToNot(BeNil())
|
||||
// A name carrying the cuda token is what the controller rejects today;
|
||||
// a bare engine name like "vllm" passes on any host and would prove
|
||||
// nothing about the union.
|
||||
Expect(env.BackendCompatible("cuda-13-vllm")).To(BeTrue())
|
||||
Expect(env.BackendCompatible("llama-cpp")).To(BeTrue())
|
||||
})
|
||||
|
||||
// The union must stay a filter, not an open door: a Linux NVIDIA fleet
|
||||
// still cannot run an Apple-only build.
|
||||
It("still rejects a backend no node in the cluster can run", func() {
|
||||
env := gallery.ClusterResolveEnv(context.Background(), controller, gib(80), []string{"nvidia-cuda-13"})
|
||||
|
||||
Expect(env.BackendCompatible("mlx")).To(BeFalse())
|
||||
})
|
||||
|
||||
It("accepts a backend that any one node in a mixed fleet can run", func() {
|
||||
env := gallery.ClusterResolveEnv(context.Background(), controller, gib(80), []string{"nvidia-cuda-13", "metal"})
|
||||
|
||||
Expect(env.BackendCompatible("mlx")).To(BeTrue())
|
||||
Expect(env.BackendCompatible("cuda-13-vllm")).To(BeTrue())
|
||||
})
|
||||
|
||||
// Ranking has to follow the hardware too, or a cluster of NVIDIA workers
|
||||
// gets offered the GGUF build over the vLLM one it should prefer.
|
||||
It("ranks engines by the workers' hardware, not the controller's", func() {
|
||||
env := gallery.ClusterResolveEnv(context.Background(), controller, gib(80), []string{"nvidia-cuda-13"})
|
||||
|
||||
Expect(env.EnginePreference).To(Equal(system.NewCapabilityState("nvidia-cuda-13").EnginePreferenceTokens()))
|
||||
})
|
||||
|
||||
// Every degradation path lands here, so it must be indistinguishable from
|
||||
// the single-node behavior that shipped before any of this existed.
|
||||
It("falls back to the host description when the cluster reports nothing", func() {
|
||||
host := gallery.HostResolveEnv(context.Background(), controller)
|
||||
env := gallery.ClusterResolveEnv(context.Background(), controller, 0, nil)
|
||||
|
||||
Expect(env.AvailableMemory).To(Equal(host.AvailableMemory))
|
||||
Expect(env.EnginePreference).To(Equal(host.EnginePreference))
|
||||
Expect(env.BackendCompatible("cuda-13-vllm")).To(Equal(host.BackendCompatible("cuda-13-vllm")))
|
||||
Expect(env.BackendCompatible("mlx")).To(Equal(host.BackendCompatible("mlx")))
|
||||
})
|
||||
|
||||
// A cluster that reports capabilities but no usable memory reading should
|
||||
// still gain the hardware view; only the size question falls back.
|
||||
It("keeps the host memory when only the memory reading is missing", func() {
|
||||
host := gallery.HostResolveEnv(context.Background(), controller)
|
||||
env := gallery.ClusterResolveEnv(context.Background(), controller, 0, []string{"nvidia-cuda-13"})
|
||||
|
||||
Expect(env.AvailableMemory).To(Equal(host.AvailableMemory))
|
||||
Expect(env.BackendCompatible("cuda-13-vllm")).To(BeTrue())
|
||||
})
|
||||
|
||||
It("keeps the probe wired so variant sizes are still measured", func() {
|
||||
env := gallery.ClusterResolveEnv(context.Background(), controller, gib(80), []string{"nvidia-cuda-13"})
|
||||
|
||||
Expect(env.ProbeMemory).ToNot(BeNil())
|
||||
Expect(env.ServingFeaturePreference).To(Equal(system.ServingFeaturePreferenceTokens()))
|
||||
})
|
||||
})
|
||||
@@ -236,8 +236,15 @@ var _ = Describe("InstallModelFromGallery with an empty base config", func() {
|
||||
Expect(install(e.Name, gallery.GalleryModel{})).To(Succeed())
|
||||
cfg := installedConfig(e.Name)
|
||||
Expect(cfg["name"]).To(Equal(e.Name))
|
||||
// The catalog's own overrides, verbatim, laid over the empty base.
|
||||
Expect(cfg["parameters"]).To(Equal(e.Overrides["parameters"]))
|
||||
// The catalog's own overrides, laid over the empty base. parameters is
|
||||
// checked key by key rather than as a whole map: the install also merges
|
||||
// the model family's inference defaults into it, and what matters here is
|
||||
// that the authored keys survive that.
|
||||
authored, ok := e.Overrides["parameters"].(map[string]any)
|
||||
Expect(ok).To(BeTrue())
|
||||
for key, want := range authored {
|
||||
Expect(cfg["parameters"]).To(HaveKeyWithValue(key, want))
|
||||
}
|
||||
Expect(cfg["known_usecases"]).To(Equal(e.Overrides["known_usecases"]))
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,101 @@
|
||||
package gallery_test
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/gallery"
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
func remarshal(value any, target any) error {
|
||||
data, err := yaml.Marshal(value)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return yaml.Unmarshal(data, target)
|
||||
}
|
||||
|
||||
var _ = Describe("EXL3 gallery entries", func() {
|
||||
It("pins the four full repositories and configures the Qwen DFlash companion", func() {
|
||||
entries, err := gallery.ReadConfigFile[[]gallery.GalleryModel](filepath.Join("..", "..", "gallery", "index.yaml"))
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
byName := make(map[string]gallery.GalleryModel, len(*entries))
|
||||
for _, entry := range *entries {
|
||||
byName[entry.Name] = entry
|
||||
}
|
||||
|
||||
expected := map[string]struct {
|
||||
repo string
|
||||
revision string
|
||||
}{
|
||||
"qwen3.8-27b-exl3-vllm-cpp": {
|
||||
repo: "Mia-AiLab/Qwen3.8-27B-EXL3-3.5bpw", revision: "19441ac874c4018295da848e250f23511361cda4",
|
||||
},
|
||||
"qwen3.8-27b-dflash2-exl3-vllm-cpp": {
|
||||
repo: "Mia-AiLab/Qwen3.8-27B-EXL3-3.5bpw", revision: "19441ac874c4018295da848e250f23511361cda4",
|
||||
},
|
||||
"deepseek-v4-flash-spark-exl3-vllm-cpp": {
|
||||
repo: "0xSero/deepseek-v4-flash-0731-spark", revision: "ce5ff0f1efb2e184aafc759d281bfae47d3a359c",
|
||||
},
|
||||
"deepseek-v4-flash-exl3-3bpw-vllm-cpp": {
|
||||
repo: "0xSero/DeepSeek-V4-Flash-0731-EXL3-3.0bpw", revision: "e0bf84ac76a5100e8790c22ad10b70b1e2d06d71",
|
||||
},
|
||||
}
|
||||
|
||||
for name, want := range expected {
|
||||
entry, found := byName[name]
|
||||
Expect(found).To(BeTrue(), "missing gallery entry %q", name)
|
||||
Expect(entry.Tags).To(ContainElements("vllm-cpp", "exl3", "gpu", "cuda"), name)
|
||||
Expect(entry.Overrides).To(HaveKeyWithValue("backend", "vllm-cpp"), name)
|
||||
cfg := config.ModelConfig{}
|
||||
Expect(remarshal(entry.Overrides, &cfg)).To(Succeed(), name)
|
||||
Expect(cfg.Artifacts).ToNot(BeEmpty(), name)
|
||||
Expect(cfg.Artifacts[0].Source.Repo).To(Equal(want.repo), name)
|
||||
Expect(cfg.Artifacts[0].Source.Revision).To(Equal(want.revision), name)
|
||||
}
|
||||
|
||||
plain := byName["qwen3.8-27b-exl3-vllm-cpp"]
|
||||
Expect(plain.Tags).ToNot(ContainElement("dflash"))
|
||||
dflashTags := 0
|
||||
for name := range expected {
|
||||
if contains(byName[name].Tags, "dflash") {
|
||||
dflashTags++
|
||||
}
|
||||
}
|
||||
Expect(dflashTags).To(Equal(1))
|
||||
|
||||
dflash := byName["qwen3.8-27b-dflash2-exl3-vllm-cpp"]
|
||||
Expect(dflash.Tags).To(ContainElement("dflash"))
|
||||
Expect(dflash.Variants).To(ConsistOf(gallery.Variant{Model: "qwen3.8-27b-exl3-vllm-cpp"}))
|
||||
cfg := config.ModelConfig{}
|
||||
Expect(remarshal(dflash.Overrides, &cfg)).To(Succeed())
|
||||
Expect(cfg.ContextSize).To(HaveValue(Equal(8192)))
|
||||
Expect(cfg.EngineArgs).To(HaveKeyWithValue("num_blocks", 2048))
|
||||
Expect(cfg.EngineArgs).To(HaveKeyWithValue("max_num_seqs", 8))
|
||||
Expect(cfg.EngineArgs).To(HaveKeyWithValue("max_num_batched_tokens", 16384))
|
||||
Expect(cfg.EngineArgs).To(HaveKeyWithValue("enable_prefix_caching", false))
|
||||
spec, ok := cfg.EngineArgs["speculative_config"].(map[string]any)
|
||||
Expect(ok).To(BeTrue())
|
||||
Expect(spec).To(HaveKeyWithValue("method", "dflash"))
|
||||
Expect(spec).To(HaveKeyWithValue("num_speculative_tokens", 7))
|
||||
Expect(cfg.Artifacts).To(HaveLen(2))
|
||||
Expect(cfg.Artifacts[1].Name).To(Equal("draft_model"))
|
||||
Expect(cfg.Artifacts[1].Target).To(Equal("companion"))
|
||||
Expect(cfg.Artifacts[1].Source.Repo).To(Equal("Mia-AiLab/Qwen3.8-27B-DFlash2-EXL3-5.0bpw"))
|
||||
Expect(cfg.Artifacts[1].Source.Revision).To(Equal("4f0436269bca761b071f05319e8e04a87cc633f9"))
|
||||
})
|
||||
})
|
||||
|
||||
func contains(values []string, target string) bool {
|
||||
for _, value := range values {
|
||||
if value == target {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
Loaded 100 of 273 files, more files were not shown because too many files have changed in this diff.
Show more
Reference in new issue
Block a user