Compare commits

..

2 Commits

Author SHA1 Message Date
localai-org-maint-bot
a4c8f4b83a Merge branch 'master' into worktree-investigate-diffusers-load-cancel-10636 2026-07-29 15:04:58 +02:00
Ettore Di Giacinto
512ec0e703 fix(backend): don't let a client disconnect cancel the model load
Image generation (and the tts/transcript/embeddings/vad/rerank/llm helpers)
pass the request context to loader.Load so distributed routing decisions
reach the request's X-LocalAI-Node holder. That context also governs
cancellation of the load, so when a client disconnects mid-load the
LoadModel RPC is aborted, stopLoadProcess tears down the backend process,
and every retry restarts from scratch. Heavy diffusers/LLM models on a slow
host (e.g. a shared-memory iGPU) take long enough to load that the request
routinely ends first, so the model never finishes loading and the UI shows
"NetworkError when attempting to fetch resource".

Wrap the load context with context.WithoutCancel: the routing holder value
still propagates, but the request's cancellation no longer aborts the load,
so it runs to completion and caches for the next request. Inference keeps the
cancellable request context, so a disconnect still stops generation.

Adds a regression spec asserting a canceled request context does not cancel
the model load while the routing holder still reaches the router.

Fixes #10636

Signed-off-by: Ettore Di Giacinto <mudler@localai.io>
Assisted-by: Claude:claude-opus-4-8 [Claude Code]
2026-07-29 09:06:51 +00:00
16 changed files with 84 additions and 1248 deletions

View File

@@ -1,95 +0,0 @@
# SPDX-License-Identifier: MIT
cmake_minimum_required(VERSION 3.20)
project(audio-cpp-grpc-server LANGUAGES CXX)
set(CMAKE_CXX_STANDARD 17)
set(CMAKE_CXX_STANDARD_REQUIRED ON)
set(CMAKE_CXX_EXTENSIONS OFF)
set(AUDIO_CPP_DIR "${CMAKE_CURRENT_SOURCE_DIR}/audio.cpp"
CACHE PATH "Path to the audio.cpp source tree")
set(LOCALAI_BACKEND_PROTO "${CMAKE_CURRENT_SOURCE_DIR}/../../backend.proto"
CACHE FILEPATH "Path to the LocalAI backend protocol")
option(ENGINE_ENABLE_CUDA "Build audio.cpp with CUDA support" OFF)
option(ENGINE_ENABLE_VULKAN "Build audio.cpp with Vulkan support" OFF)
option(ENGINE_ENABLE_METAL "Build audio.cpp with Metal support" OFF)
option(AUDIO_CPP_BUILD_TESTS "Build LocalAI audio.cpp unit tests" OFF)
option(AUDIO_CPP_BUILD_GRPC "Build the LocalAI gRPC server" ON)
find_package(Threads REQUIRED)
if(NOT EXISTS "${AUDIO_CPP_DIR}/CMakeLists.txt")
message(FATAL_ERROR
"AUDIO_CPP_DIR does not point to an audio.cpp source tree: ${AUDIO_CPP_DIR}")
endif()
add_subdirectory("${AUDIO_CPP_DIR}" "${CMAKE_CURRENT_BINARY_DIR}/audio.cpp")
add_library(localai_audio_cpp_runtime STATIC
audio_cpp_runtime.cpp
model_config.cpp)
target_include_directories(localai_audio_cpp_runtime
PUBLIC "${CMAKE_CURRENT_SOURCE_DIR}")
target_link_libraries(localai_audio_cpp_runtime PUBLIC engine_runtime)
if(AUDIO_CPP_BUILD_GRPC)
find_package(Protobuf CONFIG REQUIRED)
find_package(gRPC CONFIG REQUIRED)
find_program(PROTOC_EXECUTABLE NAMES protoc REQUIRED)
find_program(GRPC_CPP_PLUGIN_EXECUTABLE NAMES grpc_cpp_plugin REQUIRED)
get_filename_component(LOCALAI_BACKEND_PROTO_DIR
"${LOCALAI_BACKEND_PROTO}" DIRECTORY)
set(LOCALAI_PROTO_SOURCES
"${CMAKE_CURRENT_BINARY_DIR}/backend.pb.cc"
"${CMAKE_CURRENT_BINARY_DIR}/backend.grpc.pb.cc")
set(LOCALAI_PROTO_HEADERS
"${CMAKE_CURRENT_BINARY_DIR}/backend.pb.h"
"${CMAKE_CURRENT_BINARY_DIR}/backend.grpc.pb.h")
add_custom_command(
OUTPUT ${LOCALAI_PROTO_SOURCES} ${LOCALAI_PROTO_HEADERS}
COMMAND "${PROTOC_EXECUTABLE}"
ARGS
--cpp_out "${CMAKE_CURRENT_BINARY_DIR}"
--grpc_out "${CMAKE_CURRENT_BINARY_DIR}"
-I "${LOCALAI_BACKEND_PROTO_DIR}"
--plugin=protoc-gen-grpc="${GRPC_CPP_PLUGIN_EXECUTABLE}"
"${LOCALAI_BACKEND_PROTO}"
DEPENDS "${LOCALAI_BACKEND_PROTO}"
VERBATIM)
add_library(localai_backend_proto STATIC
${LOCALAI_PROTO_SOURCES}
${LOCALAI_PROTO_HEADERS})
target_include_directories(localai_backend_proto
PUBLIC "${CMAKE_CURRENT_BINARY_DIR}")
target_link_libraries(localai_backend_proto
PUBLIC protobuf::libprotobuf gRPC::grpc++)
# Task 2 replaces this generated entry point with the LocalAI service.
set(AUDIO_CPP_SERVER_PLACEHOLDER
"${CMAKE_CURRENT_BINARY_DIR}/audio-cpp-grpc-server-placeholder.cpp")
file(GENERATE OUTPUT "${AUDIO_CPP_SERVER_PLACEHOLDER}"
CONTENT "int main() { return 0; }\n")
add_executable(audio-cpp-grpc-server "${AUDIO_CPP_SERVER_PLACEHOLDER}")
target_link_libraries(audio-cpp-grpc-server PRIVATE
localai_audio_cpp_runtime
engine_runtime
localai_backend_proto
gRPC::grpc++
gRPC::grpc++_reflection)
endif()
if(AUDIO_CPP_BUILD_TESTS)
enable_testing()
add_executable(audio-cpp-runtime-test tests/runtime_tests.cpp)
target_link_libraries(audio-cpp-runtime-test PRIVATE
localai_audio_cpp_runtime
Threads::Threads)
add_test(
NAME audio-cpp-runtime-test
COMMAND audio-cpp-runtime-test "${AUDIO_CPP_DIR}")
endif()

View File

@@ -1,67 +0,0 @@
# SPDX-License-Identifier: MIT
AUDIO_CPP_VERSION?=f8fb0c19739193adfad0d9e58da99f25eda65256
AUDIO_CPP_REPO?=https://github.com/0xShug0/audio.cpp
AUDIO_CPP_SRC?=
CURRENT_MAKEFILE_DIR := $(dir $(abspath $(lastword $(MAKEFILE_LIST))))
BUILD_DIR := build
BUILD_TYPE ?=
JOBS ?= $(shell nproc 2>/dev/null || sysctl -n hw.ncpu 2>/dev/null || echo 4)
UNAME_S := $(shell uname -s)
CMAKE_ARGS ?= -DCMAKE_BUILD_TYPE=Release
CMAKE_ARGS += -DENGINE_ENABLE_CUDA=OFF
CMAKE_ARGS += -DENGINE_ENABLE_VULKAN=OFF
CMAKE_ARGS += -DENGINE_ENABLE_METAL=OFF
ifeq ($(BUILD_TYPE),cublas)
CMAKE_ARGS += -DENGINE_ENABLE_CUDA=ON
else ifeq ($(BUILD_TYPE),vulkan)
CMAKE_ARGS += -DENGINE_ENABLE_VULKAN=ON
else ifeq ($(UNAME_S),Darwin)
CMAKE_ARGS += -DENGINE_ENABLE_METAL=ON
endif
.PHONY: all grpc-server test test-unit clean purge
all: grpc-server
audio.cpp:
ifneq ($(AUDIO_CPP_SRC),)
ln -sfn $(abspath $(AUDIO_CPP_SRC)) audio.cpp
else
mkdir -p audio.cpp
cd audio.cpp && \
git init -q && \
git remote add origin $(AUDIO_CPP_REPO) && \
git fetch --depth 1 origin $(AUDIO_CPP_VERSION) && \
git checkout --detach FETCH_HEAD && \
git submodule update --init --recursive --depth 1
endif
grpc-server: audio.cpp
mkdir -p $(BUILD_DIR)
cd $(BUILD_DIR) && cmake $(CMAKE_ARGS) $(CURRENT_MAKEFILE_DIR)
cmake --build $(BUILD_DIR) --config Release \
--target audio-cpp-grpc-server -j $(JOBS)
cp $(BUILD_DIR)/audio-cpp-grpc-server grpc-server
test:
bash tests/build_contract_test.sh
test-unit: audio.cpp
mkdir -p $(BUILD_DIR)-unit
cd $(BUILD_DIR)-unit && cmake $(CMAKE_ARGS) \
-DAUDIO_CPP_BUILD_TESTS=ON -DAUDIO_CPP_BUILD_GRPC=OFF \
$(CURRENT_MAKEFILE_DIR)
cmake --build $(BUILD_DIR)-unit --config Release \
--target audio-cpp-runtime-test -j $(JOBS)
ctest --test-dir $(BUILD_DIR)-unit --output-on-failure
clean:
rm -rf $(BUILD_DIR) $(BUILD_DIR)-unit grpc-server
purge: clean
rm -rf audio.cpp

View File

@@ -1,178 +0,0 @@
// SPDX-License-Identifier: MIT
#include "audio_cpp_runtime.h"
#include <algorithm>
#include <stdexcept>
#include <string>
#include <utility>
namespace audio_cpp {
namespace {
using engine::runtime::CapabilitySet;
using engine::runtime::RunMode;
using engine::runtime::TaskSpec;
const engine::runtime::TaskCapability * find_task_capability(
const CapabilitySet & capabilities,
const TaskSpec & task) {
const auto it = std::find_if(
capabilities.supported_tasks.begin(),
capabilities.supported_tasks.end(),
[&](const engine::runtime::TaskCapability & capability) {
return capability.task == task.task;
});
return it == capabilities.supported_tasks.end() ? nullptr : &*it;
}
void validate_capability(
const engine::runtime::ILoadedVoiceModel & model,
const AudioCppModelConfig & config) {
if (model.metadata().family != config.family) {
throw std::runtime_error(
"loaded audio.cpp model family '" + model.metadata().family +
"' does not match requested family '" + config.family + "'");
}
const auto * capability = find_task_capability(
model.capabilities(),
config.task);
if (capability == nullptr) {
throw std::runtime_error(
"loaded audio.cpp model does not support requested task '" +
std::string(engine::runtime::to_string(config.task.task)) + "'");
}
if (std::find(
capability->modes.begin(),
capability->modes.end(),
config.task.mode) == capability->modes.end()) {
throw std::runtime_error(
"loaded audio.cpp model does not support requested mode '" +
std::string(engine::runtime::to_string(config.task.mode)) +
"' for task '" +
std::string(engine::runtime::to_string(config.task.task)) + "'");
}
}
void validate_session(
const engine::runtime::IVoiceTaskSession & session,
const AudioCppModelConfig & config) {
if (session.family() != config.family) {
throw std::runtime_error("audio.cpp session returned the wrong family");
}
if (session.task_kind() != config.task.task) {
throw std::runtime_error("audio.cpp session returned the wrong task");
}
if (session.run_mode() != config.task.mode) {
throw std::runtime_error("audio.cpp session returned the wrong mode");
}
if (config.task.mode == RunMode::Offline &&
dynamic_cast<const engine::runtime::IOfflineVoiceTaskSession *>(&session) == nullptr) {
throw std::runtime_error("audio.cpp session does not implement offline execution");
}
if (config.task.mode == RunMode::Streaming &&
dynamic_cast<const engine::runtime::IStreamingVoiceTaskSession *>(&session) == nullptr) {
throw std::runtime_error("audio.cpp session does not implement streaming execution");
}
}
} // namespace
AudioCppRuntime::AudioCppRuntime()
: registry_(engine::runtime::make_default_registry()) {}
AudioCppRuntime::AudioCppRuntime(engine::runtime::ModelRegistry registry)
: registry_(std::move(registry)) {}
AudioCppRuntime::~AudioCppRuntime() {
free();
}
void AudioCppRuntime::load(const AudioCppModelConfig & config) {
std::lock_guard<std::mutex> lock(mutex_);
auto candidate_model = registry_.load(config.load);
if (candidate_model == nullptr) {
throw std::runtime_error("audio.cpp registry returned a null model");
}
validate_capability(*candidate_model, config);
auto candidate_session = candidate_model->create_task_session(
config.task,
config.session);
if (candidate_session == nullptr) {
throw std::runtime_error("audio.cpp model returned a null session");
}
validate_session(*candidate_session, config);
free_locked();
model_ = std::move(candidate_model);
session_ = std::move(candidate_session);
}
void AudioCppRuntime::free() {
std::lock_guard<std::mutex> lock(mutex_);
free_locked();
}
engine::runtime::TaskResult AudioCppRuntime::run(
const engine::runtime::TaskRequest & request) {
std::lock_guard<std::mutex> lock(mutex_);
auto & session = require_session_locked();
session.prepare(engine::runtime::build_preparation_request(request));
return require_offline_locked().run(request);
}
void AudioCppRuntime::start_stream(
const engine::runtime::TaskRequest & request) {
std::lock_guard<std::mutex> lock(mutex_);
auto & session = require_session_locked();
session.prepare(engine::runtime::build_preparation_request(request));
require_streaming_locked().start_stream(request);
}
engine::runtime::StreamEvent AudioCppRuntime::process_audio_chunk(
const engine::runtime::AudioChunk & chunk) {
std::lock_guard<std::mutex> lock(mutex_);
return require_streaming_locked().process_audio_chunk(chunk);
}
engine::runtime::TaskResult AudioCppRuntime::finish_stream() {
std::lock_guard<std::mutex> lock(mutex_);
return require_streaming_locked().finish_stream();
}
engine::runtime::IVoiceTaskSession & AudioCppRuntime::require_session_locked() {
if (session_ == nullptr) {
throw std::runtime_error("audio.cpp runtime has no loaded session");
}
return *session_;
}
engine::runtime::IOfflineVoiceTaskSession &
AudioCppRuntime::require_offline_locked() {
auto * offline = dynamic_cast<engine::runtime::IOfflineVoiceTaskSession *>(
&require_session_locked());
if (offline == nullptr) {
throw std::runtime_error("loaded audio.cpp session is not offline");
}
return *offline;
}
engine::runtime::IStreamingVoiceTaskSession &
AudioCppRuntime::require_streaming_locked() {
auto * streaming = dynamic_cast<engine::runtime::IStreamingVoiceTaskSession *>(
&require_session_locked());
if (streaming == nullptr) {
throw std::runtime_error("loaded audio.cpp session is not streaming");
}
return *streaming;
}
void AudioCppRuntime::free_locked() {
session_.reset();
model_.reset();
}
} // namespace audio_cpp

View File

@@ -1,46 +0,0 @@
// SPDX-License-Identifier: MIT
#pragma once
#include "model_config.h"
#include "engine/framework/runtime/registry.h"
#include "engine/framework/runtime/session.h"
#include <memory>
#include <mutex>
namespace audio_cpp {
class AudioCppRuntime {
public:
AudioCppRuntime();
explicit AudioCppRuntime(engine::runtime::ModelRegistry registry);
~AudioCppRuntime();
AudioCppRuntime(const AudioCppRuntime &) = delete;
AudioCppRuntime & operator=(const AudioCppRuntime &) = delete;
void load(const AudioCppModelConfig & config);
void free();
engine::runtime::TaskResult run(
const engine::runtime::TaskRequest & request);
void start_stream(const engine::runtime::TaskRequest & request);
engine::runtime::StreamEvent process_audio_chunk(
const engine::runtime::AudioChunk & chunk);
engine::runtime::TaskResult finish_stream();
private:
engine::runtime::IVoiceTaskSession & require_session_locked();
engine::runtime::IOfflineVoiceTaskSession & require_offline_locked();
engine::runtime::IStreamingVoiceTaskSession & require_streaming_locked();
void free_locked();
std::mutex mutex_;
engine::runtime::ModelRegistry registry_;
std::unique_ptr<engine::runtime::ILoadedVoiceModel> model_;
std::unique_ptr<engine::runtime::IVoiceTaskSession> session_;
};
} // namespace audio_cpp

View File

@@ -1,156 +0,0 @@
// SPDX-License-Identifier: MIT
#include "model_config.h"
#include <limits>
#include <stdexcept>
#include <string>
namespace audio_cpp {
namespace {
using engine::core::BackendType;
using engine::runtime::RunMode;
using engine::runtime::VoiceTaskKind;
std::string require_option(
const std::unordered_map<std::string, std::string> & options,
const std::string & name) {
const auto it = options.find(name);
if (it == options.end() || it->second.empty()) {
throw std::invalid_argument("audio.cpp model config requires " + name);
}
return it->second;
}
VoiceTaskKind parse_task(const std::string & value) {
static const std::unordered_map<std::string, VoiceTaskKind> tasks = {
{"vad", VoiceTaskKind::Vad},
{"asr", VoiceTaskKind::Asr},
{"diarization", VoiceTaskKind::Diarization},
{"source-separation", VoiceTaskKind::SourceSeparation},
{"audio-generation", VoiceTaskKind::AudioGeneration},
{"tts", VoiceTaskKind::Tts},
{"voice-cloning", VoiceTaskKind::VoiceCloning},
{"voice-conversion", VoiceTaskKind::VoiceConversion},
{"speech-to-speech", VoiceTaskKind::SpeechToSpeech},
{"alignment", VoiceTaskKind::Alignment},
{"voice-design", VoiceTaskKind::VoiceDesign},
{"speaker-recognition", VoiceTaskKind::SpeakerRecognition},
{"svc", VoiceTaskKind::Svc},
};
const auto it = tasks.find(value);
if (it == tasks.end()) {
throw std::invalid_argument("unsupported audio.cpp task: " + value);
}
return it->second;
}
RunMode parse_mode(const std::string & value) {
if (value == "offline") {
return RunMode::Offline;
}
if (value == "streaming") {
return RunMode::Streaming;
}
throw std::invalid_argument("unsupported audio.cpp mode: " + value);
}
BackendType parse_backend(const std::string & value) {
if (value == "cpu") {
return BackendType::Cpu;
}
if (value == "cuda") {
return BackendType::Cuda;
}
if (value == "vulkan") {
return BackendType::Vulkan;
}
if (value == "metal") {
return BackendType::Metal;
}
if (value == "best") {
return BackendType::BestAvailable;
}
throw std::invalid_argument("unsupported audio.cpp backend: " + value);
}
int parse_integer(
const std::string & name,
const std::string & value,
int minimum) {
size_t parsed = 0;
long result = 0;
try {
result = std::stol(value, &parsed);
} catch (const std::exception &) {
throw std::invalid_argument("invalid audio.cpp " + name + ": " + value);
}
if (parsed != value.size() ||
result < minimum ||
result > std::numeric_limits<int>::max()) {
throw std::invalid_argument("invalid audio.cpp " + name + ": " + value);
}
return static_cast<int>(result);
}
void copy_namespaced_option(
const std::string & key,
const std::string & prefix,
const std::string & value,
std::unordered_map<std::string, std::string> & destination) {
const std::string name = key.substr(prefix.size());
if (name.empty()) {
throw std::invalid_argument("audio.cpp option namespace requires a name: " + key);
}
destination[name] = value;
}
} // namespace
AudioCppModelConfig parse_model_config(
const std::filesystem::path & model_path,
const std::unordered_map<std::string, std::string> & options) {
AudioCppModelConfig config;
config.model_path = model_path;
config.family = require_option(options, "family");
config.task.task = parse_task(require_option(options, "task"));
config.task.mode = RunMode::Offline;
config.load.model_path = model_path;
config.load.family_hint = config.family;
config.session.backend.type = BackendType::Cpu;
if (const auto it = options.find("mode"); it != options.end()) {
config.task.mode = parse_mode(it->second);
}
if (const auto it = options.find("backend"); it != options.end()) {
config.session.backend.type = parse_backend(it->second);
}
if (const auto it = options.find("device"); it != options.end()) {
config.session.backend.device = parse_integer("device", it->second, 0);
}
if (const auto it = options.find("threads"); it != options.end()) {
config.session.backend.threads = parse_integer("threads", it->second, 1);
}
if (const auto it = options.find("model_spec"); it != options.end()) {
config.load.model_spec_override = std::filesystem::path(it->second);
}
if (const auto it = options.find("config_id"); it != options.end()) {
config.load.config_id = it->second;
}
if (const auto it = options.find("weight_id"); it != options.end()) {
config.load.weight_id = it->second;
}
for (const auto & [key, value] : options) {
if (key.rfind("load.", 0) == 0) {
copy_namespaced_option(key, "load.", value, config.load.options);
} else if (key.rfind("session.", 0) == 0) {
copy_namespaced_option(key, "session.", value, config.session.options);
}
}
return config;
}
} // namespace audio_cpp

View File

@@ -1,26 +0,0 @@
// SPDX-License-Identifier: MIT
#pragma once
#include "engine/framework/runtime/model.h"
#include "engine/framework/runtime/session.h"
#include <filesystem>
#include <string>
#include <unordered_map>
namespace audio_cpp {
struct AudioCppModelConfig {
std::filesystem::path model_path;
std::string family;
engine::runtime::TaskSpec task;
engine::runtime::ModelLoadRequest load;
engine::runtime::SessionOptions session;
};
AudioCppModelConfig parse_model_config(
const std::filesystem::path & model_path,
const std::unordered_map<std::string, std::string> & options);
} // namespace audio_cpp

View File

@@ -1,164 +0,0 @@
#!/usr/bin/env bash
# SPDX-License-Identifier: MIT
set -euo pipefail
backend_dir="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)"
tmp_dir="$(mktemp -d)"
trap 'rm -rf "${tmp_dir}"' EXIT
fixture_dir="${tmp_dir}/audio.cpp"
prefix_dir="${tmp_dir}/prefix"
tools_dir="${tmp_dir}/tools"
mkdir -p "${fixture_dir}" "${prefix_dir}/lib/cmake/Protobuf" \
"${prefix_dir}/lib/cmake/gRPC" "${tools_dir}"
cat >"${fixture_dir}/engine_runtime.cpp" <<'EOF'
void audio_cpp_build_contract_fixture() {}
EOF
cat >"${fixture_dir}/CMakeLists.txt" <<'EOF'
cmake_minimum_required(VERSION 3.20)
project(AudioCppBuildContractFixture LANGUAGES CXX)
option(ENGINE_ENABLE_CUDA "Build with CUDA" OFF)
option(ENGINE_ENABLE_VULKAN "Build with Vulkan" OFF)
option(ENGINE_ENABLE_METAL "Build with Metal" OFF)
foreach(wrong_option IN ITEMS
AUDIO_CPP_ENABLE_CUDA
AUDIO_CPP_ENABLE_VULKAN
AUDIO_CPP_ENABLE_METAL
AUDIOCPP_ENABLE_CUDA
AUDIOCPP_ENABLE_VULKAN
AUDIOCPP_ENABLE_METAL
ENGINE_CUDA
ENGINE_VULKAN
ENGINE_METAL
GGML_CUDA
GGML_VULKAN
GGML_METAL)
if(DEFINED ${wrong_option})
message(FATAL_ERROR "legacy or unsupported audio.cpp option: ${wrong_option}")
endif()
endforeach()
add_library(engine_runtime STATIC engine_runtime.cpp)
EOF
cat >"${prefix_dir}/lib/cmake/Protobuf/ProtobufConfig.cmake" <<'EOF'
set(Protobuf_FOUND TRUE)
set(Protobuf_VERSION 0.0.0)
if(NOT TARGET protobuf::libprotobuf)
add_library(protobuf::libprotobuf INTERFACE IMPORTED)
endif()
EOF
cat >"${prefix_dir}/lib/cmake/gRPC/gRPCConfig.cmake" <<'EOF'
set(gRPC_FOUND TRUE)
if(NOT TARGET gRPC::grpc++)
add_library(gRPC::grpc++ INTERFACE IMPORTED)
endif()
if(NOT TARGET gRPC::grpc++_reflection)
add_library(gRPC::grpc++_reflection INTERFACE IMPORTED)
endif()
EOF
cat >"${tools_dir}/protoc" <<'EOF'
#!/usr/bin/env sh
exit 0
EOF
cat >"${tools_dir}/grpc_cpp_plugin" <<'EOF'
#!/usr/bin/env sh
exit 0
EOF
chmod +x "${tools_dir}/protoc" "${tools_dir}/grpc_cpp_plugin"
assert_cache_bool() {
local cache_file="$1"
local name="$2"
local expected="$3"
grep -q "^${name}:BOOL=${expected}$" "${cache_file}" || {
echo "expected ${name}:BOOL=${expected} in ${cache_file}" >&2
return 1
}
}
configure_case() {
local name="$1"
local cuda="$2"
local vulkan="$3"
local metal="$4"
local build_dir="${tmp_dir}/build-${name}"
PATH="${tools_dir}:${PATH}" cmake \
-S "${backend_dir}" \
-B "${build_dir}" \
-DCMAKE_PREFIX_PATH="${prefix_dir}" \
-DAUDIO_CPP_DIR="${fixture_dir}" \
-DENGINE_ENABLE_CUDA="${cuda}" \
-DENGINE_ENABLE_VULKAN="${vulkan}" \
-DENGINE_ENABLE_METAL="${metal}" \
>/dev/null
assert_cache_bool "${build_dir}/CMakeCache.txt" ENGINE_ENABLE_CUDA "${cuda}"
assert_cache_bool "${build_dir}/CMakeCache.txt" ENGINE_ENABLE_VULKAN "${vulkan}"
assert_cache_bool "${build_dir}/CMakeCache.txt" ENGINE_ENABLE_METAL "${metal}"
grep -q 'engine_runtime' \
"${build_dir}/CMakeFiles/audio-cpp-grpc-server.dir/link.txt" || {
echo "audio-cpp-grpc-server does not link engine_runtime" >&2
return 1
}
}
configure_case cpu OFF OFF OFF
configure_case cuda ON OFF OFF
configure_case vulkan OFF ON OFF
configure_case metal OFF OFF ON
if PATH="${tools_dir}:${PATH}" cmake \
-S "${backend_dir}" \
-B "${tmp_dir}/build-wrong-option" \
-DCMAKE_PREFIX_PATH="${prefix_dir}" \
-DAUDIO_CPP_DIR="${fixture_dir}" \
-DGGML_CUDA=ON \
>/dev/null 2>&1; then
echo "strict audio.cpp fixture accepted legacy GGML_CUDA option" >&2
exit 1
fi
make_database="${tmp_dir}/make-database"
make -C "${backend_dir}" -pn >"${make_database}"
audio_cpp_version="$(
sed -n 's/^AUDIO_CPP_VERSION = //p' "${make_database}" | head -n 1
)"
[[ "${audio_cpp_version}" =~ ^[0-9a-f]{40}$ ]] || {
echo "AUDIO_CPP_VERSION must be a pinned 40-character commit" >&2
exit 1
}
fetch_plan="${tmp_dir}/fetch-plan"
make -C "${backend_dir}" -Bn audio.cpp >"${fetch_plan}"
grep -q 'github.com/0xShug0/audio.cpp' "${fetch_plan}"
grep -q "${audio_cpp_version}" "${fetch_plan}"
cat >"${tools_dir}/uname" <<'EOF'
#!/usr/bin/env sh
if [ "$#" -eq 1 ] && [ "$1" = "-s" ]; then
echo Darwin
exit 0
fi
echo "build contract requires uname -s" >&2
exit 64
EOF
chmod +x "${tools_dir}/uname"
darwin_plan="${tmp_dir}/darwin-plan"
PATH="${tools_dir}:${PATH}" make -C "${backend_dir}" -n \
AUDIO_CPP_SRC="${fixture_dir}" grpc-server >"${darwin_plan}"
grep -q -- '-DENGINE_ENABLE_CUDA=OFF' "${darwin_plan}"
grep -q -- '-DENGINE_ENABLE_VULKAN=OFF' "${darwin_plan}"
grep -q -- '-DENGINE_ENABLE_METAL=ON' "${darwin_plan}"
echo "audio.cpp build contract: PASS"

View File

@@ -1,490 +0,0 @@
// SPDX-License-Identifier: MIT
#include "audio_cpp_runtime.h"
#include "model_config.h"
#include "engine/framework/runtime/model.h"
#include "engine/framework/runtime/registry.h"
#include "engine/framework/runtime/session.h"
#include <atomic>
#include <chrono>
#include <condition_variable>
#include <exception>
#include <filesystem>
#include <future>
#include <iostream>
#include <memory>
#include <mutex>
#include <stdexcept>
#include <string>
#include <unordered_map>
#include <utility>
#include <vector>
namespace {
using engine::runtime::AudioChunk;
using engine::runtime::CapabilitySet;
using engine::runtime::ILoadedVoiceModel;
using engine::runtime::IOfflineVoiceTaskSession;
using engine::runtime::IStreamingVoiceTaskSession;
using engine::runtime::IVoiceModelLoader;
using engine::runtime::IVoiceTaskSession;
using engine::runtime::ModelInspection;
using engine::runtime::ModelLoadRequest;
using engine::runtime::ModelMetadata;
using engine::runtime::RunMode;
using engine::runtime::SessionOptions;
using engine::runtime::SessionPreparationRequest;
using engine::runtime::StreamEvent;
using engine::runtime::TaskCapability;
using engine::runtime::TaskRequest;
using engine::runtime::TaskResult;
using engine::runtime::TaskSpec;
using engine::runtime::VoiceTaskKind;
void require(bool condition, const std::string & message) {
if (!condition) {
throw std::runtime_error(message);
}
}
template <typename Function>
void require_throws(Function && function, const std::string & expected) {
try {
function();
} catch (const std::exception & error) {
require(
std::string(error.what()).find(expected) != std::string::npos,
"expected error containing '" + expected + "', got '" + error.what() + "'");
return;
}
throw std::runtime_error("expected exception containing '" + expected + "'");
}
struct SessionGate {
std::mutex mutex;
std::condition_variable condition;
bool first_entered = false;
bool release_first = false;
std::atomic<int> entries{0};
};
struct FakeState {
std::mutex mutex;
std::vector<std::string> events;
CapabilitySet capabilities;
bool fail_load = false;
bool fail_session = false;
int generation = 0;
std::shared_ptr<SessionGate> gate;
void record(std::string event) {
std::lock_guard<std::mutex> lock(mutex);
events.push_back(std::move(event));
}
};
class FakeSession final
: public IOfflineVoiceTaskSession,
public IStreamingVoiceTaskSession {
public:
FakeSession(
std::shared_ptr<FakeState> state,
int generation,
TaskSpec task,
SessionOptions options)
: state_(std::move(state)),
generation_(generation),
task_(task),
options_(std::move(options)) {}
~FakeSession() override {
state_->record("session-" + std::to_string(generation_) + "-destroyed");
}
std::string family() const override { return "fake-family"; }
VoiceTaskKind task_kind() const override { return task_.task; }
RunMode run_mode() const override { return task_.mode; }
void prepare(const SessionPreparationRequest & request) override {
prepared_ = request;
}
TaskResult run(const TaskRequest &) override {
if (state_->gate != nullptr) {
const int entry = ++state_->gate->entries;
if (entry == 1) {
std::unique_lock<std::mutex> lock(state_->gate->mutex);
state_->gate->first_entered = true;
state_->gate->condition.notify_all();
state_->gate->condition.wait(
lock,
[&] { return state_->gate->release_first; });
}
}
TaskResult result;
result.text_output = engine::runtime::Transcript{
"generation-" + std::to_string(generation_),
"en",
};
return result;
}
engine::runtime::StreamingPolicy streaming_policy() const override {
engine::runtime::StreamingPolicy policy;
policy.input = engine::runtime::StreamingInputKind::AudioChunks;
policy.output = engine::runtime::StreamingOutputKind::PullEvents;
policy.preferred_audio_chunk_samples = 160;
return policy;
}
void start_stream(const TaskRequest &) override {
streaming_ = true;
}
std::optional<StreamEvent> next_stream_event() override {
return std::nullopt;
}
void set_stream_event_sink(engine::runtime::StreamEventCallback sink) override {
sink_ = std::move(sink);
}
TaskResult finish_stream() override {
streaming_ = false;
return run({});
}
void reset() override {
streaming_ = false;
}
StreamEvent process_audio_chunk(const AudioChunk & chunk) override {
require(streaming_, "stream was not started");
StreamEvent event;
event.audio_output = engine::runtime::AudioBuffer{
chunk.sample_rate,
chunk.channels,
chunk.samples,
};
if (sink_) {
sink_(event);
}
return event;
}
TaskResult finalize() override {
streaming_ = false;
return run({});
}
private:
std::shared_ptr<FakeState> state_;
int generation_;
TaskSpec task_;
SessionOptions options_;
SessionPreparationRequest prepared_;
engine::runtime::StreamEventCallback sink_;
bool streaming_ = false;
};
class FakeLoadedModel final : public ILoadedVoiceModel {
public:
FakeLoadedModel(std::shared_ptr<FakeState> state, int generation)
: state_(std::move(state)),
generation_(generation) {
metadata_.family = "fake-family";
metadata_.variant = "complete-fake";
metadata_.description = "complete test implementation";
metadata_.config_candidates = {"config.json"};
metadata_.weight_candidates = {"weights.gguf"};
}
~FakeLoadedModel() override {
state_->record("model-" + std::to_string(generation_) + "-destroyed");
}
const ModelMetadata & metadata() const noexcept override {
return metadata_;
}
const CapabilitySet & capabilities() const noexcept override {
return state_->capabilities;
}
std::unique_ptr<IVoiceTaskSession> create_task_session(
const TaskSpec & task,
const SessionOptions & options) const override {
if (state_->fail_session) {
throw std::runtime_error("session creation failed");
}
return std::make_unique<FakeSession>(state_, generation_, task, options);
}
private:
std::shared_ptr<FakeState> state_;
int generation_;
ModelMetadata metadata_;
};
class FakeLoader final : public IVoiceModelLoader {
public:
explicit FakeLoader(std::shared_ptr<FakeState> state)
: state_(std::move(state)) {}
std::string family() const override { return "fake-family"; }
bool can_load(const ModelLoadRequest & request) const override {
return request.family_hint == family();
}
ModelInspection inspect(const ModelLoadRequest & request) const override {
ModelInspection inspection;
inspection.metadata.family = family();
inspection.metadata.variant = "complete-fake";
inspection.metadata.description = "complete test loader";
inspection.metadata.config_candidates = {"config.json"};
inspection.metadata.weight_candidates = {"weights.gguf"};
inspection.capabilities = state_->capabilities;
inspection.model_root = request.model_path;
return inspection;
}
std::unique_ptr<ILoadedVoiceModel> load(
const ModelLoadRequest &) const override {
if (state_->fail_load) {
throw std::runtime_error("model load failed");
}
const int generation = ++state_->generation;
return std::make_unique<FakeLoadedModel>(state_, generation);
}
CapabilitySet advertised_capabilities() const override {
return state_->capabilities;
}
std::string advertised_instructions_policy() const override {
return "explicit";
}
std::vector<std::string> advertised_api_endpoints() const override {
return {"/v1/audio/transcriptions"};
}
private:
std::shared_ptr<FakeState> state_;
};
audio_cpp::AudioCppModelConfig offline_asr_config(
const std::filesystem::path & model_path) {
return audio_cpp::parse_model_config(
model_path,
{
{"family", "fake-family"},
{"task", "asr"},
{"mode", "offline"},
{"backend", "cpu"},
{"device", "2"},
{"threads", "3"},
{"load.cache", "memory"},
{"session.language", "en"},
});
}
std::unique_ptr<audio_cpp::AudioCppRuntime> make_runtime(
const std::shared_ptr<FakeState> & state) {
engine::runtime::ModelRegistry registry;
registry.register_loader(std::make_shared<FakeLoader>(state));
return std::make_unique<audio_cpp::AudioCppRuntime>(std::move(registry));
}
void test_model_config(const std::filesystem::path & model_path) {
const auto config = offline_asr_config(model_path);
require(config.model_path == model_path, "model path was not preserved");
require(config.family == "fake-family", "family was not parsed");
require(config.task.task == VoiceTaskKind::Asr, "task was not parsed");
require(config.task.mode == RunMode::Offline, "mode was not parsed");
require(
config.session.backend.type == engine::core::BackendType::Cpu,
"backend was not parsed");
require(config.session.backend.device == 2, "device was not parsed");
require(config.session.backend.threads == 3, "threads were not parsed");
require(config.load.options.at("cache") == "memory", "load option prefix was not stripped");
require(
config.session.options.at("language") == "en",
"session option prefix was not stripped");
require_throws(
[&] { audio_cpp::parse_model_config(model_path, {{"task", "asr"}}); },
"family");
require_throws(
[&] { audio_cpp::parse_model_config(model_path, {{"family", "fake-family"}}); },
"task");
require_throws(
[&] {
audio_cpp::parse_model_config(
model_path,
{{"family", "fake-family"}, {"task", "asr"}, {"mode", "batch"}});
},
"mode");
require_throws(
[&] {
audio_cpp::parse_model_config(
model_path,
{{"family", "fake-family"}, {"task", "asr"}, {"backend", "tpu"}});
},
"backend");
}
void test_capability_validation(const std::filesystem::path & model_path) {
auto state = std::make_shared<FakeState>();
state->capabilities.supported_tasks = {
{VoiceTaskKind::Tts, {RunMode::Offline}},
};
auto runtime = make_runtime(state);
require_throws(
[&] { runtime->load(offline_asr_config(model_path)); },
"task");
state->capabilities.supported_tasks = {
{VoiceTaskKind::Asr, {RunMode::Streaming}},
};
require_throws(
[&] { runtime->load(offline_asr_config(model_path)); },
"mode");
}
void test_atomic_replacement(const std::filesystem::path & model_path) {
auto state = std::make_shared<FakeState>();
state->capabilities.supported_tasks = {
{VoiceTaskKind::Asr, {RunMode::Offline}},
};
auto runtime = make_runtime(state);
runtime->load(offline_asr_config(model_path));
state->fail_load = true;
require_throws(
[&] { runtime->load(offline_asr_config(model_path)); },
"model load failed");
require(
runtime->run({}).text_output->text == "generation-1",
"old model was not retained after load failure");
state->fail_load = false;
state->fail_session = true;
require_throws(
[&] { runtime->load(offline_asr_config(model_path)); },
"session creation failed");
require(
runtime->run({}).text_output->text == "generation-1",
"old model was not retained after session creation failure");
}
void test_teardown_order(const std::filesystem::path & model_path) {
auto state = std::make_shared<FakeState>();
state->capabilities.supported_tasks = {
{VoiceTaskKind::Asr, {RunMode::Offline}},
};
auto runtime = make_runtime(state);
runtime->load(offline_asr_config(model_path));
runtime->free();
std::lock_guard<std::mutex> lock(state->mutex);
require(state->events.size() == 2, "expected one session and one model teardown");
require(
state->events[0] == "session-1-destroyed",
"session was not destroyed before model");
require(
state->events[1] == "model-1-destroyed",
"model teardown event was not second");
}
void test_runtime_serializes_calls(const std::filesystem::path & model_path) {
auto state = std::make_shared<FakeState>();
state->capabilities.supported_tasks = {
{VoiceTaskKind::Asr, {RunMode::Offline}},
};
state->gate = std::make_shared<SessionGate>();
auto runtime = make_runtime(state);
runtime->load(offline_asr_config(model_path));
auto first = std::async(std::launch::async, [&] { return runtime->run({}); });
{
std::unique_lock<std::mutex> lock(state->gate->mutex);
state->gate->condition.wait(
lock,
[&] { return state->gate->first_entered; });
}
std::promise<void> release_second;
std::shared_future<void> second_barrier = release_second.get_future().share();
std::promise<void> second_attempted_promise;
auto second_attempted = second_attempted_promise.get_future();
auto second = std::async(std::launch::async, [&] {
second_barrier.wait();
second_attempted_promise.set_value();
return runtime->run({});
});
release_second.set_value();
second_attempted.wait();
require(
second.wait_for(std::chrono::milliseconds(50)) == std::future_status::timeout,
"second call completed while first call held the runtime");
require(
state->gate->entries.load() == 1,
"second call entered the upstream session concurrently");
{
std::lock_guard<std::mutex> lock(state->gate->mutex);
state->gate->release_first = true;
}
state->gate->condition.notify_all();
first.get();
second.get();
require(state->gate->entries.load() == 2, "second call never reached the session");
}
void test_streaming_surface(const std::filesystem::path & model_path) {
auto state = std::make_shared<FakeState>();
state->capabilities.supported_tasks = {
{VoiceTaskKind::Asr, {RunMode::Streaming}},
};
auto runtime = make_runtime(state);
auto config = offline_asr_config(model_path);
config.task.mode = RunMode::Streaming;
runtime->load(config);
runtime->start_stream({});
const auto event = runtime->process_audio_chunk({16000, 1, 0, {0.25f}});
require(event.audio_output.has_value(), "streaming chunk result was lost");
require(
event.audio_output->samples == std::vector<float>{0.25f},
"streaming chunk samples changed");
require(
runtime->finish_stream().text_output->text == "generation-1",
"streaming final result was lost");
}
} // namespace
int main(int argc, char ** argv) {
try {
require(argc == 2, "runtime_test requires an existing model path argument");
const std::filesystem::path model_path(argv[1]);
test_model_config(model_path);
test_capability_validation(model_path);
test_atomic_replacement(model_path);
test_teardown_order(model_path);
test_runtime_serializes_calls(model_path);
test_streaming_surface(model_path);
std::cout << "audio.cpp runtime unit tests: PASS\n";
return 0;
} catch (const std::exception & error) {
std::cerr << "audio.cpp runtime unit tests: FAIL: " << error.what() << '\n';
return 1;
}
}

View File

@@ -157,6 +157,33 @@ var _ = Describe("X-LocalAI-Node ctx propagation contract", func() {
stampViaRouterCtx()
})
// Regression for #10636: a canceled request context must NOT cancel the
// model LOAD. The heavy image/audio backends bind the load to the request
// context so the routing holder reaches the SmartRouter; but a large
// diffusers/LLM model on a slow (e.g. shared-memory iGPU) host can take
// far longer to load than the client stays connected. If the request's
// cancellation propagates to the load, the LoadModel RPC is aborted, the
// backend process is torn down, and every retry restarts from scratch and
// never converges. The load must instead run to completion and cache while
// still carrying the request's routing holder value.
It("ImageGeneration does not propagate request cancellation to the model load", func() {
canceledCtx, cancel := context.WithCancel(reqCtx)
cancel() // client disconnected while the (slow) load was still running
_, err := backend.ImageGeneration(canceledCtx, 64, 64, 1, 0, "p", "", "", "/tmp/out.png", loader, modelCfg, appCfg, nil)
// The load reached the router (short-circuit sentinel), i.e. it was
// NOT aborted early by the already-canceled request context.
Expect(err).To(HaveOccurred())
Expect(err.Error()).To(ContainSubstring("router short-circuit (test)"))
routerCtx := routerCtxOf()
Expect(routerCtx).ToNot(BeNil(), "router callback must have been invoked")
Expect(routerCtx.Err()).To(BeNil(),
"a canceled request must not cancel the model load")
// The routing holder value still propagates despite the decoupling.
stampViaRouterCtx()
})
It("does NOT leak the holder when the app context is used instead", func() {
// Sanity: the bug being fixed manifests as the router getting
// appCfg.Context (no holder) instead of reqCtx (holder). A direct

View File

@@ -40,10 +40,14 @@ func (e *modelEmbedder) Embed(ctx context.Context, text string) ([]float32, erro
func ModelEmbedding(ctx context.Context, s string, tokens []int, loader *model.ModelLoader, modelConfig config.ModelConfig, appConfig *config.ApplicationConfig) (func() ([]float32, error), error) {
// model.WithContext(ctx) overrides the app-context default set in
// ModelOptions so distributed routing decisions reach the request's
// X-LocalAI-Node holder via distributedhdr.Stamp.
opts := ModelOptions(modelConfig, appConfig, model.WithContext(ctx))
// model.WithContext carries the request context into the load so distributed
// routing decisions reach the request's X-LocalAI-Node holder via
// distributedhdr.Stamp. context.WithoutCancel keeps those values but drops
// the request's cancellation, so a slow first load still completes and
// caches if the client disconnects instead of aborting the LoadModel RPC and
// tearing down the backend process (issue #10636). Inference below keeps the
// cancellable ctx, so a disconnect still stops generation.
opts := ModelOptions(modelConfig, appConfig, model.WithContext(context.WithoutCancel(ctx)))
inferenceModel, err := loader.Load(opts...)
if err != nil {

View File

@@ -13,10 +13,14 @@ import (
func ImageGeneration(ctx context.Context, height, width, step, seed int, positive_prompt, negative_prompt, src, dst string, loader *model.ModelLoader, modelConfig config.ModelConfig, appConfig *config.ApplicationConfig, refImages []string) (func() error, error) {
// model.WithContext(ctx) overrides the app-context default set in
// ModelOptions so distributed routing decisions reach the request's
// X-LocalAI-Node holder via distributedhdr.Stamp.
opts := ModelOptions(modelConfig, appConfig, model.WithContext(ctx))
// model.WithContext carries the request context into the load so distributed
// routing decisions reach the request's X-LocalAI-Node holder via
// distributedhdr.Stamp. context.WithoutCancel keeps those values but drops
// the request's cancellation, so a slow first load still completes and
// caches if the client disconnects instead of aborting the LoadModel RPC and
// tearing down the backend process (issue #10636). Inference below keeps the
// cancellable ctx, so a disconnect still stops generation.
opts := ModelOptions(modelConfig, appConfig, model.WithContext(context.WithoutCancel(ctx)))
inferenceModel, err := loader.Load(
opts...,
)

View File

@@ -133,7 +133,12 @@ func ModelInference(ctx context.Context, s string, messages schema.Messages, ima
}
ctx = distributedhdr.MaybeWithPrefixChain(ctx, c.ModelID(), chainSource)
opts := ModelOptions(*c, o, model.WithContext(ctx))
// context.WithoutCancel decouples the model load from the request's
// cancellation while preserving its routing values, so a slow load still
// completes and caches if the client disconnects instead of aborting the
// LoadModel RPC mid-load (issue #10636). Inference below keeps the
// cancellable ctx, so a disconnect still stops generation.
opts := ModelOptions(*c, o, model.WithContext(context.WithoutCancel(ctx)))
inferenceModel, err := loader.Load(opts...)
if err != nil {
recordModelLoadFailure(o, c.Name, c.Backend, err, map[string]any{"model_file": modelFile})

View File

@@ -57,10 +57,14 @@ func (r *modelReranker) Rerank(ctx context.Context, query string, documents []st
}
func Rerank(ctx context.Context, request *proto.RerankRequest, loader *model.ModelLoader, appConfig *config.ApplicationConfig, modelConfig config.ModelConfig) (*proto.RerankResult, error) {
// model.WithContext(ctx) overrides the app-context default set in
// ModelOptions so distributed routing decisions reach the request's
// X-LocalAI-Node holder via distributedhdr.Stamp.
opts := ModelOptions(modelConfig, appConfig, model.WithContext(ctx))
// model.WithContext carries the request context into the load so distributed
// routing decisions reach the request's X-LocalAI-Node holder via
// distributedhdr.Stamp. context.WithoutCancel keeps those values but drops
// the request's cancellation, so a slow first load still completes and
// caches if the client disconnects instead of aborting the LoadModel RPC and
// tearing down the backend process (issue #10636). Inference below keeps the
// cancellable ctx, so a disconnect still stops generation.
opts := ModelOptions(modelConfig, appConfig, model.WithContext(context.WithoutCancel(ctx)))
rerankModel, err := loader.Load(opts...)
if err != nil {
recordModelLoadFailure(appConfig, modelConfig.Name, modelConfig.Backend, err, nil)

View File

@@ -51,10 +51,14 @@ func loadTranscriptionModel(ctx context.Context, ml *model.ModelLoader, modelCon
if modelConfig.Backend == "" {
modelConfig.Backend = model.WhisperBackend
}
// model.WithContext(ctx) overrides the app-context default set in
// ModelOptions so distributed routing decisions reach the request's
// X-LocalAI-Node holder via distributedhdr.Stamp.
opts := ModelOptions(modelConfig, appConfig, model.WithContext(ctx))
// model.WithContext carries the request context into the load so distributed
// routing decisions reach the request's X-LocalAI-Node holder via
// distributedhdr.Stamp. context.WithoutCancel keeps those values but drops
// the request's cancellation, so a slow first load still completes and
// caches if the client disconnects instead of aborting the LoadModel RPC and
// tearing down the backend process (issue #10636). Inference below keeps the
// cancellable ctx, so a disconnect still stops generation.
opts := ModelOptions(modelConfig, appConfig, model.WithContext(context.WithoutCancel(ctx)))
transcriptionModel, err := ml.Load(opts...)
if err != nil {
recordModelLoadFailure(appConfig, modelConfig.Name, modelConfig.Backend, err, nil)

View File

@@ -57,10 +57,14 @@ func ModelTTS(
appConfig *config.ApplicationConfig,
modelConfig config.ModelConfig,
) (string, *proto.Result, error) {
// model.WithContext(ctx) overrides the app-context default set in
// ModelOptions so distributed routing decisions reach the request's
// X-LocalAI-Node holder via distributedhdr.Stamp.
opts := ModelOptions(modelConfig, appConfig, model.WithContext(ctx))
// model.WithContext carries the request context into the load so distributed
// routing decisions reach the request's X-LocalAI-Node holder via
// distributedhdr.Stamp. context.WithoutCancel keeps those values but drops
// the request's cancellation, so a slow first load still completes and
// caches if the client disconnects instead of aborting the LoadModel RPC and
// tearing down the backend process (issue #10636). Inference below keeps the
// cancellable ctx, so a disconnect still stops generation.
opts := ModelOptions(modelConfig, appConfig, model.WithContext(context.WithoutCancel(ctx)))
ttsModel, err := loader.Load(opts...)
if err != nil {
recordModelLoadFailure(appConfig, modelConfig.Name, modelConfig.Backend, err, nil)
@@ -160,7 +164,9 @@ func ModelTTSStream(
modelConfig config.ModelConfig,
audioCallback func([]byte) error,
) error {
opts := ModelOptions(modelConfig, appConfig, model.WithContext(ctx))
// See ModelTTS above: WithoutCancel decouples the load from request
// cancellation while preserving routing values (issue #10636).
opts := ModelOptions(modelConfig, appConfig, model.WithContext(context.WithoutCancel(ctx)))
ttsModel, err := loader.Load(opts...)
if err != nil {
recordModelLoadFailure(appConfig, modelConfig.Name, modelConfig.Backend, err, nil)

View File

@@ -14,10 +14,14 @@ func VAD(request *schema.VADRequest,
ml *model.ModelLoader,
appConfig *config.ApplicationConfig,
modelConfig config.ModelConfig) (*schema.VADResponse, error) {
// model.WithContext(ctx) overrides the app-context default set in
// ModelOptions so distributed routing decisions reach the request's
// X-LocalAI-Node holder via distributedhdr.Stamp.
opts := ModelOptions(modelConfig, appConfig, model.WithContext(ctx))
// model.WithContext carries the request context into the load so distributed
// routing decisions reach the request's X-LocalAI-Node holder via
// distributedhdr.Stamp. context.WithoutCancel keeps those values but drops
// the request's cancellation, so a slow first load still completes and
// caches if the client disconnects instead of aborting the LoadModel RPC and
// tearing down the backend process (issue #10636). Inference below keeps the
// cancellable ctx, so a disconnect still stops generation.
opts := ModelOptions(modelConfig, appConfig, model.WithContext(context.WithoutCancel(ctx)))
vadModel, err := ml.Load(opts...)
if err != nil {
recordModelLoadFailure(appConfig, modelConfig.Name, modelConfig.Backend, err, nil)