mirror of
https://github.com/mudler/LocalAI.git
synced 2026-10-04 20:14:43 -04:00
feat(router): route with native decision models (#12449)
* fix(schema): preserve SystemOne image inputs Assisted-by: OpenAI * test(schema): follow Ginkgo conventions for decision inputs Assisted-by: OpenAI * feat(llama-cpp): dispatch native decisions through Score Upgrade the stock dependency and reconcile Score/TTS patches. Reuse native decision parsing, tasks, formatting and response-reader cleanup; preserve ordinary scoring admission and guard older dependencies. Assisted-by: OpenAI * refactor(systemone): share request and model validation Assisted-by: OpenAI:gpt-5 * fix(systemone): preserve HTTP wire-byte validation limit Keep structural validation separate from the serialized internal request bound so HTML escaping cannot reject valid HTTP payloads. Assisted-by: OpenAI:gpt-5 * feat(systemone): bound images and account native decisions Preserve public wire limits independently from router serialization. Reject unsupported NER images, map native request/capability errors, and stamp explicit usage once. Advertise decisions for stock llama-cpp. Assisted-by: OpenAI * fix(systemone): record usage on registered native route Exercise real registration and billing with a mock native backend. Reject empty native responses, malformed image URLs, trailing JSON, and wire overflow including whitespace. Assisted-by: OpenAI * feat(router): add lazy native decision transport Bind named models through internal ModelSystemOne calls with shared validation and bounded abandoned operations. Remove request and echoed-error contents from decision traces. Assisted-by: OpenAI:gpt-5 * feat(router): classify overlapping policies with native decisions Ask independent noul questions, validate probabilities and preserve first-superset routing. Wire the central factory with config-sensitive invalidation and cancellation-safe resolution. Document native framing and bounded operation limits. Assisted-by: OpenAI:gpt-5 * feat(gallery): add pinned Julia-1 native decision model Add a separate text-only llama-cpp Q8 entry with pinned Apache-2.0 source provenance and checksum. Installed using the gallery installer and exercised choice, score and noul on CPU. Assisted-by: OpenAI * test(router): verify native decisions through central factory Add an opt-in real-model Ginkgo integration covering the native Go loader and C++ transport, token usage, independent overlapping labels, and candidate selection. Document owned-server execution and the intentionally non-quality threshold. Assisted-by: Codex:gpt-5 * fix(llama-cpp): align upstream pin and preserve decision signatures Advance to bed0a856 without losing the automated upstream bump. Detect full-request fill_task support at compile time and forward every question for Nimble framing while retaining the earlier native signature. Preserve reconciled SCORE/TTS patches; add standalone compatibility coverage. Assisted-by: Codex:gpt-5 * feat(gallery): add native decision family defaults Pin Laya, Kev-4B, lev, OpenJev and Nimble artifacts. Verify Laya/Kev/lev gallery installs and CPU contracts on both native pins; clearly mark OpenJev/Nimble runtime validation pending and their noncommercial licenses. Assisted-by: OpenAI * docs(decisions): clarify integrated Nimble prerequisite Record the exact combined backend pin while retaining pending OpenJev and Nimble installation/runtime validation status. Assisted-by: Codex:gpt-5 * fix(gallery): indent native decision model sequences Match repository yamllint indentation for Laya, Kev, lev and OpenJev list fields. Parsed gallery data is unchanged; reproduce CI gallery lint failure before the whitespace-only fix and pass the same command afterward. Assisted-by: Codex:gpt-5 * docs(decisions): record OpenJev and Nimble CPU validation Record gallery installation, checksum/metadata verification and multiquestion native smoke results on bed0a856. Retain noncommercial and text-only limitations without accuracy or deterministic-output claims. Assisted-by: OpenAI * fix(ui): expose native Decisions router classifiers Select classifier models using metadata-driven capability routing, retain tuned thresholds, and validate native decision selections before saving. Cover both native backends and create/save/reopen in the real React editor. Assisted-by: Codex:gpt-5 * fix(router): exclude aliases from native decision discovery Check the originally named config before advertising native Decisions eligibility. Retain target capability inheritance for ordinary generation aliases. Exercise the actual capabilities endpoint with native models on both backends, aliases, and disabled models. Assisted-by: Codex:gpt-5 * feat(systemone): share bounded multimodal input validation Preserve text wire limits while admitting bounded PNG/JPEG decision input. Share collection and header validation across internal and public callers and keep the native runner response budget independent. Assisted-by: OpenAI:API-assistant * fix(systemone): bound admission lifetimes and validate complete images Retain shared admission leases through actual work completion, including abandoned internal operations. Decode bounded image pixels, cap public native responses before usage stamping, and preserve oversized malformed text status precedence. Assisted-by: OpenAI:API-assistant * fix(router): classify images before media fetching Preserve ordered structured probes for native decisions. Defer OpenAI media preparation until routing selects the served model, so rejected decision URLs cannot trigger downloads before shared validation. Guard direct image collection with context-aware shared admission. Keep text classifiers and embedding caches from discarding image input. Retain fail-closed classifier configuration and runtime fallback policy. Add middleware, typed-content, admission, cancellation and cache tests. Assisted-by: OpenAI:API-assistant * fix(router): bound extraction before serialization Check probe budgets before copying text or marshaling message state. Count JSON escaping so oversized internal inputs fail before allocation. Preserve typed Anthropic blocks through selected-model conversion and fallback. Keep retry coverage in Ginkgo without global test registration. Assisted-by: OpenAI * fix(router): bound supported probe serialization Arbitrary structs can bypass the probe budget through pointer marshalers, string tags, and promoted fields. Accept concrete chat schema types and plain JSON values instead of emulating arbitrary struct serialization. Budget escaped direct prompts before marshaling so raw length cannot hide serialized expansion. Preserve runtime fallback and reject oversized input before invoking the decision runner. Add Ginkgo allocation, boundary, and marshaler invocation regressions. Six-package tests, three-package race tests, and full-T2 delta lint pass. Assisted-by: OpenAI:GPT-5 golangci-lint * feat(decisions): enable bounded OpenJev images Validate native decision images before permissive media parsing and pixel allocation. Require both decision image support and a vision projector; missing or audio-only projectors cannot silently become text decisions. Pin the OpenJev Q8 projector and document its license and disk footprint. Add native safety tests, canonical limit parity, gallery and load-option checks, and a reproducible CPU direct-RPC contrasting-image smoke. Assisted-by: OpenAI:GPT-5 * fix(decisions): reject incomplete image streams stb accepts corrupt PNG Adler checksums and truncated JPEG scans. Use bounded zlib validation and strict libjpeg decoding before parsing. Keep dimension and aggregate pixel checks ahead of decoder allocations. Wire decoder dependencies into native builds and runtime packaging. Add regressions for appended EOI and embedded marker bypasses. Assisted-by: OpenAI:GPT-5 * fix(ci): gate native decision image validation Run the decoder security tests outside the stdlib-only native suite. Fetch vendor headers at the backend pin and provision decoder dependencies. Gate Go limit parity and production CMake wiring without model downloads. Assisted-by: OpenAI:GPT-5 * test(decisions): cover multimodal public API paths Exercise shared image contracts through the registered HTTP routes and external mock backend. Add opt-in cached gallery installation and real OpenJev image decisions through SystemOne and both routing APIs. Assisted-by: Codex:gpt-5 * test(decisions): assert isolation and cache bypass Observe external RPC calls and compare complete classifier history. Winner-only and cache-miss checks could hide dropped history or cache use. Give real inference its own application and model directory so shared backend mappings and loaded processes cannot affect mixed suite order. Assisted-by: OpenAI:ChatGPT * test(decisions): isolate fixture globals Disable optional global services in the isolated HTTP fixture and register cleanup before setup assertions. Verify meter provider identity survives fixture creation and destruction. Snapshot observed usage before assertions so failures cannot retain the mutex. Require a successful usage stamp before checking error responses. Assisted-by: Codex:gpt-5 golangci-lint * fix(application): honor optional telemetry controls Skip failover gauge registration when metrics are disabled. Register against the application meter rather than looking up the global provider. Allow embedders to retain the bounded routing log without billing stats. Keep the existing default when stats are disabled. The isolated HTTP fixture uses this option without losing its native router assertions. Assisted-by: Codex:gpt-5 golangci-lint --------- Co-authored-by: Ettore Di Giacinto <mudler@localai.io>
This commit is contained in:
1 parent
f035746db9
commit
99043b442c
108 files changed
+6263
-274
No files matched your search
@@ -17,6 +17,12 @@ if [[ -n "${CUDA_DOCKER_ARCH:-}" ]]; then
|
||||
rm -rf /LocalAI/backend/cpp/llama-cpp-*-build
|
||||
fi
|
||||
|
||||
# Install here, not only in the base builder: existing prebuilt bases must
|
||||
# acquire the strict decision decoders too. Runtime libraries are collected by
|
||||
# the existing backend dependency packaging.
|
||||
sh /LocalAI/.docker/apt-mirror.sh || true
|
||||
apt-get update -qq && apt-get install -y --no-install-recommends libjpeg-dev zlib1g-dev
|
||||
|
||||
cd /LocalAI/backend/cpp/llama-cpp
|
||||
BUILD_TARGET=$(/LocalAI/.docker/llama-cpp-build-target.sh "${TARGETARCH}" "${BUILD_TYPE:-}")
|
||||
if [ "$BUILD_TARGET" = "llama-cpp-cpu-all" ]; then
|
||||
|
||||
@@ -120,7 +120,7 @@ jobs:
|
||||
# libopusshim.dylib and to locate libopus.dylib for bundling. brew's
|
||||
# pkg-config defaults its search path to the Homebrew prefix so the
|
||||
# opus.pc is found.
|
||||
brew install protobuf grpc make protoc-gen-go protoc-gen-go-grpc libomp llvm ccache blake3 fmt hiredis xxhash zstd nlohmann-json opus pkg-config
|
||||
brew install protobuf grpc make protoc-gen-go protoc-gen-go-grpc libomp llvm ccache blake3 fmt hiredis xxhash zstd nlohmann-json opus pkg-config jpeg-turbo zlib
|
||||
# Force-reinstall ccache so brew re-validates its full runtime-dep
|
||||
# closure on every run. This is the durable fix: when the upstream
|
||||
# ccache formula gains a new transitive dep (as it has multiple times
|
||||
@@ -139,7 +139,7 @@ jobs:
|
||||
# and decides "already installed" without re-linking, so on a cache-
|
||||
# hit run the formulas aren't on PATH. Force-link them; --overwrite
|
||||
# tolerates pre-existing symlinks from earlier installs.
|
||||
brew link --overwrite protobuf grpc make protoc-gen-go protoc-gen-go-grpc libomp llvm ccache blake3 fmt hiredis xxhash zstd nlohmann-json opus pkg-config 2>/dev/null || true
|
||||
brew link --overwrite protobuf grpc make protoc-gen-go protoc-gen-go-grpc libomp llvm ccache blake3 fmt hiredis xxhash zstd nlohmann-json opus pkg-config jpeg-turbo zlib 2>/dev/null || true
|
||||
|
||||
- name: Save Homebrew cache
|
||||
if: github.event_name != 'pull_request' && steps.brew-cache.outputs.cache-hit != 'true'
|
||||
|
||||
@@ -0,0 +1,63 @@
|
||||
---
|
||||
name: 'Native decision image tests'
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
paths:
|
||||
- 'backend/cpp/llama-cpp/**'
|
||||
- 'core/systemone/**'
|
||||
- 'Dockerfile'
|
||||
- 'Dockerfile.*'
|
||||
- '.github/workflows/backend_build_darwin.yml'
|
||||
- '.github/workflows/decision-images.yml'
|
||||
push:
|
||||
branches:
|
||||
- master
|
||||
paths:
|
||||
- 'backend/cpp/llama-cpp/**'
|
||||
- 'core/systemone/**'
|
||||
- 'Dockerfile'
|
||||
- 'Dockerfile.*'
|
||||
- '.github/workflows/backend_build_darwin.yml'
|
||||
- '.github/workflows/decision-images.yml'
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
concurrency:
|
||||
group: decision-images-${{ github.event.pull_request.number || github.sha }}-${{ github.repository }}
|
||||
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
|
||||
|
||||
jobs:
|
||||
decision-images:
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
defaults:
|
||||
run:
|
||||
shell: bash
|
||||
steps:
|
||||
- uses: actions/checkout@v7
|
||||
with:
|
||||
persist-credentials: false
|
||||
- name: Install test dependencies
|
||||
run: |
|
||||
sudo apt-get update
|
||||
sudo apt-get install --no-install-recommends -y g++ python3 python3-pil cmake make curl ca-certificates zlib1g-dev libjpeg-dev
|
||||
# Only two vendor headers are needed: no backend build or model download.
|
||||
- name: Fetch pinned upstream headers
|
||||
run: |
|
||||
backend=backend/cpp/llama-cpp
|
||||
pin=$(sed -n 's/^LLAMA_VERSION?=//p' "$backend/Makefile")
|
||||
[[ "$pin" =~ ^[0-9a-f]{40}$ ]]
|
||||
for header in nlohmann/json.hpp stb/stb_image.h; do
|
||||
dest="$backend/llama.cpp/vendor/$header"
|
||||
mkdir -p "$(dirname "$dest")"
|
||||
curl --fail --location --retry 3 \
|
||||
"https://raw.githubusercontent.com/ggerganov/llama.cpp/$pin/vendor/$header" \
|
||||
--output "$dest"
|
||||
done
|
||||
- name: Verify native image validation and Go limit parity
|
||||
run: bash backend/cpp/llama-cpp/tests/verify-decision-images.sh
|
||||
- name: Verify production decoder build wiring
|
||||
run: python3 backend/cpp/llama-cpp/tests/verify-image-build-wiring.py
|
||||
@@ -83,6 +83,13 @@ target_link_libraries(${TARGET} PRIVATE ${_LLAMA_COMMON_TARGET} llama mtmd ${CMA
|
||||
gRPC::${_REFLECTION}
|
||||
gRPC::${_GRPC_GRPCPP}
|
||||
protobuf::${_PROTOBUF_LIBPROTOBUF})
|
||||
# Match grpc-server.cpp's native-decision guard: older forks do not need
|
||||
# strict decision image decoders and must not acquire new dependencies.
|
||||
if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/../server/server-decision.cpp")
|
||||
find_package(ZLIB REQUIRED)
|
||||
find_package(JPEG REQUIRED)
|
||||
target_link_libraries(${TARGET} PRIVATE ZLIB::ZLIB JPEG::JPEG)
|
||||
endif()
|
||||
target_compile_features(${TARGET} PRIVATE cxx_std_11)
|
||||
if(TARGET BUILD_INFO)
|
||||
add_dependencies(${TARGET} BUILD_INFO)
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
|
||||
LLAMA_VERSION?=a868c3e3c56657f7e8a6231190dbbe90e7dd86c0
|
||||
LLAMA_VERSION?=bed0a856606ee4a24a164066f73d2379447033f5
|
||||
LLAMA_REPO?=https://github.com/ggerganov/llama.cpp
|
||||
|
||||
CMAKE_ARGS?=
|
||||
|
||||
@@ -0,0 +1,18 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
#pragma once
|
||||
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
|
||||
// Nimble framing needs every question. Detect the callable signature rather
|
||||
// than a revision number so older decision-capable forks keep working too.
|
||||
template <typename Decision, typename State, typename Questions, typename... Args>
|
||||
void localai_fill_decision_task(const Decision & decision, const State & state,
|
||||
const Questions & questions, Args &&... args) {
|
||||
if constexpr (std::is_invocable_v<decltype(&Decision::fill_task),
|
||||
const Decision &, const State &, const Questions &, Args...>) {
|
||||
decision.fill_task(state, questions, std::forward<Args>(args)...);
|
||||
} else {
|
||||
decision.fill_task(state, std::forward<Args>(args)...);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
#include "decision_compat.h"
|
||||
#include <cassert>
|
||||
#include <vector>
|
||||
|
||||
struct legacy_decision {
|
||||
void fill_task(const int & state, int question, int & result) const {
|
||||
result = state + question;
|
||||
}
|
||||
};
|
||||
struct full_request_decision {
|
||||
const std::vector<int> * expected;
|
||||
void fill_task(const int & state, const std::vector<int> & questions,
|
||||
int question, int & result) const {
|
||||
assert(&questions == expected); // no copy or singleton substitution
|
||||
assert(questions.size() == 2);
|
||||
result = state + question + questions[1];
|
||||
}
|
||||
};
|
||||
int main() {
|
||||
const std::vector<int> questions{3, 7};
|
||||
int result = 0;
|
||||
localai_fill_decision_task(legacy_decision{}, 2, questions, 3, result);
|
||||
assert(result == 5);
|
||||
localai_fill_decision_task(full_request_decision{&questions}, 2, questions, 3, result);
|
||||
assert(result == 12);
|
||||
}
|
||||
@@ -0,0 +1,210 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
#pragma once
|
||||
#include <nlohmann/json.hpp>
|
||||
#include <stdexcept>
|
||||
#include <string>
|
||||
#include <vector>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <memory>
|
||||
#include <csetjmp>
|
||||
#include <cstdio>
|
||||
#include <jpeglib.h>
|
||||
#include <zlib.h>
|
||||
#include "stb/stb_image.h"
|
||||
|
||||
// Decision-only limits, mirrored from core/systemone/images.go. The verification
|
||||
// script checks parity; these must not change ordinary chat or fork backends.
|
||||
namespace localai_decision {
|
||||
using json = nlohmann::ordered_json;
|
||||
constexpr size_t max_images = 8;
|
||||
constexpr size_t decoded_bytes = 8 << 20;
|
||||
constexpr size_t encoded_bytes = 12 << 20;
|
||||
constexpr size_t body_bytes = 16 << 20;
|
||||
constexpr size_t text_bytes = 64 << 10;
|
||||
constexpr size_t max_dimension = 4096;
|
||||
constexpr size_t max_pixels = 16000000;
|
||||
struct image_error : std::invalid_argument {
|
||||
bool too_large;
|
||||
image_error(const char * message, bool large=false) : std::invalid_argument(message), too_large(large) {}
|
||||
};
|
||||
inline bool supports_images(bool decision, bool vision) { return decision && vision; }
|
||||
inline void require(bool ok, const char * message, bool large=false) {
|
||||
if (!ok) throw image_error(message, large);
|
||||
}
|
||||
// libjpeg normally repairs premature EOF and incomplete entropy scans. Treat
|
||||
// warnings as failures as well as fatal errors; an appended EOI cannot hide a
|
||||
// short scan. Keep all mutable decoder state on the heap across longjmp.
|
||||
struct jpeg_validator {
|
||||
jpeg_decompress_struct decoder{};
|
||||
jpeg_error_mgr errors{};
|
||||
std::jmp_buf jump;
|
||||
};
|
||||
inline void jpeg_failure(j_common_ptr decoder) {
|
||||
auto * state = static_cast<jpeg_validator *>(decoder->client_data);
|
||||
std::longjmp(state->jump, 1);
|
||||
}
|
||||
inline void jpeg_message(j_common_ptr decoder, int level) {
|
||||
if (level < 0) jpeg_failure(decoder);
|
||||
}
|
||||
inline void validate_jpeg(const std::vector<unsigned char> & raw, int width, int height) {
|
||||
auto state = std::make_unique<jpeg_validator>();
|
||||
auto * decoder = &state->decoder;
|
||||
decoder->err = jpeg_std_error(&state->errors);
|
||||
state->errors.error_exit = jpeg_failure;
|
||||
state->errors.emit_message = jpeg_message;
|
||||
decoder->client_data = state.get();
|
||||
if (setjmp(state->jump)) {
|
||||
jpeg_destroy_decompress(decoder);
|
||||
throw image_error("invalid or incomplete JPEG");
|
||||
}
|
||||
jpeg_create_decompress(decoder);
|
||||
jpeg_mem_src(decoder, raw.data(), raw.size());
|
||||
jpeg_read_header(decoder, TRUE);
|
||||
// The dimension/aggregate checks in validate_url precede all pixel or
|
||||
// coefficient allocations. Verify both decoders saw the same dimensions.
|
||||
if (decoder->image_width != unsigned(width) || decoder->image_height != unsigned(height)) {
|
||||
jpeg_destroy_decompress(decoder);
|
||||
throw image_error("inconsistent JPEG dimensions");
|
||||
}
|
||||
jpeg_start_decompress(decoder);
|
||||
auto row = (*decoder->mem->alloc_sarray)(reinterpret_cast<j_common_ptr>(decoder),
|
||||
JPOOL_IMAGE, decoder->output_width * decoder->output_components, 1);
|
||||
while (decoder->output_scanline < decoder->output_height) {
|
||||
jpeg_read_scanlines(decoder, row, 1);
|
||||
}
|
||||
jpeg_finish_decompress(decoder);
|
||||
jpeg_destroy_decompress(decoder);
|
||||
}
|
||||
inline int digit(unsigned char c) {
|
||||
if (c >= 'A' && c <= 'Z') return c-'A';
|
||||
if (c >= 'a' && c <= 'z') return c-'a'+26;
|
||||
if (c >= '0' && c <= '9') return c-'0'+52;
|
||||
return c=='+' ? 62 : c=='/' ? 63 : -1;
|
||||
}
|
||||
inline uint32_t be32(const unsigned char * p) {
|
||||
return uint32_t(p[0])<<24 | uint32_t(p[1])<<16 | uint32_t(p[2])<<8 | p[3];
|
||||
}
|
||||
inline uint32_t png_crc(const unsigned char * data, size_t size) {
|
||||
uint32_t crc = 0xffffffffu;
|
||||
for (size_t i = 0; i < size; ++i) {
|
||||
crc ^= data[i];
|
||||
for (int bit = 0; bit < 8; ++bit) crc = (crc >> 1) ^ (0xedb88320u & (0u - (crc & 1)));
|
||||
}
|
||||
return crc ^ 0xffffffffu;
|
||||
}
|
||||
inline void validate_url(const std::string & url, size_t & decoded, size_t & pixels) {
|
||||
const auto comma = url.find(',');
|
||||
const auto header = url.substr(0, comma);
|
||||
bool png = header == "data:image/png;base64";
|
||||
require(comma != std::string::npos && (png || header == "data:image/jpeg;base64"), "images must be PNG/JPEG base64 data URLs");
|
||||
size_t n = url.size()-comma-1;
|
||||
require(n > 0 && n%4 == 0, "invalid base64 length");
|
||||
const char * data = url.data()+comma+1;
|
||||
size_t pad = (data[n-1]=='=') + (data[n-2]=='=');
|
||||
size_t size = n/4*3-pad;
|
||||
require(size <= decoded_bytes-decoded, "decoded image aggregate exceeds limit", true);
|
||||
std::vector<unsigned char> raw;
|
||||
raw.reserve(size);
|
||||
for (size_t i=0; i<n; i+=4) {
|
||||
int a=digit(data[i]), b=digit(data[i+1]);
|
||||
int c=data[i+2]=='=' ? 0 : digit(data[i+2]);
|
||||
int d=data[i+3]=='=' ? 0 : digit(data[i+3]);
|
||||
require(a>=0 && b>=0 && c>=0 && d>=0, "invalid base64 character");
|
||||
require((data[i+2]!='=' && data[i+3]!='=') || i+4==n, "invalid base64 padding");
|
||||
require(data[i+2]!='=' || (data[i+3]=='=' && (b&15)==0), "invalid base64 padding bits");
|
||||
require(data[i+3]!='=' || data[i+2]=='=' || (c&3)==0, "invalid base64 padding bits");
|
||||
raw.push_back((a<<2)|(b>>4));
|
||||
if (data[i+2]!='=') raw.push_back((b<<4)|(c>>2));
|
||||
if (data[i+3]!='=') raw.push_back((c<<6)|d);
|
||||
}
|
||||
decoded += raw.size();
|
||||
require(png ? raw.size()>=24 && std::memcmp(raw.data(),"\x89PNG\r\n\x1a\n",8)==0
|
||||
: raw.size()>=3 && raw[0]==255 && raw[1]==216 && raw[2]==255, "image MIME mismatch");
|
||||
int w=0,h=0,c=0;
|
||||
require(stbi_info_from_memory(raw.data(),raw.size(),&w,&h,&c)!=0 && w>0 && h>0, "invalid image header");
|
||||
require(size_t(w)<=max_dimension && size_t(h)<=max_dimension, "image dimensions exceed limit", true);
|
||||
size_t count=size_t(w)*size_t(h);
|
||||
require(count<=max_pixels-pixels, "image pixel aggregate exceeds limit", true);
|
||||
pixels += count;
|
||||
if (png) {
|
||||
// stb's PNG inflater grows independently of IHDR. Validate IDAT with a
|
||||
// fixed output buffer first, preventing small-header decompression bombs.
|
||||
// 16-bit RGBA plus Adam7 row filters fit this conservative pixel bound.
|
||||
std::vector<unsigned char> idat;
|
||||
size_t pos=8;
|
||||
bool end=false;
|
||||
while (pos+12<=raw.size()) {
|
||||
size_t len=be32(raw.data()+pos);
|
||||
require(len<=raw.size()-pos-12, "truncated PNG chunk");
|
||||
require(png_crc(raw.data()+pos+4,len+4)==be32(raw.data()+pos+8+len), "invalid PNG checksum");
|
||||
if (std::memcmp(raw.data()+pos+4,"IDAT",4)==0)
|
||||
idat.insert(idat.end(),raw.begin()+pos+8,raw.begin()+pos+8+len);
|
||||
if (std::memcmp(raw.data()+pos+4,"IEND",4)==0) { end=true; break; }
|
||||
pos+=len+12;
|
||||
}
|
||||
require(end && !idat.empty(), "incomplete PNG");
|
||||
std::vector<char> inflated(9*count+8*size_t(h)+1024);
|
||||
z_stream stream{};
|
||||
stream.next_in = idat.data();
|
||||
stream.avail_in = static_cast<uInt>(idat.size());
|
||||
stream.next_out = reinterpret_cast<Bytef *>(inflated.data());
|
||||
stream.avail_out = static_cast<uInt>(inflated.size());
|
||||
require(inflateInit(&stream) == Z_OK, "PNG inflater initialization failed");
|
||||
int result = inflate(&stream, Z_FINISH);
|
||||
bool complete = result == Z_STREAM_END && stream.avail_in == 0;
|
||||
inflateEnd(&stream);
|
||||
require(complete, "invalid or oversized PNG decompression");
|
||||
} else {
|
||||
validate_jpeg(raw, w, h);
|
||||
}
|
||||
auto * image=stbi_load_from_memory(raw.data(),raw.size(),&w,&h,&c,3);
|
||||
require(image!=nullptr, "invalid image pixels");
|
||||
stbi_image_free(image);
|
||||
}
|
||||
// Normalize only actual chat content parts, as in the canonical Go collector.
|
||||
// Upstream parse_state understands image_url but not Anthropic source objects.
|
||||
inline size_t validate(json & body, size_t wire_size) {
|
||||
require(wire_size<=body_bytes, "decision request exceeds limit", true);
|
||||
std::vector<const std::string *> urls;
|
||||
auto add=[&](const json & value) {
|
||||
require(value.is_string(), "image URL must be a string");
|
||||
require(urls.size()<max_images, "too many decision images", true);
|
||||
urls.push_back(&value.get_ref<const std::string &>());
|
||||
};
|
||||
if (body.contains("images") && !body["images"].is_null()) {
|
||||
require(body["images"].is_array(), "images must be an array");
|
||||
for (const auto & url : body["images"]) add(url);
|
||||
}
|
||||
auto state=body.find("state");
|
||||
if (state!=body.end()) {
|
||||
json * messages=&*state;
|
||||
if (state->is_object() && state->contains("messages")) messages=&(*state)["messages"];
|
||||
if (messages->is_array()) for (auto & msg : *messages) {
|
||||
if (!msg.is_object() || !msg.contains("content") || !msg["content"].is_array()) continue;
|
||||
for (auto & part : msg["content"]) {
|
||||
if (!part.is_object() || !part.contains("type")) continue;
|
||||
if (part["type"]=="image") {
|
||||
require(part.contains("source") && part["source"].is_object(), "invalid image source");
|
||||
auto & s=part["source"];
|
||||
require(s.value("type",std::string())=="base64" && s.contains("media_type") && s["media_type"].is_string() && s.contains("data") && s["data"].is_string(), "invalid image source");
|
||||
require(s["data"].get_ref<const std::string &>().size()<=encoded_bytes && s["media_type"].get_ref<const std::string &>().size()<=32, "image source exceeds limit",true);
|
||||
std::string url="data:"+s["media_type"].get<std::string>()+";base64,"+s["data"].get<std::string>();
|
||||
part=json{{"type","image_url"},{"image_url",{{"url",url}}}};
|
||||
}
|
||||
if (part["type"]=="image_url") {
|
||||
require(part.contains("image_url"), "missing image URL");
|
||||
auto & u=part["image_url"];
|
||||
if (u.is_object()) { require(u.contains("url"), "missing image URL"); add(u["url"]); }
|
||||
else add(u);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
require(!urls.empty() || wire_size<=text_bytes, "text decision request exceeds limit",true);
|
||||
size_t encoded=0,decoded=0,pixels=0;
|
||||
for (const auto * u : urls) { require(u->size()<=encoded_bytes-encoded,"encoded image aggregate exceeds limit",true); encoded+=u->size(); }
|
||||
for (const auto * u : urls) validate_url(*u,decoded,pixels);
|
||||
return urls.size();
|
||||
}
|
||||
}
|
||||
@@ -43,6 +43,12 @@
|
||||
#if __has_include("server-stream.cpp")
|
||||
#include "server-stream.cpp"
|
||||
#endif
|
||||
#if __has_include("server-decision.cpp")
|
||||
#define LOCALAI_HAS_NATIVE_DECISIONS 1
|
||||
#include "server-decision.cpp"
|
||||
#include "decision_compat.h"
|
||||
#include "decision_images.h"
|
||||
#endif
|
||||
#include "server-context.cpp"
|
||||
|
||||
// LocalAI
|
||||
@@ -3301,6 +3307,107 @@ public:
|
||||
// together in one batch, so a warm scoring call costs roughly one
|
||||
// forward pass over the new prompt tokens plus one batched pass over
|
||||
// the candidate tails.
|
||||
#ifdef LOCALAI_HAS_NATIVE_DECISIONS
|
||||
grpc::Status SystemOne(ServerContext* context, const backend::ScoreRequest* request,
|
||||
backend::ScoreResponse* response) {
|
||||
const auto & decision = ctx_server.impl->decision;
|
||||
if (decision.type == COMMON_DECISION_TYPE_NONE) {
|
||||
return grpc::Status(grpc::StatusCode::UNIMPLEMENTED, "This model is not a decision model");
|
||||
}
|
||||
try {
|
||||
if (request->prompt().size() > localai_decision::body_bytes) {
|
||||
return grpc::Status(grpc::StatusCode::RESOURCE_EXHAUSTED, "Decision request exceeds limit");
|
||||
}
|
||||
auto checked_body = localai_decision::json::parse(request->prompt());
|
||||
const auto image_count = localai_decision::validate(checked_body, request->prompt().size());
|
||||
const json body = json::parse(checked_body.dump());
|
||||
if (image_count && !localai_decision::supports_images(decision.can_use_images(),
|
||||
ctx_server.impl->mctx && mtmd_support_vision(ctx_server.impl->mctx))) {
|
||||
return grpc::Status(grpc::StatusCode::UNIMPLEMENTED,
|
||||
"This decision model requires an image-capable projector for image input");
|
||||
}
|
||||
const auto questions = decision.parse_questions(body);
|
||||
std::vector<raw_buffer> files;
|
||||
const json state = decision.parse_state(body, files);
|
||||
if (context->IsCancelled()) {
|
||||
return grpc::Status(grpc::StatusCode::CANCELLED, "Request cancelled by client");
|
||||
}
|
||||
auto rd = ctx_server.get_response_reader(); // destructor cancels outstanding tasks
|
||||
std::vector<server_task> tasks;
|
||||
size_t expected_results = 0;
|
||||
for (const auto & question : questions) {
|
||||
for (size_t variant = 0; variant < decision.n_variants(question); ++variant) {
|
||||
server_task task(SERVER_TASK_TYPE_DECISION);
|
||||
task.id = rd.get_new_id();
|
||||
localai_fill_decision_task(decision, state, questions, question, variant, files,
|
||||
ctx_server.impl->mctx, ctx_server.impl->init_opt, task);
|
||||
tasks.push_back(std::move(task));
|
||||
++expected_results;
|
||||
}
|
||||
}
|
||||
if (decision.can_share_prompt()) {
|
||||
tasks = server_decision_group_tasks(std::move(tasks), params_base.n_parallel);
|
||||
}
|
||||
rd.post_tasks(std::move(tasks));
|
||||
auto results = rd.wait_for_all([&context]() { return context->IsCancelled(); });
|
||||
if (results.is_terminated) {
|
||||
return grpc::Status(grpc::StatusCode::CANCELLED, "Request cancelled by client");
|
||||
}
|
||||
if (results.error) {
|
||||
auto code = grpc::StatusCode::INTERNAL;
|
||||
const auto * error = dynamic_cast<server_task_result_error *>(results.error.get());
|
||||
if (!error) {
|
||||
return grpc::Status(grpc::StatusCode::INTERNAL, "Unexpected decision error type");
|
||||
}
|
||||
switch (error->err_type) {
|
||||
case ERROR_TYPE_INVALID_REQUEST:
|
||||
case ERROR_TYPE_EXCEED_CONTEXT_SIZE: code = grpc::StatusCode::INVALID_ARGUMENT; break;
|
||||
case ERROR_TYPE_NOT_SUPPORTED: code = grpc::StatusCode::UNIMPLEMENTED; break;
|
||||
default: break;
|
||||
}
|
||||
return grpc::Status(code, error->err_msg);
|
||||
}
|
||||
if (results.results.size() != expected_results) {
|
||||
return grpc::Status(grpc::StatusCode::INTERNAL, "Unexpected decision result count");
|
||||
}
|
||||
json answers = json::object();
|
||||
int64_t n_tokens = 0;
|
||||
size_t index = 0;
|
||||
for (const auto & question : questions) {
|
||||
std::vector<std::vector<float>> scores;
|
||||
for (size_t variant = 0; variant < decision.n_variants(question); ++variant) {
|
||||
auto * result = dynamic_cast<server_task_result_decision *>(results.results[index++].get());
|
||||
if (!result) {
|
||||
return grpc::Status(grpc::StatusCode::INTERNAL, "Unexpected decision result type");
|
||||
}
|
||||
scores.push_back(result->scores);
|
||||
n_tokens += result->n_tokens;
|
||||
}
|
||||
answers[question.id] = decision.format_answer(question, scores);
|
||||
}
|
||||
const auto output = json{
|
||||
{"model", body.value("model", std::string())}, {"answers", answers},
|
||||
{"usage", {{"input_tokens", n_tokens}, {"output_tokens", 0}}}
|
||||
}.dump();
|
||||
if (output.size() > localai_decision::text_bytes) {
|
||||
return grpc::Status(grpc::StatusCode::RESOURCE_EXHAUSTED, "Decision response exceeds limit");
|
||||
}
|
||||
response->set_response_json(output);
|
||||
return grpc::Status::OK;
|
||||
} catch (const localai_decision::image_error & err) {
|
||||
return grpc::Status(err.too_large ? grpc::StatusCode::RESOURCE_EXHAUSTED : grpc::StatusCode::INVALID_ARGUMENT, err.what());
|
||||
} catch (const localai_decision::json::exception & err) {
|
||||
return grpc::Status(grpc::StatusCode::INVALID_ARGUMENT, err.what());
|
||||
} catch (const common_json_error & err) {
|
||||
return grpc::Status(grpc::StatusCode::INVALID_ARGUMENT, err.what());
|
||||
} catch (const std::invalid_argument & err) {
|
||||
return grpc::Status(grpc::StatusCode::INVALID_ARGUMENT, err.what());
|
||||
} catch (const std::exception & err) {
|
||||
return grpc::Status(grpc::StatusCode::INTERNAL, err.what());
|
||||
}
|
||||
}
|
||||
#endif
|
||||
|
||||
grpc::Status Score(ServerContext* context, const backend::ScoreRequest* request, backend::ScoreResponse* response) override {
|
||||
auto auth = checkAuth(context);
|
||||
if (!auth.ok()) return auth;
|
||||
@@ -3309,6 +3416,14 @@ public:
|
||||
if (params_base.model.path.empty()) {
|
||||
return grpc::Status(grpc::StatusCode::FAILED_PRECONDITION, "Model not loaded");
|
||||
}
|
||||
if (request->question_type() == "systemone") {
|
||||
#ifdef LOCALAI_HAS_NATIVE_DECISIONS
|
||||
return SystemOne(context, request, response);
|
||||
#else
|
||||
return grpc::Status(grpc::StatusCode::UNIMPLEMENTED,
|
||||
"Native decisions are unavailable in this llama.cpp fork backend");
|
||||
#endif
|
||||
}
|
||||
#ifdef LOCALAI_LLAMA_CPP_NO_SCORE_TASK
|
||||
(void) request;
|
||||
(void) response;
|
||||
|
||||
@@ -39,6 +39,13 @@ GPU_LIB_SCRIPT="${REPO_ROOT}/scripts/build/package-gpu-libs.sh"
|
||||
if [ -f "$GPU_LIB_SCRIPT" ]; then
|
||||
echo "Packaging GPU libraries for BUILD_TYPE=${BUILD_TYPE:-cpu}..."
|
||||
source "$GPU_LIB_SCRIPT" "$CURDIR/package/lib"
|
||||
# Native decision validation links zlib/libjpeg. Collect actual ELF
|
||||
# dependencies rather than assuming those libraries exist on the host.
|
||||
for binary in "$CURDIR"/package/llama-cpp-*; do
|
||||
if [ -f "$binary" ]; then
|
||||
copy_elf_deps "$binary"
|
||||
fi
|
||||
done
|
||||
package_gpu_libs
|
||||
fi
|
||||
|
||||
|
||||
@@ -1,20 +1,8 @@
|
||||
From 75220a0d74892e3315f4042274b1efa6195868d8 Mon Sep 17 00:00:00 2001
|
||||
From: Codex <codex@local>
|
||||
Date: Mon, 10 Aug 2026 23:05:52 +0000
|
||||
Subject: [PATCH 1/2] score-patch
|
||||
|
||||
---
|
||||
common/common.cpp | 6 +-
|
||||
common/common.h | 3 +
|
||||
tools/server/server-context.cpp | 358 +++++++++++++++++++++++++++++++-
|
||||
tools/server/server-task.h | 47 +++++
|
||||
4 files changed, 405 insertions(+), 9 deletions(-)
|
||||
|
||||
diff --git a/common/common.cpp b/common/common.cpp
|
||||
index 2e3f14c..0cec0dc 100644
|
||||
index aca1949..428b120 100644
|
||||
--- a/common/common.cpp
|
||||
+++ b/common/common.cpp
|
||||
@@ -1636,8 +1636,10 @@ struct llama_context_params common_context_params_to_llama(const common_params &
|
||||
@@ -1662,8 +1662,10 @@ struct llama_context_params common_context_params_to_llama(const common_params &
|
||||
auto cparams = llama_context_default_params();
|
||||
|
||||
cparams.n_ctx = params.n_ctx;
|
||||
@@ -28,10 +16,10 @@ index 2e3f14c..0cec0dc 100644
|
||||
cparams.n_outputs_max_per_seq = std::max(params.n_outputs_max_per_seq, 0);
|
||||
cparams.n_batch = params.n_batch;
|
||||
diff --git a/common/common.h b/common/common.h
|
||||
index 878534d..4001df2 100644
|
||||
index 04ffbcd..5daa023 100644
|
||||
--- a/common/common.h
|
||||
+++ b/common/common.h
|
||||
@@ -445,6 +445,9 @@ struct common_params {
|
||||
@@ -455,6 +455,9 @@ struct common_params {
|
||||
int32_t n_keep = 0; // number of tokens to keep from initial prompt
|
||||
int32_t n_chunks = -1; // max number of chunks to process (-1 = unlimited)
|
||||
int32_t n_parallel = 1; // number of parallel sequences to decode
|
||||
@@ -42,7 +30,7 @@ index 878534d..4001df2 100644
|
||||
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/server/server-context.cpp b/tools/server/server-context.cpp
|
||||
index 3b5f6a1..d0e18e6 100644
|
||||
index edb8e2d..9c88feb 100644
|
||||
--- a/tools/server/server-context.cpp
|
||||
+++ b/tools/server/server-context.cpp
|
||||
@@ -48,6 +48,13 @@ static common_speculative_output_limits server_output_limits(const common_params
|
||||
@@ -59,7 +47,7 @@ index 3b5f6a1..d0e18e6 100644
|
||||
result.total = std::max<int32_t>(1, result.total);
|
||||
result.per_seq = std::max<int32_t>(1, result.per_seq);
|
||||
return result;
|
||||
@@ -239,6 +246,26 @@ struct server_slot {
|
||||
@@ -243,6 +250,26 @@ struct server_slot {
|
||||
|
||||
std::vector<completion_token_output> generated_token_probs;
|
||||
|
||||
@@ -86,7 +74,7 @@ index 3b5f6a1..d0e18e6 100644
|
||||
bool has_next_token = true;
|
||||
bool has_new_line = false;
|
||||
bool truncated = false;
|
||||
@@ -341,6 +368,10 @@ struct server_slot {
|
||||
@@ -346,6 +373,10 @@ struct server_slot {
|
||||
}
|
||||
generated_tokens.clear();
|
||||
generated_token_probs.clear();
|
||||
@@ -97,7 +85,7 @@ index 3b5f6a1..d0e18e6 100644
|
||||
json_schema = json();
|
||||
|
||||
task_prev = std::move(task);
|
||||
@@ -2271,6 +2302,227 @@ private:
|
||||
@@ -2308,6 +2339,227 @@ private:
|
||||
queue_results.send(std::move(res));
|
||||
}
|
||||
|
||||
@@ -325,15 +313,15 @@ index 3b5f6a1..d0e18e6 100644
|
||||
//
|
||||
// Functions to process the task
|
||||
//
|
||||
@@ -2407,6 +2661,7 @@ private:
|
||||
case SERVER_TASK_TYPE_INFILL:
|
||||
@@ -2468,6 +2720,7 @@ private:
|
||||
case SERVER_TASK_TYPE_EMBEDDING:
|
||||
case SERVER_TASK_TYPE_RERANK:
|
||||
case SERVER_TASK_TYPE_DECISION:
|
||||
+ case SERVER_TASK_TYPE_SCORE:
|
||||
{
|
||||
// special case: if input is provided via CLI, tokenize it first
|
||||
// otherwise, no need to tokenize as it's already done inside the HTTP thread
|
||||
@@ -2903,6 +3158,13 @@ private:
|
||||
@@ -2986,6 +3239,13 @@ private:
|
||||
break; // stop any further processing
|
||||
}
|
||||
}
|
||||
@@ -347,7 +335,7 @@ index 3b5f6a1..d0e18e6 100644
|
||||
}
|
||||
|
||||
void pre_decode() {
|
||||
@@ -3222,6 +3484,16 @@ private:
|
||||
@@ -3325,6 +3585,16 @@ private:
|
||||
n_past = std::min(n_past, slot.alora_invocation_start - 1);
|
||||
}
|
||||
|
||||
@@ -364,7 +352,7 @@ index 3b5f6a1..d0e18e6 100644
|
||||
const auto n_cache_reuse = slot.task->params.n_cache_reuse;
|
||||
|
||||
const bool can_cache_reuse =
|
||||
@@ -3455,8 +3727,12 @@ private:
|
||||
@@ -3578,8 +3848,12 @@ private:
|
||||
|
||||
bool do_checkpoint = params_base.n_ctx_checkpoints > 0;
|
||||
|
||||
@@ -379,7 +367,7 @@ index 3b5f6a1..d0e18e6 100644
|
||||
|
||||
// make a checkpoint of the parts of the memory that cannot be rolled back.
|
||||
// checkpoints are created only if:
|
||||
@@ -3444,9 +3720,16 @@ private:
|
||||
@@ -3670,13 +3944,46 @@ private:
|
||||
// embedding requires all tokens in the batch to be output;
|
||||
// MTP also wants logits at every prompt position so the
|
||||
// streaming hook can mirror t_h_nextn into ctx_dft.
|
||||
@@ -397,7 +385,7 @@ index 3b5f6a1..d0e18e6 100644
|
||||
+ /* output = */ slot.need_embd() || need_score_logit,
|
||||
/* is_prompt = */ true);
|
||||
slot.prompt.tokens.push_back(cur_tok);
|
||||
@@ -3454,2 +3737,28 @@ private:
|
||||
|
||||
+ // score tasks: break at the shared-prompt boundary so the checkpoint
|
||||
+ // below lands exactly there — the other candidates of the same
|
||||
+ // scoring call re-process only their own tokens. Also break at the
|
||||
@@ -426,7 +414,8 @@ index 3b5f6a1..d0e18e6 100644
|
||||
+
|
||||
// break at the last user message, or at user messages at least min step past the last checkpoint
|
||||
if (do_checkpoint && spans.is_user_start(slot.prompt.n_tokens())) {
|
||||
@@ -3573,6 +3882,15 @@ private:
|
||||
const auto pos = slot.prompt.n_tokens();
|
||||
@@ -3719,6 +4026,15 @@ private:
|
||||
const bool is_user_start = spans.is_user_start(n_tokens_start);
|
||||
const bool is_last_user_message = n_tokens_start == last_user_pos;
|
||||
|
||||
@@ -442,7 +431,7 @@ index 3b5f6a1..d0e18e6 100644
|
||||
// entire prompt has been processed
|
||||
if (slot.prompt.n_tokens() == slot.task->n_tokens()) {
|
||||
slot.state = SLOT_STATE_DONE_PROMPT;
|
||||
@@ -3588,8 +3906,8 @@ private:
|
||||
@@ -3734,8 +4050,8 @@ private:
|
||||
slot.init_sampler();
|
||||
} else {
|
||||
// skip ordinary mid-prompt checkpoints, unless the batch starts a user
|
||||
@@ -453,7 +442,7 @@ index 3b5f6a1..d0e18e6 100644
|
||||
do_checkpoint = false;
|
||||
}
|
||||
}
|
||||
@@ -3606,10 +3924,10 @@ private:
|
||||
@@ -3752,10 +4068,10 @@ private:
|
||||
// do not checkpoint after mtmd chunks
|
||||
do_checkpoint = do_checkpoint && !has_mtmd;
|
||||
|
||||
@@ -466,7 +455,7 @@ index 3b5f6a1..d0e18e6 100644
|
||||
n_tokens_start > slot.prompt.checkpoints.back().n_tokens + params_base.checkpoint_min_step);
|
||||
SLT_DBG(slot, "main/do_checkpoint = %s, pos_min = %d, pos_max = %d\n", do_checkpoint ? "yes" : "no", pos_min, pos_max);
|
||||
|
||||
@@ -3772,6 +4090,13 @@ private:
|
||||
@@ -3943,6 +4259,13 @@ private:
|
||||
}
|
||||
}
|
||||
|
||||
@@ -480,7 +469,7 @@ index 3b5f6a1..d0e18e6 100644
|
||||
if (!is_inside_view(slot.i_batch)) {
|
||||
// the required token not in this sub-batch, skip
|
||||
return;
|
||||
@@ -3793,6 +4118,25 @@ private:
|
||||
@@ -3971,6 +4294,25 @@ private:
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -507,12 +496,12 @@ index 3b5f6a1..d0e18e6 100644
|
||||
|
||||
// prompt evaluated for next-token prediction
|
||||
diff --git a/tools/server/server-task.h b/tools/server/server-task.h
|
||||
index 6275ec7..5bedf19 100644
|
||||
index 8c5fa9a..4b38805 100644
|
||||
--- a/tools/server/server-task.h
|
||||
+++ b/tools/server/server-task.h
|
||||
@@ -13,10 +13,25 @@
|
||||
@@ -12,11 +12,26 @@
|
||||
#include "server-common.h"
|
||||
|
||||
using json = nlohmann::ordered_json;
|
||||
|
||||
+// SERVER_TASK_TYPE_SCORE emits one logits output per candidate token (plus
|
||||
+// the forced last-token output), and the context's output budget
|
||||
@@ -532,11 +521,12 @@ index 6275ec7..5bedf19 100644
|
||||
SERVER_TASK_TYPE_COMPLETION,
|
||||
SERVER_TASK_TYPE_EMBEDDING,
|
||||
SERVER_TASK_TYPE_RERANK,
|
||||
SERVER_TASK_TYPE_DECISION,
|
||||
+ SERVER_TASK_TYPE_SCORE,
|
||||
SERVER_TASK_TYPE_INFILL,
|
||||
SERVER_TASK_TYPE_CANCEL,
|
||||
SERVER_TASK_TYPE_CONTROL,
|
||||
@@ -153,6 +168,18 @@ struct server_task {
|
||||
@@ -156,6 +171,18 @@ struct server_task {
|
||||
task_params params;
|
||||
server_tokens tokens;
|
||||
|
||||
@@ -555,7 +545,7 @@ index 6275ec7..5bedf19 100644
|
||||
// only used by CLI, this allow tokenizing CLI inputs on server side
|
||||
// we need this because mtmd_context and vocab are not accessible outside of server_context
|
||||
bool cli = false;
|
||||
@@ -197,6 +224,7 @@ struct server_task {
|
||||
@@ -234,6 +261,7 @@ struct server_task {
|
||||
switch (type) {
|
||||
case SERVER_TASK_TYPE_COMPLETION:
|
||||
case SERVER_TASK_TYPE_INFILL:
|
||||
@@ -563,7 +553,7 @@ index 6275ec7..5bedf19 100644
|
||||
return true;
|
||||
default:
|
||||
return false;
|
||||
@@ -494,6 +522,25 @@ struct server_task_result_rerank : server_task_result {
|
||||
@@ -509,6 +537,25 @@ struct server_task_result_decision : server_task_result {
|
||||
virtual json to_json() override;
|
||||
};
|
||||
|
||||
@@ -589,5 +579,3 @@ index 6275ec7..5bedf19 100644
|
||||
struct server_task_result_error : server_task_result {
|
||||
error_type err_type = ERROR_TYPE_SERVER;
|
||||
std::string err_msg;
|
||||
--
|
||||
2.39.5
|
||||
@@ -1,5 +1,5 @@
|
||||
diff --git a/tools/mtmd/mtmd-helper-gen.cpp b/tools/mtmd/mtmd-helper-gen.cpp
|
||||
index 1c58d3ae1..196cbd433 100644
|
||||
index 5fb7ea9..b554bf2 100644
|
||||
--- a/tools/mtmd/mtmd-helper-gen.cpp
|
||||
+++ b/tools/mtmd/mtmd-helper-gen.cpp
|
||||
@@ -50,29 +50,38 @@ static llama_token find_special_token(const llama_vocab * vocab, const std::stri
|
||||
@@ -86,7 +86,7 @@ index 1c58d3ae1..196cbd433 100644
|
||||
|
||||
// the prompt above holds the whole text stream up to tts_eos, so every generated
|
||||
// frame adds tts_pad on top of the codes embedding
|
||||
@@ -302,31 +317,60 @@ public:
|
||||
@@ -301,31 +316,60 @@ public:
|
||||
}
|
||||
|
||||
int32_t get_output(int32_t * out_sample_rate, const char ** out_data, size_t * out_data_len, int64_t * out_n_samples) override {
|
||||
@@ -156,7 +156,7 @@ index 1c58d3ae1..196cbd433 100644
|
||||
private:
|
||||
bool ensure_cache() {
|
||||
if (specials_ok) {
|
||||
@@ -370,7 +414,7 @@ private:
|
||||
@@ -369,7 +413,7 @@ private:
|
||||
LOG_ERR("mtmd_helper_gen_audio: mmproj has no speaker/audio encoder\n");
|
||||
return false;
|
||||
}
|
||||
@@ -165,7 +165,7 @@ index 1c58d3ae1..196cbd433 100644
|
||||
mtmd_input_text text{ marker.c_str(), marker.size(), false, true };
|
||||
mtmd_input_chunks * chunks = mtmd_input_chunks_init();
|
||||
const mtmd_bitmap * bptr = bitmap;
|
||||
@@ -456,6 +500,9 @@ private:
|
||||
@@ -455,6 +499,9 @@ private:
|
||||
std::vector<float> h_state_buf;
|
||||
mtmd_helper_gen_audio_outtype out_type = MTMD_HELPER_GEN_AUDIO_OUTTYPE_WAV;
|
||||
std::vector<char> out_buf;
|
||||
@@ -175,7 +175,7 @@ index 1c58d3ae1..196cbd433 100644
|
||||
};
|
||||
|
||||
// settings that only live in the reference's per-pack yaml, not in the checkpoint
|
||||
@@ -1024,6 +1071,14 @@ void mtmd_helper_gen_audio_reset(mtmd_helper_gen_audio * ctx) {
|
||||
@@ -1022,6 +1069,14 @@ void mtmd_helper_gen_audio_reset(mtmd_helper_gen_audio * ctx) {
|
||||
}
|
||||
}
|
||||
|
||||
@@ -190,7 +190,7 @@ index 1c58d3ae1..196cbd433 100644
|
||||
int32_t mtmd_helper_gen_audio_set_input(mtmd_helper_gen_audio * ctx, const mtmd_helper_gen_audio_inp * inp) {
|
||||
if (!ctx->pipeline) {
|
||||
LOG_ERR("mtmd_helper_gen_audio: unsupported or missing gen-audio pipeline\n");
|
||||
@@ -1060,3 +1115,10 @@ int32_t mtmd_helper_gen_audio_get_output(mtmd_helper_gen_audio * ctx, int32_t *
|
||||
@@ -1058,3 +1113,10 @@ int32_t mtmd_helper_gen_audio_get_output(mtmd_helper_gen_audio * ctx, int32_t *
|
||||
}
|
||||
return ctx->pipeline->get_output(out_sample_rate, out_data, out_data_len, out_n_samples);
|
||||
}
|
||||
@@ -202,10 +202,10 @@ index 1c58d3ae1..196cbd433 100644
|
||||
+ return ctx->pipeline->flush();
|
||||
+}
|
||||
diff --git a/tools/mtmd/mtmd-helper.h b/tools/mtmd/mtmd-helper.h
|
||||
index 832f7171a..3eaa01aab 100644
|
||||
index 7436230..acbfefc 100644
|
||||
--- a/tools/mtmd/mtmd-helper.h
|
||||
+++ b/tools/mtmd/mtmd-helper.h
|
||||
@@ -175,6 +175,7 @@ enum mtmd_helper_gen_audio_outtype {
|
||||
@@ -204,6 +204,7 @@ enum mtmd_helper_gen_audio_outtype {
|
||||
MTMD_HELPER_GEN_AUDIO_OUTTYPE_WAV, // WAV PCM 16-bit LE, mono
|
||||
};
|
||||
struct mtmd_helper_gen_audio_inp {
|
||||
@@ -213,7 +213,7 @@ index 832f7171a..3eaa01aab 100644
|
||||
llama_seq_id seq_id;
|
||||
|
||||
const char * prompt;
|
||||
@@ -190,6 +191,8 @@ struct mtmd_helper_gen_audio_inp {
|
||||
@@ -219,6 +220,8 @@ struct mtmd_helper_gen_audio_inp {
|
||||
enum mtmd_helper_gen_audio_outtype out_type;
|
||||
};
|
||||
|
||||
@@ -222,7 +222,7 @@ index 832f7171a..3eaa01aab 100644
|
||||
MTMD_API mtmd_helper_gen_audio * mtmd_helper_gen_audio_init(
|
||||
struct llama_context * lctx,
|
||||
struct mtmd_context * mctx);
|
||||
@@ -221,6 +224,8 @@ MTMD_API int32_t mtmd_helper_gen_audio_step_gen(
|
||||
@@ -250,6 +253,8 @@ MTMD_API int32_t mtmd_helper_gen_audio_step_gen(
|
||||
|
||||
// out_data valid until next get_output() or reset() call
|
||||
// out_n_samples (optional, can be NULL) receives the number of generated PCM samples
|
||||
@@ -231,7 +231,7 @@ index 832f7171a..3eaa01aab 100644
|
||||
MTMD_API int32_t mtmd_helper_gen_audio_get_output(
|
||||
mtmd_helper_gen_audio * ctx,
|
||||
int32_t * out_sample_rate,
|
||||
@@ -228,6 +233,10 @@ MTMD_API int32_t mtmd_helper_gen_audio_get_output(
|
||||
@@ -257,6 +262,10 @@ MTMD_API int32_t mtmd_helper_gen_audio_get_output(
|
||||
size_t * out_data_len,
|
||||
int64_t * out_n_samples);
|
||||
|
||||
@@ -242,7 +242,7 @@ index 832f7171a..3eaa01aab 100644
|
||||
#ifdef __cplusplus
|
||||
} // extern "C"
|
||||
#endif
|
||||
@@ -254,8 +263,41 @@ struct mtmd_helper_gen_audio_deleter {
|
||||
@@ -283,8 +292,41 @@ struct mtmd_helper_gen_audio_deleter {
|
||||
};
|
||||
using gen_audio_ptr = std::unique_ptr<mtmd_helper_gen_audio, mtmd_helper_gen_audio_deleter>;
|
||||
struct gen_audio {
|
||||
@@ -285,7 +285,7 @@ index 832f7171a..3eaa01aab 100644
|
||||
void reset() {
|
||||
mtmd_helper_gen_audio_reset(ctx.get());
|
||||
}
|
||||
@@ -271,6 +313,9 @@ struct gen_audio {
|
||||
@@ -300,6 +342,9 @@ struct gen_audio {
|
||||
int32_t get_output(int32_t * out_sample_rate, const char ** out_data, size_t * out_data_len, int64_t * out_n_samples = nullptr) {
|
||||
return mtmd_helper_gen_audio_get_output(ctx.get(), out_sample_rate, out_data, out_data_len, out_n_samples);
|
||||
}
|
||||
@@ -296,10 +296,10 @@ index 832f7171a..3eaa01aab 100644
|
||||
|
||||
} // namespace mtmd_helper
|
||||
diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp
|
||||
index 9069463fe..b7fa1e534 100644
|
||||
index 9c88feb..064bb51 100644
|
||||
--- a/tools/server/server-context.cpp
|
||||
+++ b/tools/server/server-context.cpp
|
||||
@@ -16,6 +16,7 @@
|
||||
@@ -17,6 +17,7 @@
|
||||
#include "speculative.h"
|
||||
#include "mtmd.h"
|
||||
#include "mtmd-helper.h"
|
||||
@@ -319,7 +319,7 @@ index 9069463fe..b7fa1e534 100644
|
||||
}
|
||||
|
||||
auto result = common_speculative_get_output_limits(
|
||||
@@ -212,6 +214,30 @@ struct server_slot {
|
||||
@@ -215,6 +217,30 @@ struct server_slot {
|
||||
mtmd_context * mctx = nullptr;
|
||||
mtmd::batch_ptr mbatch = nullptr;
|
||||
|
||||
@@ -350,7 +350,7 @@ index 9069463fe..b7fa1e534 100644
|
||||
// speculative decoding
|
||||
common_speculative * spec;
|
||||
|
||||
@@ -391,6 +417,8 @@ struct server_slot {
|
||||
@@ -395,6 +421,8 @@ struct server_slot {
|
||||
|
||||
// clear multimodal state
|
||||
mbatch.reset();
|
||||
@@ -359,9 +359,9 @@ index 9069463fe..b7fa1e534 100644
|
||||
}
|
||||
|
||||
void init_sampler() const {
|
||||
@@ -829,6 +857,14 @@ public:
|
||||
mtmd_context * mctx = nullptr;
|
||||
const llama_vocab * vocab = nullptr;
|
||||
@@ -871,6 +899,14 @@ public:
|
||||
|
||||
server_decision_context decision;
|
||||
|
||||
+ bool has_cap_tts() const {
|
||||
+ return mctx != nullptr && mtmd_gen_audio_get_info(mctx).type != MTMD_GEN_AUDIO_TYPE_NONE;
|
||||
@@ -374,7 +374,7 @@ index 9069463fe..b7fa1e534 100644
|
||||
server_queue queue_tasks;
|
||||
server_response queue_results;
|
||||
|
||||
@@ -1288,6 +1324,10 @@ private:
|
||||
@@ -1350,6 +1386,10 @@ private:
|
||||
slot.mctx = mctx;
|
||||
slot.prompt.tokens.has_mtmd = mctx != nullptr;
|
||||
|
||||
@@ -385,7 +385,7 @@ index 9069463fe..b7fa1e534 100644
|
||||
SLT_TRC(slot, "new slot, n_ctx = %d\n", slot.n_ctx);
|
||||
|
||||
slot.callback_on_release = [this](int id_slot) {
|
||||
@@ -1748,6 +1788,28 @@ private:
|
||||
@@ -1830,6 +1870,28 @@ private:
|
||||
|
||||
SLT_DBG(slot, "launching slot : %s\n", safe_json_to_str(slot.to_json()).c_str());
|
||||
|
||||
@@ -414,7 +414,7 @@ index 9069463fe..b7fa1e534 100644
|
||||
// initialize samplers
|
||||
if (task.need_sampling()) {
|
||||
try {
|
||||
@@ -1765,6 +1827,9 @@ private:
|
||||
@@ -1847,6 +1909,9 @@ private:
|
||||
// TODO: getting pre sampling logits is not yet supported with backend sampling
|
||||
use_backend_sampling &= !need_pre_sample_logits;
|
||||
|
||||
@@ -424,7 +424,7 @@ index 9069463fe..b7fa1e534 100644
|
||||
// TODO: tmp until backend sampling is fully implemented
|
||||
if (use_backend_sampling) {
|
||||
llama_set_sampler(ctx_tgt, slot.id, common_sampler_get(slot.smpl.get()));
|
||||
@@ -1783,9 +1848,13 @@ private:
|
||||
@@ -1872,9 +1937,13 @@ private:
|
||||
|
||||
slot.task = std::make_unique<const server_task>(std::move(task));
|
||||
|
||||
@@ -441,7 +441,7 @@ index 9069463fe..b7fa1e534 100644
|
||||
|
||||
// reset server kill-switch counter
|
||||
n_empty_consecutive = 0;
|
||||
@@ -2050,6 +2119,18 @@ private:
|
||||
@@ -2139,6 +2208,18 @@ private:
|
||||
queue_results.send(std::move(res));
|
||||
}
|
||||
|
||||
@@ -460,15 +460,18 @@ index 9069463fe..b7fa1e534 100644
|
||||
void send_final_response(server_slot & slot) {
|
||||
auto res = std::make_unique<server_task_result_cmpl_final>();
|
||||
|
||||
@@ -2556,6 +2637,7 @@ private:
|
||||
case SERVER_TASK_TYPE_EMBEDDING:
|
||||
@@ -2721,6 +2802,7 @@ private:
|
||||
case SERVER_TASK_TYPE_RERANK:
|
||||
case SERVER_TASK_TYPE_DECISION:
|
||||
case SERVER_TASK_TYPE_SCORE:
|
||||
+ case SERVER_TASK_TYPE_TTS:
|
||||
{
|
||||
// special case: if input is provided via CLI, tokenize it first
|
||||
// otherwise, no need to tokenize as it's already done inside the HTTP thread
|
||||
@@ -3007,1 +3089,9 @@ private:
|
||||
@@ -3179,6 +3261,14 @@ private:
|
||||
return;
|
||||
}
|
||||
|
||||
+ // note: TTS slots bypass the shared batch entirely
|
||||
+ try {
|
||||
+ process_tts_slots();
|
||||
@@ -478,7 +481,9 @@ index 9069463fe..b7fa1e534 100644
|
||||
+ }
|
||||
+
|
||||
GGML_ASSERT(batch.slot_batched || batch.size() == 0);
|
||||
@@ -3074,10 +3164,77 @@ private:
|
||||
|
||||
if (batch.slot_batched) {
|
||||
@@ -3248,10 +3338,77 @@ private:
|
||||
}
|
||||
}
|
||||
|
||||
@@ -556,7 +561,7 @@ index 9069463fe..b7fa1e534 100644
|
||||
if (slot.state == SLOT_STATE_GENERATING && slot.prompt.n_tokens() + 1 >= slot.n_ctx) {
|
||||
if (!params_base.ctx_shift) {
|
||||
// this check is redundant (for good)
|
||||
@@ -3150,7 +3307,7 @@ private:
|
||||
@@ -3324,7 +3481,7 @@ private:
|
||||
|
||||
// determine which slots are generating and drafting
|
||||
iterate(slots, [&](server_slot & slot) {
|
||||
@@ -565,7 +570,7 @@ index 9069463fe..b7fa1e534 100644
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -3284,7 +3441,7 @@ private:
|
||||
@@ -3458,7 +3615,7 @@ private:
|
||||
return; // batch is full, skip remaining slots
|
||||
}
|
||||
|
||||
@@ -574,16 +579,16 @@ index 9069463fe..b7fa1e534 100644
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -4433,6 +4590,8 @@ server_context_meta server_context::get_meta() const {
|
||||
@@ -4678,6 +4835,8 @@ server_context_meta server_context::get_meta() const {
|
||||
/* has_inp_image */ impl->chat_params.allow_image,
|
||||
/* has_inp_audio */ impl->chat_params.allow_audio,
|
||||
/* has_inp_video */ impl->chat_params.allow_video,
|
||||
+ /* has_cap_chat */ impl->has_cap_chat(),
|
||||
+ /* has_cap_tts */ impl->has_cap_tts(),
|
||||
/* json_ui_settings */ impl->json_ui_settings,
|
||||
/* slot_n_ctx */ impl->get_slot_n_ctx(),
|
||||
/* slot_n_ctx */ impl->n_ctx_slot(),
|
||||
/* pooling_type */ llama_pooling_type(impl->ctx_tgt),
|
||||
@@ -4512,6 +4671,11 @@ std::unique_ptr<server_res_generator> server_routes::handle_completions_impl(
|
||||
@@ -4751,6 +4910,11 @@ std::unique_ptr<server_res_generator> server_routes::handle_completions_impl(
|
||||
|
||||
res->set_req(&req); // will also set spipe if needed
|
||||
|
||||
@@ -595,7 +600,7 @@ index 9069463fe..b7fa1e534 100644
|
||||
int32_t sse_ping_interval = params.sse_ping_interval;
|
||||
|
||||
try {
|
||||
@@ -5399,6 +5563,150 @@ void server_routes::init_routes() {
|
||||
@@ -5776,6 +5940,150 @@ void server_routes::init_routes() {
|
||||
return res;
|
||||
};
|
||||
|
||||
@@ -747,10 +752,10 @@ index 9069463fe..b7fa1e534 100644
|
||||
auto res = create_response();
|
||||
|
||||
diff --git a/tools/server/server-context.h b/tools/server/server-context.h
|
||||
index f9ab1132b..610512678 100644
|
||||
index c554bb9..f025792 100644
|
||||
--- a/tools/server/server-context.h
|
||||
+++ b/tools/server/server-context.h
|
||||
@@ -22,6 +22,8 @@ struct server_context_meta {
|
||||
@@ -23,6 +23,8 @@ struct server_context_meta {
|
||||
bool has_inp_image;
|
||||
bool has_inp_audio;
|
||||
bool has_inp_video;
|
||||
@@ -759,19 +764,19 @@ index f9ab1132b..610512678 100644
|
||||
json json_ui_settings;
|
||||
int slot_n_ctx;
|
||||
enum llama_pooling_type pooling_type;
|
||||
@@ -151,6 +153,7 @@ struct server_routes {
|
||||
@@ -152,6 +154,7 @@ struct server_routes {
|
||||
server_http_context::handler_t post_embeddings;
|
||||
server_http_context::handler_t post_embeddings_oai;
|
||||
server_http_context::handler_t post_rerank;
|
||||
+ server_http_context::handler_t post_tts;
|
||||
server_http_context::handler_t post_systemone;
|
||||
server_http_context::handler_t get_lora_adapters;
|
||||
server_http_context::handler_t post_lora_adapters;
|
||||
|
||||
diff --git a/tools/server/server-task.cpp b/tools/server/server-task.cpp
|
||||
index 1ee677553..939630b8b 100644
|
||||
index a5c33c0..7e59745 100644
|
||||
--- a/tools/server/server-task.cpp
|
||||
+++ b/tools/server/server-task.cpp
|
||||
@@ -1497,6 +1497,17 @@ json server_task_result_rerank::to_json() {
|
||||
@@ -1506,6 +1506,17 @@ json server_task_result_decision::to_json() {
|
||||
};
|
||||
}
|
||||
|
||||
@@ -790,7 +795,7 @@ index 1ee677553..939630b8b 100644
|
||||
// server_task_result_error
|
||||
//
|
||||
diff --git a/tools/server/server-task.h b/tools/server/server-task.h
|
||||
index 5bedf1987..e6ca67a65 100644
|
||||
index 4b38805..d79ea4f 100644
|
||||
--- a/tools/server/server-task.h
|
||||
+++ b/tools/server/server-task.h
|
||||
@@ -10,6 +10,7 @@
|
||||
@@ -799,9 +804,9 @@ index 5bedf1987..e6ca67a65 100644
|
||||
#include "server-common.h"
|
||||
+#include "mtmd-helper.h"
|
||||
|
||||
using json = nlohmann::ordered_json;
|
||||
|
||||
@@ -42,6 +43,7 @@ enum server_task_type {
|
||||
// SERVER_TASK_TYPE_SCORE emits one logits output per candidate token (plus
|
||||
@@ -43,6 +44,7 @@ enum server_task_type {
|
||||
SERVER_TASK_TYPE_SLOT_ERASE,
|
||||
SERVER_TASK_TYPE_GET_LORA,
|
||||
SERVER_TASK_TYPE_SET_LORA,
|
||||
@@ -809,7 +814,7 @@ index 5bedf1987..e6ca67a65 100644
|
||||
};
|
||||
|
||||
// TODO: change this to more generic "response_format" to replace the "format_response_*" in server-common
|
||||
@@ -202,6 +204,9 @@ struct server_task {
|
||||
@@ -225,6 +227,9 @@ struct server_task {
|
||||
// used by SERVER_TASK_TYPE_SET_LORA
|
||||
std::map<int, float> set_lora; // mapping adapter ID -> scale
|
||||
|
||||
@@ -819,15 +824,15 @@ index 5bedf1987..e6ca67a65 100644
|
||||
server_task() = default;
|
||||
|
||||
server_task(server_task_type type) : type(type) {}
|
||||
@@ -235,6 +240,7 @@ struct server_task {
|
||||
@@ -249,6 +254,7 @@ struct server_task {
|
||||
switch (type) {
|
||||
case SERVER_TASK_TYPE_COMPLETION:
|
||||
case SERVER_TASK_TYPE_INFILL:
|
||||
+ case SERVER_TASK_TYPE_TTS:
|
||||
return true;
|
||||
default:
|
||||
return false;
|
||||
@@ -494,5 +500,15 @@ struct server_task_result_embd : server_task_result {
|
||||
case SERVER_TASK_TYPE_DECISION:
|
||||
return !decision.labels.empty();
|
||||
@@ -521,5 +527,15 @@ struct server_task_result_embd : server_task_result {
|
||||
json to_json_oaicompat();
|
||||
};
|
||||
|
||||
|
||||
@@ -45,6 +45,8 @@ done
|
||||
|
||||
cp -r CMakeLists.txt llama.cpp/tools/grpc-server/
|
||||
cp -r grpc-server.cpp llama.cpp/tools/grpc-server/
|
||||
cp -r decision_compat.h llama.cpp/tools/grpc-server/
|
||||
cp -r decision_images.h llama.cpp/tools/grpc-server/
|
||||
# Model-load diagnostics (included by grpc-server.cpp) and their standalone
|
||||
# regression test.
|
||||
cp -r model_load_error.h llama.cpp/tools/grpc-server/
|
||||
|
||||
@@ -0,0 +1,75 @@
|
||||
# Native decision images
|
||||
|
||||
Run from the repository root after obtaining the pinned llama.cpp checkout:
|
||||
|
||||
```sh
|
||||
bash backend/cpp/llama-cpp/tests/verify-decision-images.sh
|
||||
```
|
||||
|
||||
Prerequisites: C++17, zlib and libjpeg development headers/libraries, Python 3
|
||||
with Pillow (fixture generation only). Ubuntu: `libjpeg-dev zlib1g-dev`;
|
||||
macOS: `brew install jpeg-turbo zlib`. CMake requires these libraries only when
|
||||
the checkout has native decisions; older forks retain their dependency guard.
|
||||
Both Docker builder paths install the packages in the shared compile stage,
|
||||
including builds using cached base images. Darwin CI installs the Homebrew
|
||||
packages. The llama.cpp packager collects the executable dependency closure, including
|
||||
zlib/libjpeg. Darwin uses its existing dylib collection.
|
||||
|
||||
This compiles the production validation helper with upstream stb, zlib and
|
||||
libjpeg. zlib requires a complete stream and valid Adler-32 within a fixed
|
||||
output budget. libjpeg decodes all scans with warnings treated as errors, so
|
||||
synthetic EOI recovery and short entropy scans cannot pass. Dimensions and
|
||||
aggregate pixels are checked before decoder pixel/coefficient allocation. It checks
|
||||
strict base64, MIME matching, count/byte/dimension/pixel limits, PNG decompression
|
||||
bombs, invalid Adler-32 with valid chunk CRC, baseline/progressive JPEG,
|
||||
missing EOI, truncated scans with appended EOI, embedded markers, truncation, chat-only collection, Anthropic normalization, empty-image
|
||||
text limits, capability combinations, and parity with Go's canonical limits.
|
||||
Fixtures are generated locally; no image or model downloads occur.
|
||||
|
||||
Build the native backend normally with `make backends/llama-cpp`. For a prepared
|
||||
CPU checkout and extracted distro gRPC dependencies, the existing adapter is:
|
||||
|
||||
```sh
|
||||
DEPS_ROOT=/path/to/deps CPU_BUILD=/path/to/llama.cpp/build-cpu \
|
||||
bash backend/cpp/llama-cpp/tests/build-decision-bridge.sh
|
||||
```
|
||||
|
||||
The adapter compiles the actual prepared grpc-server.cpp and links upstream CPU
|
||||
libraries. It is not a replacement implementation or mocked backend.
|
||||
|
||||
Start the resulting `grpc-server --addr=127.0.0.1:50061`, then run:
|
||||
|
||||
```sh
|
||||
PYTHONPATH=/path/to/build-decision-validation \
|
||||
python3 backend/cpp/llama-cpp/tests/decision-image-smoke.py \
|
||||
--model /models/OpenJev-Q4_K_M.gguf \
|
||||
--projector /models/mmproj-OpenJev-Q8_0.gguf \
|
||||
--fixtures backend/cpp/llama-cpp/llama.cpp/build-image-tests/fixtures.json
|
||||
```
|
||||
|
||||
The smoke uses CPU only, four threads, one slot, 8192 context, batch 512. It first
|
||||
loads without a projector and asserts explicit unsupported plus direct-RPC
|
||||
safety errors, then loads with the projector and compares red/blue probabilities.
|
||||
It requires existing weights and a checksum-verified projector; it downloads
|
||||
nothing. These are direct RPC tests, **not** public HTTP/router E2E evidence.
|
||||
|
||||
Check the production CMake dependency block and the non-decision fork guard:
|
||||
|
||||
```sh
|
||||
python3 backend/cpp/llama-cpp/tests/verify-image-build-wiring.py
|
||||
```
|
||||
|
||||
This focused check does not replace a full backend or platform build.
|
||||
|
||||
## CI gate
|
||||
|
||||
`.github/workflows/decision-images.yml` runs both verification commands on
|
||||
relevant pull requests and master pushes, or by manual dispatch. It installs
|
||||
C++17, Python/Pillow, CMake, zlib and libjpeg development dependencies and fetches
|
||||
only the two vendor headers at `LLAMA_VERSION` from the backend Makefile (not a
|
||||
floating upstream branch). No model, projector, GPU or full backend build is
|
||||
needed. The path filters include this workflow, the backend helper/tests/build
|
||||
files and upstream pin, Go limits, Dockerfiles and Darwin dependency setup.
|
||||
|
||||
This gate is separate from `backend/cpp/run-unit-tests.sh`: that stdlib-only
|
||||
runner discovers `*_test.cpp`, not the dependency-bearing `decision-images.cpp`.
|
||||
@@ -0,0 +1,96 @@
|
||||
# Native decision bridge validation
|
||||
|
||||
The stock dependency is pinned to `bed0a856606ee4a24a164066f73d2379447033f5`.
|
||||
`Score(question_type="systemone")` uses upstream decision tasks internally, not
|
||||
HTTP. Plain Score keeps its existing admission checks. Older dependencies without
|
||||
`server-decision.cpp` return gRPC `UNIMPLEMENTED` for this request type.
|
||||
|
||||
## Decision signature compatibility
|
||||
|
||||
This pin includes upstream Nimble support, in addition to OpenJev, Lev, Kev,
|
||||
and Laya. The native bridge forwards the complete parsed question collection
|
||||
when upstream's `fill_task` accepts it, as required by Nimble's schema framing.
|
||||
`decision_compat.h` detects the callable C++ signature at compile time; older
|
||||
native-decision forks still use their original single-question signature.
|
||||
Forks without native decision support retain the existing `UNIMPLEMENTED` guard.
|
||||
The standalone `decision_compat_test.cpp` checks both signatures and that the
|
||||
full collection is passed by reference, not replaced with a singleton.
|
||||
It is automatically discovered by `backend/cpp/run-unit-tests.sh`.
|
||||
|
||||
Signature and compile validation do not establish Nimble model accuracy or
|
||||
runtime support for every artifact. Nimble weights are not part of this test
|
||||
fixture. The official `ggml-org/Bespoke-Nimble-9B-v3-GGUF` model card declares
|
||||
CC-BY-NC-4.0; check its restrictions before deployment.
|
||||
|
||||
## CPU build
|
||||
|
||||
Use a fresh stock checkout at the pin in `backend/cpp/llama-cpp/llama.cpp`.
|
||||
Do not reuse a customized developer checkout. Apply patches once:
|
||||
|
||||
```sh
|
||||
cd backend/cpp/llama-cpp/llama.cpp
|
||||
git apply --check ../patches/0001-add-server-task-type-score.patch
|
||||
git apply ../patches/0001-add-server-task-type-score.patch
|
||||
git apply --check ../patches/0002-add-server-task-type-tts.patch
|
||||
git apply ../patches/0002-add-server-task-type-tts.patch
|
||||
cmake -S . -B build-cpu -DGGML_NATIVE=OFF -DLLAMA_OPENSSL=OFF \
|
||||
-DLLAMA_CURL=OFF -DBUILD_SHARED_LIBS=OFF
|
||||
cmake --build build-cpu --target llama-server -j2
|
||||
```
|
||||
|
||||
For the normal product build, start instead from a fresh unpatched checkout and
|
||||
run `make -C backend/cpp/llama-cpp grpc-server JOBS=2`; preparation applies the
|
||||
patches and stages the bridge. This requires CMake packages for gRPC, protobuf,
|
||||
and Abseil, plus `protoc` and `grpc_cpp_plugin`.
|
||||
|
||||
`build-decision-bridge.sh` is an alternative link validation for distro packages
|
||||
that lack `ProtobufConfig.cmake`. It uses the already patched CPU static libraries
|
||||
and the bridge source staged by `prepare.sh`. Do not apply patches twice when
|
||||
staging that source. Set `CPU_BUILD` to the absolute `build-cpu` directory,
|
||||
`OUT_DIR` to a scratch output directory, and optionally `DEPS_ROOT` to the root
|
||||
of **locally extracted** distro packages. It does not install or download anything.
|
||||
It also generates Python bindings, requiring `grpc_python_plugin`.
|
||||
|
||||
## Direct RPC smoke
|
||||
|
||||
The upstream test fixture `ggml-org/tinylaya-for-testing-gguf` has file
|
||||
`tinylaya-for-testing-Q8_0.gguf`, size **97,200,288 bytes**, SHA-256
|
||||
`a8b2b8f7fe6b7e10a884c55bf72362b0a8701e40dc3f332831d58246e5fa0b70`.
|
||||
This is a test model, not a production gallery recommendation.
|
||||
|
||||
Start the backend with cores disabled, using the matching library path when
|
||||
validating extracted distro dependencies:
|
||||
|
||||
```sh
|
||||
ulimit -c 0
|
||||
export LD_LIBRARY_PATH="$DEPS_ROOT/usr/lib/x86_64-linux-gnu${LD_LIBRARY_PATH:+:$LD_LIBRARY_PATH}"
|
||||
"$OUT_DIR/grpc-server" --addr 127.0.0.1:50051
|
||||
```
|
||||
|
||||
In another terminal (Python requires `grpcio` and `protobuf`):
|
||||
|
||||
```sh
|
||||
PYTHONPATH="$OUT_DIR" python3 backend/cpp/llama-cpp/tests/decision_smoke.py \
|
||||
--address 127.0.0.1:50051 --model "$MODEL_FILE"
|
||||
```
|
||||
|
||||
The smoke covers multiquestion choice/score/noul, normalized probabilities,
|
||||
positive input and explicit zero output usage, concurrent calls, invalid JSON,
|
||||
client cancellation/recovery, ordinary Score disabled/enabled, and missing
|
||||
metadata. It reloads models; do not share the backend with another test runner.
|
||||
The missing-metadata test makes a temporary equal-length metadata-key rename of
|
||||
the fixture and explicitly enables embeddings, as appropriate for this encoder.
|
||||
|
||||
Limits: immediate client cancellation does not prove interruption during active
|
||||
evaluation. The tiny fixture finishes too quickly for a deterministic timing-only
|
||||
assertion; a queue barrier or server-side instrumentation is needed for that gate.
|
||||
TTS is compiled and linked, not runtime-tested by this text-only fixture. Older
|
||||
pin compile validation does not establish every supported fork's full build.
|
||||
For bounded image/projector support and its separate runtime checks, see
|
||||
[README-decision-images.md](README-decision-images.md).
|
||||
|
||||
The metadata-stripped encoder with embeddings disabled and `-np 1` aborts
|
||||
in warmup at `llama-context.cpp`'s output-budget assertion on **clean unpatched**
|
||||
upstream at the pinned revision as well. This is a preexisting invalid-fixture
|
||||
configuration hazard, not a decision-dispatch or Score/TTS patch regression.
|
||||
Do not use that configuration as the missing-metadata test.
|
||||
@@ -0,0 +1,50 @@
|
||||
#!/usr/bin/env bash
|
||||
# SPDX-License-Identifier: MIT
|
||||
# CPU validation against an already prepared stock checkout. No downloads.
|
||||
# Run prepare.sh on a clean checkout first. DEPS_ROOT is an optional extracted
|
||||
# distro /usr tree parent, not a system installation. Output remains local.
|
||||
set -euo pipefail
|
||||
root=$(git rev-parse --show-toplevel)
|
||||
backend="$root/backend/cpp/llama-cpp"
|
||||
source="$backend/llama.cpp"
|
||||
out=${OUT_DIR:-"$source/build-decision-validation"}
|
||||
mkdir -p "$out"
|
||||
deps=${DEPS_ROOT:-/}
|
||||
export LD_LIBRARY_PATH="$deps/usr/lib/x86_64-linux-gnu${LD_LIBRARY_PATH:+:$LD_LIBRARY_PATH}"
|
||||
export PKG_CONFIG_SYSROOT_DIR="$deps"
|
||||
export PKG_CONFIG_PATH="$deps/usr/lib/x86_64-linux-gnu/pkgconfig${PKG_CONFIG_PATH:+:$PKG_CONFIG_PATH}"
|
||||
# Build upstream libraries without the optional grpc CMake subdirectory: distro
|
||||
# protobuf packages may expose FindProtobuf rather than ProtobufConfig.cmake.
|
||||
# An existing CPU build can be supplied to avoid rebuilding them.
|
||||
build=${CPU_BUILD:?set CPU_BUILD to the patched upstream CPU build directory}
|
||||
protoc -I "$root/backend" --cpp_out="$out" --grpc_out="$out" \
|
||||
--plugin=protoc-gen-grpc="$deps/usr/bin/grpc_cpp_plugin" "$root/backend/backend.proto"
|
||||
protoc -I "$root/backend" --python_out="$out" --grpc_out="$out" \
|
||||
--plugin=protoc-gen-grpc="$deps/usr/bin/grpc_python_plugin" "$root/backend/backend.proto"
|
||||
includes=(-I"$out" -I"$deps/usr/include")
|
||||
for dir in '' include common vendor ggml/include tools/mtmd tools/server; do
|
||||
includes+=(-I"$source/$dir")
|
||||
done
|
||||
cxx=${CXX:-g++}
|
||||
"$cxx" -O0 -std=c++17 -pthread "${includes[@]}" -c "$source/tools/grpc-server/grpc-server.cpp" -o "$out/grpc-server.o"
|
||||
for file in backend.pb backend.grpc.pb; do
|
||||
"$cxx" -O0 -std=c++17 -pthread "${includes[@]}" -c "$out/$file.cc" -o "$out/$file.o"
|
||||
done
|
||||
image_libs=()
|
||||
if [[ -f "$source/tools/server/server-decision.cpp" ]]; then
|
||||
image_libs=(-lz -ljpeg)
|
||||
fi
|
||||
libs=()
|
||||
for lib in common/llama-common common/llama-common-base tools/mtmd/mtmd src/llama ggml/src/ggml ggml/src/ggml-cpu ggml/src/ggml-base vendor/hash/vendor-hash vendor/cpp-httplib/cpp-httplib; do
|
||||
libs+=("$build/${lib%/*}/lib${lib##*/}.a")
|
||||
done
|
||||
# pkg-config emits a linker flag list, so intentional word splitting here.
|
||||
# shellcheck disable=SC2046
|
||||
"$cxx" -pthread "$out/grpc-server.o" "$out/backend.pb.o" "$out/backend.grpc.pb.o" "${libs[@]}" \
|
||||
-L"$deps/usr/lib/x86_64-linux-gnu" -Wl,-rpath-link,"$deps/usr/lib/x86_64-linux-gnu" \
|
||||
-lgrpc++_reflection $(pkg-config --libs grpc++) \
|
||||
-labsl_flags_parse -labsl_flags_usage -labsl_flags_usage_internal \
|
||||
-labsl_flags_commandlineflag -labsl_flags_commandlineflag_internal \
|
||||
-labsl_flags_config -labsl_flags_internal -labsl_flags_reflection \
|
||||
-labsl_flags_marshalling -lprotobuf "${image_libs[@]}" -ldl -lm -lgomp -o "$out/grpc-server"
|
||||
printf 'Built %s\n' "$out/grpc-server"
|
||||
@@ -0,0 +1,66 @@
|
||||
# SPDX-License-Identifier: MIT
|
||||
"""Direct RPC only; public API/router end-to-end tests are a separate gate."""
|
||||
import argparse
|
||||
import json
|
||||
import math
|
||||
import grpc
|
||||
import backend_pb2 as pb
|
||||
import backend_pb2_grpc as rpc
|
||||
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument('--address', default='127.0.0.1:50061')
|
||||
p.add_argument('--model', required=True)
|
||||
p.add_argument('--projector', required=True)
|
||||
p.add_argument('--fixtures', required=True)
|
||||
a = p.parse_args()
|
||||
f = json.load(open(a.fixtures))
|
||||
channel = grpc.insecure_channel(a.address, options=[('grpc.max_send_message_length', 20 << 20)])
|
||||
grpc.channel_ready_future(channel).result(timeout=20)
|
||||
s = rpc.BackendStub(channel)
|
||||
|
||||
def load(projector):
|
||||
r = s.LoadModel(pb.ModelOptions(ModelFile=a.model, MMProj=projector,
|
||||
ContextSize=8192, NBatch=512, Threads=4, NGPULayers=0,
|
||||
Options=['parallel:1']), timeout=600)
|
||||
assert r.success, r
|
||||
print('LOAD', 'vision' if projector else 'no projector', 'PASS', flush=True)
|
||||
|
||||
def body(image):
|
||||
return {'state': {}, 'images': [image], 'questions': {'color': {
|
||||
'type': 'choice', 'instructions': 'What is the dominant color of the image?',
|
||||
'criteria': {'red': None, 'blue': None}}}}
|
||||
|
||||
def score(b):
|
||||
return s.Score(pb.ScoreRequest(question_type='systemone', prompt=json.dumps(b)), timeout=600)
|
||||
|
||||
def reject(b, code):
|
||||
try:
|
||||
score(b)
|
||||
raise AssertionError('request unexpectedly accepted')
|
||||
except grpc.RpcError as e:
|
||||
assert e.code() == code, (e.code(), e.details())
|
||||
print('REJECT', code.name, e.details(), flush=True)
|
||||
|
||||
load('')
|
||||
reject(body(f['red']), grpc.StatusCode.UNIMPLEMENTED)
|
||||
for k in ['dimension', 'pixels', 'jpeg_dimension', 'jpeg_pixels']:
|
||||
reject(body(f[k]), grpc.StatusCode.RESOURCE_EXHAUSTED)
|
||||
for k in ['bomb', 'truncated', 'bad_crc', 'bad_adler', 'jpeg_missing_eoi',
|
||||
'jpeg_truncated_scan', 'jpeg_appended_eoi', 'jpeg_embedded_missing_eoi']:
|
||||
reject(body(f[k]), grpc.StatusCode.INVALID_ARGUMENT)
|
||||
reject(body('data:image/png;base64,AB=='), grpc.StatusCode.INVALID_ARGUMENT)
|
||||
reject(body('https://example.invalid/a.png'), grpc.StatusCode.INVALID_ARGUMENT)
|
||||
reject({'state': 'x' * (64 << 10)}, grpc.StatusCode.RESOURCE_EXHAUSTED)
|
||||
reject({'state': 'x' * (16 << 20)}, grpc.StatusCode.RESOURCE_EXHAUSTED)
|
||||
load(a.projector)
|
||||
results = {}
|
||||
for color in ['red', 'blue']:
|
||||
r = json.loads(score(body(f[color])).response_json)
|
||||
print('IMAGE', color, json.dumps(r), flush=True)
|
||||
assert r['usage']['input_tokens'] > 0 and r['usage']['output_tokens'] == 0, r
|
||||
probs = r['answers']['color']['probabilities']
|
||||
assert all(math.isfinite(v) for v in probs.values()) and abs(sum(probs.values())-1)<1e-4, r
|
||||
results[color] = probs
|
||||
assert results['red']['red'] > results['blue']['red'], results
|
||||
assert results['blue']['blue'] > results['red']['blue'], results
|
||||
print('CONTRASTING IMAGE EXECUTION PASS', flush=True)
|
||||
@@ -0,0 +1,49 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
#define STB_IMAGE_IMPLEMENTATION
|
||||
#include "stb/stb_image.h"
|
||||
#undef STB_IMAGE_IMPLEMENTATION
|
||||
#include "decision_images.h"
|
||||
#include <cassert>
|
||||
#include <iostream>
|
||||
#include <fstream>
|
||||
using namespace localai_decision;
|
||||
template<class F> void rejects(F f, bool large=false) {
|
||||
try { f(); assert(false); } catch (const image_error & e) { assert(e.too_large == large); }
|
||||
}
|
||||
int main(int argc, char ** argv) {
|
||||
assert(argc==2);
|
||||
std::ifstream input(argv[1]);
|
||||
json fixtures; input >> fixtures;
|
||||
assert(!supports_images(true, false));
|
||||
assert(!supports_images(false, true));
|
||||
assert(supports_images(true, true));
|
||||
auto check=[](json j) { return validate(j, j.dump().size()); };
|
||||
check(json{{"state", "text"}});
|
||||
for (auto images : {json(), json::array()}) {
|
||||
rejects([&]{check(json{{"state",std::string(text_bytes, 'x')},{"images",images}});},true);
|
||||
}
|
||||
rejects([&]{check(json{{"state",std::string(body_bytes, 'x')}});},true);
|
||||
rejects([&]{check(json{{"state",std::string(text_bytes, 'x')}});},true);
|
||||
rejects([&]{check(json{{"images", {"https://invalid/image.png"}}});});
|
||||
rejects([&]{check(json{{"images", {"data:image/png;base64,AAAA\n"}}});});
|
||||
rejects([&]{check(json{{"images", {"data:image/png;base64,AB=="}}});});
|
||||
rejects([&]{check(json{{"images", std::vector<std::string>(9,"x")}});},true);
|
||||
rejects([&]{check(json{{"images", {std::string(encoded_bytes+1,'x')}}});},true);
|
||||
rejects([&]{check(json{{"images", {"data:image/png;base64,"+std::string(12*1024*1024-24,'A')}}});},true);
|
||||
for (auto key : {"dimension", "pixels", "jpeg_dimension", "jpeg_pixels"}) rejects([&]{check(json{{"images",{fixtures[key]}}});},true);
|
||||
for (auto key : {"bomb", "truncated", "bad_crc", "bad_adler", "jpeg_missing_eoi", "jpeg_truncated_scan", "jpeg_appended_eoi", "jpeg_embedded_missing_eoi"}) rejects([&]{check(json{{"images",{fixtures[key]}}});});
|
||||
// Aggregate pixels reject even when each individual image fits.
|
||||
rejects([&]{check(json{{"images",{fixtures["aggregate"],fixtures["aggregate"]}}});},true);
|
||||
for (auto key : {"red", "blue", "jpeg", "jpeg_progressive", "jpeg_embedded_marker"}) assert(check(json{{"images",{fixtures[key]}}})==1);
|
||||
// Valid one-pixel PNG, and MIME mismatch.
|
||||
std::string png="data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNk+A8AAQUBAScY42YAAAAASUVORK5CYII=";
|
||||
assert(check(json{{"images",{png}}})==1);
|
||||
auto jpeg=png; jpeg.replace(5,9,"image/jpeg");
|
||||
rejects([&]{check(json{{"images",{jpeg}}});});
|
||||
json chat={{"state",{{"messages",json::array({{{"content",json::array({{{"type","image"},{"source",{{"type","base64"},{"media_type","image/png"},{"data",png.substr(22)}}}}})}}})}}}};
|
||||
assert(validate(chat,chat.dump().size())==1);
|
||||
assert(chat["state"]["messages"][0]["content"][0]["type"]=="image_url");
|
||||
// Domain state is not chat content.
|
||||
assert(check(json{{"state",{{"image_url","https://invalid"}}}})==0);
|
||||
std::cout << "decision image safety PASS\n";
|
||||
}
|
||||
@@ -0,0 +1,68 @@
|
||||
# SPDX-License-Identifier: MIT
|
||||
# Requires generated backend_pb2{,_grpc}.py on PYTHONPATH and a running backend.
|
||||
import argparse
|
||||
import json, grpc, os, concurrent.futures
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('--address', default='127.0.0.1:50051')
|
||||
parser.add_argument('--model', required=True)
|
||||
args = parser.parse_args()
|
||||
import backend_pb2 as pb
|
||||
import backend_pb2_grpc as rpc
|
||||
channel=grpc.insecure_channel(args.address)
|
||||
grpc.channel_ready_future(channel).result(timeout=10)
|
||||
s=rpc.BackendStub(channel)
|
||||
r=s.LoadModel(pb.ModelOptions(ModelFile=os.path.abspath(args.model),ContextSize=1024,NBatch=512,Threads=2,NGPULayers=0,Options=['parallel:2']),timeout=120)
|
||||
assert r.success,r
|
||||
print('LOAD PASS',flush=True)
|
||||
body={'model':'tinylaya','state':'I was charged twice for my order last week and nobody has replied.','questions':{
|
||||
'route':{'type':'choice','instructions':'Which team should handle this?','criteria':{'billing':'payments and refunds','shipping':None,'technical':None}},
|
||||
'urgency':{'type':'score','instructions':'How urgent is this?','criteria':['can wait','this week','today','right now']},
|
||||
'angry':{'type':'noul','instructions':'Is the customer angry?'}}}
|
||||
def run():
|
||||
r=json.loads(s.Score(pb.ScoreRequest(question_type='systemone',prompt=json.dumps(body)),timeout=60).response_json)
|
||||
assert r['usage']['input_tokens']>0 and r['usage']['output_tokens']==0,r
|
||||
a=r['answers']; assert set(a)==set(body['questions']),r
|
||||
assert 0<=a['angry']['noul']<=1,r
|
||||
for k in ['route','urgency']: assert abs(sum(a[k]['probabilities'].values())-1)<1e-4,r
|
||||
assert a['route']['choice']==max(a['route']['probabilities'],key=a['route']['probabilities'].get),r
|
||||
assert abs(a['urgency']['score']-sum(int(k)*v for k,v in a['urgency']['probabilities'].items()))<1e-4,r
|
||||
return r
|
||||
print('MULTIQUESTION',json.dumps(run()),flush=True)
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=2) as pool: list(pool.map(lambda _:run(),range(4)))
|
||||
print('CONCURRENCY PASS',flush=True)
|
||||
for payload in ['{',json.dumps({'state':'x','questions':{}})]:
|
||||
try:s.Score(pb.ScoreRequest(question_type='systemone',prompt=payload),timeout=30);raise AssertionError('expected invalid')
|
||||
except grpc.RpcError as e:assert e.code()==grpc.StatusCode.INVALID_ARGUMENT,e
|
||||
print('INVALID PASS',flush=True)
|
||||
f=s.Score.future(pb.ScoreRequest(question_type='systemone',prompt=json.dumps(body)),timeout=60);f.cancel()
|
||||
try:f.result();raise AssertionError('expected cancelled')
|
||||
except grpc.FutureCancelledError:pass
|
||||
run();print('CANCEL AND RECOVERY PASS',flush=True)
|
||||
try:s.Score(pb.ScoreRequest(prompt='hello',candidates=['world']),timeout=30);raise AssertionError('score unexpectedly enabled')
|
||||
except grpc.RpcError as e:assert e.code()==grpc.StatusCode.FAILED_PRECONDITION,e
|
||||
print('PLAIN SCORE DISABLED GUARD PASS',flush=True)
|
||||
|
||||
r=s.LoadModel(pb.ModelOptions(ModelFile=os.path.abspath(args.model),ContextSize=1024,NBatch=512,Threads=2,EnableScore=True,Options=['parallel:2']),timeout=120)
|
||||
assert r.success,r
|
||||
r=s.Score(pb.ScoreRequest(prompt='Hello',candidates=[' world',' there']),timeout=30)
|
||||
assert len(r.candidates)==2 and all(c.num_tokens>0 for c in r.candidates),r
|
||||
assert all(__import__('math').isfinite(c.log_prob) for c in r.candidates),r
|
||||
print('PLAIN SCORE ENABLED PASS',flush=True)
|
||||
|
||||
# This fixture is an encoder: keep embeddings enabled when removing decision
|
||||
# metadata, otherwise upstream warmup can exceed the one-slot output budget.
|
||||
data = Path(args.model).read_bytes()
|
||||
assert data.count(b'.decision.type') == 1, 'expected the tinylaya test fixture'
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
model = Path(directory) / 'no-decision.gguf'
|
||||
model.write_bytes(data.replace(b'.decision.type', b'.disabled.type'))
|
||||
r=s.LoadModel(pb.ModelOptions(ModelFile=str(model),ContextSize=1024,NBatch=512,Threads=2,Embeddings=True),timeout=120)
|
||||
assert r.success,r
|
||||
try:
|
||||
s.Score(pb.ScoreRequest(question_type='systemone',prompt=json.dumps(body)),timeout=30)
|
||||
raise AssertionError('missing decision metadata accepted')
|
||||
except grpc.RpcError as e:
|
||||
assert e.code()==grpc.StatusCode.UNIMPLEMENTED,e
|
||||
print('MISSING DECISION METADATA PASS',flush=True)
|
||||
@@ -0,0 +1,66 @@
|
||||
# SPDX-License-Identifier: MIT
|
||||
"""Generate small compressed fixtures, including hostile IHDR/IDAT combinations."""
|
||||
import base64
|
||||
import json
|
||||
import io
|
||||
from PIL import Image
|
||||
import struct
|
||||
import sys
|
||||
import zlib
|
||||
|
||||
def png(w, h, pixels):
|
||||
def chunk(kind, data):
|
||||
return struct.pack('>I', len(data)) + kind + data + struct.pack('>I', zlib.crc32(kind + data))
|
||||
return b'\x89PNG\r\n\x1a\n' + chunk(b'IHDR', struct.pack('>IIBBBBB', w, h, 8, 2, 0, 0, 0)) + chunk(b'IDAT', zlib.compress(pixels)) + chunk(b'IEND', b'')
|
||||
|
||||
def url(raw):
|
||||
return 'data:image/png;base64,' + base64.b64encode(raw).decode()
|
||||
|
||||
fixtures = {
|
||||
'aggregate': url(png(3000, 3000, (b'\0' * 9001)*3000)),
|
||||
'dimension': url(png(4097, 1, b'\0'*12292)),
|
||||
'pixels': url(png(4096, 4096, b'\0')),
|
||||
'bomb': url(png(1, 1, b'\0'*1000000)),
|
||||
'truncated': url(png(1, 1, b'\0'*4)[:-15]),
|
||||
'red': url(png(64, 64, (b'\0'+b'\xff\0\0'*64)*64)),
|
||||
'blue': url(png(64, 64, (b'\0'+b'\0\0\xff'*64)*64)),
|
||||
}
|
||||
bad = bytearray(png(1, 1, b'\0'*4))
|
||||
bad[29] ^= 1
|
||||
fixtures['bad_crc'] = url(bad)
|
||||
# CRC-valid IDAT with an invalid zlib Adler-32 checksum.
|
||||
bad = bytearray(base64.b64decode(fixtures['red'].split(',')[1]))
|
||||
pos = bad.index(b'IDAT')
|
||||
n = struct.unpack('>I', bad[pos-4:pos])[0]
|
||||
bad[pos+4+n-1] ^= 1
|
||||
bad[pos+4+n:pos+8+n] = struct.pack('>I', zlib.crc32(bad[pos:pos+4+n]))
|
||||
try:
|
||||
zlib.decompress(bad[pos+4:pos+4+n])
|
||||
raise AssertionError('invalid Adler-32 accepted')
|
||||
except zlib.error:
|
||||
pass
|
||||
fixtures['bad_adler'] = url(bad)
|
||||
|
||||
def jpeg(w, h, progressive=False):
|
||||
out = io.BytesIO()
|
||||
Image.new('RGB', (w, h), 'red').save(out, format='JPEG', progressive=progressive)
|
||||
return out.getvalue()
|
||||
|
||||
def jpg_url(raw):
|
||||
return 'data:image/jpeg;base64,' + base64.b64encode(raw).decode()
|
||||
|
||||
jpg = jpeg(64, 64)
|
||||
for key, raw in {
|
||||
'jpeg': jpg,
|
||||
'jpeg_progressive': jpeg(64, 64, True),
|
||||
# EOI inside a comment is data, not an end marker.
|
||||
'jpeg_embedded_marker': jpg[:2] + b'\xff\xfe\x00\x04\xff\xd9' + jpg[2:],
|
||||
'jpeg_missing_eoi': jpg[:-2],
|
||||
'jpeg_truncated_scan': jpg[:-30],
|
||||
'jpeg_appended_eoi': jpg[:-30] + b'\xff\xd9',
|
||||
'jpeg_embedded_missing_eoi': jpg[:2] + b'\xff\xfe\x00\x04\xff\xd9' + jpg[2:-2],
|
||||
'jpeg_dimension': jpeg(4097, 1),
|
||||
'jpeg_pixels': jpeg(4096, 4096),
|
||||
}.items():
|
||||
fixtures[key] = jpg_url(raw)
|
||||
json.dump(fixtures, open(sys.argv[1], 'w'))
|
||||
+22
@@ -0,0 +1,22 @@
|
||||
#!/usr/bin/env bash
|
||||
# SPDX-License-Identifier: MIT
|
||||
set -euo pipefail
|
||||
root=$(git rev-parse --show-toplevel)
|
||||
b="$root/backend/cpp/llama-cpp"
|
||||
out="$b/llama.cpp/build-image-tests"
|
||||
mkdir -p "$out"
|
||||
${CXX:-g++} -std=c++17 -Wall -Wextra -I"$b" -I"$b/llama.cpp/vendor" "$b/tests/decision-images.cpp" -lz -ljpeg -o "$out/decision-images"
|
||||
python3 "$b/tests/image-fixtures.py" "$out/fixtures.json"
|
||||
"$out/decision-images" "$out/fixtures.json"
|
||||
# Keep the native boundary in lockstep with canonical Go limits.
|
||||
python3 - "$root" <<'PY'
|
||||
import pathlib, re, sys
|
||||
root=pathlib.Path(sys.argv[1])
|
||||
go=(root/'core/systemone/images.go').read_text()
|
||||
cpp=(root/'backend/cpp/llama-cpp/decision_images.h').read_text()
|
||||
for g,c in [('MaxImages','max_images'),('MaxImageDecodedBytes','decoded_bytes'),('MaxImageEncodedBytes','encoded_bytes'),('MaxImageBodyBytes','body_bytes'),('MaxImageDimension','max_dimension'),('MaxImagePixels','max_pixels'),('MaxResponseBytes','text_bytes')]:
|
||||
gv=re.search(r'\b'+g+r'\s*=\s*([^\n]+)',go)[1]
|
||||
cv=re.search(r'\b'+c+r'\s*=\s*([^;]+)',cpp)[1]
|
||||
assert eval(gv)==eval(cv),(g,c)
|
||||
print('Go/native limit parity PASS')
|
||||
PY
|
||||
@@ -0,0 +1,38 @@
|
||||
# SPDX-License-Identifier: MIT
|
||||
"""Exercise the production CMake decoder dependency block, including old forks."""
|
||||
import pathlib
|
||||
import subprocess
|
||||
import tempfile
|
||||
|
||||
backend = pathlib.Path(__file__).resolve().parents[1]
|
||||
cmake = (backend / 'CMakeLists.txt').read_text()
|
||||
start = cmake.index('if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/../server/server-decision.cpp")')
|
||||
block = cmake[start:cmake.index('endif()', start) + len('endif()')]
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
root = pathlib.Path(tmp)
|
||||
source = root / 'grpc-server'
|
||||
source.mkdir()
|
||||
(root / 'server').mkdir()
|
||||
(source / 'main.cpp').write_text('int main() {}\n')
|
||||
(source / 'CMakeLists.txt').write_text('''cmake_minimum_required(VERSION 3.15)
|
||||
project(decoder_wiring LANGUAGES CXX)
|
||||
set(TARGET grpc-server)
|
||||
add_executable(${TARGET} main.cpp)
|
||||
''' + block + '''
|
||||
get_target_property(libs ${TARGET} LINK_LIBRARIES)
|
||||
if(EXPECT_DECODERS)
|
||||
if(NOT "${libs}" STREQUAL "ZLIB::ZLIB;JPEG::JPEG")
|
||||
message(FATAL_ERROR "Missing decoder links: ${libs}")
|
||||
endif()
|
||||
elseif(libs)
|
||||
message(FATAL_ERROR "Old fork acquired decoder dependencies: ${libs}")
|
||||
endif()
|
||||
''')
|
||||
subprocess.run(['cmake', '-S', str(source), '-B', str(root / 'fork'),
|
||||
'-DCMAKE_DISABLE_FIND_PACKAGE_ZLIB=TRUE',
|
||||
'-DCMAKE_DISABLE_FIND_PACKAGE_JPEG=TRUE'], check=True)
|
||||
(root / 'server/server-decision.cpp').touch()
|
||||
subprocess.run(['cmake', '-S', str(source), '-B', str(root / 'native'),
|
||||
'-DEXPECT_DECODERS=ON'], check=True)
|
||||
subprocess.run(['cmake', '--build', str(root / 'native')], check=True)
|
||||
print('Production CMake decoder links and old-fork guard PASS')
|
||||
@@ -474,8 +474,8 @@ func (a *Application) MITMHostOwners() map[string]string {
|
||||
}
|
||||
|
||||
// RouterDecisions returns the routing decision store. nil when stats
|
||||
// are disabled (--disable-stats); the RouteModel middleware skips the
|
||||
// log write in that case but still rewrites requests.
|
||||
// are disabled (--disable-stats), unless WithRouterDecisionLog explicitly
|
||||
// retains the log. A nil store skips logging but still rewrites requests.
|
||||
func (a *Application) RouterDecisions() router.DecisionStore {
|
||||
return a.routerDecisions
|
||||
}
|
||||
|
||||
@@ -0,0 +1,10 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package application
|
||||
|
||||
import "github.com/mudler/LocalAI/core/backend"
|
||||
|
||||
// DecisionRunner lazily resolves the named model on every native decision call.
|
||||
// Construction does not load weights or guess an unavailable model's usecase.
|
||||
func (a *Application) DecisionRunner(modelName string) backend.DecisionRunner {
|
||||
return backend.NewDecisionRunner(modelName, a.adapterConfig, a.modelLoader, a.applicationConfig)
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package application
|
||||
|
||||
import (
|
||||
"context"
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/services/monitoring"
|
||||
"github.com/mudler/LocalAI/pkg/system"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"go.opentelemetry.io/otel/metric"
|
||||
"go.opentelemetry.io/otel/metric/noop"
|
||||
)
|
||||
|
||||
var _ = Describe("optional startup services", func() {
|
||||
DescribeTable("router log without billing stats", func(enabled bool) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
DeferCleanup(cancel)
|
||||
state, err := system.GetSystemState(system.WithModelPath(GinkgoT().TempDir()), system.WithBackendPath(GinkgoT().TempDir()))
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
app, err := New(config.WithContext(ctx), config.WithSystemState(state), config.DisableMetricsEndpoint, config.WithDisableLocalAIAssistant(true), config.WithDisableStats(true), config.WithRouterDecisionLog(enabled))
|
||||
if app != nil {
|
||||
DeferCleanup(func() { Expect(app.Shutdown()).To(Succeed()) })
|
||||
}
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(app.StatsRecorder()).To(BeNil())
|
||||
if enabled {
|
||||
Expect(app.RouterDecisions()).NotTo(BeNil())
|
||||
} else {
|
||||
Expect(app.RouterDecisions()).To(BeNil())
|
||||
}
|
||||
}, Entry("retained by explicit opt-in", true), Entry("disabled by default", false))
|
||||
})
|
||||
|
||||
type registrationMeter struct {
|
||||
metric.Meter
|
||||
gauges int
|
||||
}
|
||||
|
||||
func (m *registrationMeter) Int64ObservableGauge(name string, opts ...metric.Int64ObservableGaugeOption) (metric.Int64ObservableGauge, error) {
|
||||
m.gauges++
|
||||
return m.Meter.Int64ObservableGauge(name, opts...)
|
||||
}
|
||||
|
||||
var _ = Describe("optional failover metrics", func() {
|
||||
It("registers only on the application's enabled meter", func() {
|
||||
meter := ®istrationMeter{Meter: noop.NewMeterProvider().Meter("test")}
|
||||
app := &Application{applicationConfig: &config.ApplicationConfig{DisableMetrics: true}, metricsService: &monitoring.LocalAIMetricsService{Meter: meter}}
|
||||
app.registerFailoverMetrics()
|
||||
Expect(meter.gauges).To(BeZero())
|
||||
app.applicationConfig.DisableMetrics = false
|
||||
app.registerFailoverMetrics()
|
||||
Expect(meter.gauges).To(Equal(1))
|
||||
})
|
||||
})
|
||||
@@ -5,7 +5,9 @@ import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"encoding/json"
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/pkg/model"
|
||||
"github.com/mudler/LocalAI/pkg/system"
|
||||
|
||||
@@ -62,6 +64,22 @@ var _ = Describe("router_factories lazy config resolution", func() {
|
||||
Expect(os.Remove(filepath.Join(tmpDir, name+".yaml"))).To(Succeed())
|
||||
}
|
||||
|
||||
Context("DecisionRunner", func() {
|
||||
It("constructs lazily and observes subsequent model removal", func() {
|
||||
runner := app.DecisionRunner("dec-test")
|
||||
Expect(runner).NotTo(BeNil())
|
||||
req := &schema.SystemOneRequest{State: json.RawMessage(`"text"`), Questions: map[string]schema.SystemOneQuestion{"q": {Type: "noul"}}}
|
||||
_, err := runner.Decide(context.Background(), req)
|
||||
Expect(err).To(MatchError(ContainSubstring("no longer available")))
|
||||
writeCfg("dec-test", "vllm-cpp")
|
||||
_, err = runner.Decide(context.Background(), req)
|
||||
Expect(err).To(MatchError(ContainSubstring("explicitly declare")))
|
||||
removeCfg("dec-test")
|
||||
_, err = runner.Decide(context.Background(), req)
|
||||
Expect(err).To(MatchError(ContainSubstring("no longer available")))
|
||||
})
|
||||
})
|
||||
|
||||
Context("Embedder", func() {
|
||||
It("returns nil at construction for an unknown model", func() {
|
||||
Expect(app.Embedder("missing")).To(BeNil())
|
||||
|
||||
@@ -243,8 +243,9 @@ func New(opts ...config.AppOption) (*Application, error) {
|
||||
|
||||
// Wire the routing decision log. Always-on when stats are enabled —
|
||||
// the per-router admin page reads this as the live activity feed
|
||||
// and as input to drift checks for subsystem 5.
|
||||
if !options.DisableStats {
|
||||
// and as input to drift checks for subsystem 5. Embedders may retain this
|
||||
// bounded log independently without enabling billing stats.
|
||||
if !options.DisableStats || options.RouterDecisionLog {
|
||||
application.routerDecisions = router.NewMemoryDecisionStore(0)
|
||||
}
|
||||
// Process-wide classifier cache shared across all route middlewares so
|
||||
@@ -569,7 +570,7 @@ func New(opts ...config.AppOption) (*Application, error) {
|
||||
// Start the failover scheduler: it syncs chains from config, runs
|
||||
// liveness/recovery probes and dwell-based fail-back. Run is the only
|
||||
// caller of Sync in production so onWarm callbacks stay ordered.
|
||||
failover.RegisterMetrics(application.failoverManager)
|
||||
application.registerFailoverMetrics()
|
||||
go application.failoverManager.Run(options.Context)
|
||||
|
||||
// Watch the configuration directory
|
||||
@@ -746,3 +747,12 @@ func migrateDataFiles(srcDir, dstDir string) {
|
||||
xlog.Info("Data migration complete", "from", srcDir, "to", dstDir)
|
||||
}
|
||||
}
|
||||
|
||||
// registerFailoverMetrics must not retain a manager on the global provider
|
||||
// when this application has metrics disabled.
|
||||
func (a *Application) registerFailoverMetrics() {
|
||||
if a.applicationConfig.DisableMetrics || a.metricsService == nil {
|
||||
return
|
||||
}
|
||||
failover.RegisterMetrics(a.failoverManager, a.metricsService.Meter)
|
||||
}
|
||||
@@ -0,0 +1,24 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package backend
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
var _ = Describe("decision projector load options", func() {
|
||||
It("resolves the YAML projector under the model directory for the native loader", func() {
|
||||
var cfg config.ModelConfig
|
||||
Expect(yaml.Unmarshal([]byte("name: openjev\nbackend: llama-cpp\nthreads: 1\nmmproj: mmproj-OpenJev-Q8_0.gguf\nparameters:\n model: OpenJev-Q4_K_M.gguf\n"), &cfg)).To(Succeed())
|
||||
opts := grpcModelOpts(cfg, "/models")
|
||||
Expect(opts.MMProj).To(Equal(filepath.Join("/models", "mmproj-OpenJev-Q8_0.gguf")))
|
||||
})
|
||||
It("does not invent a projector for text-only configurations", func() {
|
||||
threads := 1
|
||||
Expect(grpcModelOpts(config.ModelConfig{Threads: &threads}, "/models").MMProj).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,137 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package backend
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/systemone"
|
||||
"github.com/mudler/LocalAI/pkg/model"
|
||||
)
|
||||
|
||||
// DecisionRunner is the typed native decision transport. It never uses HTTP,
|
||||
// generation or NER as a substitute for the model's decision pipeline.
|
||||
type DecisionRunner interface {
|
||||
Decide(context.Context, *schema.SystemOneRequest) (*schema.SystemOneResponse, error)
|
||||
}
|
||||
|
||||
// NewDecisionRunner binds a name, not a snapshot: configuration is resolved on
|
||||
// every call so edits and removal cannot leave a cached adapter using old policy.
|
||||
func NewDecisionRunner(name string, lookup func(string) *config.ModelConfig, loader *model.ModelLoader, app *config.ApplicationConfig) DecisionRunner {
|
||||
return &decisionRunner{modelName: name, lookup: lookup, load: func(body string, cfg config.ModelConfig) (func(context.Context) (string, error), error) {
|
||||
return ModelSystemOne(body, loader, cfg, app)
|
||||
}}
|
||||
}
|
||||
|
||||
type decisionRunner struct {
|
||||
modelName string
|
||||
lookup func(string) *config.ModelConfig
|
||||
load func(string, config.ModelConfig) (func(context.Context) (string, error), error)
|
||||
}
|
||||
|
||||
// Load has no context API. Bound abandoned work process-wide (including across
|
||||
// registry replacements), retaining the permit until the underlying operation
|
||||
// finishes. Cancellation releases the caller, not the loader or backend itself.
|
||||
// Saturation fails promptly rather than spawning an unbounded goroutine queue.
|
||||
const maxDecisionOperations = 8
|
||||
|
||||
var decisionOperations = make(chan struct{}, maxDecisionOperations)
|
||||
|
||||
func (r *decisionRunner) Decide(ctx context.Context, req *schema.SystemOneRequest) (*schema.SystemOneResponse, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
release, err := systemone.AcquireAdmission(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
transferred := false
|
||||
defer func() {
|
||||
if !transferred {
|
||||
release()
|
||||
}
|
||||
}()
|
||||
if req == nil {
|
||||
return nil, systemone.ValidateRequest(nil)
|
||||
}
|
||||
if req.Model != "" && req.Model != r.modelName {
|
||||
return nil, fmt.Errorf("decision request model does not match bound model")
|
||||
}
|
||||
copied := *req
|
||||
copied.Model = r.modelName
|
||||
if err := systemone.ValidateRequest(&copied); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Marshal before launching work so no goroutine retains caller-owned data.
|
||||
body, err := json.Marshal(&copied)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
select {
|
||||
case decisionOperations <- struct{}{}:
|
||||
default:
|
||||
return nil, fmt.Errorf("native decision operation capacity reached")
|
||||
}
|
||||
type result struct {
|
||||
response *schema.SystemOneResponse
|
||||
err error
|
||||
}
|
||||
done := make(chan result, 1)
|
||||
transferred = true
|
||||
go func() {
|
||||
defer release()
|
||||
defer func() { <-decisionOperations }()
|
||||
if err := ctx.Err(); err != nil {
|
||||
done <- result{err: err}
|
||||
return
|
||||
}
|
||||
cfg := r.lookup(r.modelName)
|
||||
if cfg == nil {
|
||||
done <- result{err: fmt.Errorf("decision model %q no longer available", r.modelName)}
|
||||
return
|
||||
}
|
||||
if err := systemone.ValidateDecisionModel(*cfg); err != nil {
|
||||
done <- result{err: err}
|
||||
return
|
||||
}
|
||||
fn, err := r.load(string(body), *cfg)
|
||||
if err == nil {
|
||||
err = ctx.Err()
|
||||
}
|
||||
if err != nil {
|
||||
done <- result{err: err}
|
||||
return
|
||||
}
|
||||
raw, err := fn(ctx)
|
||||
if err != nil {
|
||||
done <- result{err: err}
|
||||
return
|
||||
}
|
||||
if len(raw) > systemone.MaxResponseBytes {
|
||||
done <- result{err: fmt.Errorf("decision response exceeds 64 KiB")}
|
||||
return
|
||||
}
|
||||
var response schema.SystemOneResponse
|
||||
if err := json.Unmarshal([]byte(raw), &response); err != nil {
|
||||
done <- result{err: fmt.Errorf("invalid decision response JSON")}
|
||||
return
|
||||
}
|
||||
if response.Answers == nil {
|
||||
done <- result{err: fmt.Errorf("decision response has no answers")}
|
||||
return
|
||||
}
|
||||
done <- result{response: &response}
|
||||
}()
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
case res := <-done:
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return res.response, res.err
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,138 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package backend
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/systemone"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("internal decision runner", func() {
|
||||
var runner *decisionRunner
|
||||
var calls atomic.Int32
|
||||
var cfg config.ModelConfig
|
||||
var req *schema.SystemOneRequest
|
||||
BeforeEach(func() {
|
||||
Eventually(func() int { return len(decisionOperations) }).Should(BeZero())
|
||||
calls.Store(0)
|
||||
cfg = config.ModelConfig{}
|
||||
cfg.Name = "native"
|
||||
cfg.Backend = "vllm-cpp"
|
||||
flags := config.FLAG_DECISIONS
|
||||
cfg.KnownUsecases = &flags
|
||||
req = &schema.SystemOneRequest{State: json.RawMessage(`"text"`), Questions: map[string]schema.SystemOneQuestion{"q": {Type: "noul"}}}
|
||||
runner = &decisionRunner{modelName: "native", lookup: func(string) *config.ModelConfig { return &cfg }, load: func(body string, c config.ModelConfig) (func(context.Context) (string, error), error) {
|
||||
defer GinkgoRecover()
|
||||
calls.Add(1)
|
||||
var sent schema.SystemOneRequest
|
||||
Expect(json.Unmarshal([]byte(body), &sent)).To(Succeed())
|
||||
Expect(sent.Model).To(Equal("native"))
|
||||
return func(context.Context) (string, error) { return `{"answers":{"q":{"type":"noul","noul":0.75}}}`, nil }, nil
|
||||
}}
|
||||
})
|
||||
|
||||
It("rejects admission saturation before inspecting caller data", func() {
|
||||
var releases []func()
|
||||
defer func() {
|
||||
for _, r := range releases {
|
||||
r()
|
||||
}
|
||||
}()
|
||||
for i := 0; i < systemone.MaxAdmissions; i++ {
|
||||
r, err := systemone.AcquireAdmission(context.Background())
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
releases = append(releases, r)
|
||||
}
|
||||
req.State = json.RawMessage(`invalid`)
|
||||
_, err := runner.Decide(context.Background(), req)
|
||||
Expect(err).To(MatchError(systemone.ErrAdmissionCapacity))
|
||||
Expect(calls.Load()).To(BeZero())
|
||||
})
|
||||
It("uses a named internal call and preserves numeric noul without probabilities", func() {
|
||||
result, err := runner.Decide(context.Background(), req)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(*result.Answers["q"].Noul).To(Equal(.75))
|
||||
Expect(req.Model).To(BeEmpty())
|
||||
Expect(calls.Load()).To(Equal(int32(1)))
|
||||
})
|
||||
It("revalidates model config and bounds requests before loading", func() {
|
||||
cfg.KnownUsecases = nil
|
||||
_, err := runner.Decide(context.Background(), req)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(calls.Load()).To(BeZero())
|
||||
flags := config.FLAG_DECISIONS
|
||||
cfg.KnownUsecases = &flags
|
||||
req.State, _ = json.Marshal(strings.Repeat("a", 65536))
|
||||
_, err = runner.Decide(context.Background(), req)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(calls.Load()).To(BeZero())
|
||||
})
|
||||
It("returns cancellation while an uncancellable load is still running", func() {
|
||||
started := make(chan struct{})
|
||||
finish := make(chan struct{})
|
||||
defer close(finish)
|
||||
runner.load = func(string, config.ModelConfig) (func(context.Context) (string, error), error) {
|
||||
close(started)
|
||||
<-finish
|
||||
return func(context.Context) (string, error) { calls.Add(1); return `{}`, nil }, nil
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
done := make(chan error, 1)
|
||||
go func() { _, err := runner.Decide(ctx, req); done <- err }()
|
||||
Eventually(started, time.Second).Should(BeClosed())
|
||||
cancel()
|
||||
Eventually(done, time.Second).Should(Receive(Equal(context.Canceled)))
|
||||
Expect(calls.Load()).To(BeZero())
|
||||
})
|
||||
It("bounds abandoned loads across new runner instances", func() {
|
||||
Eventually(func() int { return len(decisionOperations) }).Should(BeZero())
|
||||
finish := make(chan struct{})
|
||||
defer close(finish)
|
||||
started := make(chan struct{}, cap(decisionOperations))
|
||||
load := func(string, config.ModelConfig) (func(context.Context) (string, error), error) {
|
||||
started <- struct{}{}
|
||||
<-finish
|
||||
return func(context.Context) (string, error) { return `{}`, nil }, nil
|
||||
}
|
||||
// Cancellation does not free the global permit until Load actually ends.
|
||||
for i := 0; i < cap(decisionOperations); i++ {
|
||||
copyRunner := *runner
|
||||
copyRunner.load = load
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
done := make(chan error, 1)
|
||||
go func() { _, err := copyRunner.Decide(ctx, req); done <- err }()
|
||||
Eventually(started).Should(Receive())
|
||||
cancel()
|
||||
Eventually(done).Should(Receive(Equal(context.Canceled)))
|
||||
}
|
||||
_, err := runner.Decide(context.Background(), req)
|
||||
Expect(err).To(MatchError(ContainSubstring("capacity reached")))
|
||||
Expect(calls.Load()).To(BeZero())
|
||||
})
|
||||
|
||||
It("does not load for an already cancelled parent", func() {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
_, err := runner.Decide(ctx, req)
|
||||
Expect(err).To(MatchError(context.Canceled))
|
||||
Expect(calls.Load()).To(BeZero())
|
||||
})
|
||||
It("rejects malformed and oversized responses", func() {
|
||||
for _, body := range []string{`null`, `[]`, `{}`, `{"answers":null}`, `{"answers":{}} trailing`, `{"answers":{"q":{"type":"noul","noul":"0.5"}}}`, `{"answers":{"q":{"type":"noul","noul":true}}}`, `{"answers":{"q":{"type":"noul","noul":1e999}}}`, strings.Repeat("x", 65537)} {
|
||||
runner.load = func(string, config.ModelConfig) (func(context.Context) (string, error), error) {
|
||||
return func(context.Context) (string, error) { return body, nil }, nil
|
||||
}
|
||||
_, err := runner.Decide(context.Background(), req)
|
||||
Expect(err).To(HaveOccurred())
|
||||
}
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,74 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package backend_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"github.com/mudler/LocalAI/core/backend"
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/trace"
|
||||
grpcPkg "github.com/mudler/LocalAI/pkg/grpc"
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
"github.com/mudler/LocalAI/pkg/model"
|
||||
"github.com/mudler/LocalAI/pkg/system"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
ggrpc "google.golang.org/grpc"
|
||||
)
|
||||
|
||||
type nativeDecisionBackend struct {
|
||||
grpcPkg.Backend
|
||||
request *pb.ScoreRequest
|
||||
fail bool
|
||||
}
|
||||
|
||||
func (b *nativeDecisionBackend) HealthCheck(context.Context) (bool, error) { return true, nil }
|
||||
func (b *nativeDecisionBackend) IsBusy() bool { return false }
|
||||
|
||||
func (b *nativeDecisionBackend) Score(_ context.Context, r *pb.ScoreRequest, _ ...ggrpc.CallOption) (*pb.ScoreResponse, error) {
|
||||
b.request = r
|
||||
if b.fail {
|
||||
return nil, fmt.Errorf("backend echoed secret-text")
|
||||
}
|
||||
return &pb.ScoreResponse{ResponseJson: `{"answers":{"q":{"type":"noul","noul":0.8}}}`}, nil
|
||||
}
|
||||
|
||||
var _ = Describe("native decision transport", func() {
|
||||
It("uses ModelSystemOne Score directly and never records prompt or echoed error", func() {
|
||||
rec := &nativeDecisionBackend{}
|
||||
state := &system.SystemState{}
|
||||
loader := model.NewModelLoader(state)
|
||||
loader.SetModelRouter(func(_ context.Context, id string, _, _, _, _ string, _ *pb.ModelOptions, _ bool) (*model.Model, error) {
|
||||
return model.NewModelWithClient(id, "test://native", rec), nil
|
||||
})
|
||||
cfg := identityModelCfg()
|
||||
cfg.Backend = "vllm-cpp"
|
||||
cfg.Model = "native-test.gguf"
|
||||
flags := config.FLAG_DECISIONS
|
||||
cfg.KnownUsecases = &flags
|
||||
app := config.NewApplicationConfig(config.WithSystemState(state))
|
||||
app.EnableTracing = true
|
||||
trace.ClearBackendTraces()
|
||||
defer trace.ClearBackendTraces()
|
||||
runner := backend.NewDecisionRunner(cfg.Name, func(string) *config.ModelConfig { return &cfg }, loader, app)
|
||||
req := &schema.SystemOneRequest{State: json.RawMessage(`"secret-text"`), Questions: map[string]schema.SystemOneQuestion{"q": {Type: "noul"}}}
|
||||
response, err := runner.Decide(context.Background(), req)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(*response.Answers["q"].Noul).To(Equal(.8))
|
||||
Expect(rec.request.QuestionType).To(Equal("systemone"))
|
||||
Expect(rec.request.ModelIdentity).To(Equal(cfg.Model))
|
||||
Expect(rec.request.Prompt).To(ContainSubstring("secret-text"))
|
||||
rec.fail = true
|
||||
_, err = runner.Decide(context.Background(), req)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Eventually(func() int { return len(trace.GetBackendTraces()) }).Should(Equal(2))
|
||||
traces := trace.GetBackendTraces()
|
||||
Expect(traces).NotTo(BeEmpty())
|
||||
for _, tr := range traces {
|
||||
Expect(tr.Summary).NotTo(ContainSubstring("secret-text"))
|
||||
Expect(tr.Error).NotTo(ContainSubstring("secret-text"))
|
||||
}
|
||||
})
|
||||
})
|
||||
@@ -28,6 +28,9 @@ func ModelSystemOne(requestJSON string, loader *model.ModelLoader, modelConfig c
|
||||
return nil, fmt.Errorf("systemone not supported by backend %q", modelConfig.Backend)
|
||||
}
|
||||
return func(ctx context.Context) (string, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
release, err := AcquireGlobalBackendSlot()
|
||||
if err != nil {
|
||||
return "", err
|
||||
@@ -43,7 +46,7 @@ func ModelSystemOne(requestJSON string, loader *model.ModelLoader, modelConfig c
|
||||
Type: trace.BackendTraceScore,
|
||||
ModelName: modelConfig.Name,
|
||||
Backend: modelConfig.Backend,
|
||||
Summary: trace.TruncateString(requestJSON, 200),
|
||||
Summary: "native decision request",
|
||||
})
|
||||
}
|
||||
defer trace.CancelBackendTrace(traceID)
|
||||
@@ -55,7 +58,7 @@ func ModelSystemOne(requestJSON string, loader *model.ModelLoader, modelConfig c
|
||||
if appConfig.EnableTracing {
|
||||
errStr := ""
|
||||
if err != nil {
|
||||
errStr = err.Error()
|
||||
errStr = "native decision backend failed"
|
||||
}
|
||||
trace.RecordBackendTrace(trace.BackendTrace{
|
||||
ID: traceID,
|
||||
@@ -64,7 +67,7 @@ func ModelSystemOne(requestJSON string, loader *model.ModelLoader, modelConfig c
|
||||
Type: trace.BackendTraceScore,
|
||||
ModelName: modelConfig.Name,
|
||||
Backend: modelConfig.Backend,
|
||||
Summary: trace.TruncateString(requestJSON, 200),
|
||||
Summary: "native decision request",
|
||||
Error: errStr,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -88,6 +88,10 @@ type ApplicationConfig struct {
|
||||
// touch disk or memory.
|
||||
DisableStats bool
|
||||
|
||||
// RouterDecisionLog retains the bounded in-memory routing log even when
|
||||
// billing stats are disabled. It does not enable token usage recording.
|
||||
RouterDecisionLog bool
|
||||
|
||||
// MITMListen is the address (host:port) the cloudproxy MITM
|
||||
// listener binds on. Empty disables the MITM proxy entirely.
|
||||
// Use case: redacting PII from Claude Code / Codex CLI traffic
|
||||
@@ -1224,3 +1228,8 @@ func (o *ApplicationConfig) ApplyRuntimeSettings(settings *RuntimeSettings) (req
|
||||
// o.Metrics = meter
|
||||
// }
|
||||
// }
|
||||
|
||||
// WithRouterDecisionLog retains routing decisions independently of billing stats.
|
||||
func WithRouterDecisionLog(enabled bool) AppOption {
|
||||
return func(o *ApplicationConfig) { o.RouterDecisionLog = enabled }
|
||||
}
|
||||
@@ -309,7 +309,7 @@ var BackendCapabilities = map[string]BackendCapability{
|
||||
// llama-cpp models in the gallery are text LLMs that clone nothing.
|
||||
"llama-cpp": {
|
||||
GRPCMethods: []GRPCMethod{MethodPredict, MethodPredictStream, MethodEmbedding, MethodTokenizeString, MethodScore, MethodTTS, MethodTTSStream},
|
||||
PossibleUsecases: []string{UsecaseChat, UsecaseCompletion, UsecaseEdit, UsecaseEmbeddings, UsecaseTokenize, UsecaseVision, UsecaseScore, UsecaseTTS},
|
||||
PossibleUsecases: []string{UsecaseChat, UsecaseCompletion, UsecaseEdit, UsecaseEmbeddings, UsecaseTokenize, UsecaseVision, UsecaseScore, UsecaseDecisions, UsecaseTTS},
|
||||
DefaultUsecases: []string{UsecaseChat},
|
||||
AcceptsImages: true, // requires mmproj
|
||||
VoiceCloning: referenceVoiceCloning(),
|
||||
|
||||
@@ -114,6 +114,7 @@ func applyOverride(f *FieldMeta, o FieldMetaOverride) {
|
||||
if o.Options != nil {
|
||||
f.Options = o.Options
|
||||
}
|
||||
f.AutocompleteBy = o.AutocompleteBy
|
||||
if o.AutocompleteProvider != "" {
|
||||
f.AutocompleteProvider = o.AutocompleteProvider
|
||||
}
|
||||
@@ -132,4 +133,3 @@ func applyOverride(f *FieldMeta, o FieldMetaOverride) {
|
||||
func BuildForTest(modelConfigType reflect.Type, registry map[string]FieldMetaOverride) *ConfigMetadata {
|
||||
return buildConfigMetadataUncached(modelConfigType, registry)
|
||||
}
|
||||
|
||||
@@ -1146,10 +1146,11 @@ func DefaultRegistry() map[string]FieldMetaOverride {
|
||||
"router.classifier": {
|
||||
Section: "router",
|
||||
Label: "Classifier",
|
||||
Description: "How the router picks labels for a prompt. \"score\" asks the classifier_model to rank each policy label and reads off the softmax; \"colbert\" reranks policy descriptions against the prompt via a reranker model; \"knn\" votes over a curated corpus of labelled example prompts (seeded via the corpus API) and routes to the fallback when the prompt is unlike all corpus entries. Empty defaults to \"score\".",
|
||||
Description: "How the router picks labels for a prompt. Decisions returns independent label probabilities, not an exclusive choice. \"score\" asks the classifier_model to rank each policy label and reads off the softmax; \"colbert\" reranks policy descriptions against the prompt via a reranker model; \"knn\" votes over a curated corpus of labelled example prompts (seeded via the corpus API) and routes to the fallback when the prompt is unlike all corpus entries. Empty defaults to \"score\".",
|
||||
Component: "select",
|
||||
Options: []FieldOption{
|
||||
{Value: "score", Label: "Score (Arch-Router-style)"},
|
||||
{Value: "decisions", Label: "Decisions (native probabilities)"},
|
||||
{Value: "colbert", Label: "Colbert (reranker)"},
|
||||
{Value: "knn", Label: "KNN (labelled corpus)"},
|
||||
},
|
||||
@@ -1158,9 +1159,10 @@ func DefaultRegistry() map[string]FieldMetaOverride {
|
||||
"router.classifier_model": {
|
||||
Section: "router",
|
||||
Label: "Classifier Model",
|
||||
Description: "Loaded LocalAI model the score classifier asks to rank each policy label as a continuation (for colbert: the reranker model). Must support the Score gRPC primitive (today: llama-cpp, vLLM) and use the ChatML template. Arch-Router-1.5B Q4_K_M is the canonical choice; any small ChatML instruct model also works at a higher activation_threshold. Not used by the knn classifier.",
|
||||
Description: "Installed classifier model. Score uses a ChatML continuation model; Colbert uses a reranker. Decisions uses a native decision model explicitly declaring known_usecases: [decisions] on a Score-capable backend, with no ChatML template required. Not used by KNN.",
|
||||
Component: "model-select",
|
||||
AutocompleteProvider: ProviderModelsScore,
|
||||
AutocompleteBy: &ConditionalProvider{Field: "router.classifier", Providers: map[string]string{"decisions": "models:decisions", "colbert": "models:rerank", "knn": ""}},
|
||||
Order: 231,
|
||||
},
|
||||
"router.fallback": {
|
||||
@@ -1174,7 +1176,7 @@ func DefaultRegistry() map[string]FieldMetaOverride {
|
||||
"router.activation_threshold": {
|
||||
Section: "router",
|
||||
Label: "Activation Threshold",
|
||||
Description: "Softmax-probability floor a policy must clear to join the active label set for a request. Higher → single-label dominant routes; lower → more multi-label activations. 0 picks the package default (0.15). On Arch-Router-1.5B a value around 0.40 keeps the dominant label clean without losing genuine compound activations.",
|
||||
Description: "For Decisions, use 0.5 as a starting threshold for independent label probabilities (0 selects its default of 0.5). Switching classifiers preserves your threshold. For Score: softmax-probability floor a policy must clear to join the active label set for a request. Higher → single-label dominant routes; lower → more multi-label activations. 0 picks the package default (0.15). On Arch-Router-1.5B a value around 0.40 keeps the dominant label clean without losing genuine compound activations.",
|
||||
Component: "slider",
|
||||
Min: f64(0),
|
||||
Max: f64(1),
|
||||
|
||||
@@ -68,3 +68,20 @@ var _ = Describe("MCP field metadata", func() {
|
||||
Entry("stdio servers", "mcp.stdio", "MCP STDIO Servers", "local commands"),
|
||||
)
|
||||
})
|
||||
|
||||
var _ = Describe("Decisions router metadata", func() {
|
||||
It("offers decisions alongside existing classifiers", func() {
|
||||
f := meta.DefaultRegistry()["router.classifier"]
|
||||
Expect(f.Options).To(ContainElement(meta.FieldOption{Value: "decisions", Label: "Decisions (native probabilities)"}))
|
||||
md := meta.BuildForTest(reflect.TypeOf(config.ModelConfig{}), meta.DefaultRegistry())
|
||||
for _, field := range md.Fields {
|
||||
if field.Path == "router.classifier_model" {
|
||||
Expect(field.AutocompleteBy).NotTo(BeNil())
|
||||
Expect(field.AutocompleteBy.Field).To(Equal("router.classifier"))
|
||||
Expect(field.AutocompleteBy.Providers).To(HaveKeyWithValue("decisions", "models:decisions"))
|
||||
Expect(field.AutocompleteBy.Providers).To(HaveKeyWithValue("colbert", "models:rerank"))
|
||||
Expect(field.AutocompleteProvider).To(Equal(meta.ProviderModelsScore))
|
||||
}
|
||||
}
|
||||
})
|
||||
})
|
||||
@@ -19,10 +19,11 @@ type FieldMeta struct {
|
||||
Step *float64 `json:"step,omitempty"`
|
||||
Options []FieldOption `json:"options,omitempty"`
|
||||
|
||||
AutocompleteProvider string `json:"autocomplete_provider,omitempty"` // "backends", "models:chat", etc.
|
||||
VRAMImpact bool `json:"vram_impact,omitempty"`
|
||||
Advanced bool `json:"advanced,omitempty"`
|
||||
Order int `json:"order"`
|
||||
AutocompleteBy *ConditionalProvider `json:"autocomplete_by,omitempty"`
|
||||
AutocompleteProvider string `json:"autocomplete_provider,omitempty"` // "backends", "models:chat", etc.
|
||||
VRAMImpact bool `json:"vram_impact,omitempty"`
|
||||
Advanced bool `json:"advanced,omitempty"`
|
||||
Order int `json:"order"`
|
||||
}
|
||||
|
||||
// FieldOption represents a choice in a select/enum field.
|
||||
@@ -59,6 +60,7 @@ type FieldMetaOverride struct {
|
||||
Max *float64
|
||||
Step *float64
|
||||
Options []FieldOption
|
||||
AutocompleteBy *ConditionalProvider
|
||||
AutocompleteProvider string
|
||||
VRAMImpact bool
|
||||
Advanced bool
|
||||
@@ -89,3 +91,9 @@ func DefaultSections() []Section {
|
||||
{ID: "other", Label: "Other", Icon: "more-horizontal", Order: 100},
|
||||
}
|
||||
}
|
||||
|
||||
// ConditionalProvider selects a provider from a sibling field without changing saved values.
|
||||
type ConditionalProvider struct {
|
||||
Field string `json:"field"`
|
||||
Providers map[string]string `json:"providers"`
|
||||
}
|
||||
@@ -143,6 +143,8 @@ func (c *ModelConfig) Capabilities() []string {
|
||||
}
|
||||
}
|
||||
|
||||
add(c.NativeDecisionsEligible(), UsecaseDecisions)
|
||||
add(c.HasUsecases(FLAG_SCORE), UsecaseScore)
|
||||
add(chat, UsecaseChat)
|
||||
add(completion, UsecaseCompletion)
|
||||
add(c.HasUsecases(FLAG_EDIT), UsecaseEdit)
|
||||
@@ -239,3 +241,20 @@ func (c *ModelConfig) OutputModalities() []string {
|
||||
modalities[Modality3D] = modalities[Modality3D] || threeDOut
|
||||
return orderedModalities(modalities)
|
||||
}
|
||||
|
||||
// NativeDecisionsEligible excludes generation/NER heuristics and dispatchers.
|
||||
func (c *ModelConfig) NativeDecisionsEligible() bool {
|
||||
if c.IsDisabled() || c.HasRouter() || c.Router.Classifier != "" || c.IsAlias() || c.KnownUsecases == nil || (*c.KnownUsecases&FLAG_DECISIONS) == 0 {
|
||||
return false
|
||||
}
|
||||
backend := GetBackendCapability(c.Backend)
|
||||
if backend == nil {
|
||||
return false
|
||||
}
|
||||
for _, method := range backend.GRPCMethods {
|
||||
if method == MethodScore {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -352,7 +352,7 @@ func (c *ModelConfig) IsCloudProxyBackendPassthrough() bool {
|
||||
// config load to keep the dispatch graph acyclic and predictable. The
|
||||
// middleware also asserts depth ≤ 1 at runtime as a defensive check.
|
||||
type RouterConfig struct {
|
||||
// Classifier picks the implementation. Only "score" ships today:
|
||||
// Classifier selects score, colbert, knn, or native decisions. For score:
|
||||
// it asks the classifier model to score every Policy label as a
|
||||
// continuation of the routing prompt and reads off the
|
||||
// distribution. Empty defaults to "score".
|
||||
@@ -390,6 +390,8 @@ type RouterConfig struct {
|
||||
// 0 disables the cache. Default 1024.
|
||||
ClassifierCacheSize int `yaml:"classifier_cache_size,omitempty" json:"classifier_cache_size,omitempty"`
|
||||
|
||||
// For decisions, ActivationThreshold is an independent P(true) floor;
|
||||
// zero selects the default 0.5.
|
||||
// ActivationThreshold is the softmax-probability floor a policy
|
||||
// must clear to be considered "active" for the request. 0
|
||||
// defaults to a sensible value (~0.15) inside the classifier.
|
||||
@@ -1814,6 +1816,29 @@ func (c *ModelConfig) Validate() (bool, error) {
|
||||
return false, fmt.Errorf("router: unknown score_normalization %q (expected %q or %q)",
|
||||
c.Router.ScoreNormalization, ScoreNormalizationRaw, ScoreNormalizationMean)
|
||||
}
|
||||
|
||||
if c.Router.Classifier == "decisions" {
|
||||
t := c.Router.ActivationThreshold
|
||||
if math.IsNaN(t) || math.IsInf(t, 0) || t < 0 || t > 1 {
|
||||
return false, fmt.Errorf("router.decisions activation_threshold must be finite and in [0,1]")
|
||||
}
|
||||
if c.Router.ClassifierModel == "" {
|
||||
return false, fmt.Errorf("router.decisions requires classifier_model")
|
||||
}
|
||||
if len(c.Router.Policies) == 0 || len(c.Router.Policies) > 64 {
|
||||
return false, fmt.Errorf("router.decisions requires 1 to 64 policies")
|
||||
}
|
||||
seen := map[string]bool{}
|
||||
for _, p := range c.Router.Policies {
|
||||
if strings.TrimSpace(p.Label) == "" || strings.TrimSpace(p.Description) == "" || seen[p.Label] {
|
||||
return false, fmt.Errorf("router.decisions requires unique nonblank labels and descriptions")
|
||||
}
|
||||
seen[p.Label] = true
|
||||
}
|
||||
if c.Router.KNN != nil || c.Router.EmbeddingCache != nil {
|
||||
return false, fmt.Errorf("router.decisions does not support knn or embedding_cache composition")
|
||||
}
|
||||
}
|
||||
if c.Router.KNN != nil {
|
||||
if err := c.Router.KNN.Validate(); err != nil {
|
||||
return false, err
|
||||
|
||||
@@ -0,0 +1,24 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package config_test
|
||||
|
||||
import (
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"math"
|
||||
)
|
||||
|
||||
var _ = Describe("decisions router validation", func() {
|
||||
It("rejects invalid thresholds and duplicate policies at config validation", func() {
|
||||
c := config.ModelConfig{Name: "route", Router: config.RouterConfig{Classifier: "decisions", ClassifierModel: "native", Policies: []config.RouterPolicy{{Label: "x", Description: "x"}}}}
|
||||
for _, v := range []float64{-.1, 1.1, math.NaN(), math.Inf(1)} {
|
||||
c.Router.ActivationThreshold = v
|
||||
_, err := c.Validate()
|
||||
Expect(err).To(HaveOccurred())
|
||||
}
|
||||
c.Router.ActivationThreshold = .5
|
||||
c.Router.Policies = append(c.Router.Policies, c.Router.Policies[0])
|
||||
_, err := c.Validate()
|
||||
Expect(err).To(HaveOccurred())
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,69 @@
|
||||
package gallery_test
|
||||
|
||||
import (
|
||||
"github.com/mudler/LocalAI/core/gallery"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"gopkg.in/yaml.v3"
|
||||
"os"
|
||||
)
|
||||
|
||||
var _ = Describe("Julia native decision gallery", func() {
|
||||
It("pins a separately named text-only stock backend artifact", func() {
|
||||
data, err := os.ReadFile("../../gallery/index.yaml")
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
var entries []gallery.GalleryModel
|
||||
Expect(yaml.Unmarshal(data, &entries)).To(Succeed())
|
||||
var matches []gallery.GalleryModel
|
||||
for _, entry := range entries {
|
||||
if entry.Name == "julia-1-llama-cpp" {
|
||||
matches = append(matches, entry)
|
||||
}
|
||||
}
|
||||
Expect(matches).To(HaveLen(1))
|
||||
entry := matches[0]
|
||||
Expect(entry.License).To(Equal("apache-2.0"))
|
||||
Expect(entry.Overrides["backend"]).To(Equal("llama-cpp"))
|
||||
Expect(entry.Overrides["known_usecases"]).To(ConsistOf("decisions"))
|
||||
Expect(entry.AdditionalFiles).To(HaveLen(1))
|
||||
Expect(entry.AdditionalFiles[0].URI).To(Equal("https://huggingface.co/ggml-org/Julia-1-GGUF/resolve/16fee17949206fbf58da9347daea44d792a81211/Julia-1-Q8_0.gguf"))
|
||||
Expect(entry.AdditionalFiles[0].SHA256).To(Equal("1ea6a7e87156eeeda88cb7a36a61265b37ba7b993897b7289b99aea5b5e47069"))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("native decision family defaults", func() {
|
||||
It("pins all six defaults and enables only OpenJev with its exact projector", func() {
|
||||
data, err := os.ReadFile("../../gallery/index.yaml")
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
var entries []gallery.GalleryModel
|
||||
Expect(yaml.Unmarshal(data, &entries)).To(Succeed())
|
||||
expected := map[string]string{"julia-1-llama-cpp": "apache-2.0", "laya-llama-cpp": "apache-2.0", "kev-4b-llama-cpp": "apache-2.0", "lev-llama-cpp": "apache-2.0", "openjev-llama-cpp": "cc-by-nc-4.0", "nimble-9b-v3-llama-cpp": "cc-by-nc-4.0"}
|
||||
found := map[string]int{}
|
||||
for _, entry := range entries {
|
||||
license, ok := expected[entry.Name]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
found[entry.Name]++
|
||||
Expect(entry.License).To(Equal(license))
|
||||
Expect(entry.Overrides["backend"]).To(Equal("llama-cpp"))
|
||||
Expect(entry.Overrides["known_usecases"]).To(ConsistOf("decisions"))
|
||||
if entry.Name == "openjev-llama-cpp" {
|
||||
Expect(entry.AdditionalFiles).To(HaveLen(2))
|
||||
Expect(entry.Overrides["mmproj"]).To(Equal("mmproj-OpenJev-Q8_0.gguf"))
|
||||
Expect(entry.Overrides["context_size"]).To(Equal(8192))
|
||||
Expect(entry.AdditionalFiles[1].URI).To(Equal("https://huggingface.co/ggml-org/OpenJev-GGUF/resolve/10840f375658dea7afc5ff4711127bca8218b560/mmproj-OpenJev-Q8_0.gguf"))
|
||||
Expect(entry.AdditionalFiles[1].SHA256).To(Equal("e372cdbf59fdd6bd2504cb64c988b31c7a42ac406a8f711df4b7a7acd9216f1e"))
|
||||
} else {
|
||||
Expect(entry.AdditionalFiles).To(HaveLen(1))
|
||||
Expect(entry.Overrides).NotTo(HaveKey("mmproj"))
|
||||
}
|
||||
Expect(entry.AdditionalFiles[0].URI).To(MatchRegexp(`^https://huggingface.co/ggml-org/[^/]+/resolve/[0-9a-f]{40}/[^/]+\.gguf$`))
|
||||
Expect(entry.AdditionalFiles[0].SHA256).To(MatchRegexp(`^[0-9a-f]{64}$`))
|
||||
Expect(entry.Tags).NotTo(ContainElement("vision"))
|
||||
}
|
||||
for name := range expected {
|
||||
Expect(found[name]).To(Equal(1), name)
|
||||
}
|
||||
})
|
||||
})
|
||||
@@ -871,7 +871,7 @@ func convertAnthropicToOpenAIMessages(input *schema.AnthropicRequest) []schema.M
|
||||
}
|
||||
|
||||
// Handle content (can be string or array of content blocks)
|
||||
switch content := msg.Content.(type) {
|
||||
switch content := anthropicContentBlocks(msg.Content).(type) {
|
||||
case string:
|
||||
openAIMsg.StringContent = content
|
||||
openAIMsg.Content = content
|
||||
@@ -942,13 +942,15 @@ func convertAnthropicToOpenAIMessages(input *schema.AnthropicRequest) []schema.M
|
||||
// For now, we'll add it as text content
|
||||
toolUseID, _ := blockMap["tool_use_id"].(string)
|
||||
isError := false
|
||||
if isErrorPtr, ok := blockMap["is_error"].(*bool); ok && isErrorPtr != nil {
|
||||
if value, ok := blockMap["is_error"].(bool); ok {
|
||||
isError = value
|
||||
} else if isErrorPtr, ok := blockMap["is_error"].(*bool); ok && isErrorPtr != nil {
|
||||
isError = *isErrorPtr
|
||||
}
|
||||
|
||||
var resultText string
|
||||
if resultContent, ok := blockMap["content"]; ok {
|
||||
switch rc := resultContent.(type) {
|
||||
switch rc := anthropicContentBlocks(resultContent).(type) {
|
||||
case string:
|
||||
resultText = rc
|
||||
case []any:
|
||||
@@ -1049,3 +1051,21 @@ func forwardCloudProxyAnthropicViaBackend(c echo.Context, cfg *config.ModelConfi
|
||||
}
|
||||
return cloudproxy.ForwardViaBackend(c, cfg, body, ml, appConfig)
|
||||
}
|
||||
|
||||
// anthropicContentBlocks normalizes typed blocks without a JSON round trip or
|
||||
// mutating the caller's payload. Both representations use the same conversion.
|
||||
func anthropicContentBlocks(content any) any {
|
||||
blocks, ok := content.([]schema.AnthropicContentBlock)
|
||||
if !ok {
|
||||
return content
|
||||
}
|
||||
result := make([]any, 0, len(blocks))
|
||||
for _, b := range blocks {
|
||||
m := map[string]any{"type": b.Type, "text": b.Text, "thinking": b.Thinking, "id": b.ID, "name": b.Name, "input": b.Input, "tool_use_id": b.ToolUseID, "content": b.Content, "is_error": b.IsError}
|
||||
if b.Source != nil {
|
||||
m["source"] = map[string]any{"type": b.Source.Type, "media_type": b.Source.MediaType, "data": b.Source.Data}
|
||||
}
|
||||
result = append(result, m)
|
||||
}
|
||||
return result
|
||||
}
|
||||
@@ -110,3 +110,15 @@ func indexOf(items []string, target string) int {
|
||||
}
|
||||
return -1
|
||||
}
|
||||
|
||||
var _ = Describe("Typed inbound conversion", func() {
|
||||
It("preserves ordered typed text images and tool errors without mutation", func() {
|
||||
yes := true
|
||||
blocks := []schema.AnthropicContentBlock{{Type: "text", Text: "first"}, {Type: "image", Source: &schema.AnthropicImageSource{Type: "base64", MediaType: "image/png", Data: "AA=="}}, {Type: "text", Text: "second"}, {Type: "image", Source: &schema.AnthropicImageSource{Type: "base64", MediaType: "image/jpeg", Data: "AQ=="}}, {Type: "tool_result", ToolUseID: "id", Content: "failed", IsError: &yes}}
|
||||
req := &schema.AnthropicRequest{Messages: []schema.AnthropicMessage{{Role: "user", Content: blocks}}}
|
||||
msgs := convertAnthropicToOpenAIMessages(req)
|
||||
Expect(msgs[0].StringContent).To(Equal("firstsecond\n[Tool Result for id]: Error: failed"))
|
||||
Expect(msgs[0].StringImages).To(Equal([]string{"data:image/png;base64,AA==", "data:image/jpeg;base64,AQ=="}))
|
||||
Expect(req.Messages[0].Content).To(Equal(blocks))
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,92 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package anthropic
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"image"
|
||||
"image/png"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/labstack/echo/v4"
|
||||
"github.com/mudler/LocalAI/core/backend"
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/http/middleware"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/services/routing/router"
|
||||
"github.com/mudler/LocalAI/pkg/system"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
type conversionRunner func(context.Context, *schema.SystemOneRequest) (*schema.SystemOneResponse, error)
|
||||
|
||||
func (f conversionRunner) Decide(c context.Context, r *schema.SystemOneRequest) (*schema.SystemOneResponse, error) {
|
||||
return f(c, r)
|
||||
}
|
||||
|
||||
var _ = Describe("Routed typed Anthropic endpoint conversion", func() {
|
||||
for _, fallback := range []bool{false, true} {
|
||||
It("preserves original content after native selection or runtime fallback", func() {
|
||||
dir := GinkgoT().TempDir()
|
||||
for _, name := range []string{"chosen", "fallback"} {
|
||||
Expect(os.WriteFile(filepath.Join(dir, name+".yaml"), []byte("name: "+name+"\nbackend: mock-backend\n"), 0600)).To(Succeed())
|
||||
}
|
||||
cfg := &config.ModelConfig{Name: "route", Router: config.RouterConfig{Classifier: "decisions", ClassifierModel: "native", Fallback: "fallback", Policies: []config.RouterPolicy{{Label: "visual", Description: "image"}}, Candidates: []config.RouterCandidate{{Model: "chosen", Labels: []string{"visual"}}}}}
|
||||
var b bytes.Buffer
|
||||
Expect(png.Encode(&b, image.NewGray(image.Rect(0, 0, 1, 1)))).To(Succeed())
|
||||
data := base64.StdEncoding.EncodeToString(b.Bytes())
|
||||
blocks := []schema.AnthropicContentBlock{{Type: "text", Text: "first"}, {Type: "image", Source: &schema.AnthropicImageSource{Type: "base64", MediaType: "image/png", Data: data}}, {Type: "text", Text: "second"}, {Type: "image", Source: &schema.AnthropicImageSource{Type: "base64", MediaType: "image/png", Data: data}}}
|
||||
req := &schema.AnthropicRequest{Model: "route", Messages: []schema.AnthropicMessage{{Role: "user", Content: blocks}}}
|
||||
before, err := json.Marshal(req.Messages)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
app := &config.ApplicationConfig{Context: context.Background(), SystemState: &system.SystemState{Model: system.Model{ModelsPath: dir}}}
|
||||
c := echo.New().NewContext(httptest.NewRequest("POST", "/v1/messages", nil), httptest.NewRecorder())
|
||||
c.Set(middleware.CONTEXT_LOCALS_KEY_MODEL_CONFIG, cfg)
|
||||
c.Set(middleware.CONTEXT_LOCALS_KEY_LOCALAI_REQUEST, req)
|
||||
called := false
|
||||
deps := middleware.ClassifierDeps{ModelLookup: func(string) *config.ModelConfig {
|
||||
u := config.FLAG_DECISIONS
|
||||
return &config.ModelConfig{Backend: "llama-cpp", KnownUsecases: &u}
|
||||
}, Decisions: func(string) backend.DecisionRunner {
|
||||
return conversionRunner(func(_ context.Context, r *schema.SystemOneRequest) (*schema.SystemOneResponse, error) {
|
||||
called = true
|
||||
if fallback {
|
||||
return nil, errors.New("runtime failure")
|
||||
}
|
||||
answers := map[string]schema.SystemOneAnswer{}
|
||||
for k := range r.Questions {
|
||||
v := 1.0
|
||||
answers[k] = schema.SystemOneAnswer{Type: "noul", Noul: &v}
|
||||
}
|
||||
return &schema.SystemOneResponse{Answers: answers}, nil
|
||||
})
|
||||
}}
|
||||
var converted []schema.Message
|
||||
handler := middleware.RouteModel(config.NewModelConfigLoader(dir), app, nil, nil, middleware.AnthropicProbe, router.SourceAnthropic, deps)(func(c echo.Context) error {
|
||||
// Invoke the actual conversion used by the messages endpoint, not a
|
||||
// stub which only checks that middleware called next.
|
||||
converted = convertAnthropicToOpenAIMessages(c.Get(middleware.CONTEXT_LOCALS_KEY_LOCALAI_REQUEST).(*schema.AnthropicRequest))
|
||||
return nil
|
||||
})
|
||||
Expect(handler(c)).To(Succeed())
|
||||
Expect(called).To(BeTrue())
|
||||
if fallback {
|
||||
Expect(req.Model).To(Equal("fallback"))
|
||||
} else {
|
||||
Expect(req.Model).To(Equal("chosen"))
|
||||
}
|
||||
Expect(converted).To(HaveLen(1))
|
||||
Expect(converted[0].StringContent).To(Equal("firstsecond"))
|
||||
Expect(converted[0].StringImages).To(Equal([]string{"data:image/png;base64," + data, "data:image/png;base64," + data}))
|
||||
after, err := json.Marshal(req.Messages)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(after).To(Equal(before))
|
||||
})
|
||||
}
|
||||
})
|
||||
@@ -124,6 +124,8 @@ func AutocompleteEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, a
|
||||
filterFn = config.BuildUsecaseFilterFn(config.FLAG_VAD)
|
||||
case config.UsecaseTranscript:
|
||||
filterFn = config.BuildUsecaseFilterFn(config.FLAG_TRANSCRIPT)
|
||||
case config.UsecaseDecisions:
|
||||
filterFn = func(_ string, c *config.ModelConfig) bool { return c.NativeDecisionsEligible() }
|
||||
case "score": // router classifier usecase (FLAG_SCORE); not in UsecaseInfoMap
|
||||
filterFn = config.BuildUsecaseFilterFn(config.FLAG_SCORE)
|
||||
case config.UsecaseTokenClassify: // PII NER detector usecase (FLAG_TOKEN_CLASSIFY)
|
||||
|
||||
@@ -65,6 +65,36 @@ var _ = Describe("Config Metadata Endpoints", func() {
|
||||
os.RemoveAll(tempDir)
|
||||
})
|
||||
|
||||
It("lists native decisions on both Score backends but not NER, routers or disabled models", func() {
|
||||
for name, body := range map[string]string{
|
||||
"llama-decision": "backend: llama-cpp\nknown_usecases: [decisions]\n",
|
||||
"vllm-decision": "backend: vllm-cpp\nknown_usecases: [decisions]\n",
|
||||
"ner": "backend: vllm-cpp\nknown_usecases: [token_classify]\n",
|
||||
"chat": "backend: llama-cpp\nknown_usecases: [chat]\n",
|
||||
"unsupported": "backend: piper\nknown_usecases: [decisions]\n",
|
||||
"disabled": "backend: llama-cpp\nknown_usecases: [decisions]\ndisabled: true\n",
|
||||
"router": "backend: llama-cpp\nknown_usecases: [decisions]\nrouter:\n classifier: decisions\n classifier_model: llama-decision\n",
|
||||
} {
|
||||
Expect(os.WriteFile(filepath.Join(tempDir, name+".yaml"), []byte("name: "+name+"\n"+body), 0600)).To(Succeed())
|
||||
}
|
||||
Expect(configLoader.LoadModelConfigsFromPath(tempDir)).To(Succeed())
|
||||
rec := httptest.NewRecorder()
|
||||
app.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/models/config-metadata/autocomplete/models:decisions", nil))
|
||||
Expect(rec.Code).To(Equal(http.StatusOK))
|
||||
var result struct {
|
||||
Values []string `json:"values"`
|
||||
}
|
||||
Expect(json.Unmarshal(rec.Body.Bytes(), &result)).To(Succeed())
|
||||
Expect(result.Values).To(ConsistOf("llama-decision", "vllm-decision"))
|
||||
for _, cfg := range configLoader.GetAllModelsConfigs() {
|
||||
if cfg.Name == "llama-decision" || cfg.Name == "vllm-decision" {
|
||||
Expect(cfg.Capabilities()).To(ContainElement(config.UsecaseDecisions))
|
||||
} else {
|
||||
Expect(cfg.Capabilities()).NotTo(ContainElement(config.UsecaseDecisions))
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
Context("GET /api/models/config-metadata", func() {
|
||||
It("should return section index when no section param", func() {
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/models/config-metadata", nil)
|
||||
|
||||
@@ -52,6 +52,26 @@ var _ = Describe("RouterDecideEndpoint", func() {
|
||||
_ = os.RemoveAll(modelDir)
|
||||
})
|
||||
|
||||
It("routes overlapping native decisions through the oracle factory", func() {
|
||||
cfg := config.ModelConfig{Name: "native-router", Router: config.RouterConfig{Classifier: "decisions", ClassifierModel: "native", Policies: []config.RouterPolicy{{Label: "code", Description: "coding"}, {Label: "private", Description: "private data"}}, Candidates: []config.RouterCandidate{{Model: "small-model", Labels: []string{"code"}}, {Model: "big-model", Labels: []string{"code", "private"}}}}}
|
||||
b, err := yaml.Marshal(cfg)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(os.WriteFile(filepath.Join(modelDir, "native-router.yaml"), b, 0600)).To(Succeed())
|
||||
writeBareModel(modelDir, "small-model")
|
||||
writeBareModel(modelDir, "big-model")
|
||||
flags := config.FLAG_DECISIONS
|
||||
native := &config.ModelConfig{Name: "native", Backend: "vllm-cpp", KnownUsecases: &flags}
|
||||
calls := 0
|
||||
d := middleware.ClassifierDeps{Registry: router.NewRegistry(), ModelLookup: func(string) *config.ModelConfig { return native }, Decisions: func(name string) backend.DecisionRunner {
|
||||
Expect(name).To(Equal("native"))
|
||||
return oracleDecisionRunner{calls: &calls}
|
||||
}}
|
||||
rec, body := invokeDecide(loader, appConfig, d, `{"router":"native-router","input":"private coding task"}`)
|
||||
Expect(rec.Code).To(Equal(http.StatusOK))
|
||||
Expect(body.Candidate).To(Equal("big-model"))
|
||||
Expect(calls).To(Equal(1))
|
||||
})
|
||||
|
||||
It("rejects requests with no router field", func() {
|
||||
rec, _ := invokeDecide(loader, appConfig, deps(nil), `{"input":"hello"}`)
|
||||
Expect(rec.Code).To(Equal(http.StatusBadRequest))
|
||||
@@ -246,3 +266,11 @@ func writeBareModel(modelDir, name string) {
|
||||
body := "name: " + name + "\nbackend: mock-backend\n"
|
||||
Expect(os.WriteFile(filepath.Join(modelDir, name+".yaml"), []byte(body), 0o644)).To(Succeed())
|
||||
}
|
||||
|
||||
type oracleDecisionRunner struct{ calls *int }
|
||||
|
||||
func (r oracleDecisionRunner) Decide(_ context.Context, req *schema.SystemOneRequest) (*schema.SystemOneResponse, error) {
|
||||
*r.calls++
|
||||
a, b := .8, .9
|
||||
return &schema.SystemOneResponse{Answers: map[string]schema.SystemOneAnswer{"p0": {Type: "noul", Noul: &a}, "p1": {Type: "noul", Noul: &b}}}, nil
|
||||
}
|
||||
@@ -1,13 +1,15 @@
|
||||
package localai
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"math"
|
||||
"math/rand"
|
||||
"net/http"
|
||||
"slices"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
@@ -17,9 +19,81 @@ import (
|
||||
"github.com/mudler/LocalAI/core/application"
|
||||
"github.com/mudler/LocalAI/core/backend"
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/http/middleware"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/systemone"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
)
|
||||
|
||||
// systemOneBackendStatus preserves the distinction between invalid input and
|
||||
// unsupported capabilities; unknown/load failures remain server errors.
|
||||
func systemOneBackendStatus(err error) int {
|
||||
switch status.Code(err) {
|
||||
case codes.InvalidArgument:
|
||||
return http.StatusBadRequest
|
||||
case codes.Unimplemented:
|
||||
return http.StatusNotImplemented
|
||||
default:
|
||||
return http.StatusInternalServerError
|
||||
}
|
||||
}
|
||||
|
||||
func stampSystemOneUsage(c echo.Context, model, response string) error {
|
||||
var result struct {
|
||||
Usage *struct {
|
||||
Input *int `json:"input_tokens"`
|
||||
Output *int `json:"output_tokens"`
|
||||
} `json:"usage"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(response), &result); err != nil {
|
||||
return err
|
||||
}
|
||||
if result.Usage == nil {
|
||||
return nil
|
||||
}
|
||||
if result.Usage.Input == nil && result.Usage.Output == nil {
|
||||
return nil
|
||||
}
|
||||
if result.Usage.Input == nil || result.Usage.Output == nil {
|
||||
return fmt.Errorf("incomplete decision usage")
|
||||
}
|
||||
input, output := *result.Usage.Input, *result.Usage.Output
|
||||
if input < 0 || output < 0 || input > int(^uint(0)>>1)-output {
|
||||
return fmt.Errorf("invalid decision usage counts")
|
||||
}
|
||||
middleware.StampUsage(c, model, input, output)
|
||||
return nil
|
||||
}
|
||||
|
||||
// respondSystemOne is the native route's response path. UsageMiddleware records
|
||||
// the stamp once and does not also parse the response body.
|
||||
func respondSystemOne(c echo.Context, model string, run func(context.Context) (string, error)) error {
|
||||
response, err := run(c.Request().Context())
|
||||
if err != nil {
|
||||
return systemOneError(c, systemOneBackendStatus(err), err.Error())
|
||||
}
|
||||
if len(response) > systemone.MaxResponseBytes {
|
||||
return systemOneError(c, http.StatusInternalServerError, "decision response exceeds 64 KiB")
|
||||
}
|
||||
var envelope struct {
|
||||
Answers map[string]json.RawMessage `json:"answers"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(response), &envelope); err != nil || len(envelope.Answers) == 0 {
|
||||
return systemOneError(c, http.StatusInternalServerError, "invalid decision response: answers required")
|
||||
}
|
||||
for _, answer := range envelope.Answers {
|
||||
var fields map[string]json.RawMessage
|
||||
if err := json.Unmarshal(answer, &fields); err != nil || len(fields) == 0 {
|
||||
return systemOneError(c, http.StatusInternalServerError, "invalid decision answer")
|
||||
}
|
||||
}
|
||||
if err := stampSystemOneUsage(c, model, response); err != nil {
|
||||
return systemOneError(c, http.StatusInternalServerError, "invalid decision response")
|
||||
}
|
||||
return c.JSON(http.StatusOK, json.RawMessage(response))
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Helpers — ported from kev/api.py (render, r2, choice_confidence,
|
||||
// score_confidence, softmax) and mirrored in vllm.cpp api_server.cpp.
|
||||
@@ -198,6 +272,13 @@ type parsedSystemOne struct {
|
||||
}
|
||||
|
||||
func parseSystemOneRequest(req *schema.SystemOneRequest) (*parsedSystemOne, error) {
|
||||
images, err := systemOneImages(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(images) > 0 {
|
||||
return nil, errSystemOneImagesUnsupported
|
||||
}
|
||||
p := &parsedSystemOne{
|
||||
model: req.Model,
|
||||
threshold: 0.5,
|
||||
@@ -377,13 +458,7 @@ func systemOneError(c echo.Context, status int, msg string) error {
|
||||
// declares no usecases predates the flag and stays allowed, and a
|
||||
// token_classify model is allowed because the NER path serves it.
|
||||
func systemOneModelAllowed(cfg config.ModelConfig) error {
|
||||
if cfg.KnownUsecases == nil {
|
||||
return nil
|
||||
}
|
||||
if *cfg.KnownUsecases&(config.FLAG_DECISIONS|config.FLAG_TOKEN_CLASSIFY) != 0 {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("model %q does not declare the decisions usecase (known_usecases: [decisions])", cfg.Name)
|
||||
return systemone.ModelAllowed(cfg)
|
||||
}
|
||||
|
||||
// checkSystemOneModel applies systemOneModelAllowed to a model looked up by
|
||||
@@ -409,31 +484,14 @@ func checkSystemOneModel(app *application.Application, modelName string) error {
|
||||
// decision pipeline, which is what setups that predate the decisions usecase
|
||||
// relied on.
|
||||
func systemOneUsesDecisionPipeline(cfg config.ModelConfig) bool {
|
||||
if !backendSupportsScore(cfg.Backend) {
|
||||
return false
|
||||
}
|
||||
if cfg.KnownUsecases == nil {
|
||||
return true
|
||||
}
|
||||
declared := *cfg.KnownUsecases
|
||||
if declared&config.FLAG_DECISIONS != 0 {
|
||||
return true
|
||||
}
|
||||
return declared&config.FLAG_TOKEN_CLASSIFY == 0
|
||||
return systemone.UsesDecisionPipeline(cfg)
|
||||
}
|
||||
|
||||
// systemOneNERAllowed guards /permute and /separate, which always run the NER
|
||||
// path. A decision model cannot serve them: the backend's NER entry point
|
||||
// refuses its architecture, and the caller would see a backend error.
|
||||
func systemOneNERAllowed(cfg config.ModelConfig) error {
|
||||
if cfg.KnownUsecases == nil {
|
||||
return nil
|
||||
}
|
||||
declared := *cfg.KnownUsecases
|
||||
if declared&config.FLAG_DECISIONS != 0 && declared&config.FLAG_TOKEN_CLASSIFY == 0 {
|
||||
return fmt.Errorf("model %q is a decision model: /permute and /separate use the NER path, use POST /v1/systemone instead", cfg.Name)
|
||||
}
|
||||
return nil
|
||||
return systemone.NERAllowed(cfg)
|
||||
}
|
||||
|
||||
// checkSystemOneNERModel applies systemOneNERAllowed to a model looked up by
|
||||
@@ -450,21 +508,58 @@ func checkSystemOneNERModel(app *application.Application, modelName string) erro
|
||||
return systemOneNERAllowed(cfg)
|
||||
}
|
||||
|
||||
// systemOneMaxBody and systemOneMaxQuestions bound one request. They keep a
|
||||
// single call from pinning a decision model on an unbounded prompt, and match
|
||||
// the limits Ollama documents for the same wire contract, so a client written
|
||||
// for one server behaves the same on the other. The engine enforces any
|
||||
// per-model option cap (letter-answer models refuse more than 26 options).
|
||||
// Text wire requests retain their original cap independently of image requests.
|
||||
const (
|
||||
systemOneMaxBody = 64 << 10
|
||||
systemOneMaxQuestions = 64
|
||||
systemOneMaxBody = systemone.MaxBodyBytes // public raw-wire cap, independent of internal serialized cap
|
||||
systemOneMaxQuestions = systemone.MaxQuestions
|
||||
)
|
||||
|
||||
// systemOneBind binds the JSON body with a size cap. Bind reads the whole body
|
||||
// first, so the cap has to be on the reader.
|
||||
func systemOneBind(c echo.Context, v any) error {
|
||||
c.Request().Body = http.MaxBytesReader(c.Response(), c.Request().Body, systemOneMaxBody)
|
||||
return c.Bind(v)
|
||||
data, err := io.ReadAll(http.MaxBytesReader(c.Response(), c.Request().Body, systemone.MaxImageBodyBytes))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// Unmarshal consumes the whole payload: trailing JSON or garbage is invalid,
|
||||
// and trailing whitespace is included in the raw wire-byte budget above.
|
||||
if !json.Valid(data) {
|
||||
if len(data) > systemOneMaxBody {
|
||||
return &http.MaxBytesError{Limit: int64(systemOneMaxBody)}
|
||||
}
|
||||
return fmt.Errorf("invalid request body")
|
||||
}
|
||||
c.Request().Body = io.NopCloser(bytes.NewReader(data))
|
||||
if err := c.Bind(v); err != nil {
|
||||
if len(data) > systemOneMaxBody {
|
||||
return &http.MaxBytesError{Limit: int64(systemOneMaxBody)}
|
||||
}
|
||||
return err
|
||||
}
|
||||
var req *schema.SystemOneRequest
|
||||
switch value := v.(type) {
|
||||
case *schema.SystemOneRequest:
|
||||
req = value
|
||||
case *schema.SystemOnePermuteRequest:
|
||||
req = &value.Request
|
||||
}
|
||||
limit := systemOneMaxBody
|
||||
if req != nil {
|
||||
var err error
|
||||
limit, err = systemone.RequestBodyLimit(req)
|
||||
if err != nil {
|
||||
// Invalid/missing state cannot opt a text request into the image
|
||||
// budget. Preserve raw-wire overflow precedence, including spaces.
|
||||
if len(data) > systemOneMaxBody {
|
||||
return &http.MaxBytesError{Limit: int64(systemOneMaxBody)}
|
||||
}
|
||||
return err
|
||||
}
|
||||
}
|
||||
if len(data) > limit {
|
||||
return &http.MaxBytesError{Limit: int64(limit)}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// systemOneBindStatus maps a bind failure to its status: 413 when the body
|
||||
@@ -479,7 +574,9 @@ func systemOneBindStatus(err error) int {
|
||||
|
||||
func systemOneBindMessage(err error) string {
|
||||
if systemOneBindStatus(err) == http.StatusRequestEntityTooLarge {
|
||||
return fmt.Sprintf("request body exceeds %d KiB", systemOneMaxBody>>10)
|
||||
var tooLarge *http.MaxBytesError
|
||||
errors.As(err, &tooLarge)
|
||||
return fmt.Sprintf("request body exceeds %d KiB", tooLarge.Limit>>10)
|
||||
}
|
||||
return "invalid request body"
|
||||
}
|
||||
@@ -489,70 +586,8 @@ func systemOneBindMessage(err error) string {
|
||||
// forwarded path never sees parseSystemOneRequest, so without this a malformed
|
||||
// question would surface as a backend error instead of a 400.
|
||||
func validateSystemOneRequest(req *schema.SystemOneRequest) error {
|
||||
if len(req.State) == 0 || string(req.State) == "null" {
|
||||
return fmt.Errorf("state is required")
|
||||
}
|
||||
var state any
|
||||
if err := json.Unmarshal(req.State, &state); err != nil {
|
||||
return fmt.Errorf("state is not valid JSON: %w", err)
|
||||
}
|
||||
if s, ok := state.(string); ok && strings.TrimSpace(s) == "" {
|
||||
return fmt.Errorf("state is required")
|
||||
}
|
||||
if len(req.Questions) == 0 {
|
||||
return fmt.Errorf("questions is required and must contain at least one question")
|
||||
}
|
||||
if len(req.Questions) > systemOneMaxQuestions {
|
||||
return fmt.Errorf("questions must contain at most %d questions", systemOneMaxQuestions)
|
||||
}
|
||||
qids := make([]string, 0, len(req.Questions))
|
||||
for id := range req.Questions {
|
||||
qids = append(qids, id)
|
||||
}
|
||||
sort.Strings(qids)
|
||||
for _, id := range qids {
|
||||
if strings.TrimSpace(id) == "" {
|
||||
return fmt.Errorf("question ids must not be blank")
|
||||
}
|
||||
q := req.Questions[id]
|
||||
switch q.Type {
|
||||
case "choice":
|
||||
var criteria map[string]json.RawMessage
|
||||
if err := json.Unmarshal(q.Criteria, &criteria); err != nil {
|
||||
return fmt.Errorf("question %q (choice) requires a criteria object", id)
|
||||
}
|
||||
if len(criteria) < 2 {
|
||||
return fmt.Errorf("question %q (choice) requires at least 2 options", id)
|
||||
}
|
||||
for k := range criteria {
|
||||
if strings.TrimSpace(k) == "" {
|
||||
return fmt.Errorf("question %q (choice) has a blank option key", id)
|
||||
}
|
||||
}
|
||||
case "score":
|
||||
var criteria []json.RawMessage
|
||||
if err := json.Unmarshal(q.Criteria, &criteria); err != nil {
|
||||
return fmt.Errorf("question %q (score) requires a criteria array", id)
|
||||
}
|
||||
if len(criteria) < 2 {
|
||||
return fmt.Errorf("question %q (score) requires at least 2 levels", id)
|
||||
}
|
||||
case "noul":
|
||||
if len(q.Criteria) == 0 || string(q.Criteria) == "null" {
|
||||
continue
|
||||
}
|
||||
var criteria map[string]json.RawMessage
|
||||
if err := json.Unmarshal(q.Criteria, &criteria); err != nil {
|
||||
return fmt.Errorf("question %q (noul) criteria must be an object with \"false\" and \"true\" descriptions", id)
|
||||
}
|
||||
for k := range criteria {
|
||||
if k != "false" && k != "true" {
|
||||
return fmt.Errorf("question %q (noul) criteria may only have \"false\" and \"true\" keys", id)
|
||||
}
|
||||
}
|
||||
default:
|
||||
return fmt.Errorf("question %q has unknown type: %s", id, q.Type)
|
||||
}
|
||||
if err := systemone.ValidateRequestStructure(req); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -562,11 +597,7 @@ func validateSystemOneRequest(req *schema.SystemOneRequest) error {
|
||||
// scoring via the unified vllm_decide C ABI); other backends fall through to
|
||||
// the NER-based path.
|
||||
func backendSupportsScore(backendName string) bool {
|
||||
cap := config.GetBackendCapability(backendName)
|
||||
if cap == nil {
|
||||
return false
|
||||
}
|
||||
return slices.Contains(cap.GRPCMethods, config.MethodScore)
|
||||
return systemone.BackendSupportsScore(backendName)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
@@ -587,6 +618,12 @@ func backendSupportsScore(backendName string) bool {
|
||||
// @Router /v1/systemone [post]
|
||||
func SystemOneEndpoint(app *application.Application) echo.HandlerFunc {
|
||||
return func(c echo.Context) error {
|
||||
release, err := systemone.AcquireAdmission(c.Request().Context())
|
||||
if err != nil {
|
||||
return systemOneError(c, http.StatusServiceUnavailable, err.Error())
|
||||
}
|
||||
defer release()
|
||||
|
||||
var req schema.SystemOneRequest
|
||||
if err := systemOneBind(c, &req); err != nil {
|
||||
return systemOneError(c, systemOneBindStatus(err), systemOneBindMessage(err))
|
||||
@@ -598,7 +635,7 @@ func SystemOneEndpoint(app *application.Application) echo.HandlerFunc {
|
||||
return systemOneError(c, http.StatusBadRequest, err.Error())
|
||||
}
|
||||
if err := validateSystemOneRequest(&req); err != nil {
|
||||
return systemOneError(c, http.StatusBadRequest, err.Error())
|
||||
return systemOneError(c, systemOneInputStatus(err), err.Error())
|
||||
}
|
||||
// vllm-cpp models (kev/laya) implement the decision pipeline natively
|
||||
// via the vllm_decide C ABI. Forward the raw request JSON through the
|
||||
@@ -614,17 +651,13 @@ func SystemOneEndpoint(app *application.Application) echo.HandlerFunc {
|
||||
if err != nil {
|
||||
return systemOneError(c, http.StatusInternalServerError, err.Error())
|
||||
}
|
||||
respJSON, err := fn(c.Request().Context())
|
||||
if err != nil {
|
||||
return systemOneError(c, http.StatusInternalServerError, err.Error())
|
||||
}
|
||||
return c.JSON(http.StatusOK, json.RawMessage(respJSON))
|
||||
return respondSystemOne(c, req.Model, fn)
|
||||
}
|
||||
}
|
||||
// NER-based path (GLiNER2.5 zero-shot NER).
|
||||
parsed, err := parseSystemOneRequest(&req)
|
||||
if err != nil {
|
||||
return systemOneError(c, http.StatusBadRequest, err.Error())
|
||||
return systemOneError(c, systemOneInputStatus(err), err.Error())
|
||||
}
|
||||
classifier, err := resolveClassifier(app, req.Model, parsed.threshold)
|
||||
if err != nil {
|
||||
@@ -659,6 +692,12 @@ func SystemOneEndpoint(app *application.Application) echo.HandlerFunc {
|
||||
// @Router /v1/systemone/permute [post]
|
||||
func SystemOnePermuteEndpoint(app *application.Application) echo.HandlerFunc {
|
||||
return func(c echo.Context) error {
|
||||
release, err := systemone.AcquireAdmission(c.Request().Context())
|
||||
if err != nil {
|
||||
return systemOneError(c, http.StatusServiceUnavailable, err.Error())
|
||||
}
|
||||
defer release()
|
||||
|
||||
var req schema.SystemOnePermuteRequest
|
||||
if err := systemOneBind(c, &req); err != nil {
|
||||
return systemOneError(c, systemOneBindStatus(err), systemOneBindMessage(err))
|
||||
@@ -673,14 +712,14 @@ func SystemOnePermuteEndpoint(app *application.Application) echo.HandlerFunc {
|
||||
return systemOneError(c, http.StatusBadRequest, err.Error())
|
||||
}
|
||||
if err := validateSystemOneRequest(&req.Request); err != nil {
|
||||
return systemOneError(c, http.StatusBadRequest, err.Error())
|
||||
return systemOneError(c, systemOneInputStatus(err), err.Error())
|
||||
}
|
||||
if req.Question == "" {
|
||||
return systemOneError(c, http.StatusBadRequest, "question is required")
|
||||
}
|
||||
parsed, err := parseSystemOneRequest(&req.Request)
|
||||
if err != nil {
|
||||
return systemOneError(c, http.StatusBadRequest, err.Error())
|
||||
return systemOneError(c, systemOneInputStatus(err), err.Error())
|
||||
}
|
||||
var target *parsedQuestion
|
||||
for i := range parsed.questions {
|
||||
@@ -804,6 +843,12 @@ func SystemOnePermuteEndpoint(app *application.Application) echo.HandlerFunc {
|
||||
// @Router /v1/systemone/separate [post]
|
||||
func SystemOneSeparateEndpoint(app *application.Application) echo.HandlerFunc {
|
||||
return func(c echo.Context) error {
|
||||
release, err := systemone.AcquireAdmission(c.Request().Context())
|
||||
if err != nil {
|
||||
return systemOneError(c, http.StatusServiceUnavailable, err.Error())
|
||||
}
|
||||
defer release()
|
||||
|
||||
var req schema.SystemOneRequest
|
||||
if err := systemOneBind(c, &req); err != nil {
|
||||
return systemOneError(c, systemOneBindStatus(err), systemOneBindMessage(err))
|
||||
@@ -818,11 +863,11 @@ func SystemOneSeparateEndpoint(app *application.Application) echo.HandlerFunc {
|
||||
return systemOneError(c, http.StatusBadRequest, err.Error())
|
||||
}
|
||||
if err := validateSystemOneRequest(&req); err != nil {
|
||||
return systemOneError(c, http.StatusBadRequest, err.Error())
|
||||
return systemOneError(c, systemOneInputStatus(err), err.Error())
|
||||
}
|
||||
parsed, err := parseSystemOneRequest(&req)
|
||||
if err != nil {
|
||||
return systemOneError(c, http.StatusBadRequest, err.Error())
|
||||
return systemOneError(c, systemOneInputStatus(err), err.Error())
|
||||
}
|
||||
classifier, err := resolveClassifier(app, req.Model, parsed.threshold)
|
||||
if err != nil {
|
||||
|
||||
@@ -72,3 +72,12 @@ var _ = Describe("systemone routing by model kind", func() {
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("native decision capabilities", func() {
|
||||
It("advertises stock llama-cpp decisions without advertising older forks", func() {
|
||||
Expect(config.PossibleUsecasesForBackend("llama-cpp")).To(ContainElement(config.UsecaseDecisions))
|
||||
for _, name := range []string{"bonsai", "turboquant", "ik-llama-cpp"} {
|
||||
Expect(config.PossibleUsecasesForBackend(name)).NotTo(ContainElement(config.UsecaseDecisions))
|
||||
}
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,39 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package localai
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/systemone"
|
||||
"net/http"
|
||||
)
|
||||
|
||||
const systemOneMaxImageBytes = systemone.MaxImageEncodedBytes
|
||||
|
||||
var errSystemOneImagesUnsupported = errors.New("image input is not supported by the NER decision path")
|
||||
|
||||
func systemOneInputStatus(err error) int {
|
||||
if errors.Is(err, errSystemOneImagesUnsupported) {
|
||||
return http.StatusNotImplemented
|
||||
}
|
||||
var validation *systemone.ValidationError
|
||||
if errors.As(err, &validation) {
|
||||
switch validation.Kind {
|
||||
case systemone.InputTooLarge:
|
||||
return http.StatusRequestEntityTooLarge
|
||||
case systemone.UnsupportedBackend:
|
||||
return http.StatusNotImplemented
|
||||
}
|
||||
}
|
||||
return http.StatusBadRequest
|
||||
}
|
||||
func systemOneImages(req *schema.SystemOneRequest) ([]string, error) {
|
||||
return systemone.CollectImages(req)
|
||||
}
|
||||
func validateSystemOneImages(req *schema.SystemOneRequest) error {
|
||||
images, err := systemone.CollectImages(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return systemone.ValidateImages(images)
|
||||
}
|
||||
@@ -0,0 +1,118 @@
|
||||
package localai
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"github.com/labstack/echo/v4"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/systemone"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
)
|
||||
|
||||
var _ = Describe("SystemOne image admission", func() {
|
||||
It("rejects top-level and structured-state images on the NER path", func() {
|
||||
for _, body := range []string{
|
||||
`{"state":"x","images":["data:image/png;base64,AA=="],"questions":{"q":{"type":"noul","instructions":"x"}}}`,
|
||||
`{"state":[{"role":"user","content":[{"type":"image_url","image_url":{"url":"data:image/png;base64,AA=="}}]}],"questions":{"q":{"type":"noul","instructions":"x"}}}`,
|
||||
} {
|
||||
var req schema.SystemOneRequest
|
||||
Expect(json.Unmarshal([]byte(body), &req)).To(Succeed())
|
||||
_, err := parseSystemOneRequest(&req)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(systemOneInputStatus(err)).To(Equal(http.StatusNotImplemented))
|
||||
}
|
||||
})
|
||||
It("bounds image count and aggregate encoded image bytes without changing the router cap", func() {
|
||||
req := schema.SystemOneRequest{State: json.RawMessage(`"x"`)}
|
||||
req.Images = json.RawMessage(`[` + strings.TrimSuffix(strings.Repeat(`"data:image/png;base64,AA==",`, 9), ",") + `]`)
|
||||
Expect(validateSystemOneImages(&req)).NotTo(Succeed())
|
||||
req.Images = json.RawMessage(` ["data:image/png;base64,` + strings.Repeat("A", systemOneMaxImageBytes) + `"]`)
|
||||
Expect(systemOneInputStatus(validateSystemOneImages(&req))).To(Equal(http.StatusRequestEntityTooLarge))
|
||||
})
|
||||
It("retains the raw wire cap and rejects overflow before binding", func() {
|
||||
e := echo.New()
|
||||
r := httptest.NewRequest(http.MethodPost, "/v1/systemone", strings.NewReader(`{"state":"`+strings.Repeat("x", systemOneMaxBody)+`"}`))
|
||||
r.Header.Set(echo.HeaderContentType, echo.MIMEApplicationJSON)
|
||||
c := e.NewContext(r, httptest.NewRecorder())
|
||||
var req schema.SystemOneRequest
|
||||
Expect(systemOneBindStatus(systemOneBind(c, &req))).To(Equal(http.StatusRequestEntityTooLarge))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("SystemOne bounded wire and image parsing", func() {
|
||||
It("rejects trailing JSON and counts trailing whitespace toward the wire cap", func() {
|
||||
for _, item := range []struct {
|
||||
body string
|
||||
code int
|
||||
}{{`{} {}`, 400}, {`{"state":` + strings.Repeat(" ", systemOneMaxBody), 413}, {`{"questions":5,"state":"` + strings.Repeat("x", systemOneMaxBody) + `"}`, 413}, {`{}` + strings.Repeat(" ", systemOneMaxBody), 413}} {
|
||||
e := echo.New()
|
||||
r := httptest.NewRequest(http.MethodPost, "/v1/systemone", strings.NewReader(item.body))
|
||||
r.Header.Set(echo.HeaderContentType, echo.MIMEApplicationJSON)
|
||||
var req schema.SystemOneRequest
|
||||
Expect(systemOneBindStatus(systemOneBind(e.NewContext(r, httptest.NewRecorder()), &req))).To(Equal(item.code))
|
||||
}
|
||||
})
|
||||
It("rejects empty MIME subtype or image data", func() {
|
||||
for _, image := range []string{"data:image/;base64,", "data:image/png;base64,", "data:image/;base64,AA=="} {
|
||||
data, err := json.Marshal([]string{image})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(validateSystemOneImages(&schema.SystemOneRequest{State: json.RawMessage(`"x"`), Images: data})).NotTo(Succeed())
|
||||
}
|
||||
})
|
||||
})
|
||||
|
||||
type unreadDecisionBody struct{}
|
||||
|
||||
func (unreadDecisionBody) Read([]byte) (int, error) {
|
||||
Fail("saturated admission read request body")
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
var _ = Describe("HTTP decision admission", func() {
|
||||
It("rejects saturation before buffering on all decision handlers", func() {
|
||||
var releases []func()
|
||||
defer func() {
|
||||
for _, r := range releases {
|
||||
r()
|
||||
}
|
||||
}()
|
||||
for i := 0; i < systemone.MaxAdmissions; i++ {
|
||||
r, err := systemone.AcquireAdmission(context.Background())
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
releases = append(releases, r)
|
||||
}
|
||||
for _, handler := range []echo.HandlerFunc{SystemOneEndpoint(nil), SystemOnePermuteEndpoint(nil), SystemOneSeparateEndpoint(nil)} {
|
||||
e := echo.New()
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodPost, "/v1/systemone", unreadDecisionBody{})
|
||||
Expect(handler(e.NewContext(req, w))).To(Succeed())
|
||||
Expect(w.Code).To(Equal(503))
|
||||
}
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("Exact public decision wire budgets", func() {
|
||||
It("accepts exact text/image wire limits but rejects one more byte", func() {
|
||||
for _, item := range []struct {
|
||||
body string
|
||||
limit int
|
||||
}{{`{"state":"x"}`, systemone.MaxBodyBytes}, {`{"state":{},"images":["data:image/png;base64,AA=="]}`, systemone.MaxImageBodyBytes}} {
|
||||
for _, extra := range []int{0, 1} {
|
||||
body := item.body + strings.Repeat(" ", item.limit-len(item.body)+extra)
|
||||
req := httptest.NewRequest(http.MethodPost, "/v1/systemone", strings.NewReader(body))
|
||||
req.Header.Set(echo.HeaderContentType, echo.MIMEApplicationJSON)
|
||||
var value schema.SystemOneRequest
|
||||
err := systemOneBind(echo.New().NewContext(req, httptest.NewRecorder()), &value)
|
||||
if extra == 0 {
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
} else {
|
||||
Expect(systemOneBindStatus(err)).To(Equal(413))
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,84 @@
|
||||
package localai
|
||||
|
||||
import (
|
||||
"context"
|
||||
"github.com/labstack/echo/v4"
|
||||
"github.com/mudler/LocalAI/core/http/auth"
|
||||
"github.com/mudler/LocalAI/core/http/middleware"
|
||||
"github.com/mudler/LocalAI/core/services/routing/billing"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
)
|
||||
|
||||
var _ = Describe("SystemOne native response handling", func() {
|
||||
It("maps backend request and capability errors without changing unknown failures", func() {
|
||||
Expect(systemOneBackendStatus(status.Error(codes.InvalidArgument, "bad question"))).To(Equal(http.StatusBadRequest))
|
||||
Expect(systemOneBackendStatus(status.Error(codes.Unimplemented, "no decision metadata"))).To(Equal(http.StatusNotImplemented))
|
||||
Expect(systemOneBackendStatus(status.Error(codes.Internal, "failed"))).To(Equal(http.StatusInternalServerError))
|
||||
})
|
||||
It("stamps explicit zeros but not missing usage", func() {
|
||||
for _, body := range []string{`{"usage":{"input_tokens":0,"output_tokens":0}}`, `{"usage":{"input_tokens":12,"output_tokens":0}}`} {
|
||||
c := echo.New().NewContext(httptest.NewRequest(http.MethodPost, "/v1/systemone", nil), httptest.NewRecorder())
|
||||
Expect(stampSystemOneUsage(c, "decision", body)).To(Succeed())
|
||||
Expect(c.Get(middleware.ContextKeyCompletionTokens)).To(Equal(int64(0)))
|
||||
Expect(c.Get(middleware.ContextKeyResponseModel)).To(Equal("decision"))
|
||||
}
|
||||
for _, body := range []string{`{}`, `{"usage":{}}`} {
|
||||
c := echo.New().NewContext(httptest.NewRequest(http.MethodPost, "/v1/systemone", nil), httptest.NewRecorder())
|
||||
Expect(stampSystemOneUsage(c, "decision", body)).To(Succeed())
|
||||
Expect(c.Get(middleware.ContextKeyPromptTokens)).To(BeNil())
|
||||
}
|
||||
})
|
||||
It("rejects negative or invalid usage instead of recording it", func() {
|
||||
c := echo.New().NewContext(httptest.NewRequest(http.MethodPost, "/v1/systemone", nil), httptest.NewRecorder())
|
||||
Expect(stampSystemOneUsage(c, "decision", `{"usage":{"input_tokens":-1,"output_tokens":0}}`)).NotTo(Succeed())
|
||||
Expect(c.Get(middleware.ContextKeyPromptTokens)).To(BeNil())
|
||||
})
|
||||
})
|
||||
|
||||
type decisionUsageCapture struct{ records []*auth.UsageRecord }
|
||||
|
||||
func (b *decisionUsageCapture) Record(_ context.Context, r *auth.UsageRecord) error {
|
||||
b.records = append(b.records, r)
|
||||
return nil
|
||||
}
|
||||
func (*decisionUsageCapture) Aggregate(context.Context, billing.AggregateQuery) ([]auth.UsageBucket, error) {
|
||||
return nil, nil
|
||||
}
|
||||
func (*decisionUsageCapture) Close() error { return nil }
|
||||
|
||||
var _ = Describe("SystemOne native route accounting", func() {
|
||||
DescribeTable("records exactly once and never fabricates output", func(body string, wantStatus, wantRecords int) {
|
||||
capture := &decisionUsageCapture{}
|
||||
e := echo.New()
|
||||
e.POST("/v1/systemone", func(c echo.Context) error {
|
||||
return respondSystemOne(c, "decision", func(context.Context) (string, error) { return body, nil })
|
||||
}, middleware.UsageMiddleware(billing.NewRecorder(capture), &auth.User{ID: "test"}))
|
||||
w := httptest.NewRecorder()
|
||||
e.ServeHTTP(w, httptest.NewRequest(http.MethodPost, "/v1/systemone", nil))
|
||||
Expect(w.Code).To(Equal(wantStatus))
|
||||
Expect(capture.records).To(HaveLen(wantRecords))
|
||||
if wantRecords > 0 {
|
||||
Expect(capture.records[0].CompletionTokens).To(Equal(int64(0)))
|
||||
}
|
||||
},
|
||||
Entry("positive input zero output", `{"answers":{"q":{"type":"noul","noul":0.5}},"usage":{"input_tokens":12,"output_tokens":0}}`, 200, 1),
|
||||
Entry("explicit zeros", `{"answers":{"q":{"type":"noul","noul":0.5}},"usage":{"input_tokens":0,"output_tokens":0}}`, 200, 1),
|
||||
Entry("missing", `{"answers":{"q":{"type":"noul","noul":0.5}}}`, 200, 0),
|
||||
Entry("incomplete", `{"answers":{"q":{"type":"noul","noul":0.5}},"usage":{"input_tokens":12}}`, 500, 0),
|
||||
Entry("negative", `{"answers":{"q":{"type":"noul","noul":0.5}},"usage":{"input_tokens":-1,"output_tokens":0}}`, 500, 0),
|
||||
)
|
||||
DescribeTable("maps RPC errors at the HTTP route", func(code codes.Code, want int) {
|
||||
e := echo.New()
|
||||
e.POST("/v1/systemone", func(c echo.Context) error {
|
||||
return respondSystemOne(c, "decision", func(context.Context) (string, error) { return "", status.Error(code, "backend error") })
|
||||
})
|
||||
w := httptest.NewRecorder()
|
||||
e.ServeHTTP(w, httptest.NewRequest(http.MethodPost, "/v1/systemone", nil))
|
||||
Expect(w.Code).To(Equal(want))
|
||||
}, Entry("invalid", codes.InvalidArgument, 400), Entry("unsupported", codes.Unimplemented, 501))
|
||||
})
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
|
||||
"github.com/labstack/echo/v4"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/systemone"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
@@ -76,6 +77,23 @@ var _ = Describe("systemOneBind", func() {
|
||||
return http.StatusOK, nil
|
||||
}
|
||||
|
||||
It("accepts HTTP text whose JSON escaping exceeds the internal serialized bound", func() {
|
||||
body := `{"state":"` + strings.Repeat("<", 11000) + `","questions":{"q":{"type":"noul"}}}`
|
||||
Expect(len(body)).To(BeNumerically("<", systemOneMaxBody))
|
||||
e := echo.New()
|
||||
r := httptest.NewRequest(http.MethodPost, "/v1/systemone", strings.NewReader(body))
|
||||
r.Header.Set("Content-Type", "application/json")
|
||||
c := e.NewContext(r, httptest.NewRecorder())
|
||||
var out schema.SystemOneRequest
|
||||
Expect(systemOneBind(c, &out)).To(Succeed())
|
||||
serialized, err := json.Marshal(out)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(len(serialized)).To(BeNumerically(">", systemOneMaxBody))
|
||||
Expect(validateSystemOneRequest(&out)).To(Succeed())
|
||||
// The internal transport still bounds its actual serialized representation.
|
||||
Expect(systemone.ValidateRequest(&out)).To(MatchError(ContainSubstring("exceeds 64 KiB")))
|
||||
})
|
||||
|
||||
It("binds a normal body", func() {
|
||||
status, err := bind(`{"model":"m","state":"x","questions":{}}`)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
model "github.com/mudler/LocalAI/pkg/model"
|
||||
"gorm.io/gorm"
|
||||
"slices"
|
||||
)
|
||||
|
||||
// ListModelCapabilitiesEndpoint is a LocalAI-specific extension of the OpenAI
|
||||
@@ -43,6 +44,14 @@ func ListModelCapabilitiesEndpoint(bcl *config.ModelConfigLoader, ml *model.Mode
|
||||
cfg.ContextSize = &appConfig.ContextSize
|
||||
}
|
||||
entry.Capabilities = cfg.Capabilities()
|
||||
// Generation aliases inherit target capabilities, but the router's
|
||||
// native classifier loads the named config directly (no alias resolution).
|
||||
original, exists := bcl.GetModelConfig(m)
|
||||
if !exists || !original.NativeDecisionsEligible() {
|
||||
entry.Capabilities = slices.DeleteFunc(entry.Capabilities, func(capability string) bool {
|
||||
return capability == config.UsecaseDecisions
|
||||
})
|
||||
}
|
||||
entry.ThreeDOperations = cfg.ThreeDOperations()
|
||||
entry.InputModalities = cfg.InputModalities()
|
||||
entry.OutputModalities = cfg.OutputModalities()
|
||||
|
||||
@@ -76,6 +76,28 @@ var _ = Describe("ListModelCapabilitiesEndpoint", func() {
|
||||
return nil
|
||||
}
|
||||
|
||||
It("does not inherit native decisions eligibility through aliases", func() {
|
||||
writeConfig("native", "name: native\nbackend: llama-cpp\nknown_usecases: [decisions]\n")
|
||||
writeConfig("native-vllm", "name: native-vllm\nbackend: vllm-cpp\nknown_usecases: [decisions]\n")
|
||||
writeConfig("decision-alias", "name: decision-alias\nalias: native\n")
|
||||
writeConfig("disabled-native", "name: disabled-native\nbackend: llama-cpp\nknown_usecases: [decisions]\ndisabled: true\n")
|
||||
writeConfig("disabled-alias", "name: disabled-alias\nalias: disabled-native\n")
|
||||
resp := call()
|
||||
for _, name := range []string{"native", "native-vllm"} {
|
||||
entry := entryFor(resp, name)
|
||||
Expect(entry).NotTo(BeNil())
|
||||
Expect(entry.Capabilities).To(ContainElement(config.UsecaseDecisions))
|
||||
}
|
||||
alias := entryFor(resp, "decision-alias")
|
||||
Expect(alias).NotTo(BeNil())
|
||||
Expect(alias.Capabilities).NotTo(ContainElement(config.UsecaseDecisions))
|
||||
for _, name := range []string{"disabled-native", "disabled-alias"} {
|
||||
if entry := entryFor(resp, name); entry != nil {
|
||||
Expect(entry.Capabilities).NotTo(ContainElement(config.UsecaseDecisions))
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
It("returns the list envelope even with no models", func() {
|
||||
resp := call()
|
||||
Expect(resp.Object).To(Equal("list"))
|
||||
|
||||
@@ -0,0 +1,76 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package middleware_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
"github.com/mudler/LocalAI/core/backend"
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
. "github.com/mudler/LocalAI/core/http/middleware"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/services/routing/router"
|
||||
"github.com/mudler/LocalAI/pkg/model"
|
||||
"github.com/mudler/LocalAI/pkg/system"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("native decisions integration", Label("real-models"), func() {
|
||||
It("routes overlapping labels through the central factory and native runner", func() {
|
||||
weights, addr := os.Getenv("LOCALAI_DECISIONS_TEST_MODEL"), os.Getenv("LOCALAI_DECISIONS_TEST_GRPC")
|
||||
if weights == "" || addr == "" {
|
||||
Skip("requires an existing decision GGUF and a dedicated llama-cpp gRPC server")
|
||||
}
|
||||
abs, err := filepath.Abs(weights)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
_, err = os.Stat(abs)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
state := &system.SystemState{Model: system.Model{ModelsPath: filepath.Dir(abs)}}
|
||||
loader := model.NewModelLoader(state)
|
||||
app := config.NewApplicationConfig(config.WithSystemState(state), config.WithExternalBackend("llama-cpp", addr))
|
||||
flags, size, threads := config.FLAG_DECISIONS, 2048, 2
|
||||
native := &config.ModelConfig{Name: "native-integration", Backend: "llama-cpp", KnownUsecases: &flags}
|
||||
native.ContextSize = &size
|
||||
native.Model = filepath.Base(abs)
|
||||
native.Threads = &threads
|
||||
lookup := func(name string) *config.ModelConfig {
|
||||
if name == native.Name {
|
||||
return native
|
||||
}
|
||||
return nil
|
||||
}
|
||||
runner := backend.NewDecisionRunner(native.Name, lookup, loader, app)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
|
||||
defer cancel()
|
||||
// This separately checks the actual request/usage contract, not just labels.
|
||||
response, err := runner.Decide(ctx, &schema.SystemOneRequest{State: json.RawMessage(`"I was charged twice and need a refund today."`), Questions: map[string]schema.SystemOneQuestion{"refund": {Type: "noul", Instructions: json.RawMessage(`"Does the user request a refund?"`), Criteria: json.RawMessage(`{"false":"No refund requested","true":"Refund requested"}`)}}})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(response.Answers).To(HaveKey("refund"))
|
||||
Expect(response.Answers["refund"].Noul).NotTo(BeNil())
|
||||
Expect(response.Usage.InputTokens).To(BeNumerically(">", 0))
|
||||
Expect(response.Usage.OutputTokens).To(Equal(0))
|
||||
raw, err := json.Marshal(response)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
GinkgoWriter.Printf("native contract: %s\n", raw)
|
||||
// A deliberately low positive threshold checks overlapping-label plumbing,
|
||||
// not model quality. Both policies are independent binary questions.
|
||||
cfg := &config.ModelConfig{Name: "integration-router", Router: config.RouterConfig{Classifier: "decisions", ClassifierModel: native.Name, ActivationThreshold: 0.000001, Policies: []config.RouterPolicy{{Label: "billing", Description: "The user discusses a charge or payment."}, {Label: "refund", Description: "The user requests a refund."}}, Candidates: []config.RouterCandidate{{Model: "billing-only", Labels: []string{"billing"}}, {Model: "combined-target", Labels: []string{"billing", "refund"}}}}}
|
||||
classifier, err := GetOrBuildClassifier(router.NewRegistry(), cfg, ClassifierDeps{ModelLookup: lookup, Decisions: func(string) backend.DecisionRunner { return runner }})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
result, err := router.Resolve(ctx, cfg, classifier, func(name string) (*config.ModelConfig, error) { return &config.ModelConfig{Name: name}, nil }, router.Probe{Prompt: "I was charged twice and need a refund today."})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(result.UsedFallback).To(BeFalse())
|
||||
Expect(result.Labels).To(ConsistOf("billing", "refund"))
|
||||
Expect(result.Decision.LabelScores).To(HaveLen(2))
|
||||
for _, score := range result.Decision.LabelScores {
|
||||
Expect(score.Score).To(BeNumerically(">=", 0))
|
||||
Expect(score.Score).To(BeNumerically("<=", 1))
|
||||
}
|
||||
Expect(result.ChosenModel).To(Equal("combined-target"))
|
||||
GinkgoWriter.Printf("native routing: scores=%+v labels=%v candidate=%s\n", result.Decision.LabelScores, result.Labels, result.ChosenModel)
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,50 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package middleware_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"github.com/mudler/LocalAI/core/backend"
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
. "github.com/mudler/LocalAI/core/http/middleware"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/services/routing/router"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
type nativeRunner struct{}
|
||||
|
||||
func (nativeRunner) Decide(context.Context, *schema.SystemOneRequest) (*schema.SystemOneResponse, error) {
|
||||
v := .8
|
||||
return &schema.SystemOneResponse{Answers: map[string]schema.SystemOneAnswer{"p0": {Type: "noul", Noul: &v}}}, nil
|
||||
}
|
||||
|
||||
var _ = Describe("decisions central factory", func() {
|
||||
It("requires explicit usecase and invalidates on native model edits", func() {
|
||||
cfg := &config.ModelConfig{Name: "router", Router: config.RouterConfig{Classifier: "decisions", ClassifierModel: "native", Policies: []config.RouterPolicy{{Label: "code", Description: "coding"}}, Candidates: []config.RouterCandidate{{Model: "target", Labels: []string{"code"}}}}}
|
||||
flags := config.FLAG_DECISIONS
|
||||
model := &config.ModelConfig{Name: "native", Backend: "vllm-cpp", KnownUsecases: &flags}
|
||||
deps := ClassifierDeps{Decisions: func(string) backend.DecisionRunner { return nativeRunner{} }, ModelLookup: func(string) *config.ModelConfig { return model }}
|
||||
registry := router.NewRegistry()
|
||||
first, err := GetOrBuildClassifier(registry, cfg, deps)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(first.Name()).To(Equal("decisions"))
|
||||
same, err := GetOrBuildClassifier(registry, cfg, deps)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(same).To(BeIdenticalTo(first))
|
||||
model.Model = "new-weights.gguf"
|
||||
changed, err := GetOrBuildClassifier(registry, cfg, deps)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(changed).NotTo(BeIdenticalTo(first))
|
||||
Expect(model.StampPersistedConfigRevision()).To(Succeed())
|
||||
revised, err := GetOrBuildClassifier(registry, cfg, deps)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(revised).NotTo(BeIdenticalTo(changed))
|
||||
model.KnownUsecases = nil
|
||||
_, err = GetOrBuildClassifier(registry, cfg, deps)
|
||||
Expect(err).To(HaveOccurred())
|
||||
deps.ModelLookup = nil
|
||||
_, err = GetOrBuildClassifier(router.NewRegistry(), cfg, deps)
|
||||
Expect(err).To(HaveOccurred())
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,259 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package middleware_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"github.com/mudler/LocalAI/core/backend"
|
||||
"github.com/mudler/LocalAI/core/systemone"
|
||||
"github.com/mudler/LocalAI/pkg/system"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"image"
|
||||
"image/png"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"reflect"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
|
||||
"github.com/labstack/echo/v4"
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
. "github.com/mudler/LocalAI/core/http/middleware"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/services/routing/router"
|
||||
)
|
||||
|
||||
var _ = It("TestRetryRemoteImageMiddleware", func() {
|
||||
var calls atomic.Int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { calls.Add(1); _, _ = w.Write([]byte("image")) }))
|
||||
defer server.Close()
|
||||
GinkgoT().Setenv("HTTP_PROXY", server.URL)
|
||||
GinkgoT().Setenv("NO_PROXY", "")
|
||||
dir := GinkgoT().TempDir()
|
||||
cfg := newScoreRouterModel(dir, "smart-router")
|
||||
cfg.Router.Fallback = ""
|
||||
cfg.Router.Classifier = router.ClassifierDecisions
|
||||
app := &config.ApplicationConfig{Context: context.Background(), SystemState: &system.SystemState{Model: system.Model{ModelsPath: dir}}}
|
||||
loader := config.NewModelConfigLoader(dir)
|
||||
req := openAIChat("")
|
||||
req.Messages[0].Content = []any{map[string]any{"type": "image_url", "image_url": map[string]any{"url": "http://93.184.215.14/decision.png"}}}
|
||||
e := echo.New()
|
||||
c := e.NewContext(httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil), httptest.NewRecorder())
|
||||
c.Set(CONTEXT_LOCALS_KEY_LOCALAI_REQUEST, req)
|
||||
c.Set(CONTEXT_LOCALS_KEY_MODEL_CONFIG, cfg)
|
||||
re := NewRequestExtractor(loader, nil, app)
|
||||
// Same ordered schema parsing and routing middleware as registered chat.
|
||||
if err := re.SetOpenAIRequest(c); err != nil {
|
||||
Fail(fmt.Sprint(err))
|
||||
}
|
||||
defer req.Cancel()
|
||||
h := RouteModel(loader, app, nil, nil, OpenAIProbe, router.SourceChat, retryDecisionDeps())(func(echo.Context) error { Fail(fmt.Sprint("rejected image reached endpoint")); return nil })
|
||||
if err := h(c); err == nil || !strings.Contains(err.Error(), "images must be PNG or JPEG base64 data URLs") {
|
||||
Fail(fmt.Sprintf("expected shared image validation, got %v", err))
|
||||
}
|
||||
if got := calls.Load(); got != 0 {
|
||||
Fail(fmt.Sprintf("decision input performed %d network calls", got))
|
||||
}
|
||||
})
|
||||
|
||||
var _ = It("TestRetryTypedOpenAIText", func() {
|
||||
var typed []schema.Content
|
||||
raw := []byte(`[{"type":"text","text":"first"},{"type":"image_url","image_url":{"url":"data:image/png;base64,AA=="}},{"type":"text","text":"second"}]`)
|
||||
if err := json.Unmarshal(raw, &typed); err != nil {
|
||||
Fail(fmt.Sprint(err))
|
||||
}
|
||||
var untyped []any
|
||||
Expect(json.Unmarshal(raw, &untyped)).To(Succeed())
|
||||
a := OpenAIProbeFromRequest(&schema.OpenAIRequest{Messages: []schema.Message{{Role: "user", Content: typed}}})
|
||||
b := OpenAIProbeFromRequest(&schema.OpenAIRequest{Messages: []schema.Message{{Role: "user", Content: untyped}}})
|
||||
if a.Prompt != b.Prompt || a.Prompt == "" {
|
||||
Fail(fmt.Sprintf("typed %q untyped %q", a.Prompt, b.Prompt))
|
||||
}
|
||||
})
|
||||
|
||||
type retryDecisionFunc func(context.Context, *schema.SystemOneRequest) (*schema.SystemOneResponse, error)
|
||||
|
||||
func (f retryDecisionFunc) Decide(ctx context.Context, r *schema.SystemOneRequest) (*schema.SystemOneResponse, error) {
|
||||
return f(ctx, r)
|
||||
}
|
||||
func retryDecisionDeps() ClassifierDeps {
|
||||
return ClassifierDeps{
|
||||
ModelLookup: func(string) *config.ModelConfig {
|
||||
u := config.FLAG_DECISIONS
|
||||
return &config.ModelConfig{Backend: "llama-cpp", KnownUsecases: &u}
|
||||
},
|
||||
Decisions: func(string) backend.DecisionRunner {
|
||||
return retryDecisionFunc(func(context.Context, *schema.SystemOneRequest) (*schema.SystemOneResponse, error) {
|
||||
Fail(fmt.Sprint("invalid image reached runner"))
|
||||
return nil, nil
|
||||
})
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
var _ = It("TestRetryConstructionFailsClosed", func() {
|
||||
dir := GinkgoT().TempDir()
|
||||
cfg := newScoreRouterModel(dir, "smart-router")
|
||||
writeCandidate(dir, cfg.Router.Fallback)
|
||||
req := openAIChat("original")
|
||||
rec, err := runRouterWithDeps(config.NewModelConfigLoader(dir), &config.ApplicationConfig{Context: context.Background(), SystemState: &system.SystemState{Model: system.Model{ModelsPath: dir}}}, nil, cfg, req, ClassifierDeps{})
|
||||
if err == nil || !strings.Contains(err.Error(), "no scorer factory") || rec.Body.Len() != 0 {
|
||||
Fail(fmt.Sprintf("construction must fail closed: %v", err))
|
||||
}
|
||||
if req.Messages[0].Content != "original" {
|
||||
Fail(fmt.Sprint("payload changed"))
|
||||
}
|
||||
})
|
||||
|
||||
var _ = It("TestRetryMaterialization", func() {
|
||||
for _, mode := range []string{"ordinary", "success", "fallback"} {
|
||||
By(mode)
|
||||
func() {
|
||||
dir := GinkgoT().TempDir()
|
||||
cfg := newScoreRouterModel(dir, "smart-router")
|
||||
writeCandidate(dir, "small-model")
|
||||
writeCandidate(dir, cfg.Router.Fallback)
|
||||
if mode == "ordinary" {
|
||||
cfg.Router = config.RouterConfig{}
|
||||
}
|
||||
req := openAIChat("")
|
||||
req.Messages[0].Content = []any{map[string]any{"type": "text", "text": "original"}, map[string]any{"type": "image_url", "image_url": map[string]any{"url": "data:image/png;base64,AA=="}}}
|
||||
expectedImage := "AA=="
|
||||
if mode == "success" {
|
||||
var b bytes.Buffer
|
||||
if err := png.Encode(&b, image.NewGray(image.Rect(0, 0, 1, 1))); err != nil {
|
||||
Fail(fmt.Sprint(err))
|
||||
}
|
||||
expectedImage = base64.StdEncoding.EncodeToString(b.Bytes())
|
||||
req.Messages[0].Content.([]any)[1].(map[string]any)["image_url"] = map[string]any{"url": "data:image/png;base64," + expectedImage}
|
||||
raw, _ := json.Marshal(req.Messages[0].Content)
|
||||
var typed []schema.Content
|
||||
Expect(json.Unmarshal(raw, &typed)).To(Succeed())
|
||||
req.Messages[0].Content = typed
|
||||
}
|
||||
before, _ := json.Marshal(req.Messages[0].Content)
|
||||
app := &config.ApplicationConfig{Context: context.Background(), SystemState: &system.SystemState{Model: system.Model{ModelsPath: dir}}}
|
||||
loader := config.NewModelConfigLoader(dir)
|
||||
c := echo.New().NewContext(httptest.NewRequest(http.MethodPost, "/", nil), httptest.NewRecorder())
|
||||
c.Set(CONTEXT_LOCALS_KEY_LOCALAI_REQUEST, req)
|
||||
c.Set(CONTEXT_LOCALS_KEY_MODEL_CONFIG, cfg)
|
||||
if err := NewRequestExtractor(loader, nil, app).SetOpenAIRequest(c); err != nil {
|
||||
Fail(fmt.Sprint(err))
|
||||
}
|
||||
defer req.Cancel()
|
||||
if mode != "ordinary" && len(req.Messages[0].StringImages) != 0 {
|
||||
Fail(fmt.Sprint("premature image preparation"))
|
||||
}
|
||||
deps := ClassifierDeps{Scorer: func(string) backend.Scorer { return &stubScorer{} }}
|
||||
extractor := OpenAIProbe
|
||||
if mode == "success" {
|
||||
cfg.Router.Classifier = router.ClassifierDecisions
|
||||
deps = retryDecisionDeps()
|
||||
deps.Decisions = func(string) backend.DecisionRunner {
|
||||
return retryDecisionFunc(func(_ context.Context, r *schema.SystemOneRequest) (*schema.SystemOneResponse, error) {
|
||||
if !bytes.Contains(r.State, []byte(expectedImage)) {
|
||||
Fail(fmt.Sprint("image missing from native probe"))
|
||||
}
|
||||
answers := map[string]schema.SystemOneAnswer{}
|
||||
for key := range r.Questions {
|
||||
v := 0.0
|
||||
if key == "p1" {
|
||||
v = 1
|
||||
}
|
||||
answers[key] = schema.SystemOneAnswer{Type: "noul", Noul: &v}
|
||||
}
|
||||
return &schema.SystemOneResponse{Answers: answers}, nil
|
||||
})
|
||||
}
|
||||
}
|
||||
reached := false
|
||||
err := RouteModel(loader, app, nil, nil, extractor, router.SourceChat, deps)(func(echo.Context) error { reached = true; return nil })(c)
|
||||
if err != nil || !reached {
|
||||
Fail(fmt.Sprintf("dispatch: %v", err))
|
||||
}
|
||||
if len(req.Messages[0].StringImages) != 1 || req.Messages[0].StringImages[0] != expectedImage {
|
||||
Fail(fmt.Sprintf("images lost: %#v", req.Messages[0]))
|
||||
}
|
||||
after, _ := json.Marshal(req.Messages[0].Content)
|
||||
if string(before) != string(after) {
|
||||
Fail(fmt.Sprint("original content changed"))
|
||||
}
|
||||
}()
|
||||
}
|
||||
})
|
||||
|
||||
var _ = It("TestRetryOrderedProbes", func() {
|
||||
for _, api := range []string{"openai", "anthropic"} {
|
||||
raw := `[{"type":"text","text":"first"},{"type":"image_url","image_url":{"url":"data:image/png;base64,AA=="}},{"type":"text","text":"second"},{"type":"image_url","image_url":{"url":"data:image/png;base64,AQ=="}}]`
|
||||
if api == "anthropic" {
|
||||
raw = `[{"type":"text","text":"first"},{"type":"image","source":{"type":"base64","media_type":"image/png","data":"AA=="}},{"type":"text","text":"second"},{"type":"image","source":{"type":"base64","media_type":"image/png","data":"AQ=="}}]`
|
||||
}
|
||||
var untyped []any
|
||||
if err := json.Unmarshal([]byte(raw), &untyped); err != nil {
|
||||
Fail(fmt.Sprint(err))
|
||||
}
|
||||
var typed any
|
||||
if api == "openai" {
|
||||
var blocks []schema.Content
|
||||
Expect(json.Unmarshal([]byte(raw), &blocks)).To(Succeed())
|
||||
typed = blocks
|
||||
} else {
|
||||
var blocks []schema.AnthropicContentBlock
|
||||
Expect(json.Unmarshal([]byte(raw), &blocks)).To(Succeed())
|
||||
typed = blocks
|
||||
}
|
||||
for _, content := range []any{typed, untyped} {
|
||||
var p router.Probe
|
||||
if api == "openai" {
|
||||
p = OpenAIProbeFromRequest(&schema.OpenAIRequest{Messages: []schema.Message{{Role: "user", Content: content}}})
|
||||
} else {
|
||||
p, _ = AnthropicProbe(&schema.AnthropicRequest{Messages: []schema.AnthropicMessage{{Role: "user", Content: content}}})
|
||||
}
|
||||
if p.InputError != nil || p.Prompt != "first\nsecond\n" {
|
||||
Fail(fmt.Sprintf("%s: %#v", api, p))
|
||||
}
|
||||
images, err := systemone.CollectImages(&schema.SystemOneRequest{State: p.State})
|
||||
if err != nil || !reflect.DeepEqual(images, []string{"data:image/png;base64,AA==", "data:image/png;base64,AQ=="}) {
|
||||
Fail(fmt.Sprintf("%s order: %v %v", api, images, err))
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
var _ = It("TestRetryAnthropicAndMarshalFallback", func() {
|
||||
for _, bad := range []bool{false, true} {
|
||||
By(fmt.Sprint(bad))
|
||||
func() {
|
||||
dir := GinkgoT().TempDir()
|
||||
cfg := newScoreRouterModel(dir, "smart-router")
|
||||
writeCandidate(dir, cfg.Router.Fallback)
|
||||
content := []any{map[string]any{"type": "image", "source": map[string]any{"type": "base64", "media_type": "image/png", "data": "AA=="}}}
|
||||
if bad {
|
||||
content = append(content, map[string]any{"type": "text", "text": "keep", "unencodable": make(chan int)})
|
||||
}
|
||||
req := &schema.AnthropicRequest{Model: cfg.Name, Messages: []schema.AnthropicMessage{{Role: "user", Content: content}}}
|
||||
original := req.Messages[0].Content
|
||||
probe, _ := AnthropicProbe(req)
|
||||
if bad && probe.InputError == nil {
|
||||
Fail(fmt.Sprint("marshal error lost"))
|
||||
}
|
||||
app := &config.ApplicationConfig{Context: context.Background(), SystemState: &system.SystemState{Model: system.Model{ModelsPath: dir}}}
|
||||
c := echo.New().NewContext(httptest.NewRequest(http.MethodPost, "/v1/messages", nil), httptest.NewRecorder())
|
||||
c.Set(CONTEXT_LOCALS_KEY_LOCALAI_REQUEST, req)
|
||||
c.Set(CONTEXT_LOCALS_KEY_MODEL_CONFIG, cfg)
|
||||
reached := false
|
||||
err := RouteModel(config.NewModelConfigLoader(dir), app, nil, nil, AnthropicProbe, router.SourceAnthropic, ClassifierDeps{Scorer: func(string) backend.Scorer { return &stubScorer{} }})(func(echo.Context) error { reached = true; return nil })(c)
|
||||
if err != nil || !reached || req.Model != cfg.Router.Fallback {
|
||||
Fail(fmt.Sprintf("fallback: %v", err))
|
||||
}
|
||||
if !reflect.DeepEqual(original, req.Messages[0].Content) {
|
||||
Fail(fmt.Sprint("payload changed"))
|
||||
}
|
||||
}()
|
||||
}
|
||||
})
|
||||
@@ -0,0 +1,97 @@
|
||||
package middleware_test
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
. "github.com/mudler/LocalAI/core/http/middleware"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/systemone"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"runtime"
|
||||
"strings"
|
||||
)
|
||||
|
||||
var _ = Describe("Probe extraction bounds", func() {
|
||||
for _, api := range []string{"openai", "anthropic"} {
|
||||
for _, typed := range []bool{false, true} {
|
||||
for _, field := range []string{"text", "image"} {
|
||||
It(api+" bounds allocations before copying "+field, func() {
|
||||
huge := strings.Repeat("x", systemone.MaxImageBodyBytes+1)
|
||||
raw := `[{"type":"text","text":"` + huge + `"}]`
|
||||
if field == "image" {
|
||||
raw = `[{"type":"image_url","image_url":{"url":"` + huge + `"}}]`
|
||||
if api == "anthropic" {
|
||||
raw = `[{"type":"image","source":{"type":"base64","data":"` + huge + `"}}]`
|
||||
}
|
||||
}
|
||||
var content any
|
||||
if typed {
|
||||
if api == "openai" {
|
||||
var v []schema.Content
|
||||
Expect(json.Unmarshal([]byte(raw), &v)).To(Succeed())
|
||||
content = v
|
||||
} else {
|
||||
var v []schema.AnthropicContentBlock
|
||||
Expect(json.Unmarshal([]byte(raw), &v)).To(Succeed())
|
||||
content = v
|
||||
}
|
||||
} else {
|
||||
Expect(json.Unmarshal([]byte(raw), &content)).To(Succeed())
|
||||
}
|
||||
runtime.GC()
|
||||
var before, after runtime.MemStats
|
||||
runtime.ReadMemStats(&before)
|
||||
if api == "openai" {
|
||||
p := OpenAIProbeFromRequest(&schema.OpenAIRequest{Messages: []schema.Message{{Content: content}}})
|
||||
Expect(p.InputError).To(HaveOccurred())
|
||||
} else {
|
||||
p, _ := AnthropicProbe(&schema.AnthropicRequest{Messages: []schema.AnthropicMessage{{Content: content}}})
|
||||
Expect(p.InputError).To(HaveOccurred())
|
||||
}
|
||||
runtime.ReadMemStats(&after)
|
||||
Expect(after.TotalAlloc - before.TotalAlloc).To(BeNumerically("<", systemone.MaxBodyBytes))
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
It("counts JSON escaping before allocation", func() {
|
||||
req := &schema.OpenAIRequest{Messages: []schema.Message{{Content: strings.Repeat("\x00", systemone.MaxImageBodyBytes/len(`\u0000`)+1)}}}
|
||||
runtime.GC()
|
||||
var a, b runtime.MemStats
|
||||
runtime.ReadMemStats(&a)
|
||||
p := OpenAIProbeFromRequest(req)
|
||||
runtime.ReadMemStats(&b)
|
||||
Expect(p.InputError).To(HaveOccurred())
|
||||
Expect(b.TotalAlloc - a.TotalAlloc).To(BeNumerically("<", systemone.MaxBodyBytes))
|
||||
})
|
||||
It("preserves text at the ordinary text budget", func() {
|
||||
p := OpenAIProbeFromRequest(&schema.OpenAIRequest{Messages: []schema.Message{{Content: strings.Repeat("\x00", systemone.MaxBodyBytes)}}})
|
||||
Expect(p.InputError).NotTo(HaveOccurred())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("Probe escaping boundaries", func() {
|
||||
It("accepts exactly the serialized limit and rejects the next escaped byte", func() {
|
||||
const escapeBytes = len(`\u0000`)
|
||||
req := &schema.OpenAIRequest{Messages: []schema.Message{{Content: ""}}}
|
||||
empty, err := json.Marshal(req.Messages)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
available := systemone.MaxImageBodyBytes - len(empty)
|
||||
text := strings.Repeat("\x00", available/escapeBytes) + strings.Repeat("a", available%escapeBytes)
|
||||
req.Messages[0].Content = text
|
||||
p := OpenAIProbeFromRequest(req)
|
||||
Expect(p.InputError).NotTo(HaveOccurred())
|
||||
Expect(p.State).To(HaveLen(systemone.MaxImageBodyBytes))
|
||||
req.Messages[0].Content = text + "\x00"
|
||||
p = OpenAIProbeFromRequest(req)
|
||||
Expect(p.InputError).To(HaveOccurred())
|
||||
Expect(p.State).To(BeEmpty())
|
||||
Expect(p.Prompt).To(BeEmpty())
|
||||
})
|
||||
It("bounds recursive direct internal values", func() {
|
||||
cycle := map[string]any{}
|
||||
cycle["self"] = cycle
|
||||
p := OpenAIProbeFromRequest(&schema.OpenAIRequest{Messages: []schema.Message{{Content: cycle}}})
|
||||
Expect(p.InputError).To(HaveOccurred())
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,198 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"encoding"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"strings"
|
||||
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/systemone"
|
||||
)
|
||||
|
||||
// probeBudget checks the serialized size before either extracting text or
|
||||
// invoking encoding/json (whose Encoder also buffers a complete value). It
|
||||
// walks caller-owned values without copying strings or allocating map keys.
|
||||
// Unknown custom marshalers fail closed: their allocation cannot be bounded.
|
||||
func probeBudget(value any) error {
|
||||
remaining := systemone.MaxImageBodyBytes
|
||||
if !probeValueBudget(reflect.ValueOf(value), &remaining, 0) {
|
||||
return fmt.Errorf("router state exceeds decision image request budget or contains unsupported values")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
const (
|
||||
probeMaxDepth = 100 // Bound recursion for cyclic direct/internal requests.
|
||||
jsonNumberBytes = 32 // Upper bound for JSON's built-in numeric representations.
|
||||
)
|
||||
|
||||
func probeSpend(left *int, n int) bool {
|
||||
if n > *left {
|
||||
return false
|
||||
}
|
||||
*left -= n
|
||||
return true
|
||||
}
|
||||
|
||||
func probeValueBudget(v reflect.Value, left *int, depth int) bool {
|
||||
if depth > probeMaxDepth {
|
||||
return false
|
||||
}
|
||||
if !v.IsValid() {
|
||||
return probeSpend(left, len("null"))
|
||||
}
|
||||
if v.Kind() == reflect.Interface {
|
||||
if v.IsNil() {
|
||||
return probeSpend(left, len("null"))
|
||||
}
|
||||
return probeValueBudget(v.Elem(), left, depth+1)
|
||||
}
|
||||
// Only these concrete schema structs and plain JSON representations are
|
||||
// accepted. Do not emulate arbitrary encoding/json field promotion or tags.
|
||||
if v.Type() == reflect.TypeFor[*json.RawMessage]() && !v.IsNil() {
|
||||
return probeValueBudget(v.Elem(), left, depth+1)
|
||||
}
|
||||
if v.Type() == reflect.TypeFor[json.RawMessage]() {
|
||||
if v.IsNil() {
|
||||
return probeSpend(left, 4)
|
||||
}
|
||||
// RawMessage is compacted and HTML-escaped by encoding/json. Count raw
|
||||
// whitespace too; this intentionally overestimates, without executing it.
|
||||
if v.Len() > *left {
|
||||
return false
|
||||
}
|
||||
for _, b := range v.Bytes() {
|
||||
n := 1
|
||||
if b == '<' || b == '>' || b == '&' || b >= 0x80 {
|
||||
n = 6
|
||||
}
|
||||
if !probeSpend(left, n) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
typ := v.Type()
|
||||
if typ.Implements(reflect.TypeFor[json.Marshaler]()) || typ.Implements(reflect.TypeFor[encoding.TextMarshaler]()) ||
|
||||
reflect.PointerTo(typ).Implements(reflect.TypeFor[json.Marshaler]()) || reflect.PointerTo(typ).Implements(reflect.TypeFor[encoding.TextMarshaler]()) {
|
||||
return false
|
||||
}
|
||||
if typ.Name() != "" && typ.PkgPath() != "" && !probeSchemaType(typ) && typ != reflect.TypeFor[json.Number]() {
|
||||
return false
|
||||
}
|
||||
|
||||
switch v.Kind() {
|
||||
case reflect.Pointer:
|
||||
if v.IsNil() {
|
||||
return probeSpend(left, len("null"))
|
||||
}
|
||||
return probeValueBudget(v.Elem(), left, depth+1)
|
||||
case reflect.String:
|
||||
return systemone.SpendJSONString(v.String(), left)
|
||||
case reflect.Bool:
|
||||
return probeSpend(left, len("false"))
|
||||
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64, reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64, reflect.Float32, reflect.Float64:
|
||||
return probeSpend(left, jsonNumberBytes)
|
||||
case reflect.Slice, reflect.Array:
|
||||
if v.Kind() == reflect.Slice && v.IsNil() {
|
||||
return probeSpend(left, len("null"))
|
||||
}
|
||||
// encoding/json base64-encodes byte slices.
|
||||
if v.Kind() == reflect.Slice && v.Type().Elem().Kind() == reflect.Uint8 {
|
||||
return v.Type().Elem() == reflect.TypeFor[byte]() && v.Len() <= *left && probeSpend(left, 2+4*((v.Len()+2)/3))
|
||||
}
|
||||
if !probeSpend(left, 2) || v.Len() > *left {
|
||||
return false
|
||||
}
|
||||
for i := 0; i < v.Len(); i++ {
|
||||
if i > 0 && !probeSpend(left, 1) {
|
||||
return false
|
||||
}
|
||||
if !probeValueBudget(v.Index(i), left, depth+1) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
case reflect.Map:
|
||||
if v.IsNil() {
|
||||
return probeSpend(left, len("null"))
|
||||
}
|
||||
if v.Type().Key() != reflect.TypeFor[string]() || !probeSpend(left, 2) || v.Len() > *left {
|
||||
return false
|
||||
}
|
||||
iter := v.MapRange()
|
||||
first := true
|
||||
for iter.Next() {
|
||||
if !first && !probeSpend(left, 1) {
|
||||
return false
|
||||
}
|
||||
first = false
|
||||
if !systemone.SpendJSONString(iter.Key().String(), left) || !probeSpend(left, 1) || !probeValueBudget(iter.Value(), left, depth+1) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
case reflect.Struct:
|
||||
if !probeSchemaType(v.Type()) {
|
||||
return false
|
||||
}
|
||||
if !probeSpend(left, 2) {
|
||||
return false
|
||||
}
|
||||
first := true
|
||||
typ := v.Type()
|
||||
for i := 0; i < v.NumField(); i++ {
|
||||
field := typ.Field(i)
|
||||
if field.PkgPath != "" {
|
||||
continue
|
||||
}
|
||||
tag := field.Tag.Get("json")
|
||||
name, opts, _ := strings.Cut(tag, ",")
|
||||
if field.Anonymous || strings.Contains(opts, "string") {
|
||||
return false
|
||||
}
|
||||
if name == "-" {
|
||||
continue
|
||||
}
|
||||
if name == "" {
|
||||
name = field.Name
|
||||
}
|
||||
fv := v.Field(i)
|
||||
if strings.Contains(opts, "omitempty") && probeEmpty(fv) {
|
||||
continue
|
||||
}
|
||||
if !first && !probeSpend(left, 1) {
|
||||
return false
|
||||
}
|
||||
first = false
|
||||
if !systemone.SpendJSONString(name, left) || !probeSpend(left, 1) || !probeValueBudget(fv, left, depth+1) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func probeEmpty(v reflect.Value) bool {
|
||||
switch v.Kind() {
|
||||
case reflect.Array, reflect.Map, reflect.Slice, reflect.String:
|
||||
return v.Len() == 0
|
||||
default:
|
||||
return v.IsZero()
|
||||
}
|
||||
}
|
||||
|
||||
func probeSchemaType(t reflect.Type) bool {
|
||||
switch t {
|
||||
case reflect.TypeFor[schema.Message](), reflect.TypeFor[schema.Messages](),
|
||||
reflect.TypeFor[schema.Content](), reflect.TypeFor[schema.ContentURL](), reflect.TypeFor[schema.InputAudio](),
|
||||
reflect.TypeFor[schema.ToolCall](), reflect.TypeFor[schema.FunctionCall](),
|
||||
reflect.TypeFor[schema.AnthropicMessage](), reflect.TypeFor[schema.AnthropicContentBlock](), reflect.TypeFor[schema.AnthropicImageSource]():
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package middleware_test
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"runtime"
|
||||
"strings"
|
||||
|
||||
. "github.com/mudler/LocalAI/core/http/middleware"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/systemone"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
type pointerJSON struct {
|
||||
Calls *int `json:"-"`
|
||||
}
|
||||
|
||||
func (p *pointerJSON) MarshalJSON() ([]byte, error) { *p.Calls++; return []byte(`"custom"`), nil }
|
||||
|
||||
type pointerText struct {
|
||||
Calls *int `json:"-"`
|
||||
}
|
||||
|
||||
func (p *pointerText) MarshalText() ([]byte, error) { *p.Calls++; return []byte(`custom`), nil }
|
||||
|
||||
type quotedContent struct {
|
||||
Text string `json:"text,string"`
|
||||
}
|
||||
type promotedContent struct {
|
||||
Text string `json:"text"`
|
||||
}
|
||||
type embeddedContent struct{ promotedContent }
|
||||
|
||||
var _ = Describe("Probe accepted representations", func() {
|
||||
It("rejects pointer marshalers without invoking them", func() {
|
||||
calls := 0
|
||||
for _, content := range []any{[]pointerJSON{{&calls}}, []pointerText{{&calls}}} {
|
||||
p := OpenAIProbeFromRequest(&schema.OpenAIRequest{Messages: []schema.Message{{Content: content}}})
|
||||
Expect(calls).To(BeZero())
|
||||
Expect(p.InputError).To(HaveOccurred())
|
||||
}
|
||||
})
|
||||
for _, kind := range []string{"quoted", "promoted"} {
|
||||
It("rejects unsupported "+kind+" structs before allocating", func() {
|
||||
var content any = quotedContent{strings.Repeat("\x00", systemone.MaxImageBodyBytes/6-100)}
|
||||
if kind == "promoted" {
|
||||
content = embeddedContent{promotedContent{strings.Repeat("x", systemone.MaxImageBodyBytes+1)}}
|
||||
}
|
||||
runtime.GC()
|
||||
var a, b runtime.MemStats
|
||||
runtime.ReadMemStats(&a)
|
||||
p := OpenAIProbeFromRequest(&schema.OpenAIRequest{Messages: []schema.Message{{Content: content}}})
|
||||
runtime.ReadMemStats(&b)
|
||||
Expect(p.InputError).To(HaveOccurred())
|
||||
Expect(b.TotalAlloc - a.TotalAlloc).To(BeNumerically("<", systemone.MaxBodyBytes))
|
||||
})
|
||||
}
|
||||
It("preserves raw tool JSON and schema tool content", func() {
|
||||
p := OpenAIProbeFromRequest(&schema.OpenAIRequest{Messages: []schema.Message{{Content: json.RawMessage(`[{"type":"text","text":"hello"}]`), FunctionCall: schema.FunctionCall{Name: "tool", Arguments: `{}`}, ToolCalls: []schema.ToolCall{{FunctionCall: schema.FunctionCall{Arguments: `{}`}}}}}})
|
||||
Expect(p.InputError).NotTo(HaveOccurred())
|
||||
Expect(json.Valid(p.State)).To(BeTrue())
|
||||
})
|
||||
})
|
||||
@@ -273,6 +273,11 @@ func (re *RequestExtractor) SetOpenAIRequest(c echo.Context) error {
|
||||
input.Context = ctxWithCorrelationID
|
||||
input.Cancel = cancel
|
||||
|
||||
// Router inputs must be classified before any candidate-specific media fetch.
|
||||
if cfg.HasRouter() {
|
||||
return nil
|
||||
}
|
||||
|
||||
err := mergeOpenAIRequestAndModelConfig(cfg, input)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -478,10 +483,17 @@ func mergeOpenAIRequestAndModelConfig(config *config.ModelConfig, input *schema.
|
||||
switch content := m.Content.(type) {
|
||||
case string:
|
||||
input.Messages[i].StringContent = content
|
||||
case []any:
|
||||
dat, _ := json.Marshal(content)
|
||||
c := []schema.Content{}
|
||||
json.Unmarshal(dat, &c)
|
||||
case []any, []schema.Content:
|
||||
c, typed := content.([]schema.Content)
|
||||
if !typed {
|
||||
dat, err := json.Marshal(content)
|
||||
if err != nil {
|
||||
return fmt.Errorf("message content: %w", err)
|
||||
}
|
||||
if err := json.Unmarshal(dat, &c); err != nil {
|
||||
return fmt.Errorf("message content: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
textContent := ""
|
||||
// we will template this at the end
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"hash/fnv"
|
||||
"strconv"
|
||||
@@ -17,6 +18,7 @@ import (
|
||||
"github.com/mudler/LocalAI/core/http/auth"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/services/routing/router"
|
||||
"github.com/mudler/LocalAI/core/systemone"
|
||||
"github.com/mudler/LocalAI/core/templates"
|
||||
"github.com/mudler/xlog"
|
||||
"gopkg.in/yaml.v3"
|
||||
@@ -102,8 +104,9 @@ func (s *reseedingVectorStore) Search(ctx context.Context, vec []float32) (float
|
||||
// score classifier runs unwrapped and the embedding-cache YAML is
|
||||
// ignored with a warning.
|
||||
type ClassifierDeps struct {
|
||||
Scorer ScorerFactory
|
||||
Embedder EmbedderFactory
|
||||
Decisions func(string) backend.DecisionRunner
|
||||
Scorer ScorerFactory
|
||||
Embedder EmbedderFactory
|
||||
// EmbedderFingerprint identifies the weights/config behind Embedder so
|
||||
// KNN corpus vectors cannot be queried across embedding spaces.
|
||||
EmbedderFingerprint EmbedderFingerprintFactory
|
||||
@@ -152,6 +155,7 @@ type ClassifierDeps struct {
|
||||
func NewClassifierDeps(app *application.Application) ClassifierDeps {
|
||||
return ClassifierDeps{
|
||||
Scorer: app.Scorer,
|
||||
Decisions: app.DecisionRunner,
|
||||
Corpus: app.RouterCorpus(),
|
||||
TokenCounter: app.TokenCounter,
|
||||
Embedder: app.Embedder,
|
||||
@@ -185,9 +189,9 @@ type ProbeExtractor func(parsed any) (router.Probe, bool)
|
||||
// 3. Invokes the classifier matching cfg.Router.Classifier
|
||||
// ("score" or "colbert"). If the classifier can't be built —
|
||||
// missing classifier_model, misconfigured policies, etc. — the
|
||||
// request fails with 503. cfg.Router.Fallback only catches
|
||||
// Classify-time errors and label-coverage misses, not config
|
||||
// bugs that would otherwise be silent.
|
||||
// request fails with 503. Invalid configuration fails closed; only
|
||||
// Classify-time errors and label-coverage misses use the
|
||||
// configured fallback.
|
||||
// 4. Resolves the chosen candidate to its model name. Reloads the
|
||||
// ModelConfig for that model and asserts depth-1 (the candidate
|
||||
// must NOT itself have a Router). Violation returns 500 — config
|
||||
@@ -257,6 +261,12 @@ func RouteModel(loader *config.ModelConfigLoader, appConfig *config.ApplicationC
|
||||
req.ModelName(&chosen)
|
||||
}
|
||||
|
||||
// Materialize the original payload only after choosing the served model.
|
||||
if req, ok := parsed.(*schema.OpenAIRequest); ok {
|
||||
if err := mergeOpenAIRequestAndModelConfig(result.ChosenConfig, req); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
c.Set(CONTEXT_LOCALS_KEY_MODEL_CONFIG, result.ChosenConfig)
|
||||
// Preserve an upstream requested model (e.g. an alias that points
|
||||
// at this router model) so accounting keeps the name the client
|
||||
@@ -346,6 +356,19 @@ func routerConfigFingerprint(rc config.RouterConfig, classifierCfg *config.Model
|
||||
h := fnv.New64a()
|
||||
h.Write(bytes)
|
||||
if classifierCfg != nil {
|
||||
if rc.Classifier == router.ClassifierDecisions {
|
||||
native, err := yaml.Marshal(classifierCfg)
|
||||
if err != nil {
|
||||
return uint64(time.Now().UnixNano())
|
||||
}
|
||||
h.Write(native)
|
||||
h.Write([]byte(classifierCfg.PersistedConfigRevision()))
|
||||
if classifierCfg.KnownUsecases != nil {
|
||||
h.Write([]byte(fmt.Sprintf("usecases:%d", *classifierCfg.KnownUsecases)))
|
||||
} else {
|
||||
h.Write([]byte("usecases:nil"))
|
||||
}
|
||||
}
|
||||
// Narrow projection: only the fields buildClassifier reads (renderer,
|
||||
// stop tokens, context_size → MaxContextTokens). Hashing the whole
|
||||
// ModelConfig would invalidate the cache on irrelevant changes;
|
||||
@@ -406,6 +429,21 @@ func buildClassifier(cfg *config.ModelConfig, deps ClassifierDeps) (router.Class
|
||||
|
||||
var inner router.Classifier
|
||||
switch name {
|
||||
case router.ClassifierDecisions:
|
||||
if rc.ClassifierModel == "" || deps.Decisions == nil || deps.ModelLookup == nil {
|
||||
return nil, fmt.Errorf("decisions requires classifier_model, native factory and model lookup")
|
||||
}
|
||||
if rc.KNN != nil || rc.EmbeddingCache != nil {
|
||||
return nil, fmt.Errorf("decisions does not support knn or embedding_cache composition")
|
||||
}
|
||||
modelCfg := deps.ModelLookup(rc.ClassifierModel)
|
||||
if modelCfg == nil {
|
||||
return nil, fmt.Errorf("decision model not available")
|
||||
}
|
||||
if err := systemone.ValidateDecisionModel(*modelCfg); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return router.NewDecisionsClassifier(policies, deps.Decisions(rc.ClassifierModel), rc.ActivationThreshold)
|
||||
case router.ClassifierScore:
|
||||
if rc.ClassifierModel == "" {
|
||||
return nil, fmt.Errorf("router classifier score requires classifier_model")
|
||||
@@ -764,8 +802,7 @@ func newDecisionID() string {
|
||||
// OpenAIProbe extracts a router.Probe from a parsed *schema.OpenAIRequest.
|
||||
// Concatenates message contents (string-form or text blocks of the
|
||||
// structured `[]any` content) so the classifier sees a single corpus
|
||||
// for length and content-shape rules. Image blocks are skipped — a
|
||||
// future multimodal classifier can take a different route.
|
||||
// for text classifiers, retaining complete chat state for native decisions.
|
||||
func OpenAIProbe(parsed any) (router.Probe, bool) {
|
||||
req, ok := parsed.(*schema.OpenAIRequest)
|
||||
if !ok || req == nil {
|
||||
@@ -777,6 +814,26 @@ func OpenAIProbe(parsed any) (router.Probe, bool) {
|
||||
// messageText flattens a chat message's Content to plain text: string content
|
||||
// verbatim; []any structured content contributes only its "text" blocks.
|
||||
func messageText(content any) string {
|
||||
// Typed API content is handled directly, avoiding a lossy JSON round trip.
|
||||
switch blocks := content.(type) {
|
||||
case []schema.Content:
|
||||
var texts []string
|
||||
for _, block := range blocks {
|
||||
if block.Type == "text" && block.Text != "" {
|
||||
texts = append(texts, block.Text)
|
||||
}
|
||||
}
|
||||
return strings.Join(texts, "\n")
|
||||
case []schema.AnthropicContentBlock:
|
||||
var texts []string
|
||||
for _, block := range blocks {
|
||||
if block.Type == "text" && block.Text != "" {
|
||||
texts = append(texts, block.Text)
|
||||
}
|
||||
}
|
||||
return strings.Join(texts, "\n")
|
||||
}
|
||||
|
||||
switch ct := content.(type) {
|
||||
case string:
|
||||
return ct
|
||||
@@ -817,6 +874,14 @@ func OpenAIProbeFromRequest(req *schema.OpenAIRequest) router.Probe {
|
||||
if req == nil {
|
||||
return router.Probe{}
|
||||
}
|
||||
release, admissionErr := systemone.AcquireAdmission(context.Background())
|
||||
if admissionErr != nil {
|
||||
return router.Probe{InputError: admissionErr}
|
||||
}
|
||||
defer release()
|
||||
if err := probeBudget(req.Messages); err != nil {
|
||||
return router.Probe{InputError: err}
|
||||
}
|
||||
texts := make([]string, len(req.Messages))
|
||||
for i := range req.Messages {
|
||||
texts[i] = messageText(req.Messages[i].Content)
|
||||
@@ -825,7 +890,12 @@ func OpenAIProbeFromRequest(req *schema.OpenAIRequest) router.Probe {
|
||||
// Prompt carries the full conversation; each classifier trims it to its own
|
||||
// model's context (see modelTokenTrim). Messages preserves the per-turn
|
||||
// split the trimmer drops oldest-first.
|
||||
return router.Probe{Prompt: router.JoinTurns(parts), Messages: parts}
|
||||
state, err := json.Marshal(req.Messages)
|
||||
if len(state) > systemone.MaxImageBodyBytes {
|
||||
state = nil
|
||||
err = fmt.Errorf("router state exceeds decision image request budget")
|
||||
}
|
||||
return router.Probe{Prompt: router.JoinTurns(parts), Messages: parts, State: state, InputError: err}
|
||||
}
|
||||
|
||||
// AnthropicProbe is the AnthropicRequest analogue of OpenAIProbe.
|
||||
@@ -834,10 +904,23 @@ func AnthropicProbe(parsed any) (router.Probe, bool) {
|
||||
if !ok || req == nil {
|
||||
return router.Probe{}, false
|
||||
}
|
||||
release, admissionErr := systemone.AcquireAdmission(context.Background())
|
||||
if admissionErr != nil {
|
||||
return router.Probe{InputError: admissionErr}, true
|
||||
}
|
||||
defer release()
|
||||
if err := probeBudget(req.Messages); err != nil {
|
||||
return router.Probe{InputError: err}, true
|
||||
}
|
||||
texts := make([]string, len(req.Messages))
|
||||
for i := range req.Messages {
|
||||
texts[i] = messageText(req.Messages[i].Content)
|
||||
}
|
||||
parts := messageProbeParts(texts)
|
||||
return router.Probe{Prompt: router.JoinTurns(parts), Messages: parts}, true
|
||||
state, err := json.Marshal(req.Messages)
|
||||
if len(state) > systemone.MaxImageBodyBytes {
|
||||
state = nil
|
||||
err = fmt.Errorf("router state exceeds decision image request budget")
|
||||
}
|
||||
return router.Probe{Prompt: router.JoinTurns(parts), Messages: parts, State: state, InputError: err}, true
|
||||
}
|
||||
@@ -2,6 +2,7 @@ package middleware_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
@@ -52,6 +53,18 @@ var _ = Describe("RouteModel middleware (score classifier)", func() {
|
||||
_ = os.RemoveAll(modelDir)
|
||||
})
|
||||
|
||||
It("uses configured fallback for image input on text classifiers without changing content", func() {
|
||||
cfg := newScoreRouterModel(modelDir, "smart-router")
|
||||
writeCandidate(modelDir, "qwen3-0.6b")
|
||||
req := openAIChat("")
|
||||
req.Messages[0].Content = []any{map[string]any{"type": "image_url", "image_url": map[string]string{"url": "data:image/png;base64,AA=="}}}
|
||||
before, _ := json.Marshal(req.Messages)
|
||||
rec, err := runRouter(loader, appConfig, store, cfg, req, func(string) backend.Scorer { return &stubScorer{} })
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(rec.Body.String()).To(Equal("served:qwen3-0.6b"))
|
||||
after, _ := json.Marshal(req.Messages)
|
||||
Expect(after).To(Equal(before))
|
||||
})
|
||||
It("routes to a candidate whose labels cover the active set", func() {
|
||||
// 3 policies, 2 candidates. Small model has [casual-chat],
|
||||
// bigger has [code-generation, math-reasoning, casual-chat].
|
||||
@@ -129,11 +142,7 @@ var _ = Describe("RouteModel middleware (score classifier)", func() {
|
||||
"math-reasoning": -4.0,
|
||||
}}
|
||||
_, err := runRouter(loader, appConfig, store, routerCfg, openAIChat("debug something"), stubScorerFactory(s))
|
||||
// Build-time config bugs (here: a candidate referencing a
|
||||
// label not declared in policies) must surface to the client
|
||||
// — the previous silent-fallback behaviour hid the broken
|
||||
// config and left operators wondering why traces never showed
|
||||
// the classifier model running.
|
||||
// Invalid classifier configuration must fail closed even with a fallback.
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(err.Error()).To(ContainSubstring("unknown label"))
|
||||
})
|
||||
@@ -822,3 +831,21 @@ var _ = Describe("RouteModel middleware (knn classifier)", func() {
|
||||
Expect(store.records[0].Cached).To(BeFalse())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("Multimodal router probes", func() {
|
||||
It("preserves image-only OpenAI content without mutating the request", func() {
|
||||
r := &schema.OpenAIRequest{Messages: []schema.Message{{Role: "user", Content: []any{map[string]any{"type": "image_url", "image_url": map[string]string{"url": "data:image/png;base64,AA=="}}}}}}
|
||||
before, _ := json.Marshal(r)
|
||||
p := OpenAIProbeFromRequest(r)
|
||||
Expect(string(p.State)).To(ContainSubstring("data:image/png;base64,AA=="))
|
||||
Expect(p.Prompt).To(BeEmpty())
|
||||
after, _ := json.Marshal(r)
|
||||
Expect(after).To(Equal(before))
|
||||
})
|
||||
It("preserves typed Anthropic base64 source", func() {
|
||||
r := &schema.AnthropicRequest{Messages: []schema.AnthropicMessage{{Role: "user", Content: []schema.AnthropicContentBlock{{Type: "image", Source: &schema.AnthropicImageSource{Type: "base64", MediaType: "image/png", Data: "AA=="}}}}}}
|
||||
p, ok := AnthropicProbe(r)
|
||||
Expect(ok).To(BeTrue())
|
||||
Expect(string(p.State)).To(ContainSubstring(`"media_type":"image/png"`))
|
||||
})
|
||||
})
|
||||
@@ -1,4 +1,5 @@
|
||||
import { test, expect } from '@playwright/test'
|
||||
import { test, expect } from './coverage-fixtures'
|
||||
import YAML from 'yaml'
|
||||
|
||||
// Router template + structured editor regression tests.
|
||||
//
|
||||
@@ -8,7 +9,7 @@ import { test, expect } from '@playwright/test'
|
||||
// of a string ("(intermediate value).split is not a function").
|
||||
//
|
||||
// The current schema is also covered:
|
||||
// - classifier=score is the only shipped classifier
|
||||
// - classifier choices preserve the existing router creation flow
|
||||
// - router.policies surfaces in its own structured editor (label +
|
||||
// description rows with duplicate detection)
|
||||
// - router.candidates is the structured {model, labels[]} editor;
|
||||
@@ -29,13 +30,14 @@ const ROUTER_METADATA = {
|
||||
{
|
||||
path: 'router.classifier', yaml_key: 'classifier', go_type: 'string', ui_type: 'string',
|
||||
section: 'other', label: 'Classifier', component: 'select',
|
||||
options: [{ value: 'score', label: 'Score (Arch-Router-style)' }],
|
||||
options: [{ value: 'score', label: 'Score (Arch-Router-style)' }, { value: 'decisions', label: 'Decisions (native probabilities)' }, { value: 'colbert', label: 'Colbert (reranker)' }, { value: 'knn', label: 'KNN (labelled corpus)' }],
|
||||
description: 'Picks a candidate by scoring every policy label against the prompt. Only "score" is shipped today.',
|
||||
order: 230,
|
||||
},
|
||||
{
|
||||
path: 'router.classifier_model', yaml_key: 'classifier_model', go_type: 'string', ui_type: 'string',
|
||||
section: 'other', label: 'Classifier Model', component: 'model-select', autocomplete_provider: 'models:chat',
|
||||
section: 'other', label: 'Classifier Model', component: 'model-select', autocomplete_provider: 'models:score',
|
||||
autocomplete_by: { field: 'router.classifier', providers: { decisions: 'models:decisions', colbert: 'models:rerank', knn: '' } },
|
||||
description: 'Loaded LocalAI model the score classifier asks to rank each policy label.',
|
||||
order: 231,
|
||||
},
|
||||
@@ -151,11 +153,11 @@ test.describe('Router template — create flow', () => {
|
||||
await expect(page.getByText('Activation Threshold').first()).toBeVisible()
|
||||
})
|
||||
|
||||
test('Classifier select offers only the score option', async ({ page }) => {
|
||||
test('Classifier defaults to score', async ({ page }) => {
|
||||
await page.goto('/app/model-editor?template=router')
|
||||
|
||||
// SearchableSelect renders the current option's *label* inside the
|
||||
// trigger button. After the schema cleanup the only option is
|
||||
// trigger button. The initial option is
|
||||
// "Score (Arch-Router-style)", pre-selected by the template.
|
||||
await expect(page.getByText('Score (Arch-Router-style)').first()).toBeVisible({ timeout: 10_000 })
|
||||
})
|
||||
@@ -216,4 +218,60 @@ test.describe('Router template — create flow', () => {
|
||||
page.locator('input[title="Duplicate label — candidates won\'t be able to distinguish them"]').first()
|
||||
).toBeVisible()
|
||||
})
|
||||
test('Decisions picker creates, saves and reopens exact model and tuned threshold', async ({ page }) => {
|
||||
const models = [
|
||||
{ id: 'openjev-llama', capabilities: ['decisions'] },
|
||||
{ id: 'LayaGLiNERDecide-vllm', capabilities: ['decisions'] },
|
||||
{ id: 'TevKev-vllm', capabilities: ['decisions'] },
|
||||
{ id: 'NimbleCLM-vllm', capabilities: ['decisions'] },
|
||||
{ id: 'generic-ner', capabilities: ['token_classify'] },
|
||||
{ id: 'chat-target', capabilities: ['chat'] },
|
||||
]
|
||||
await page.route('**/v1/models/capabilities', route => route.fulfill({ json: { data: models } }))
|
||||
await page.route('**/api/models/capabilities', route => route.fulfill({ json: { data: [{ id: 'arch-score', capabilities: ['FLAG_SCORE'] }] } }))
|
||||
let saved
|
||||
await page.route('**/models/import', async route => {
|
||||
saved = route.request().postDataJSON()
|
||||
await route.fulfill({ json: { success: true } })
|
||||
})
|
||||
await page.route('**/api/models/edit/smart-router', route => route.fulfill({ json: { config: YAML.stringify(saved) } }))
|
||||
await page.goto('/app/model-editor?template=router')
|
||||
const modelRow = page.locator('.form-row').filter({ has: page.locator('.form-row__label-text', { hasText: /^Classifier Model$/ }) })
|
||||
const threshold = page.locator('.form-row').filter({ hasText: 'Activation Threshold' }).locator('input[type="range"]')
|
||||
await modelRow.locator('input').fill('arch-score')
|
||||
await threshold.focus()
|
||||
for (let i = 0; i < 5; i++) await threshold.press('ArrowRight')
|
||||
await page.getByText('Score (Arch-Router-style)', { exact: true }).click()
|
||||
await page.getByRole('option', { name: 'Decisions (native probabilities)' }).click()
|
||||
await expect(modelRow.locator('input')).toHaveValue('')
|
||||
await expect(threshold).toHaveValue('0.65')
|
||||
await modelRow.locator('input').click()
|
||||
for (const id of ['openjev-llama', 'LayaGLiNERDecide-vllm', 'TevKev-vllm', 'NimbleCLM-vllm']) {
|
||||
await expect(page.getByRole('option', { name: id })).toBeVisible()
|
||||
}
|
||||
await expect(page.getByRole('option', { name: 'generic-ner' })).toHaveCount(0)
|
||||
await expect(page.getByRole('option', { name: 'chat-target' })).toHaveCount(0)
|
||||
await page.getByRole('option', { name: 'NimbleCLM-vllm' }).click()
|
||||
await page.getByRole('button', { name: /Create Model$/ }).click()
|
||||
await expect(page).toHaveURL(/model-editor\/smart-router/)
|
||||
expect(saved.router).toMatchObject({ classifier: 'decisions', classifier_model: 'NimbleCLM-vllm', activation_threshold: 0.65 })
|
||||
await page.reload()
|
||||
await expect(page.getByText('Decisions (native probabilities)', { exact: true })).toBeVisible()
|
||||
await expect(modelRow.locator('input')).toHaveValue('NimbleCLM-vllm')
|
||||
await expect(threshold).toHaveValue('0.65')
|
||||
await page.getByText('Decisions (native probabilities)', { exact: true }).click()
|
||||
await page.getByRole('option', { name: 'Colbert (reranker)' }).click()
|
||||
await expect(modelRow.locator('input')).toHaveValue('')
|
||||
await expect(threshold).toHaveValue('0.65')
|
||||
await page.getByText('Colbert (reranker)', { exact: true }).click()
|
||||
await page.getByRole('option', { name: 'KNN (labelled corpus)' }).click()
|
||||
await expect(modelRow).toHaveCount(0)
|
||||
await expect(threshold).toHaveValue('0.65')
|
||||
await page.getByText('KNN (labelled corpus)', { exact: true }).click()
|
||||
await page.getByRole('option', { name: 'Decisions (native probabilities)' }).click()
|
||||
await modelRow.locator('input').fill('generic-ner')
|
||||
await page.getByRole('button', { name: /Save Changes$/ }).click()
|
||||
await expect(page.getByText('Save failed: Select an eligible native Decisions classifier model')).toBeVisible()
|
||||
})
|
||||
|
||||
})
|
||||
@@ -1,3 +1,4 @@
|
||||
import { useFormContext } from '../contexts/FormContext'
|
||||
import { useState } from 'react'
|
||||
import SettingRow from './SettingRow'
|
||||
import Toggle from './Toggle'
|
||||
@@ -20,6 +21,8 @@ const PROVIDER_TO_CAPABILITY = {
|
||||
'models:transcript': 'FLAG_TRANSCRIPT',
|
||||
'models:vad': 'FLAG_VAD',
|
||||
'models:score': 'FLAG_SCORE',
|
||||
'models:decisions': 'decisions',
|
||||
'models:rerank': 'FLAG_RERANK',
|
||||
'models:token_classify': 'FLAG_TOKEN_CLASSIFY',
|
||||
}
|
||||
|
||||
@@ -160,6 +163,9 @@ function FieldLabel({ field }) {
|
||||
}
|
||||
|
||||
export default function ConfigFieldRenderer({ field, value, onChange, onRemove, annotation }) {
|
||||
const context = useFormContext()
|
||||
const conditional = field.autocomplete_by
|
||||
const provider = conditional?.providers[context?.formData?.[conditional.field]] ?? field.autocomplete_provider
|
||||
const handleChange = (raw) => {
|
||||
onChange(coerceValue(raw, field.ui_type))
|
||||
}
|
||||
@@ -191,11 +197,13 @@ export default function ConfigFieldRenderer({ field, value, onChange, onRemove,
|
||||
}
|
||||
|
||||
// Model-select
|
||||
if (component === 'model-select' && provider === '') return null
|
||||
if (component === 'model-select') {
|
||||
const cap = PROVIDER_TO_CAPABILITY[field.autocomplete_provider] || undefined
|
||||
const cap = PROVIDER_TO_CAPABILITY[provider] || undefined
|
||||
return (
|
||||
<SettingRow label={<FieldLabel field={field} />} description={description}>
|
||||
<SearchableModelSelect
|
||||
key={provider}
|
||||
value={value || ''}
|
||||
onChange={handleChange}
|
||||
capability={cap}
|
||||
@@ -402,7 +410,7 @@ export default function ConfigFieldRenderer({ field, value, onChange, onRemove,
|
||||
// PII detectors — a capability-filtered multi-select of token_classify
|
||||
// models (the consuming model's pii.detectors list).
|
||||
if (component === 'model-multi-select') {
|
||||
const cap = PROVIDER_TO_CAPABILITY[field.autocomplete_provider] || undefined
|
||||
const cap = PROVIDER_TO_CAPABILITY[provider] || undefined
|
||||
return (
|
||||
<div className="list-row">
|
||||
<div className="hstack hstack--between mb-xs">
|
||||
|
||||
@@ -186,7 +186,7 @@ export default function SearchableModelSelect({ value, onChange, capability, pla
|
||||
<span className="sms-hint">{hints[m.id]}</span>
|
||||
)}
|
||||
{isEnterTarget && (
|
||||
<span style={{ color: 'var(--color-text-muted)', fontSize: '0.75rem', flexShrink: 0 }}>↵</span>
|
||||
<span className="text-meta shrink-0">↵</span>
|
||||
)}
|
||||
</div>
|
||||
)
|
||||
|
||||
+4
-4
@@ -9,11 +9,11 @@ export function useModels(capability) {
|
||||
const fetchModels = useCallback(async ({ silent = false } = {}) => {
|
||||
try {
|
||||
if (!silent) setLoading(true)
|
||||
const data = await modelsApi.listCapabilities()
|
||||
let items = data?.data || []
|
||||
const data = await (capability === 'decisions' ? modelsApi.listNativeCapabilities() : modelsApi.listCapabilities())
|
||||
let items = (data?.data || []).filter(m => !capability || !m.disabled)
|
||||
if (capability) {
|
||||
items = items.filter(m =>
|
||||
m.capabilities?.includes(capability) ||
|
||||
m.capabilities?.includes(capability) || m.capabilities?.includes(capability.replace(/^FLAG_/, '').toLowerCase()) ||
|
||||
// Models without config (loose files) have no capabilities — show them only when no filter
|
||||
false
|
||||
)
|
||||
@@ -24,7 +24,7 @@ export function useModels(capability) {
|
||||
// Fallback to /v1/models if capabilities endpoint unavailable
|
||||
try {
|
||||
const data = await modelsApi.listV1()
|
||||
setModels((data?.data || []).map(m => ({ id: m.id, capabilities: [] })))
|
||||
setModels(capability ? [] : (data?.data || []).map(m => ({ id: m.id, capabilities: [] })))
|
||||
setError(null)
|
||||
} catch (err) {
|
||||
setError(err.message)
|
||||
|
||||
@@ -292,6 +292,12 @@ export default function ModelEditor() {
|
||||
for (const path of activeFieldPaths) {
|
||||
if (path in values) patchFlat[path] = values[path]
|
||||
}
|
||||
if (patchFlat['router.classifier'] === 'decisions') {
|
||||
const available = await modelsApi.listNativeCapabilities()
|
||||
if (!available?.data?.some(m => m.id === patchFlat['router.classifier_model'] && m.capabilities?.includes('decisions'))) {
|
||||
throw new Error('Select an eligible native Decisions classifier model')
|
||||
}
|
||||
}
|
||||
const config = unflattenConfig(patchFlat)
|
||||
|
||||
if (isCreateMode) {
|
||||
@@ -407,7 +413,16 @@ export default function ModelEditor() {
|
||||
}
|
||||
|
||||
const handleFieldChange = (path, val) => {
|
||||
setValues(prev => ({ ...prev, [path]: val }))
|
||||
setValues(prev => {
|
||||
const next = { ...prev, [path]: val }
|
||||
// A classifier change invalidates its dependent model, not a tuned threshold.
|
||||
if (prev[path] !== val) {
|
||||
for (const field of fields) {
|
||||
if (field.autocomplete_by?.field === path) next[field.path] = ''
|
||||
}
|
||||
}
|
||||
return next
|
||||
})
|
||||
}
|
||||
|
||||
const toggleSection = (id) => {
|
||||
|
||||
Vendored
+1
@@ -83,6 +83,7 @@ export async function streamChat(body, signal) {
|
||||
export const modelsApi = {
|
||||
list: (params) => fetchJSON(buildUrl(API_CONFIG.endpoints.models, params)),
|
||||
listV1: () => fetchJSON(API_CONFIG.endpoints.modelsList),
|
||||
listNativeCapabilities: () => fetchJSON('/v1/models/capabilities'),
|
||||
listCapabilities: () => fetchJSON(API_CONFIG.endpoints.modelsCapabilities),
|
||||
listAliases: () => fetchJSON(API_CONFIG.endpoints.modelsAliases),
|
||||
// variant is optional. Omitting it lets the server auto-select the best
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"github.com/labstack/echo/v4"
|
||||
"github.com/mudler/LocalAI/core/application"
|
||||
"github.com/mudler/LocalAI/core/http/endpoints/localai"
|
||||
"github.com/mudler/LocalAI/core/http/middleware"
|
||||
)
|
||||
|
||||
// RegisterSystemOneRoutes wires the kev-compatible SystemOne endpoints.
|
||||
@@ -14,7 +15,7 @@ import (
|
||||
// under n_perm option orders; POST /v1/systemone/separate answers each
|
||||
// question in its own NER pass.
|
||||
func RegisterSystemOneRoutes(e *echo.Echo, app *application.Application) {
|
||||
e.POST("/v1/systemone", localai.SystemOneEndpoint(app))
|
||||
e.POST("/v1/systemone", localai.SystemOneEndpoint(app), middleware.UsageMiddleware(app.StatsRecorder(), app.FallbackUser()))
|
||||
e.POST("/v1/systemone/permute", localai.SystemOnePermuteEndpoint(app))
|
||||
e.POST("/v1/systemone/separate", localai.SystemOneSeparateEndpoint(app))
|
||||
}
|
||||
@@ -0,0 +1,74 @@
|
||||
package routes_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"github.com/labstack/echo/v4"
|
||||
"github.com/mudler/LocalAI/core/application"
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/http/routes"
|
||||
"github.com/mudler/LocalAI/core/services/routing/billing"
|
||||
grpcpkg "github.com/mudler/LocalAI/pkg/grpc"
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
"github.com/mudler/LocalAI/pkg/model"
|
||||
"github.com/mudler/LocalAI/pkg/system"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
ggrpc "google.golang.org/grpc"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type nativeDecisionFixture struct {
|
||||
grpcpkg.Backend
|
||||
body string
|
||||
}
|
||||
|
||||
func (*nativeDecisionFixture) HealthCheck(context.Context) (bool, error) { return true, nil }
|
||||
func (*nativeDecisionFixture) IsBusy() bool { return false }
|
||||
func (*nativeDecisionFixture) Free(context.Context) error { return nil }
|
||||
func (b *nativeDecisionFixture) Score(context.Context, *pb.ScoreRequest, ...ggrpc.CallOption) (*pb.ScoreResponse, error) {
|
||||
return &pb.ScoreResponse{ResponseJson: b.body}, nil
|
||||
}
|
||||
|
||||
var _ = Describe("registered SystemOne billing", func() {
|
||||
DescribeTable("records validated native responses exactly once", func(body string, code int, count int64) {
|
||||
root := GinkgoT().TempDir()
|
||||
app, err := application.New(config.WithDataPath(root), config.WithDisableLocalAIAssistant(true), config.WithDisableCSRF(true), config.WithSystemState(&system.SystemState{Model: system.Model{ModelsPath: root}, Backend: system.Backend{BackendsPath: root}}))
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
DeferCleanup(func() { Expect(app.Shutdown()).To(Succeed()) })
|
||||
cfg := config.ModelConfig{Name: "decision", Backend: "llama-cpp", KnownUsecases: config.GetUsecasesFromYAML([]string{"decisions"})}
|
||||
cfg.SetDefaults()
|
||||
cfg.Model = "fixture.gguf"
|
||||
app.ModelConfigLoader().ReplaceModelConfigs([]config.ModelConfig{cfg})
|
||||
fixture := &nativeDecisionFixture{body: body}
|
||||
app.ModelLoader().SetModelRouter(func(_ context.Context, id string, _, _, _, _ string, _ *pb.ModelOptions, _ bool) (*model.Model, error) {
|
||||
return model.NewModelWithClient(id, "test://decision", fixture), nil
|
||||
})
|
||||
e := echo.New()
|
||||
routes.RegisterSystemOneRoutes(e, app)
|
||||
req := httptest.NewRequest(http.MethodPost, "/v1/systemone", strings.NewReader(`{"model":"decision","state":"x","questions":{"q":{"type":"noul","instructions":"yes?"}}}`))
|
||||
req.Header.Set(echo.HeaderContentType, echo.MIMEApplicationJSON)
|
||||
w := httptest.NewRecorder()
|
||||
e.ServeHTTP(w, req)
|
||||
Expect(w.Code).To(Equal(code), w.Body.String())
|
||||
buckets, err := app.StatsRecorder().Aggregate(context.Background(), billing.AggregateQuery{})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
var actual int64
|
||||
for _, b := range buckets {
|
||||
actual += b.RequestCount
|
||||
Expect(b.CompletionTokens).To(Equal(int64(0)))
|
||||
}
|
||||
Expect(actual).To(Equal(count))
|
||||
},
|
||||
Entry("positive input", `{"answers":{"q":{"type":"noul","noul":0.5}},"usage":{"input_tokens":12,"output_tokens":0}}`, 200, int64(1)),
|
||||
Entry("explicit zeros", `{"answers":{"q":{"type":"noul","noul":0.5}},"usage":{"input_tokens":0,"output_tokens":0}}`, 200, int64(1)),
|
||||
Entry("absent usage", `{"answers":{"q":{"type":"noul","noul":0.5}}}`, 200, int64(0)),
|
||||
Entry("negative", `{"answers":{"q":{"type":"noul","noul":0.5}},"usage":{"input_tokens":-1,"output_tokens":0}}`, 500, int64(0)),
|
||||
Entry("incomplete", `{"answers":{"q":{"type":"noul","noul":0.5}},"usage":{"input_tokens":12}}`, 500, int64(0)),
|
||||
Entry("null", "null", 500, int64(0)),
|
||||
Entry("oversized valid JSON", `{"answers":{"q":{"type":"noul","noul":0.5}},"usage":{"input_tokens":12,"output_tokens":0},"padding":"`+strings.Repeat("x", 64<<10)+`"}`, 500, int64(0)),
|
||||
Entry("empty", `{}`, 500, int64(0)),
|
||||
Entry("missing answers", `{"usage":{"input_tokens":12,"output_tokens":0}}`, 500, int64(0)),
|
||||
)
|
||||
})
|
||||
@@ -0,0 +1,110 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package routes_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"github.com/labstack/echo/v4"
|
||||
"github.com/mudler/LocalAI/core/application"
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/http/routes"
|
||||
"github.com/mudler/LocalAI/core/services/routing/billing"
|
||||
"github.com/mudler/LocalAI/core/systemone"
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
"github.com/mudler/LocalAI/pkg/model"
|
||||
"github.com/mudler/LocalAI/pkg/system"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"image"
|
||||
"image/jpeg"
|
||||
"image/png"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
)
|
||||
|
||||
var _ = Describe("registered SystemOne multimodal admission", func() {
|
||||
DescribeTable("enforces shared limits without billing errors", func(input string, code int, count int64) {
|
||||
root := GinkgoT().TempDir()
|
||||
app, err := application.New(config.WithDataPath(root), config.WithDisableLocalAIAssistant(true), config.WithDisableCSRF(true), config.WithSystemState(&system.SystemState{Model: system.Model{ModelsPath: root}, Backend: system.Backend{BackendsPath: root}}))
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
DeferCleanup(func() { Expect(app.Shutdown()).To(Succeed()) })
|
||||
cfg := config.ModelConfig{Name: "decision", Backend: "llama-cpp", KnownUsecases: config.GetUsecasesFromYAML([]string{"decisions"})}
|
||||
cfg.SetDefaults()
|
||||
cfg.Model = "fixture.gguf"
|
||||
app.ModelConfigLoader().ReplaceModelConfigs([]config.ModelConfig{cfg})
|
||||
fixture := &nativeDecisionFixture{body: `{"answers":{"q":{"type":"noul","noul":0.5}},"usage":{"input_tokens":12,"output_tokens":0}}`}
|
||||
app.ModelLoader().SetModelRouter(func(_ context.Context, id string, _, _, _, _ string, _ *pb.ModelOptions, _ bool) (*model.Model, error) {
|
||||
return model.NewModelWithClient(id, "test://decision", fixture), nil
|
||||
})
|
||||
e := echo.New()
|
||||
routes.RegisterSystemOneRoutes(e, app)
|
||||
req := httptest.NewRequest(http.MethodPost, "/v1/systemone", strings.NewReader(input))
|
||||
req.Header.Set(echo.HeaderContentType, echo.MIMEApplicationJSON)
|
||||
w := httptest.NewRecorder()
|
||||
e.ServeHTTP(w, req)
|
||||
Expect(w.Code).To(Equal(code), w.Body.String())
|
||||
buckets, err := app.StatsRecorder().Aggregate(context.Background(), billing.AggregateQuery{})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
var actual int64
|
||||
for _, b := range buckets {
|
||||
actual += b.RequestCount
|
||||
Expect(b.CompletionTokens).To(Equal(int64(0)))
|
||||
}
|
||||
Expect(actual).To(Equal(count))
|
||||
},
|
||||
Entry("valid image-only", imageRequest("{}", []string{routePNG(1, 1)}), 200, int64(1)),
|
||||
Entry("image-bearing above text limit", imageRequest(`{"text":"`+strings.Repeat("x", systemone.MaxBodyBytes)+`"}`, []string{routePNG(1, 1)}), 200, int64(1)),
|
||||
Entry("valid JPEG", imageRequest("{}", []string{routeJPEG(false)}), 200, int64(1)),
|
||||
Entry("truncated JPEG", imageRequest("{}", []string{routeJPEG(true)}), 400, int64(0)),
|
||||
Entry("corrupt PNG pixels", imageRequest("{}", []string{corruptRoutePNG()}), 400, int64(0)),
|
||||
Entry("truncated PNG", imageRequest("{}", []string{truncatedRoutePNG()}), 400, int64(0)),
|
||||
Entry("invalid header", imageRequest("{}", []string{"data:image/png;base64,AA=="}), 400, int64(0)),
|
||||
Entry("remote URL", imageRequest("{}", []string{"https://example.org/x.png"}), 400, int64(0)),
|
||||
Entry("dimensions", imageRequest("{}", []string{routePNG(4097, 1)}), 413, int64(0)),
|
||||
Entry("aggregate pixels", imageRequest("{}", []string{routePNG(3000, 3000), routePNG(3000, 3000)}), 413, int64(0)),
|
||||
Entry("count", imageRequest("{}", strings.Split(strings.Repeat(routePNG(1, 1)+" ", 9), " ")[:9]), 413, int64(0)),
|
||||
Entry("encoded", imageRequest("{}", []string{strings.Repeat("A", systemone.MaxImageEncodedBytes+1)}), 413, int64(0)),
|
||||
Entry("decoded", imageRequest("{}", []string{"data:image/png;base64," + strings.Repeat("A", ((systemone.MaxImageDecodedBytes+3)/3)*4)}), 413, int64(0)),
|
||||
Entry("empty images keep text cap", imageRequest(`"`+strings.Repeat("x", systemone.MaxBodyBytes)+`"`, []string{}), 413, int64(0)),
|
||||
Entry("escaping is not a second wire cap", imageRequest(`"`+strings.Repeat("<", 20000)+`"`, nil), 200, int64(1)),
|
||||
Entry("image body cap", imageRequest("{}", []string{routePNG(1, 1)})+strings.Repeat(" ", systemone.MaxImageBodyBytes), 413, int64(0)),
|
||||
)
|
||||
})
|
||||
|
||||
func routePNG(w, h int) string {
|
||||
var b bytes.Buffer
|
||||
if err := png.Encode(&b, image.NewGray(image.Rect(0, 0, w, h))); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return "data:image/png;base64," + base64.StdEncoding.EncodeToString(b.Bytes())
|
||||
}
|
||||
func imageRequest(state string, images []string) string {
|
||||
raw, _ := json.Marshal(images)
|
||||
return `{"model":"decision","state":` + state + `,"images":` + string(raw) + `,"questions":{"q":{"type":"noul"}}}`
|
||||
}
|
||||
|
||||
func truncatedRoutePNG() string {
|
||||
raw, _ := base64.StdEncoding.DecodeString(strings.SplitN(routePNG(8, 8), ",", 2)[1])
|
||||
return "data:image/png;base64," + base64.StdEncoding.EncodeToString(raw[:33])
|
||||
}
|
||||
|
||||
func routeJPEG(truncated bool) string {
|
||||
var b bytes.Buffer
|
||||
if err := jpeg.Encode(&b, image.NewGray(image.Rect(0, 0, 8, 8)), nil); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
raw := b.Bytes()
|
||||
if truncated {
|
||||
raw = raw[:len(raw)-10]
|
||||
}
|
||||
return "data:image/jpeg;base64," + base64.StdEncoding.EncodeToString(raw)
|
||||
}
|
||||
func corruptRoutePNG() string {
|
||||
raw, _ := base64.StdEncoding.DecodeString(strings.SplitN(routePNG(8, 8), ",", 2)[1])
|
||||
at := bytes.Index(raw, []byte("IDAT"))
|
||||
raw[at+5] ^= 255
|
||||
return "data:image/png;base64," + base64.StdEncoding.EncodeToString(raw)
|
||||
}
|
||||
@@ -9,6 +9,10 @@ type SystemOneRequest struct {
|
||||
// State is the text (or any JSON value) to extract from. A non-string
|
||||
// value is rendered to its JSON representation before NER.
|
||||
State json.RawMessage `json:"state"`
|
||||
// Images holds PNG/JPEG data URLs validated by the shared decision contract. RawMessage
|
||||
// distinguishes absence from explicit null and an empty array. Unsupported
|
||||
// image input must be rejected, never silently discarded.
|
||||
Images json.RawMessage `json:"images,omitempty"`
|
||||
// Questions maps question IDs to their definitions.
|
||||
Questions map[string]SystemOneQuestion `json:"questions"`
|
||||
// Model names the NER model to use. Optional.
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
package schema_test
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("SystemOne input preservation", func() {
|
||||
It("preserves structured input and absent, null, or empty images semantically", func() {
|
||||
for _, input := range []string{
|
||||
`{"state":{"messages":[{"content":[{"type":"image_url","image_url":{"url":"data:image/png;base64,AA=="}}]}]},"questions":{"q":{"type":"choice","instructions":{"text":"choose"},"criteria":{"yes":null,"no":"negative"}}},"images":["data:image/png;base64,AA=="]}`,
|
||||
`{"state":"text","questions":{},"images":[]}`,
|
||||
`{"state":"text","questions":{},"images":null}`,
|
||||
`{"state":"text","questions":{}}`,
|
||||
} {
|
||||
var req schema.SystemOneRequest
|
||||
Expect(json.Unmarshal([]byte(input), &req)).To(Succeed())
|
||||
output, err := json.Marshal(req)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
var want, got any
|
||||
Expect(json.Unmarshal([]byte(input), &want)).To(Succeed())
|
||||
Expect(json.Unmarshal(output, &got)).To(Succeed())
|
||||
Expect(got).To(Equal(want))
|
||||
}
|
||||
})
|
||||
It("preserves images in permute requests semantically", func() {
|
||||
input := `{"request":{"state":"text","questions":{},"images":[{"url":"data:image/png;base64,AA=="}]},"question":"q"}`
|
||||
var req schema.SystemOnePermuteRequest
|
||||
Expect(json.Unmarshal([]byte(input), &req)).To(Succeed())
|
||||
output, err := json.Marshal(req)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
var want, got any
|
||||
Expect(json.Unmarshal([]byte(input), &want)).To(Succeed())
|
||||
Expect(json.Unmarshal(output, &got)).To(Succeed())
|
||||
Expect(got).To(Equal(want))
|
||||
})
|
||||
})
|
||||
@@ -37,8 +37,7 @@ func recordSwitch(ev Event) {
|
||||
|
||||
// RegisterMetrics exports target health as a gauge. The application calls it
|
||||
// once for its manager; tests create many managers and skip it.
|
||||
func RegisterMetrics(m *Manager) {
|
||||
meter := otel.Meter("github.com/mudler/LocalAI")
|
||||
func RegisterMetrics(m *Manager, meter metric.Meter) {
|
||||
_, _ = meter.Int64ObservableGauge("localai_failover_target_up",
|
||||
metric.WithDescription("1 when a failover target is healthy, 0 otherwise"),
|
||||
metric.WithInt64Callback(func(_ context.Context, o metric.Int64Observer) error {
|
||||
|
||||
@@ -3,13 +3,14 @@ package failover
|
||||
import (
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"go.opentelemetry.io/otel/metric/noop"
|
||||
)
|
||||
|
||||
var _ = Describe("metrics", func() {
|
||||
It("registers and records without a meter provider", func() {
|
||||
src := newFakeSource(remote("a"), local("b"), chainCfg("chain", nil, t("a"), t("b")))
|
||||
m := New(src, WithClock(newFakeClock()))
|
||||
Expect(func() { RegisterMetrics(m) }).ToNot(Panic())
|
||||
Expect(func() { RegisterMetrics(m, noop.NewMeterProvider().Meter("test")) }).ToNot(Panic())
|
||||
Expect(func() { m.ReportFailure("a", errBoom) }).ToNot(Panic())
|
||||
})
|
||||
It("records an attempt trace only when enabled", func() {
|
||||
|
||||
@@ -138,6 +138,17 @@ func (c *EmbeddingCacheClassifier) Stats() EmbeddingCacheStats {
|
||||
}
|
||||
|
||||
func (c *EmbeddingCacheClassifier) Classify(ctx context.Context, p Probe) (Decision, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return Decision{}, err
|
||||
}
|
||||
images, err := p.HasImages(ctx)
|
||||
if err != nil {
|
||||
return Decision{}, err
|
||||
}
|
||||
// Text embeddings cannot distinguish images; neither read nor populate cache.
|
||||
if images {
|
||||
return c.inner.Classify(ctx, p)
|
||||
}
|
||||
start := time.Now()
|
||||
|
||||
vec, err := c.embedder.Embed(ctx, trimmedProbeText(p, c.budget, identityRender))
|
||||
|
||||
@@ -390,3 +390,43 @@ var _ = Describe("EmbeddingCache latency", func() {
|
||||
Expect(d.Latency).To(BeNumerically("<", time.Second), "Latency unreasonably high for an in-memory hit")
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("Multimodal embedding cache isolation", func() {
|
||||
It("never embeds, trims, reads or writes cache for image probes", func() {
|
||||
inner := &stubInner{name: "decisions", decision: router.Decision{Labels: []string{"visual"}, Score: .9}}
|
||||
embedder := &countingImageEmbedder{}
|
||||
store := &countingImageStore{}
|
||||
cache := router.NewEmbeddingCacheClassifier(inner, embedder, store, .9, .5)
|
||||
for _, data := range []string{"AA==", "AQ=="} {
|
||||
p := router.Probe{Prompt: "identical", Messages: []string{"old image turn", "new text"}, State: json.RawMessage(`{}`), Images: json.RawMessage(`["data:image/png;base64,` + data + `"]`)}
|
||||
d, err := cache.Classify(context.Background(), p)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(d.Cached).To(BeFalse())
|
||||
}
|
||||
Expect(inner.calls).To(Equal(2))
|
||||
Expect(embedder.calls).To(BeZero())
|
||||
Expect(store.reads).To(BeZero())
|
||||
Expect(store.writes).To(BeZero())
|
||||
})
|
||||
})
|
||||
|
||||
// Return a usable vector and applicable cached decision: disabling the image
|
||||
// guard must fail even when the cache infrastructure succeeds.
|
||||
type countingImageEmbedder struct{ calls int }
|
||||
|
||||
func (e *countingImageEmbedder) Embed(context.Context, string) ([]float32, error) {
|
||||
e.calls++
|
||||
return []float32{1}, nil
|
||||
}
|
||||
|
||||
type countingImageStore struct{ reads, writes int }
|
||||
|
||||
func (s *countingImageStore) Search(context.Context, []float32) (float64, []byte, bool, error) {
|
||||
s.reads++
|
||||
return 1, []byte(`{"labels":["stale-text"],"score":1}`), true, nil
|
||||
}
|
||||
func (s *countingImageStore) SearchK(context.Context, []float32, int) ([]backend.Neighbor, error) {
|
||||
s.reads++
|
||||
return nil, nil
|
||||
}
|
||||
func (s *countingImageStore) Insert(context.Context, []float32, []byte) error { s.writes++; return nil }
|
||||
@@ -0,0 +1,67 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package router
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/systemone"
|
||||
)
|
||||
|
||||
func (p Probe) decisionRequest() (*schema.SystemOneRequest, error) {
|
||||
if p.InputError != nil {
|
||||
return nil, p.InputError
|
||||
}
|
||||
if err := p.collectionBounds(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
state := p.State
|
||||
if len(state) == 0 {
|
||||
left := systemone.MaxImageBodyBytes
|
||||
if !systemone.SpendJSONString(p.Prompt, &left) {
|
||||
return nil, fmt.Errorf("router prompt exceeds serialized decision budget")
|
||||
}
|
||||
state, _ = json.Marshal(p.Prompt)
|
||||
}
|
||||
return &schema.SystemOneRequest{State: state, Images: p.Images}, nil
|
||||
}
|
||||
|
||||
// HasImages uses the canonical collector, including malformed image parts.
|
||||
// Callers must not treat collection failures as text-only input.
|
||||
func (p Probe) HasImages(ctx context.Context) (bool, error) {
|
||||
release, err := systemone.AcquireAdmission(ctx)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
defer release()
|
||||
|
||||
r, err := p.decisionRequest()
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
images, err := systemone.CollectImages(r)
|
||||
return len(images) > 0, err
|
||||
}
|
||||
func requireTextProbe(ctx context.Context, p Probe) error {
|
||||
images, err := p.HasImages(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if images {
|
||||
return &systemone.ValidationError{Kind: systemone.UnsupportedBackend, Err: errTextImages}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
var errTextImages = errors.New("text-only router classifier does not support image input")
|
||||
|
||||
// Check byte lengths before JSON conversion allocates a structured copy.
|
||||
func (p Probe) collectionBounds() error {
|
||||
if len(p.State) > systemone.MaxImageBodyBytes || len(p.Images) > systemone.MaxImageEncodedBytes || len(p.Prompt) > systemone.MaxImageBodyBytes {
|
||||
return fmt.Errorf("router probe exceeds decision collection budget")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,172 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package router
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"github.com/mudler/LocalAI/core/backend"
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/systemone"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
"image"
|
||||
"image/png"
|
||||
"reflect"
|
||||
"sync"
|
||||
)
|
||||
|
||||
var _ = It("TestRetryCollectionAdmission", func() {
|
||||
for i := 0; i < systemone.MaxAdmissions; i++ {
|
||||
release, err := systemone.AcquireAdmission(context.Background())
|
||||
if err != nil {
|
||||
Fail(fmt.Sprint(err))
|
||||
}
|
||||
defer release()
|
||||
}
|
||||
_, err := (Probe{Prompt: "text"}).HasImages(context.Background())
|
||||
if !errors.Is(err, systemone.ErrAdmissionCapacity) {
|
||||
Fail(fmt.Sprintf("unguarded collection: %v", err))
|
||||
}
|
||||
})
|
||||
|
||||
var _ = It("TestRetryCollectionCancellationAndRelease", func() {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
if _, err := (Probe{Prompt: "text"}).HasImages(ctx); !errors.Is(err, context.Canceled) {
|
||||
Fail(fmt.Sprintf("cancel: %v", err))
|
||||
}
|
||||
for i := 0; i < systemone.MaxAdmissions*2; i++ {
|
||||
if _, err := (Probe{State: []byte(`{`)}).HasImages(context.Background()); err == nil {
|
||||
Fail(fmt.Sprint("invalid JSON accepted"))
|
||||
}
|
||||
}
|
||||
for i := 0; i < systemone.MaxAdmissions; i++ {
|
||||
release, err := systemone.AcquireAdmission(context.Background())
|
||||
if err != nil {
|
||||
Fail(fmt.Sprintf("leaked lease: %v", err))
|
||||
}
|
||||
defer release()
|
||||
}
|
||||
})
|
||||
|
||||
var _ = It("TestRetryConcurrentCollectionBound", func() {
|
||||
// Leave one shared slot. Concurrent direct callers must never exceed it,
|
||||
// and must recover after collection errors and cancellations.
|
||||
for i := 0; i < systemone.MaxAdmissions-1; i++ {
|
||||
release, err := systemone.AcquireAdmission(context.Background())
|
||||
if err != nil {
|
||||
Fail(fmt.Sprint(err))
|
||||
}
|
||||
defer release()
|
||||
}
|
||||
start := make(chan struct{})
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < 32; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer GinkgoRecover()
|
||||
defer wg.Done()
|
||||
<-start
|
||||
for j := 0; j < 10; j++ {
|
||||
_, err := (Probe{State: []byte(`{"messages":[]}`)}).HasImages(context.Background())
|
||||
if err != nil && !errors.Is(err, systemone.ErrAdmissionCapacity) {
|
||||
Fail(fmt.Sprintf("collection: %v", err))
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
close(start)
|
||||
wg.Wait()
|
||||
release, err := systemone.AcquireAdmission(context.Background())
|
||||
if err != nil {
|
||||
Fail(fmt.Sprint(err))
|
||||
}
|
||||
defer release()
|
||||
if _, err = (Probe{Prompt: "text"}).HasImages(context.Background()); !errors.Is(err, systemone.ErrAdmissionCapacity) {
|
||||
Fail(fmt.Sprintf("bound: %v", err))
|
||||
}
|
||||
})
|
||||
|
||||
type retryInner struct {
|
||||
want Probe
|
||||
}
|
||||
|
||||
func (*retryInner) Name() string { return "spy" }
|
||||
func (s *retryInner) Classify(_ context.Context, p Probe) (Decision, error) {
|
||||
if !reflect.DeepEqual(s.want, p) {
|
||||
Fail(fmt.Sprintf("inner probe changed: %#v", p))
|
||||
}
|
||||
return Decision{Score: 1}, nil
|
||||
}
|
||||
|
||||
type retryNoEmbed struct{}
|
||||
|
||||
func (retryNoEmbed) Embed(context.Context, string) ([]float32, error) { panic("image cache embed") }
|
||||
|
||||
type retryNoStore struct{ backend.VectorStore }
|
||||
|
||||
func (retryNoStore) Search(context.Context, []float32) (float64, []byte, bool, error) {
|
||||
panic("image cache search")
|
||||
}
|
||||
func (retryNoStore) Insert(context.Context, []float32, []byte) error { panic("image cache insert") }
|
||||
|
||||
var _ = It("TestRetryCacheCompleteProbe", func() {
|
||||
for _, image := range []string{"AA==", "AQ=="} {
|
||||
p := Probe{Prompt: "identical", Messages: []string{"first", "second"}, State: []byte(`{}`), Images: []byte(`["data:image/png;base64,` + image + `"]`)}
|
||||
c := NewEmbeddingCacheClassifier(&retryInner{p}, retryNoEmbed{}, retryNoStore{}, .9, .5).WithTokenTrim(func(string) (int, error) { panic("image cache trim") }, 10)
|
||||
if _, err := c.Classify(context.Background(), p); err != nil {
|
||||
Fail(fmt.Sprint(err))
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
type retryRunner func(context.Context, *schema.SystemOneRequest) (*schema.SystemOneResponse, error)
|
||||
|
||||
func (f retryRunner) Decide(ctx context.Context, r *schema.SystemOneRequest) (*schema.SystemOneResponse, error) {
|
||||
return f(ctx, r)
|
||||
}
|
||||
|
||||
var _ = It("TestRetryImageCancellationNoFallback", func() {
|
||||
var b bytes.Buffer
|
||||
if err := png.Encode(&b, image.NewGray(image.Rect(0, 0, 1, 1))); err != nil {
|
||||
Fail(fmt.Sprint(err))
|
||||
}
|
||||
images, _ := json.Marshal([]string{"data:image/png;base64," + base64.StdEncoding.EncodeToString(b.Bytes())})
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
called := false
|
||||
c, err := NewDecisionsClassifier([]ScorePolicy{{Label: "visual", Description: "image"}}, retryRunner(func(context.Context, *schema.SystemOneRequest) (*schema.SystemOneResponse, error) {
|
||||
called = true
|
||||
cancel()
|
||||
return nil, context.Canceled
|
||||
}), 0)
|
||||
if err != nil {
|
||||
Fail(fmt.Sprint(err))
|
||||
}
|
||||
cfg := &config.ModelConfig{Name: "route", Router: config.RouterConfig{Classifier: "decisions", Fallback: "fallback", Candidates: []config.RouterCandidate{{Model: "candidate", Labels: []string{"visual"}}}}}
|
||||
_, err = Resolve(ctx, cfg, c, func(string) (*config.ModelConfig, error) { Fail(fmt.Sprint("cancel loaded fallback")); return nil, nil }, Probe{State: []byte(`{}`), Images: images})
|
||||
if !called || !errors.Is(err, context.Canceled) {
|
||||
Fail(fmt.Sprintf("cancel: called=%v err=%v", called, err))
|
||||
}
|
||||
})
|
||||
|
||||
var _ = It("TestRetryOversizedDirectProbe", func() {
|
||||
p := Probe{State: make([]byte, systemone.MaxImageBodyBytes+1)}
|
||||
if _, err := p.HasImages(context.Background()); err == nil {
|
||||
Fail(fmt.Sprint("oversized collection accepted"))
|
||||
}
|
||||
c, err := NewDecisionsClassifier([]ScorePolicy{{Label: "x", Description: "x"}}, retryRunner(func(context.Context, *schema.SystemOneRequest) (*schema.SystemOneResponse, error) {
|
||||
Fail(fmt.Sprint("oversized runner called"))
|
||||
return nil, nil
|
||||
}), 0)
|
||||
if err != nil {
|
||||
Fail(fmt.Sprint(err))
|
||||
}
|
||||
if _, err = c.Classify(context.Background(), p); err == nil {
|
||||
Fail(fmt.Sprint("oversized native probe accepted"))
|
||||
}
|
||||
})
|
||||
@@ -105,6 +105,13 @@ func (c *KNNClassifier) WithTokenTrim(tokenize func(string) (int, error), maxCon
|
||||
func (c *KNNClassifier) Name() string { return ClassifierKNN }
|
||||
|
||||
func (c *KNNClassifier) Classify(ctx context.Context, p Probe) (Decision, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return Decision{}, err
|
||||
}
|
||||
if err := requireTextProbe(ctx, p); err != nil {
|
||||
return Decision{}, err
|
||||
}
|
||||
|
||||
start := time.Now()
|
||||
|
||||
vec, err := c.embedder.Embed(ctx, trimmedProbeText(p, c.budget, identityRender))
|
||||
|
||||
@@ -0,0 +1,108 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package router
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"math"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/mudler/LocalAI/core/backend"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/systemone"
|
||||
)
|
||||
|
||||
// DecisionsClassifier asks independent native binary questions. Probabilities
|
||||
// are not normalized across policies: a prompt may activate every label.
|
||||
type DecisionsClassifier struct {
|
||||
runner backend.DecisionRunner
|
||||
labels []string
|
||||
questions map[string]schema.SystemOneQuestion
|
||||
threshold float64
|
||||
}
|
||||
|
||||
func NewDecisionsClassifier(policies []ScorePolicy, runner backend.DecisionRunner, threshold float64) (*DecisionsClassifier, error) {
|
||||
if runner == nil {
|
||||
return nil, fmt.Errorf("decisions runner is required")
|
||||
}
|
||||
if len(policies) == 0 || len(policies) > systemone.MaxQuestions {
|
||||
return nil, fmt.Errorf("decisions requires 1 to %d policies", systemone.MaxQuestions)
|
||||
}
|
||||
if math.IsNaN(threshold) || math.IsInf(threshold, 0) || threshold < 0 || threshold > 1 {
|
||||
return nil, fmt.Errorf("decisions activation_threshold must be finite and in [0,1]")
|
||||
}
|
||||
if threshold == 0 {
|
||||
threshold = .5
|
||||
}
|
||||
c := &DecisionsClassifier{runner: runner, threshold: threshold, questions: make(map[string]schema.SystemOneQuestion, len(policies))}
|
||||
seen := map[string]bool{}
|
||||
for i, p := range policies {
|
||||
if strings.TrimSpace(p.Label) == "" || strings.TrimSpace(p.Description) == "" || seen[p.Label] {
|
||||
return nil, fmt.Errorf("decisions policies require unique nonblank labels and descriptions")
|
||||
}
|
||||
seen[p.Label] = true
|
||||
criteria, _ := json.Marshal(map[string]string{"false": "The state does not match this policy: " + p.Description, "true": "The state matches this policy: " + p.Description})
|
||||
instruction, _ := json.Marshal("Does the state match this policy? " + p.Description)
|
||||
c.questions["p"+strconv.Itoa(i)] = schema.SystemOneQuestion{Type: "noul", Instructions: instruction, Criteria: criteria}
|
||||
c.labels = append(c.labels, p.Label)
|
||||
}
|
||||
if err := systemone.ValidateRequest(&schema.SystemOneRequest{State: json.RawMessage(`"x"`), Questions: c.questions}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return c, nil
|
||||
}
|
||||
func (c *DecisionsClassifier) Name() string { return ClassifierDecisions }
|
||||
func (c *DecisionsClassifier) Classify(ctx context.Context, p Probe) (Decision, error) {
|
||||
start := time.Now()
|
||||
if err := ctx.Err(); err != nil {
|
||||
return Decision{}, err
|
||||
}
|
||||
release, err := systemone.AcquireAdmission(ctx)
|
||||
if err != nil {
|
||||
return Decision{}, err
|
||||
}
|
||||
req, err := p.decisionRequest()
|
||||
if err != nil {
|
||||
release()
|
||||
return Decision{}, err
|
||||
}
|
||||
req.Questions = c.questions
|
||||
err = systemone.ValidateRequest(req)
|
||||
release()
|
||||
if err != nil {
|
||||
return Decision{}, err
|
||||
}
|
||||
// Native question framing is engine-owned. Raw JSON token counts are not a
|
||||
// context budget: do not trim or claim a fit; propagate native rejection.
|
||||
response, err := c.runner.Decide(ctx, req)
|
||||
if ctx.Err() != nil {
|
||||
return Decision{}, ctx.Err()
|
||||
}
|
||||
if err != nil {
|
||||
return Decision{}, err
|
||||
}
|
||||
if response == nil || len(response.Answers) != len(c.labels) {
|
||||
return Decision{}, fmt.Errorf("decisions response must answer exactly the requested questions")
|
||||
}
|
||||
d := Decision{ActivationThreshold: c.threshold}
|
||||
for i, label := range c.labels {
|
||||
a, ok := response.Answers["p"+strconv.Itoa(i)]
|
||||
if !ok || a.Type != "noul" || a.Noul == nil {
|
||||
return Decision{}, fmt.Errorf("decisions response requires numeric noul answers")
|
||||
}
|
||||
v := *a.Noul
|
||||
if math.IsNaN(v) || math.IsInf(v, 0) || v < 0 || v > 1 {
|
||||
return Decision{}, fmt.Errorf("decisions probability must be finite and in [0,1]")
|
||||
}
|
||||
d.LabelScores = append(d.LabelScores, LabelScore{Label: label, Score: v})
|
||||
if v >= c.threshold {
|
||||
d.Labels = append(d.Labels, label)
|
||||
}
|
||||
d.Score = math.Max(d.Score, v)
|
||||
}
|
||||
d.Latency = time.Since(start)
|
||||
return d, nil
|
||||
}
|
||||
@@ -0,0 +1,194 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package router_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"image"
|
||||
"image/png"
|
||||
"math"
|
||||
"strings"
|
||||
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/services/routing/router"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
type decisionFunc func(context.Context, *schema.SystemOneRequest) (*schema.SystemOneResponse, error)
|
||||
|
||||
func (f decisionFunc) Decide(ctx context.Context, r *schema.SystemOneRequest) (*schema.SystemOneResponse, error) {
|
||||
return f(ctx, r)
|
||||
}
|
||||
func noul(v float64) schema.SystemOneAnswer { return schema.SystemOneAnswer{Type: "noul", Noul: &v} }
|
||||
|
||||
var _ = Describe("native decisions classifier", func() {
|
||||
policies := []router.ScorePolicy{{Label: "code", Description: "writing code"}, {Label: "private", Description: "private information"}}
|
||||
var answers map[string]schema.SystemOneAnswer
|
||||
var runner decisionFunc
|
||||
BeforeEach(func() {
|
||||
answers = map[string]schema.SystemOneAnswer{"p0": noul(.8), "p1": noul(.7)}
|
||||
runner = func(_ context.Context, req *schema.SystemOneRequest) (*schema.SystemOneResponse, error) {
|
||||
Expect(string(req.State)).To(Equal(`"prompt"`))
|
||||
Expect(req.Questions).To(HaveLen(2))
|
||||
for _, q := range req.Questions {
|
||||
Expect(q.Type).To(Equal("noul"))
|
||||
var criteria map[string]string
|
||||
Expect(json.Unmarshal(q.Criteria, &criteria)).To(Succeed())
|
||||
Expect(criteria).To(HaveKey("false"))
|
||||
Expect(criteria).To(HaveKey("true"))
|
||||
}
|
||||
return &schema.SystemOneResponse{Answers: answers}, nil
|
||||
}
|
||||
})
|
||||
It("preserves overlapping independent probabilities and policy ordering", func() {
|
||||
c, err := router.NewDecisionsClassifier(policies, runner, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
d, err := c.Classify(context.Background(), router.Probe{Prompt: "prompt"})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(d.Labels).To(Equal([]string{"code", "private"}))
|
||||
Expect(d.Score).To(Equal(.8))
|
||||
Expect(d.LabelScores).To(Equal([]router.LabelScore{{Label: "code", Score: .8}, {Label: "private", Score: .7}}))
|
||||
Expect(d.ActivationThreshold).To(Equal(.5))
|
||||
})
|
||||
It("accepts zero and threshold equality; abstains without top-one", func() {
|
||||
c, err := router.NewDecisionsClassifier(policies, runner, .5)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
answers["p0"] = noul(0)
|
||||
answers["p1"] = noul(.5)
|
||||
d, err := c.Classify(context.Background(), router.Probe{Prompt: "prompt"})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(d.Labels).To(Equal([]string{"private"}))
|
||||
answers["p1"] = noul(.49)
|
||||
d, err = c.Classify(context.Background(), router.Probe{Prompt: "prompt"})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(d.Labels).To(BeEmpty())
|
||||
})
|
||||
It("rejects missing extra null wrong-type nonfinite and out-of-range answers", func() {
|
||||
c, err := router.NewDecisionsClassifier(policies, runner, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
for _, a := range []map[string]schema.SystemOneAnswer{
|
||||
nil, {"p0": noul(.8)}, {"p0": noul(.8), "p1": noul(.7), "extra": noul(.9)},
|
||||
{"p0": {Type: "noul"}, "p1": noul(.7)}, {"p0": {Type: "score", Noul: noul(.5).Noul}, "p1": noul(.7)},
|
||||
{"p0": noul(math.NaN()), "p1": noul(.7)}, {"p0": noul(math.Inf(1)), "p1": noul(.7)}, {"p0": noul(-.1), "p1": noul(.7)}, {"p0": noul(1.1), "p1": noul(.7)},
|
||||
} {
|
||||
answers = a
|
||||
_, err = c.Classify(context.Background(), router.Probe{Prompt: "prompt"})
|
||||
Expect(err).To(HaveOccurred())
|
||||
}
|
||||
})
|
||||
It("distinguishes wire null and missing noul from zero and rejects substituted IDs", func() {
|
||||
c, err := router.NewDecisionsClassifier(policies, runner, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
for _, raw := range []string{`{"p0":null,"p1":{"type":"noul","noul":0.7}}`, `{"p0":{"type":"noul","noul":null},"p1":{"type":"noul","noul":0.7}}`, `{"p0":{"type":"noul","noul":0.8},"other":{"type":"noul","noul":0.7}}`} {
|
||||
Expect(json.Unmarshal([]byte(raw), &answers)).To(Succeed())
|
||||
_, err = c.Classify(context.Background(), router.Probe{Prompt: "prompt"})
|
||||
Expect(err).To(HaveOccurred())
|
||||
answers = nil
|
||||
}
|
||||
})
|
||||
|
||||
It("validates labels thresholds questions and request bounds", func() {
|
||||
for _, v := range []float64{-.1, 1.1, math.NaN(), math.Inf(1)} {
|
||||
_, err := router.NewDecisionsClassifier(policies, runner, v)
|
||||
Expect(err).To(HaveOccurred())
|
||||
}
|
||||
for _, p := range [][]router.ScorePolicy{nil, {{Label: "same", Description: "a"}, {Label: "same", Description: "b"}}, {{Label: " ", Description: "a"}}, {{Label: "x", Description: " "}}, make([]router.ScorePolicy, 65)} {
|
||||
_, err := router.NewDecisionsClassifier(p, runner, 0)
|
||||
Expect(err).To(HaveOccurred())
|
||||
}
|
||||
c, err := router.NewDecisionsClassifier(policies, runner, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
_, err = c.Classify(context.Background(), router.Probe{Prompt: strings.Repeat("x", 65536)})
|
||||
Expect(err).To(HaveOccurred())
|
||||
})
|
||||
It("does not call the adapter after parent cancellation", func() {
|
||||
c, err := router.NewDecisionsClassifier(policies, decisionFunc(func(context.Context, *schema.SystemOneRequest) (*schema.SystemOneResponse, error) {
|
||||
Fail("called")
|
||||
return nil, nil
|
||||
}), 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
_, err = c.Classify(ctx, router.Probe{Prompt: "prompt"})
|
||||
Expect(err).To(MatchError(context.Canceled))
|
||||
})
|
||||
It("uses first-superset and falls back on overload, context rejection or abstention, not parent cancellation", func() {
|
||||
cfg := &config.ModelConfig{Name: "route", Router: config.RouterConfig{Classifier: "decisions", Fallback: "fallback", Candidates: []config.RouterCandidate{{Model: "small", Labels: []string{"code"}}, {Model: "both", Labels: []string{"code", "private"}}}}}
|
||||
loads := 0
|
||||
load := func(name string) (*config.ModelConfig, error) { loads++; return &config.ModelConfig{Name: name}, nil }
|
||||
c, err := router.NewDecisionsClassifier(policies, runner, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
r, err := router.Resolve(context.Background(), cfg, c, load, router.Probe{Prompt: "prompt"})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(r.ChosenModel).To(Equal("both"))
|
||||
for _, failure := range []error{errors.New("native decision operation capacity reached"), errors.New("native context exceeded")} {
|
||||
c, err = router.NewDecisionsClassifier(policies, decisionFunc(func(context.Context, *schema.SystemOneRequest) (*schema.SystemOneResponse, error) {
|
||||
return nil, failure
|
||||
}), 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
r, err = router.Resolve(context.Background(), cfg, c, load, router.Probe{Prompt: "prompt"})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(r.UsedFallback).To(BeTrue())
|
||||
}
|
||||
answers["p0"] = noul(0)
|
||||
answers["p1"] = noul(0)
|
||||
c, err = router.NewDecisionsClassifier(policies, runner, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
r, err = router.Resolve(context.Background(), cfg, c, load, router.Probe{Prompt: "prompt"})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(r.UsedFallback).To(BeTrue())
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
c, err = router.NewDecisionsClassifier(policies, decisionFunc(func(context.Context, *schema.SystemOneRequest) (*schema.SystemOneResponse, error) {
|
||||
cancel()
|
||||
return nil, context.Canceled
|
||||
}), 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
loads = 0
|
||||
_, err = router.Resolve(ctx, cfg, c, load, router.Probe{Prompt: "prompt"})
|
||||
Expect(err).To(MatchError(context.Canceled))
|
||||
Expect(loads).To(BeZero())
|
||||
cancel()
|
||||
loads = 0
|
||||
_, err = router.Resolve(ctx, cfg, c, load, router.Probe{})
|
||||
Expect(err).To(MatchError(context.Canceled))
|
||||
_, err = router.Resolve(ctx, cfg, nil, load, router.Probe{})
|
||||
Expect(err).To(MatchError(context.Canceled))
|
||||
Expect(loads).To(BeZero())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("multimodal native routing", func() {
|
||||
It("passes image-only structured state unchanged and rejects invalid input before runner", func() {
|
||||
// Portable valid 1x1 PNG, generated with the standard encoder below.
|
||||
var b bytes.Buffer
|
||||
Expect(png.Encode(&b, image.NewGray(image.Rect(0, 0, 1, 1)))).To(Succeed())
|
||||
u := "data:image/png;base64," + base64.StdEncoding.EncodeToString(b.Bytes())
|
||||
state, _ := json.Marshal([]any{map[string]any{"role": "user", "content": []any{map[string]any{"type": "image_url", "image_url": map[string]string{"url": u}}}}})
|
||||
calls := 0
|
||||
c, err := router.NewDecisionsClassifier([]router.ScorePolicy{{Label: "visual", Description: "visual content"}}, decisionFunc(func(_ context.Context, r *schema.SystemOneRequest) (*schema.SystemOneResponse, error) {
|
||||
calls++
|
||||
Expect(r.State).To(Equal(json.RawMessage(state)))
|
||||
Expect(r.Images).To(BeEmpty())
|
||||
return &schema.SystemOneResponse{Answers: map[string]schema.SystemOneAnswer{"p0": noul(.9)}}, nil
|
||||
}), 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
d, err := c.Classify(context.Background(), router.Probe{State: state})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(d.Labels).To(Equal([]string{"visual"}))
|
||||
_, err = c.Classify(context.Background(), router.Probe{State: json.RawMessage(`{}`), Images: json.RawMessage(`["https://example.org/x.png"]`)})
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(calls).To(Equal(1))
|
||||
})
|
||||
It("rejects images in every text classifier before cache trimming or model use", func() {
|
||||
p := router.Probe{Prompt: "same text", State: json.RawMessage(`{}`), Images: json.RawMessage(`["data:image/png;base64,AA=="]`)}
|
||||
for _, c := range []router.Classifier{&router.ScoreClassifier{}, &router.RerankClassifier{}, &router.KNNClassifier{}} {
|
||||
_, err := c.Classify(context.Background(), p)
|
||||
Expect(err).To(MatchError(ContainSubstring("does not support image")))
|
||||
}
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,75 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package router
|
||||
|
||||
import (
|
||||
"context"
|
||||
"runtime"
|
||||
"strings"
|
||||
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/systemone"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("Direct prompt escape bounds", func() {
|
||||
for _, native := range []bool{false, true} {
|
||||
It("rejects escaped expansion before allocation", func() {
|
||||
p := Probe{Prompt: strings.Repeat("\x00", systemone.MaxImageBodyBytes)}
|
||||
called := false
|
||||
c, err := NewDecisionsClassifier([]ScorePolicy{{Label: "a", Description: "a"}}, retryRunner(func(context.Context, *schema.SystemOneRequest) (*schema.SystemOneResponse, error) {
|
||||
called = true
|
||||
return nil, nil
|
||||
}), 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
runtime.GC()
|
||||
var a, b runtime.MemStats
|
||||
runtime.ReadMemStats(&a)
|
||||
if native {
|
||||
_, err = c.Classify(context.Background(), p)
|
||||
} else {
|
||||
_, err = p.HasImages(context.Background())
|
||||
}
|
||||
runtime.ReadMemStats(&b)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(called).To(BeFalse())
|
||||
Expect(b.TotalAlloc - a.TotalAlloc).To(BeNumerically("<", systemone.MaxBodyBytes))
|
||||
})
|
||||
}
|
||||
It("accepts the serialized boundary and ordinary text limit", func() {
|
||||
available := systemone.MaxImageBodyBytes - 2
|
||||
text := strings.Repeat("\x00", available/6) + strings.Repeat("a", available%6)
|
||||
r, err := (Probe{Prompt: text}).decisionRequest()
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(r.State).To(HaveLen(systemone.MaxImageBodyBytes))
|
||||
_, err = (Probe{Prompt: text + "a"}).decisionRequest()
|
||||
Expect(err).To(HaveOccurred())
|
||||
images, err := (Probe{Prompt: strings.Repeat("a", systemone.MaxBodyBytes)}).HasImages(context.Background())
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(images).To(BeFalse())
|
||||
})
|
||||
It("uses only configured runtime fallback for oversized prompts", func() {
|
||||
c, err := NewDecisionsClassifier([]ScorePolicy{{Label: "a", Description: "a"}}, retryRunner(func(context.Context, *schema.SystemOneRequest) (*schema.SystemOneResponse, error) {
|
||||
Fail("runner called")
|
||||
return nil, nil
|
||||
}), 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
cfg := &config.ModelConfig{Name: "route", Router: config.RouterConfig{Classifier: "decisions", Candidates: []config.RouterCandidate{{Model: "candidate", Labels: []string{"a"}}}}}
|
||||
loaded := []string{}
|
||||
loader := func(name string) (*config.ModelConfig, error) {
|
||||
loaded = append(loaded, name)
|
||||
return &config.ModelConfig{Name: name}, nil
|
||||
}
|
||||
p := Probe{Prompt: strings.Repeat("\x00", systemone.MaxImageBodyBytes)}
|
||||
_, err = Resolve(context.Background(), cfg, c, loader, p)
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(loaded).To(BeEmpty())
|
||||
cfg.Router.Fallback = "fallback"
|
||||
result, err := Resolve(context.Background(), cfg, c, loader, p)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(result.UsedFallback).To(BeTrue())
|
||||
Expect(loaded).To(Equal([]string{"fallback"}))
|
||||
})
|
||||
|
||||
})
|
||||
@@ -81,6 +81,13 @@ func (c *RerankClassifier) WithTokenTrim(tokenize func(string) (int, error), max
|
||||
func (c *RerankClassifier) Name() string { return ClassifierColbert }
|
||||
|
||||
func (c *RerankClassifier) Classify(ctx context.Context, p Probe) (Decision, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return Decision{}, err
|
||||
}
|
||||
if err := requireTextProbe(ctx, p); err != nil {
|
||||
return Decision{}, err
|
||||
}
|
||||
|
||||
start := time.Now()
|
||||
query := trimmedProbeText(p, c.budget, identityRender)
|
||||
key := cacheKey(query)
|
||||
|
||||
@@ -77,12 +77,19 @@ func Resolve(ctx context.Context, routerCfg *config.ModelConfig, classifier Clas
|
||||
return nil, fmt.Errorf("router.Resolve: config has no router block")
|
||||
}
|
||||
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if classifier == nil {
|
||||
return resolveFallback(routerCfg, loader, Decision{}, LabelFallback, "classifier unavailable")
|
||||
}
|
||||
|
||||
start := time.Now()
|
||||
decision, err := classifier.Classify(ctx, probe)
|
||||
if ctx.Err() != nil {
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
if err != nil {
|
||||
return resolveFallback(routerCfg, loader, Decision{Latency: time.Since(start)}, classifier.Name(), "classifier error: "+err.Error())
|
||||
}
|
||||
|
||||
@@ -309,6 +309,13 @@ func (c *ScoreClassifier) SlotFillPrompt(p Probe, label, firstSlot string) (stri
|
||||
}
|
||||
|
||||
func (c *ScoreClassifier) Classify(ctx context.Context, p Probe) (Decision, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return Decision{}, err
|
||||
}
|
||||
if err := requireTextProbe(ctx, p); err != nil {
|
||||
return Decision{}, err
|
||||
}
|
||||
|
||||
start := time.Now()
|
||||
|
||||
prompt, userText, err := c.renderProbe(p)
|
||||
|
||||
@@ -19,6 +19,7 @@ package router
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"time"
|
||||
)
|
||||
|
||||
@@ -27,6 +28,12 @@ import (
|
||||
// middleware does the schema-shape extraction); the classifier never
|
||||
// inspects the original request struct.
|
||||
type Probe struct {
|
||||
// State preserves ordered chat content, including image-only turns. Images
|
||||
// holds optional top-level data URLs; embedded images are not duplicated here.
|
||||
State json.RawMessage
|
||||
Images json.RawMessage
|
||||
InputError error
|
||||
|
||||
// Prompt is the merged user-visible text. For chat completions it
|
||||
// is the concatenation of message contents (separated by newlines);
|
||||
// for plain completions it is the raw prompt.
|
||||
@@ -47,7 +54,7 @@ type Probe struct {
|
||||
// surrounding middleware picks the first candidate whose Labels
|
||||
// superset the active label set; that lets one prompt activate multiple
|
||||
// policies and route to a model capable of all of them. Score is the
|
||||
// softmax probability of the top label — kept for the decision log so
|
||||
// maximum label score (independent P(true) for decisions) — kept for the decision log so
|
||||
// admins can spot uncertain calls.
|
||||
type Decision struct {
|
||||
Labels []string `json:"labels"`
|
||||
@@ -57,7 +64,8 @@ type Decision struct {
|
||||
// LabelScores carries the full per-label score distribution that
|
||||
// fed the threshold check, in policy-declaration order. Score
|
||||
// classifier emits softmax probabilities (sum to 1.0); rerank
|
||||
// emits independent relevance in [0, 1]. Empty on cache hits —
|
||||
// emits independent relevance and decisions emits independent P(true)
|
||||
// in [0, 1]. Empty on cache hits —
|
||||
// the cache stores only the final label set, not the distribution.
|
||||
LabelScores []LabelScore `json:"label_scores,omitempty"`
|
||||
|
||||
@@ -144,7 +152,8 @@ const (
|
||||
// model (Arch-Router-style) to score each policy label as a
|
||||
// continuation of the routing prompt. See router/score.go for
|
||||
// the full rationale.
|
||||
ClassifierScore = "score"
|
||||
ClassifierScore = "score"
|
||||
ClassifierDecisions = "decisions"
|
||||
|
||||
// ClassifierColbert picks labels by reranking each policy's
|
||||
// description against the prompt via LocalAI's rerankers
|
||||
@@ -169,7 +178,7 @@ const (
|
||||
// available_classifiers field both derive from it, so a new classifier
|
||||
// added to the buildClassifier switch shows up on every surface by
|
||||
// extending this one slice (colbert once drifted out of both).
|
||||
var AllClassifiers = []string{ClassifierScore, ClassifierColbert, ClassifierKNN}
|
||||
var AllClassifiers = []string{ClassifierScore, ClassifierColbert, ClassifierKNN, ClassifierDecisions}
|
||||
|
||||
// LabelFallback is the synthetic label written to the decision
|
||||
// store when the middleware uses cfg.Router.Fallback rather than a
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package systemone
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// MaxAdmissions bounds decision-owned buffering, decoding and retained request
|
||||
// bodies across HTTP and internal callers. Leases last until actual work ends,
|
||||
// not merely until validation or a cancelled caller returns. This is separate
|
||||
// from the backend's abandoned-operation ceiling, and does not bound model RSS
|
||||
// or memory already owned by callers/upstream middleware.
|
||||
const MaxAdmissions = 8
|
||||
|
||||
var admissions = make(chan struct{}, MaxAdmissions)
|
||||
var ErrAdmissionCapacity = errors.New("decision admission capacity reached")
|
||||
|
||||
// AcquireAdmission fails promptly on saturation: no unbounded waiter queue.
|
||||
// Validation helpers do not acquire leases, avoiding nested acquisition.
|
||||
func AcquireAdmission(ctx context.Context) (func(), error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
select {
|
||||
case admissions <- struct{}{}:
|
||||
default:
|
||||
return nil, ErrAdmissionCapacity
|
||||
}
|
||||
var once sync.Once
|
||||
release := func() { once.Do(func() { <-admissions }) }
|
||||
if err := ctx.Err(); err != nil {
|
||||
release()
|
||||
return nil, err
|
||||
}
|
||||
return release, nil
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package systemone
|
||||
|
||||
import (
|
||||
"context"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("Decision admission leases", func() {
|
||||
It("fails promptly on saturation, respects cancellation and releases once", func() {
|
||||
var releases []func()
|
||||
defer func() {
|
||||
for _, r := range releases {
|
||||
r()
|
||||
}
|
||||
}()
|
||||
for i := 0; i < MaxAdmissions; i++ {
|
||||
r, err := AcquireAdmission(context.Background())
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
releases = append(releases, r)
|
||||
}
|
||||
_, err := AcquireAdmission(context.Background())
|
||||
Expect(err).To(MatchError(ErrAdmissionCapacity))
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
_, err = AcquireAdmission(ctx)
|
||||
Expect(err).To(MatchError(context.Canceled))
|
||||
releases[0]()
|
||||
releases[0]()
|
||||
r, err := AcquireAdmission(context.Background())
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
r()
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,163 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package systemone
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"image"
|
||||
_ "image/jpeg"
|
||||
_ "image/png"
|
||||
"strings"
|
||||
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
)
|
||||
|
||||
// Limits are shared by internal and public decision callers, not global HTTP limits.
|
||||
const (
|
||||
MaxImages = 8
|
||||
MaxImageDecodedBytes = 8 << 20
|
||||
MaxImageEncodedBytes = 12 << 20
|
||||
MaxImageBodyBytes = 16 << 20
|
||||
MaxImageDimension = 4096
|
||||
MaxImagePixels = 16_000_000
|
||||
MaxResponseBytes = 64 << 10
|
||||
)
|
||||
|
||||
const InputTooLarge ErrorKind = "input_too_large"
|
||||
|
||||
func imageError(kind ErrorKind, message string) error {
|
||||
return &ValidationError{kind, fmt.Errorf("%s", message)}
|
||||
}
|
||||
|
||||
// CollectImages preserves URL bytes and order (top-level first, then chat parts).
|
||||
// Only message content has image semantics; arbitrary JSON domain state does not.
|
||||
// Collection never fetches or decodes data. ValidateImages must precede inference.
|
||||
func CollectImages(req *schema.SystemOneRequest) ([]string, error) {
|
||||
if req == nil {
|
||||
return nil, imageError(InvalidRequest, "request is required")
|
||||
}
|
||||
var images []string
|
||||
if len(req.Images) > 0 {
|
||||
if err := json.Unmarshal(req.Images, &images); err != nil {
|
||||
return nil, imageError(InvalidRequest, "images must be an array of data URLs")
|
||||
}
|
||||
}
|
||||
var state any
|
||||
if err := json.Unmarshal(req.State, &state); err != nil {
|
||||
return nil, imageError(InvalidRequest, "state is not valid JSON")
|
||||
}
|
||||
if wrapped, ok := state.(map[string]any); ok {
|
||||
state = wrapped["messages"]
|
||||
}
|
||||
messages, _ := state.([]any)
|
||||
for _, message := range messages {
|
||||
msg, _ := message.(map[string]any)
|
||||
content, _ := msg["content"].([]any)
|
||||
for _, value := range content {
|
||||
part, _ := value.(map[string]any)
|
||||
switch part["type"] {
|
||||
case "image_url":
|
||||
value := part["image_url"]
|
||||
if obj, ok := value.(map[string]any); ok {
|
||||
value = obj["url"]
|
||||
}
|
||||
url, ok := value.(string)
|
||||
if !ok {
|
||||
return nil, imageError(InvalidRequest, "image_url must contain a URL")
|
||||
}
|
||||
images = append(images, url)
|
||||
case "image":
|
||||
source, _ := part["source"].(map[string]any)
|
||||
mime, mok := source["media_type"].(string)
|
||||
data, dok := source["data"].(string)
|
||||
if source["type"] != "base64" || !mok || !dok {
|
||||
return nil, imageError(InvalidRequest, "image source must contain base64 data and media_type")
|
||||
}
|
||||
images = append(images, "data:"+mime+";base64,"+data)
|
||||
}
|
||||
}
|
||||
}
|
||||
return images, nil
|
||||
}
|
||||
|
||||
// ValidateImages bounds encoded and decoded allocation before reading headers.
|
||||
// DecodeConfig reads dimensions without allocating a pixel buffer. Native decoders
|
||||
// must enforce the same bounds independently for direct RPC callers.
|
||||
func ValidateImages(images []string) error {
|
||||
if len(images) > MaxImages {
|
||||
return imageError(InputTooLarge, "too many decision images")
|
||||
}
|
||||
encoded := 0
|
||||
for _, url := range images {
|
||||
if len(url) > MaxImageEncodedBytes-encoded {
|
||||
return imageError(InputTooLarge, "decision images exceed encoded aggregate limit")
|
||||
}
|
||||
encoded += len(url)
|
||||
}
|
||||
decoded, pixels := 0, int64(0)
|
||||
for _, url := range images {
|
||||
header, data, ok := strings.Cut(url, ",")
|
||||
format := ""
|
||||
switch header {
|
||||
case "data:image/png;base64":
|
||||
format = "png"
|
||||
case "data:image/jpeg;base64":
|
||||
format = "jpeg"
|
||||
}
|
||||
if !ok || format == "" || data == "" || len(data)%4 != 0 {
|
||||
return imageError(InvalidRequest, "images must be PNG or JPEG base64 data URLs")
|
||||
}
|
||||
// Go's base64 decoder tolerates CR/LF even in Strict mode; the wire contract does not.
|
||||
for i := 0; i < len(data); i++ {
|
||||
c := data[i]
|
||||
if !(c >= 'A' && c <= 'Z' || c >= 'a' && c <= 'z' || c >= '0' && c <= '9' || c == '+' || c == '/' || c == '=') {
|
||||
return imageError(InvalidRequest, "invalid base64 image data")
|
||||
}
|
||||
}
|
||||
size := base64.StdEncoding.DecodedLen(len(data))
|
||||
if strings.HasSuffix(data, "==") {
|
||||
size -= 2
|
||||
} else if strings.HasSuffix(data, "=") {
|
||||
size--
|
||||
}
|
||||
if size > MaxImageDecodedBytes-decoded {
|
||||
return imageError(InputTooLarge, "decision images exceed decoded aggregate limit")
|
||||
}
|
||||
raw, err := base64.StdEncoding.Strict().DecodeString(data)
|
||||
if err != nil {
|
||||
return imageError(InvalidRequest, "invalid base64 image data")
|
||||
}
|
||||
decoded += len(raw)
|
||||
cfg, actual, err := image.DecodeConfig(bytes.NewReader(raw))
|
||||
if err != nil || actual != format || cfg.Width <= 0 || cfg.Height <= 0 {
|
||||
return imageError(InvalidRequest, "invalid image header or MIME mismatch")
|
||||
}
|
||||
if cfg.Width > MaxImageDimension || cfg.Height > MaxImageDimension {
|
||||
return imageError(InputTooLarge, "decision image dimensions exceed limit")
|
||||
}
|
||||
pixels += int64(cfg.Width) * int64(cfg.Height)
|
||||
if pixels > MaxImagePixels {
|
||||
return imageError(InputTooLarge, "decision images exceed aggregate pixel limit")
|
||||
}
|
||||
// Only allocate pixels after header bounds. DecodeConfig alone accepts
|
||||
// truncated streams and corrupt pixel payloads.
|
||||
if _, _, err := image.Decode(bytes.NewReader(raw)); err != nil {
|
||||
return imageError(InvalidRequest, "invalid image pixel data")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// RequestBodyLimit only selects a budget; it does not replace validation.
|
||||
func RequestBodyLimit(req *schema.SystemOneRequest) (int, error) {
|
||||
images, err := CollectImages(req)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if len(images) > 0 {
|
||||
return MaxImageBodyBytes, nil
|
||||
}
|
||||
return MaxBodyBytes, nil
|
||||
}
|
||||
@@ -0,0 +1,146 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package systemone
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"image"
|
||||
"image/jpeg"
|
||||
"image/png"
|
||||
"strings"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
func pngURL(w, h int) string {
|
||||
var b bytes.Buffer
|
||||
Expect(png.Encode(&b, image.NewGray(image.Rect(0, 0, w, h)))).To(Succeed())
|
||||
return "data:image/png;base64," + base64.StdEncoding.EncodeToString(b.Bytes())
|
||||
}
|
||||
|
||||
var _ = Describe("Canonical decision images", func() {
|
||||
It("validates headers, not just base64", func() {
|
||||
for _, u := range []string{"data:image/png;base64,AA==", "data:image/gif;base64,AA==", "https://example.org/x.png", strings.Replace(pngURL(1, 1), "image/png", "image/jpeg", 1)} {
|
||||
r := validRequest()
|
||||
r.Images, _ = json.Marshal([]string{u})
|
||||
Expect(ValidateRequestStructure(r)).NotTo(Succeed())
|
||||
}
|
||||
})
|
||||
It("accepts image-only structured state and preserves input", func() {
|
||||
r := validRequest()
|
||||
r.State = json.RawMessage(`{}`)
|
||||
r.Images, _ = json.Marshal([]string{pngURL(1, 1)})
|
||||
before, _ := json.Marshal(r)
|
||||
Expect(ValidateRequest(r)).To(Succeed())
|
||||
after, _ := json.Marshal(r)
|
||||
Expect(after).To(Equal(before))
|
||||
})
|
||||
It("rejects excessive dimensions and aggregate pixels", func() {
|
||||
for _, urls := range [][]string{{pngURL(4097, 1)}, {pngURL(3000, 3000), pngURL(3000, 3000)}} {
|
||||
r := validRequest()
|
||||
r.Images, _ = json.Marshal(urls)
|
||||
Expect(ValidateRequestStructure(r)).NotTo(Succeed())
|
||||
}
|
||||
})
|
||||
It("accepts image-bearing bodies above the text cap but not empty image arrays", func() {
|
||||
r := validRequest()
|
||||
r.State, _ = json.Marshal(map[string]string{"text": strings.Repeat("a", MaxBodyBytes)})
|
||||
r.Images = json.RawMessage(`[]`)
|
||||
Expect(ValidateRequest(r)).NotTo(Succeed())
|
||||
r.Images, _ = json.Marshal([]string{pngURL(1, 1)})
|
||||
Expect(ValidateRequest(r)).To(Succeed())
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("Decision image boundaries", func() {
|
||||
It("collects both chat formats in order without interpreting domain objects", func() {
|
||||
u := pngURL(1, 1)
|
||||
data := strings.SplitN(u, ",", 2)[1]
|
||||
r := validRequest()
|
||||
r.Images, _ = json.Marshal([]string{u})
|
||||
r.State, _ = json.Marshal(map[string]any{"messages": []any{map[string]any{"role": "user", "content": []any{map[string]any{"type": "image_url", "image_url": map[string]string{"url": u}}, map[string]any{"type": "image", "source": map[string]string{"type": "base64", "media_type": "image/png", "data": data}}}}}})
|
||||
urls, err := CollectImages(r)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(urls).To(Equal([]string{u, u, u}))
|
||||
Expect(ValidateRequest(r)).To(Succeed())
|
||||
r.Images = nil
|
||||
r.State = json.RawMessage(`{"type":"image","source":"domain data"}`)
|
||||
urls, err = CollectImages(r)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(urls).To(BeEmpty())
|
||||
})
|
||||
It("rejects malformed base64, headers and non-array images", func() {
|
||||
for _, u := range []string{"data:image/png;base64,AB==", "data:image/png;base64,AA\n=", "data:image/png;charset=utf8;base64,AA==", "data:image/png;base64,====", "file:///x", "data:image/png;base64,AAA"} {
|
||||
Expect(ValidateImages([]string{u})).NotTo(Succeed())
|
||||
}
|
||||
for _, raw := range []string{`""`, `{}`, `[null]`, `[1]`} {
|
||||
r := validRequest()
|
||||
r.Images = json.RawMessage(raw)
|
||||
Expect(ValidateRequestStructure(r)).NotTo(Succeed())
|
||||
}
|
||||
})
|
||||
It("bounds count, encoded and decoded aggregate before image decoding", func() {
|
||||
u := pngURL(1, 1)
|
||||
Expect(ValidateImages([]string{u, u, u, u, u, u, u, u})).To(Succeed())
|
||||
for _, urls := range [][]string{{u, u, u, u, u, u, u, u, u}, {strings.Repeat("x", MaxImageEncodedBytes+1)}, {"data:image/png;base64," + strings.Repeat("A", ((MaxImageDecodedBytes+3)/3)*4)}} {
|
||||
err := ValidateImages(urls)
|
||||
Expect(err).To(BeAssignableToTypeOf(&ValidationError{}))
|
||||
Expect(err.(*ValidationError).Kind).To(Equal(InputTooLarge))
|
||||
}
|
||||
Expect(ValidateImages([]string{pngURL(4000, 4000)})).To(Succeed())
|
||||
})
|
||||
It("keeps absent/null/empty image budgets and missing-state rejection", func() {
|
||||
for _, raw := range []json.RawMessage{nil, json.RawMessage(`null`), json.RawMessage(`[]`)} {
|
||||
r := validRequest()
|
||||
r.Images = raw
|
||||
limit, err := RequestBodyLimit(r)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(limit).To(Equal(MaxBodyBytes))
|
||||
}
|
||||
for _, raw := range []json.RawMessage{nil, json.RawMessage(`null`), json.RawMessage(`" "`)} {
|
||||
r := validRequest()
|
||||
r.State = raw
|
||||
r.Images, _ = json.Marshal([]string{pngURL(1, 1)})
|
||||
Expect(ValidateRequestStructure(r)).NotTo(Succeed())
|
||||
}
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("Complete image validation", func() {
|
||||
It("rejects header-only PNG and JPEG while accepting valid JPEG", func() {
|
||||
var b bytes.Buffer
|
||||
Expect(jpeg.Encode(&b, image.NewGray(image.Rect(0, 0, 8, 8)), nil)).To(Succeed())
|
||||
jpegBytes := b.Bytes()
|
||||
Expect(ValidateImages([]string{"data:image/jpeg;base64," + base64.StdEncoding.EncodeToString(jpegBytes)})).To(Succeed())
|
||||
raw, err := base64.StdEncoding.DecodeString(strings.SplitN(pngURL(8, 8), ",", 2)[1])
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
for _, item := range []struct {
|
||||
mime string
|
||||
data []byte
|
||||
}{{"png", raw[:33]}, {"jpeg", jpegBytes[:len(jpegBytes)-10]}} {
|
||||
Expect(ValidateImages([]string{"data:image/" + item.mime + ";base64," + base64.StdEncoding.EncodeToString(item.data)})).NotTo(Succeed())
|
||||
}
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("Exact decision byte boundaries", func() {
|
||||
It("accepts exactly the decoded budget and rejects the next byte", func() {
|
||||
raw, _ := base64.StdEncoding.DecodeString(strings.SplitN(pngURL(1, 1), ",", 2)[1])
|
||||
raw = append(raw, make([]byte, MaxImageDecodedBytes-len(raw))...)
|
||||
url := func(b []byte) string { return "data:image/png;base64," + base64.StdEncoding.EncodeToString(b) }
|
||||
Expect(ValidateImages([]string{url(raw)})).To(Succeed())
|
||||
err := ValidateImages([]string{url(append(raw, 0))})
|
||||
Expect(err.(*ValidationError).Kind).To(Equal(InputTooLarge))
|
||||
})
|
||||
It("preserves distinct image ordering and bytes", func() {
|
||||
a, b, c := pngURL(1, 1), pngURL(2, 1), pngURL(3, 1)
|
||||
r := validRequest()
|
||||
r.Images, _ = json.Marshal([]string{a, b})
|
||||
r.State, _ = json.Marshal([]any{map[string]any{"content": []any{map[string]any{"type": "image_url", "image_url": c}}}})
|
||||
images, err := CollectImages(r)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(images).To(Equal([]string{a, b, c}))
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,47 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package systemone
|
||||
|
||||
import "unicode/utf8"
|
||||
|
||||
func spendJSONBytes(left *int, n int) bool {
|
||||
if n > *left {
|
||||
return false
|
||||
}
|
||||
*left -= n
|
||||
return true
|
||||
}
|
||||
|
||||
// SpendJSONString budgets encoding/json string escaping without allocating output.
|
||||
func SpendJSONString(s string, left *int) bool {
|
||||
if !spendJSONBytes(left, 2) || len(s) > *left {
|
||||
return false
|
||||
}
|
||||
for i := 0; i < len(s); {
|
||||
c := s[i]
|
||||
n := 1
|
||||
switch {
|
||||
case c == '"' || c == '\\':
|
||||
n = 2
|
||||
case c == '\n' || c == '\r' || c == '\t' || c == '\b' || c == '\f':
|
||||
n = 2
|
||||
case c < 0x20 || c == '<' || c == '>' || c == '&':
|
||||
n = 6
|
||||
case c >= utf8.RuneSelf:
|
||||
r, size := utf8.DecodeRuneInString(s[i:])
|
||||
if r == utf8.RuneError && size == 1 {
|
||||
n = 6
|
||||
} else if r == '\u2028' || r == '\u2029' {
|
||||
n = 6
|
||||
i += size - 1
|
||||
} else {
|
||||
n = size
|
||||
i += size - 1
|
||||
}
|
||||
}
|
||||
if !spendJSONBytes(left, n) {
|
||||
return false
|
||||
}
|
||||
i++
|
||||
}
|
||||
return true
|
||||
}
|
||||
@@ -0,0 +1,199 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
// Package systemone shares decision request validation across internal and HTTP callers.
|
||||
package systemone
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"slices"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
)
|
||||
|
||||
const (
|
||||
MaxBodyBytes = 64 << 10
|
||||
MaxQuestions = 64
|
||||
)
|
||||
|
||||
type ErrorKind string
|
||||
|
||||
const (
|
||||
InvalidRequest ErrorKind = "invalid_request"
|
||||
UnsupportedBackend ErrorKind = "unsupported_backend"
|
||||
)
|
||||
|
||||
// ValidationError separates invalid input from unsupported backend capabilities.
|
||||
// HTTP callers may map these kinds to 400 and 501 respectively.
|
||||
type ValidationError struct {
|
||||
Kind ErrorKind
|
||||
Err error
|
||||
}
|
||||
|
||||
func (e *ValidationError) Error() string { return e.Err.Error() }
|
||||
func (e *ValidationError) Unwrap() error { return e.Err }
|
||||
|
||||
// ValidateRequestStructure checks request semantics without re-encoding its
|
||||
// fields. HTTP callers bound the original wire bytes before binding; escaping
|
||||
// during serialization must not impose a second, different HTTP size limit.
|
||||
func ValidateRequestStructure(req *schema.SystemOneRequest) error {
|
||||
if req == nil {
|
||||
return &ValidationError{InvalidRequest, fmt.Errorf("request is required")}
|
||||
}
|
||||
if err := validateRequest(req); err != nil {
|
||||
return &ValidationError{InvalidRequest, err}
|
||||
}
|
||||
images, err := CollectImages(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return ValidateImages(images)
|
||||
}
|
||||
|
||||
// ValidateRequest additionally bounds the serialized internal transport body.
|
||||
// Internal callers have no HTTP reader on which to enforce the wire limit.
|
||||
func ValidateRequest(req *schema.SystemOneRequest) error {
|
||||
if err := ValidateRequestStructure(req); err != nil {
|
||||
return err
|
||||
}
|
||||
body, err := json.Marshal(req)
|
||||
if err != nil {
|
||||
return &ValidationError{InvalidRequest, err}
|
||||
}
|
||||
limit, err := RequestBodyLimit(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if len(body) > limit {
|
||||
return &ValidationError{InvalidRequest, fmt.Errorf("request body exceeds %d KiB", limit>>10)}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ValidateDecisionModel is stricter than legacy HTTP admission: the router
|
||||
// never guesses a model's usecase and never falls back to NER or generation.
|
||||
func ValidateDecisionModel(cfg config.ModelConfig) error {
|
||||
if cfg.KnownUsecases == nil || *cfg.KnownUsecases&config.FLAG_DECISIONS == 0 {
|
||||
return &ValidationError{InvalidRequest, fmt.Errorf("model %q must explicitly declare known_usecases: [decisions]", cfg.Name)}
|
||||
}
|
||||
if !BackendSupportsScore(cfg.Backend) {
|
||||
return &ValidationError{UnsupportedBackend, fmt.Errorf("backend %q does not support decisions", cfg.Backend)}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func validateRequest(req *schema.SystemOneRequest) error {
|
||||
if len(req.State) == 0 || string(req.State) == "null" {
|
||||
return fmt.Errorf("state is required")
|
||||
}
|
||||
var state any
|
||||
if err := json.Unmarshal(req.State, &state); err != nil {
|
||||
return fmt.Errorf("state is not valid JSON: %w", err)
|
||||
}
|
||||
if state == nil {
|
||||
return fmt.Errorf("state is required")
|
||||
}
|
||||
if s, ok := state.(string); ok && strings.TrimSpace(s) == "" {
|
||||
return fmt.Errorf("state is required")
|
||||
}
|
||||
if len(req.Questions) == 0 {
|
||||
return fmt.Errorf("questions is required and must contain at least one question")
|
||||
}
|
||||
if len(req.Questions) > MaxQuestions {
|
||||
return fmt.Errorf("questions must contain at most %d questions", MaxQuestions)
|
||||
}
|
||||
qids := make([]string, 0, len(req.Questions))
|
||||
for id := range req.Questions {
|
||||
qids = append(qids, id)
|
||||
}
|
||||
sort.Strings(qids)
|
||||
for _, id := range qids {
|
||||
if strings.TrimSpace(id) == "" {
|
||||
return fmt.Errorf("question ids must not be blank")
|
||||
}
|
||||
q := req.Questions[id]
|
||||
switch q.Type {
|
||||
case "choice":
|
||||
var criteria map[string]json.RawMessage
|
||||
if err := json.Unmarshal(q.Criteria, &criteria); err != nil {
|
||||
return fmt.Errorf("question %q (choice) requires a criteria object", id)
|
||||
}
|
||||
if len(criteria) < 2 {
|
||||
return fmt.Errorf("question %q (choice) requires at least 2 options", id)
|
||||
}
|
||||
for k := range criteria {
|
||||
if strings.TrimSpace(k) == "" {
|
||||
return fmt.Errorf("question %q (choice) has a blank option key", id)
|
||||
}
|
||||
}
|
||||
case "score":
|
||||
var criteria []json.RawMessage
|
||||
if err := json.Unmarshal(q.Criteria, &criteria); err != nil {
|
||||
return fmt.Errorf("question %q (score) requires a criteria array", id)
|
||||
}
|
||||
if len(criteria) < 2 {
|
||||
return fmt.Errorf("question %q (score) requires at least 2 levels", id)
|
||||
}
|
||||
case "noul":
|
||||
if len(q.Criteria) == 0 || string(q.Criteria) == "null" {
|
||||
continue
|
||||
}
|
||||
var criteria map[string]json.RawMessage
|
||||
if err := json.Unmarshal(q.Criteria, &criteria); err != nil {
|
||||
return fmt.Errorf("question %q (noul) criteria must be an object with \"false\" and \"true\" descriptions", id)
|
||||
}
|
||||
for k := range criteria {
|
||||
if k != "false" && k != "true" {
|
||||
return fmt.Errorf("question %q (noul) criteria may only have \"false\" and \"true\" keys", id)
|
||||
}
|
||||
}
|
||||
default:
|
||||
return fmt.Errorf("question %q has unknown type: %s", id, q.Type)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func ModelAllowed(cfg config.ModelConfig) error {
|
||||
if cfg.KnownUsecases == nil {
|
||||
return nil
|
||||
}
|
||||
if *cfg.KnownUsecases&(config.FLAG_DECISIONS|config.FLAG_TOKEN_CLASSIFY) != 0 {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("model %q does not declare the decisions usecase (known_usecases: [decisions])", cfg.Name)
|
||||
}
|
||||
|
||||
func NERAllowed(cfg config.ModelConfig) error {
|
||||
if cfg.KnownUsecases == nil {
|
||||
return nil
|
||||
}
|
||||
declared := *cfg.KnownUsecases
|
||||
if declared&config.FLAG_DECISIONS != 0 && declared&config.FLAG_TOKEN_CLASSIFY == 0 {
|
||||
return fmt.Errorf("model %q is a decision model: /permute and /separate use the NER path, use POST /v1/systemone instead", cfg.Name)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func UsesDecisionPipeline(cfg config.ModelConfig) bool {
|
||||
if !BackendSupportsScore(cfg.Backend) {
|
||||
return false
|
||||
}
|
||||
if cfg.KnownUsecases == nil {
|
||||
return true
|
||||
}
|
||||
declared := *cfg.KnownUsecases
|
||||
if declared&config.FLAG_DECISIONS != 0 {
|
||||
return true
|
||||
}
|
||||
return declared&config.FLAG_TOKEN_CLASSIFY == 0
|
||||
}
|
||||
|
||||
func BackendSupportsScore(backendName string) bool {
|
||||
cap := config.GetBackendCapability(backendName)
|
||||
if cap == nil {
|
||||
return false
|
||||
}
|
||||
return slices.Contains(cap.GRPCMethods, config.MethodScore)
|
||||
}
|
||||
@@ -0,0 +1,78 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
package systemone
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
)
|
||||
|
||||
func validRequest() *schema.SystemOneRequest {
|
||||
return &schema.SystemOneRequest{State: json.RawMessage(`{"text":"hello"}`), Questions: map[string]schema.SystemOneQuestion{"q": {Type: "noul", Criteria: json.RawMessage(`{"false":"absent","true":"present"}`)}}}
|
||||
}
|
||||
func TestSystemOne(t *testing.T) {
|
||||
RegisterFailHandler(Fail)
|
||||
RunSpecs(t, "SystemOne validation")
|
||||
}
|
||||
|
||||
var _ = Describe("Shared validation", func() {
|
||||
It("preserves structured request fields", func() {
|
||||
req := validRequest()
|
||||
before, _ := json.Marshal(req)
|
||||
Expect(ValidateRequest(req)).To(Succeed())
|
||||
after, _ := json.Marshal(req)
|
||||
Expect(after).To(Equal(before))
|
||||
})
|
||||
It("rejects malformed and unbounded requests with typed errors", func() {
|
||||
cases := []func(*schema.SystemOneRequest){
|
||||
func(r *schema.SystemOneRequest) { r.State = nil },
|
||||
func(r *schema.SystemOneRequest) { r.State = json.RawMessage(` null `) },
|
||||
func(r *schema.SystemOneRequest) { r.State = json.RawMessage(`" "`) },
|
||||
func(r *schema.SystemOneRequest) { r.Questions = nil },
|
||||
func(r *schema.SystemOneRequest) { r.Questions["q"] = schema.SystemOneQuestion{Type: "other"} },
|
||||
func(r *schema.SystemOneRequest) {
|
||||
r.Questions["q"] = schema.SystemOneQuestion{Type: "noul", Criteria: json.RawMessage(`{"yes":"yes"}`)}
|
||||
},
|
||||
func(r *schema.SystemOneRequest) { r.State, _ = json.Marshal(strings.Repeat("a", MaxBodyBytes)) },
|
||||
func(r *schema.SystemOneRequest) {
|
||||
for i := 0; i < MaxQuestions; i++ {
|
||||
r.Questions[strings.Repeat("q", i+2)] = schema.SystemOneQuestion{Type: "noul"}
|
||||
}
|
||||
},
|
||||
}
|
||||
for _, mutate := range cases {
|
||||
req := validRequest()
|
||||
mutate(req)
|
||||
err := ValidateRequest(req)
|
||||
var typed *ValidationError
|
||||
Expect(errors.As(err, &typed)).To(BeTrue())
|
||||
Expect(typed.Kind).To(Equal(InvalidRequest))
|
||||
}
|
||||
Expect(ValidateRequest(nil)).NotTo(Succeed())
|
||||
})
|
||||
It("requires explicit native usecase but preserves legacy HTTP admission", func() {
|
||||
c := config.ModelConfig{}
|
||||
c.Name = "decision"
|
||||
c.Backend = "vllm-cpp"
|
||||
Expect(ModelAllowed(c)).To(Succeed())
|
||||
Expect(ValidateDecisionModel(c)).NotTo(Succeed())
|
||||
flags := config.FLAG_DECISIONS
|
||||
c.KnownUsecases = &flags
|
||||
Expect(ValidateDecisionModel(c)).To(Succeed())
|
||||
c.Backend = "does-not-exist"
|
||||
var typed *ValidationError
|
||||
Expect(errors.As(ValidateDecisionModel(c), &typed)).To(BeTrue())
|
||||
Expect(typed.Kind).To(Equal(UnsupportedBackend))
|
||||
flags = config.FLAG_TOKEN_CLASSIFY
|
||||
c.Backend = "vllm-cpp"
|
||||
Expect(ValidateDecisionModel(c)).NotTo(Succeed())
|
||||
Expect(ModelAllowed(c)).To(Succeed())
|
||||
Expect(UsesDecisionPipeline(c)).To(BeFalse())
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,18 @@
|
||||
---
|
||||
title: "Routing logs in embedded applications"
|
||||
---
|
||||
|
||||
Go applications that embed LocalAI can retain the bounded in-memory router
|
||||
decision log without enabling billing statistics:
|
||||
|
||||
```go
|
||||
config.WithDisableStats(true),
|
||||
config.WithRouterDecisionLog(true),
|
||||
```
|
||||
|
||||
The explicit routing-log option retains classifier, selected-model, cache and
|
||||
score diagnostics. It does not create a billing recorder or record token usage.
|
||||
Without that opt-in, disabling stats still disables both logs as before.
|
||||
|
||||
`config.DisableMetricsEndpoint` also prevents failover health-gauge registration;
|
||||
failover routing itself continues to run.
|
||||
@@ -136,7 +136,7 @@ The same model also answers `/v1/chat/completions` requests.
|
||||
|
||||
A request is refused with `400` (or `413` for the body size) when:
|
||||
|
||||
- the body is larger than 64 KiB,
|
||||
- a text-only body is larger than 64 KiB (image-bearing bodies have the bounded budget below),
|
||||
- `state` is missing or blank,
|
||||
- there are no questions, or more than 64,
|
||||
- a question id is blank,
|
||||
@@ -170,3 +170,208 @@ common case. These behaviors differ:
|
||||
When authentication is on, the three routes need the `decisions` feature. It is
|
||||
on by default for every user, like the other API features, and an administrator
|
||||
can turn it off per user.
|
||||
|
||||
## Native llama.cpp decisions
|
||||
|
||||
The stock `llama-cpp` backend supports text-only decision GGUFs carrying upstream
|
||||
SystemOne metadata. Declare `known_usecases: [decisions]`; ordinary `score` need
|
||||
not be enabled. Requests use the existing internal Score RPC, not a backend HTTP
|
||||
server. Choice, score, and noul questions may be combined in one request. Structured
|
||||
state and questions are forwarded without NER rendering. llama.cpp score questions
|
||||
accept 2–10 levels; this backend-specific limit does not constrain vllm-cpp.
|
||||
Older forks without native decision support return 501. Missing decision metadata
|
||||
also returns 501, while backend invalid requests return 400.
|
||||
|
||||
### Bounded image input
|
||||
|
||||
The public routes and internal decision validator share an image contract. Supply
|
||||
PNG or JPEG base64 data URLs in `images`, OpenAI `image_url` message content, or
|
||||
Anthropic `image` content with a `base64` source and `media_type`. Remote URLs and
|
||||
file paths are never fetched. MIME must match the decoded image header; malformed
|
||||
base64, unsupported formats and invalid headers return 400.
|
||||
|
||||
Limits per request are **8 images**, **12 MiB aggregate encoded data-URL bytes**,
|
||||
**8 MiB aggregate decoded bytes**, **4096 pixels per dimension**, and
|
||||
**16 million aggregate pixels**. Exceeding these limits returns 413. Headers are
|
||||
checked before full pixel decoding, which rejects truncated or corrupt images; native decoders must independently
|
||||
protect direct RPC inputs.
|
||||
|
||||
Image-bearing request bodies may use up to **16 MiB**. Text-only requests retain
|
||||
the **64 KiB raw-wire limit**, including whitespace; JSON escaping during internal
|
||||
serialization does not impose a second HTTP limit. Absent, `null`, or empty
|
||||
`images` do not enable the larger budget. Native decision responses retain a
|
||||
separate **64 KiB** limit, independent of the request budget. These limits do not
|
||||
raise any global HTTP limit.
|
||||
|
||||
A shared **8-request admission ceiling** covers public decision handlers and
|
||||
internal decision runners before body buffering, image decoding or serialization.
|
||||
Saturation fails promptly (HTTP 503); cancellation before admission does not take
|
||||
a slot. A slot remains held through inference/response handling, and an internal
|
||||
cancelled call retains its slot until its underlying worker actually ends. This
|
||||
bounds concurrent decision-owned allocation and retained request bodies, not total
|
||||
process memory: caller-owned inputs, upstream middleware buffers and model/backend
|
||||
memory are outside this budget. Full decoding is sequential per admitted request,
|
||||
with each pixel buffer limited by the checked dimensions and aggregate pixel
|
||||
budget (up to 16 million pixels; decoded byte storage depends on pixel format).
|
||||
Garbage collection timing is not an RSS guarantee. The separate 8-operation backend ceiling still
|
||||
bounds abandoned native operations. Validation helpers do not acquire nested slots.
|
||||
|
||||
For image-only input, provide explicit structured state such as `"state": {}`
|
||||
alongside `images`, or a message containing image content. Missing, null or blank
|
||||
string state remains invalid. Arbitrary domain JSON is preserved, not interpreted
|
||||
as image content outside message content parts. Images are never replaced by
|
||||
invented text.
|
||||
|
||||
Admission is not a promise of model image capability: the NER path and text-only
|
||||
native decision models reject images with 501 rather than silently dropping them.
|
||||
OpenJev requires its vision projector (see below). Failed requests are not billed.
|
||||
Router image probes preserve the canonical request; full public API image
|
||||
validation across installed models is separate from native RPC validation.
|
||||
|
||||
Native responses report backend input/output usage, including zero generated
|
||||
tokens. LocalAI records supplied usage once; explicit zero counts are distinct
|
||||
from missing usage. Missing counts are not estimated, and invalid negative counts
|
||||
are rejected rather than billed.
|
||||
|
||||
### Julia-1 CPU example
|
||||
|
||||
Install the separate stock llama.cpp entry (existing vllm-cpp entries are unchanged):
|
||||
|
||||
```sh
|
||||
local-ai models install julia-1-llama-cpp
|
||||
```
|
||||
|
||||
Julia-1 is a 144.3M-parameter multilingual text decision model. The gallery pins
|
||||
`ggml-org/Julia-1-GGUF` revision `16fee17949206fbf58da9347daea44d792a81211`,
|
||||
file `Julia-1-Q8_0.gguf` (168,166,496 bytes, about 160.4 MiB), SHA-256
|
||||
`1ea6a7e87156eeeda88cb7a36a61265b37ba7b993897b7289b99aea5b5e47069`.
|
||||
The source model and GGUF publisher declare Apache-2.0. Source provenance:
|
||||
`SupersonicLabs/Julia-1` revision `a85b127321d580d65176c89ced8273f305745d85`,
|
||||
based on `jhu-clsp/mmBERT-small`. This is a real model, not the upstream tiny test
|
||||
fixture; assess its accuracy for your own tasks.
|
||||
|
||||
```sh
|
||||
curl http://localhost:8080/v1/systemone \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{"model":"julia-1-llama-cpp","state":"I was charged twice and need a refund.","questions":{"route":{"type":"choice","instructions":"Which team should handle this?","criteria":{"billing":"payments and refunds","shipping":"delivery problems","technical":"software issues"}},"refund":{"type":"noul","instructions":"Does the customer request a refund?"}}}'
|
||||
```
|
||||
|
||||
The pinned artifact was installed through the gallery installer, checksum-verified,
|
||||
and tested on CPU with the native Score RPC using choice, score, and noul in one
|
||||
request. That smoke returned 97 input tokens and zero output tokens; token counts
|
||||
vary with the request. This does not establish broad model accuracy or image
|
||||
support.
|
||||
|
||||
### Native family defaults
|
||||
|
||||
Each family has a separately named default; these do not replace vllm-cpp entries.
|
||||
All entries except OpenJev are text-only and omit projectors. OpenJev includes
|
||||
its Q8 projector. Download size is not a RAM estimate.
|
||||
|
||||
| Gallery entry | Quantization | Artifact bytes | License | Validation status |
|
||||
|---|---|---:|---|---|
|
||||
| `julia-1-llama-cpp` | Q8_0 | 168,166,496 | Apache-2.0 | Gallery install and CPU request verified |
|
||||
| `laya-llama-cpp` | Q8_0 | 449,397,600 | Apache-2.0 | Gallery install and CPU choice/score/noul verified |
|
||||
| `kev-4b-llama-cpp` | Q4_K_M | 3,033,489,824 | Apache-2.0 | Gallery install and CPU choice/score/noul verified |
|
||||
| `lev-llama-cpp` | Q4_K_M | 3,011,777,440 | Apache-2.0 | Gallery install and CPU choice/score/noul verified |
|
||||
| `openjev-llama-cpp` | Q4_K_M + Q8 projector | 19,603,119,520 | **CC-BY-NC-4.0** | Gallery install and CPU choice/score/noul verified |
|
||||
| `nimble-9b-v3-llama-cpp` | Q4_K_M | 6,324,185,632 | **CC-BY-NC-4.0** | Gallery install and CPU choice/score/noul verified |
|
||||
|
||||
OpenJev and Nimble are noncommercial models. OpenJev's upstream multimodal
|
||||
capability does **not** imply LocalAI decision-image support. Nimble requires the
|
||||
native Nimble integration included in this source tree's llama.cpp pin
|
||||
`bed0a856606ee4a24a164066f73d2379447033f5`; older installed backends must be
|
||||
updated before serving it. This source prerequisite is integrated, but the
|
||||
OpenJev and Nimble installation/runtime checks remain pending as listed above.
|
||||
The published entries pin revisions and SHA-256 checksums, but metadata verification
|
||||
alone is not a runtime test. No model-quality guarantee follows from these smoke
|
||||
tests. Laya, Kev-4B, and lev were also retested against the newer native backend
|
||||
with 1- and 11-level score requests correctly rejected.
|
||||
|
||||
OpenJev and Nimble validation used the native backend at llama.cpp revision
|
||||
`bed0a856606ee4a24a164066f73d2379447033f5`, CPU-only with two threads, a 2048-token
|
||||
context, and batch size 512. Each artifact was installed through the gallery,
|
||||
SHA-256 verified, and checked for its decision metadata and SystemOne template.
|
||||
Each request included choice, score, and noul questions together, including the
|
||||
full question set required by Nimble. Response-shape, probability-normalization,
|
||||
and noul-bound assertions passed; 1- and 11-level score requests were rejected.
|
||||
The test requests reported 224 input tokens for OpenJev and 933 for Nimble, with
|
||||
explicit zero output tokens for both. No projector was installed or tested.
|
||||
|
||||
These are bounded text contract smoke tests, not accuracy benchmarks or
|
||||
performance guarantees. Floating-point probabilities can vary with hardware and
|
||||
build settings; tests do not require exact answer probabilities or token counts.
|
||||
Neither image support nor interruption during active evaluation is established
|
||||
by these tests. CC-BY-NC-4.0's noncommercial restriction still applies.
|
||||
|
||||
### Multimodal router probes
|
||||
|
||||
The `decisions` router classifier preserves ordered OpenAI message content and
|
||||
Anthropic base64 image sources as structured state, including image-only turns.
|
||||
It uses the internal decision runner, not a loopback HTTP request. The shared
|
||||
image limits above are validated before model loading; no URL is fetched by the
|
||||
classifier. Original message content is not rewritten when selecting a candidate
|
||||
or the configured fallback.
|
||||
|
||||
Score, rerank and KNN classifiers are text-only: image input produces an explicit
|
||||
classifier error and follows the existing configured fallback policy, rather
|
||||
than classifying an image-stripped prompt. Without a fallback, routing fails.
|
||||
Parent cancellation remains terminal and does not select a fallback. Image probes
|
||||
bypass text embedding caches and are not trimmed to text-only turns. Native
|
||||
context overflow is reported by the backend rather than silently dropping images.
|
||||
These transport guarantees do not establish installed projector capability or
|
||||
real-model image accuracy; those require separate native and end-to-end validation.
|
||||
|
||||
OpenAI chat routing classifies the original structured message before preparing
|
||||
media for the selected model. Remote image URLs are not downloaded as decision
|
||||
inputs. After selection (including a configured fallback), the served model's
|
||||
normal media preparation runs without replacing the original content blocks.
|
||||
Invalid classifier configuration fails closed even with a configured fallback.
|
||||
Runtime classification and input errors follow the configured fallback policy.
|
||||
Cancellation never selects a fallback.
|
||||
|
||||
Router probe extraction checks the shared 16 MiB state budget before copying
|
||||
text or serializing messages, including JSON escaping expansion. This applies
|
||||
to typed and untyped internal requests as well as parsed API requests; it does
|
||||
not add a limit to non-router inference. Direct internal probes containing
|
||||
custom structs, JSON/text marshalers, or excessively nested values fail
|
||||
extraction rather than executing their serialization. Supported internal values
|
||||
are the chat schema message/content/tool types and plain JSON values (including
|
||||
`json.RawMessage`, conservatively budgeted for escaping). Direct prompt-only
|
||||
probes also check serialized escaping before allocation. The separate 64 KiB text-only Decisions
|
||||
request limit is unchanged. Anthropic conversion preserves typed content blocks
|
||||
through both native selection and fallback, including ordered text and images.
|
||||
|
||||
### OpenJev image decisions
|
||||
|
||||
The `openjev-llama-cpp` gallery entry installs OpenJev Q4_K_M and its pinned
|
||||
Q8 vision projector (`mmproj-OpenJev-Q8_0.gguf`). Both artifacts come from
|
||||
`ggml-org/OpenJev-GGUF` revision `10840f375658dea7afc5ff4711127bca8218b560`.
|
||||
The weights occupy 18,973,872,288 bytes and the projector 629,247,232 bytes:
|
||||
19,603,119,520 bytes total (about 19.61 GB decimal), excluding runtime memory,
|
||||
KV cache and backend files. The model is **CC-BY-NC-4.0, noncommercial only**;
|
||||
LocalAI's software license does not override the model license.
|
||||
|
||||
For a manually installed model, set `mmproj: mmproj-OpenJev-Q8_0.gguf` alongside
|
||||
`parameters.model`, not inside `parameters`. LocalAI resolves that filename
|
||||
relative to the model directory and forwards it to llama.cpp's projector loader.
|
||||
The gallery uses an 8192-token context to leave room for image tokens and
|
||||
question framing. This is not a guarantee that all eight maximum-sized images
|
||||
fit; context overflow remains an error. Size the context for the actual workload.
|
||||
|
||||
Native image decisions require both a decision format that accepts images and a
|
||||
loaded projector that supports **vision input**. A missing projector, an
|
||||
audio-only projector, or a text-only decision model does not silently fall back
|
||||
to a text decision. Other decision gallery entries remain text-only.
|
||||
|
||||
Direct native RPC callers receive the same image count, encoded/decoded byte,
|
||||
dimension and aggregate pixel bounds as public callers. Validation precedes
|
||||
llama.cpp's permissive media parsing and full pixel decode; PNG decompression
|
||||
is independently bounded before stb decodes pixels. This validation applies
|
||||
only to native decision tasks, not ordinary chat or legacy scoring.
|
||||
|
||||
Native decision image validation rejects PNG streams with invalid checksums and
|
||||
incomplete JPEG scans, including truncated scans with an appended end marker.
|
||||
Source builds with native decision support require zlib and libjpeg development
|
||||
packages (`zlib1g-dev libjpeg-dev` on Ubuntu; `zlib jpeg-turbo` on Homebrew).
|
||||
Packaged backends include the required runtime libraries.
|
||||
@@ -353,21 +353,85 @@ silent-bypass.
|
||||
|
||||
### Available classifiers
|
||||
|
||||
LocalAI ships three classifier implementations. Pick one with `classifier:`
|
||||
LocalAI ships four classifier implementations. Pick one with `classifier:`
|
||||
in the router YAML:
|
||||
|
||||
| Classifier | When to use | Underlying primitive |
|
||||
|---|---|---|
|
||||
| `score` (default) | Small classifier-tuned LM (Arch-Router-style). Best when label vocabulary is well-covered by next-token continuation. | `Score` gRPC primitive (llama-cpp, vLLM). |
|
||||
| `colbert` | When label descriptions are abstract or short and a next-token classifier produces flat distributions. Robust on long-form policy descriptions. | rerankers backend in ColBERT mode (e.g. `bge-m3-colbert` from the gallery). |
|
||||
| `decisions` | Independent, overlapping policy decisions from a native decision model. | Internal SystemOne pipeline via Score; numeric noul P(true). |
|
||||
| `knn` | When you have (or can generate) labelled example prompts — including outcome-labelled production traffic. Deterministic, auditable, cheapest per request, and the only classifier with an explicit out-of-distribution fallback. | embeddings backend + local-store KNN over a persisted, curated corpus. |
|
||||
|
||||
All three share `policies`, `candidates`, `fallback`, and
|
||||
`classifier_cache_size`. `score` and `colbert` take a
|
||||
All four share `policies`, `candidates`, and `fallback`. The existing
|
||||
`score`, `colbert`, and `knn` classifiers also use `classifier_cache_size`. `score` and `colbert` take a
|
||||
`classifier_model` (+ `activation_threshold`, optional
|
||||
`embedding_cache`); `knn` instead takes a `knn:` block and a corpus
|
||||
seeded through the API.
|
||||
|
||||
### Native decision models (`decisions`)
|
||||
|
||||
Use `classifier: decisions` with an installed native decision model explicitly
|
||||
configured with `known_usecases: [decisions]` and a backend supporting the Score
|
||||
RPC. A chat model with ordinary continuation scoring is not a substitute.
|
||||
The router calls the internal `ModelSystemOne` adapter, never a loopback HTTP
|
||||
endpoint, NER extractor, or text-generation fallback.
|
||||
|
||||
```yaml
|
||||
name: policy-router
|
||||
router:
|
||||
classifier: decisions
|
||||
classifier_model: my-decision-model
|
||||
activation_threshold: 0.5
|
||||
policies:
|
||||
- label: code
|
||||
description: Writing or debugging code
|
||||
- label: private
|
||||
description: Handling private information
|
||||
candidates:
|
||||
- model: coding-model
|
||||
labels: [code]
|
||||
- model: private-coding-model
|
||||
labels: [code, private]
|
||||
fallback: general-model
|
||||
```
|
||||
|
||||
Each policy becomes a separate `noul` question with explicit `false` and `true`
|
||||
criteria. Stable question IDs map answers back to policy declaration order.
|
||||
The native numeric `noul` value is **P(true)**; a probabilities map is not
|
||||
required. `label_scores` are independent values in [0,1], not a distribution
|
||||
normalized across policies. Every label with probability **>= threshold** is
|
||||
active, and `score` is the maximum probability. The default threshold is 0.5
|
||||
(omitted or zero); positive configured thresholds must be finite and <=1.
|
||||
There is no exclusive-choice mode or top-one rescue. All-below-threshold means
|
||||
abstention. The first candidate covering all active labels wins; abstention,
|
||||
malformed answers, native context rejection, and backend errors use the normal
|
||||
configured fallback. Parent request cancellation returns without selecting a
|
||||
fallback.
|
||||
|
||||
Limits: 1–64 unique nonblank policies, nonblank descriptions, and 64 KiB each
|
||||
for serialized internal request and raw response. Missing/null/non-numeric,
|
||||
wrong-type, unknown-ID and out-of-range answers fail classification; numeric
|
||||
zero is valid. The native engine owns question framing and context enforcement.
|
||||
The router does not estimate native context from raw JSON token counts or trim
|
||||
text on that assumption. A request within the byte limit can still exceed the
|
||||
model's context and fall back.
|
||||
|
||||
The native adapter has a separate **process-wide eight-operation ceiling**,
|
||||
covering healthy calls as well as abandoned work, independently of the general
|
||||
backend admission limit (default 1024). Saturation fails classification and uses
|
||||
the configured fallback; it does not queue. Model loading has no cancellable
|
||||
API: cancellation releases the waiting caller, **not underlying resources**.
|
||||
Eight stuck loads therefore prevent further native decisions until underlying
|
||||
work finishes. This bound is a resource-safety measure, not a throughput claim.
|
||||
|
||||
The classifier registry invalidates on router configuration, resolved native
|
||||
model configuration/usecase, or persisted model revision changes. Decisions do
|
||||
not memoize prompt results or support `embedding_cache`/composite `knn` blocks;
|
||||
these combinations are rejected. `classifier_cache_size` has no effect on this
|
||||
classifier. The same central factory serves chat, Anthropic, realtime and the
|
||||
existing `/router/decide` oracle. Decision traces contain no prompt text.
|
||||
|
||||
### The Score classifier
|
||||
|
||||
The `score` classifier works like this:
|
||||
@@ -815,3 +879,18 @@ with `POST /models/reload` to pick up YAML edits without restarting.
|
||||
required for mutating endpoints and the `/app/middleware` page; in
|
||||
no-auth single-user mode the synthetic local user has admin role
|
||||
automatically.
|
||||
|
||||
### Creating a Decisions router in the UI
|
||||
|
||||
In **Middleware → Routing → Create routing model**, select **Decisions (native
|
||||
probabilities)** under Classifier. The Classifier Model picker lists installed,
|
||||
enabled native decision models explicitly declaring `known_usecases: [decisions]`
|
||||
on a backend supporting Score, including `llama-cpp` and `vllm-cpp`. NER-only
|
||||
models and routing dispatchers are not eligible. Selecting a model saves its exact
|
||||
configured name; it does not install weights.
|
||||
|
||||
Decisions needs no ChatML template and returns independent label probabilities,
|
||||
not an exclusive choice. Start with an activation threshold of **0.5** (zero uses
|
||||
the Decisions default of 0.5). Changing classifiers clears the dependent model
|
||||
selection but preserves your threshold, including the template's initial 0.40;
|
||||
set it deliberately before saving. Existing saved selections reopen unchanged.
|
||||
Loaded 100 of 108 files, more files were not shown because too many files have changed in this diff.
Show more
Reference in new issue
Block a user