mirror of
https://github.com/mudler/LocalAI.git
synced 2026-09-10 04:57:51 -04:00
Compare commits
19
Commits
No files matched your search
@@ -240,7 +240,7 @@ test-ci-scripts:
|
||||
## pure stdlib on purpose so they run without any backend venv; the list is
|
||||
## explicit because their siblings (model_identity_test) import grpc and the
|
||||
## generated protobufs, which only exist inside a built backend.
|
||||
PYTHON_HELPER_TESTS?=python_utils_test vllm_utils_test model_utils_test mlx_utils_test parent_watch_test
|
||||
PYTHON_HELPER_TESTS?=python_utils_test vllm_utils_test model_utils_test mlx_utils_test parent_watch_test temp_utils_test
|
||||
test-python-helpers:
|
||||
cd backend/python/common && python3 -m unittest $(PYTHON_HELPER_TESTS)
|
||||
|
||||
|
||||
@@ -9,7 +9,7 @@
|
||||
# recipe is a make target (not a prepare.sh) so 'make purge && make' is a clean
|
||||
# rebuild and so the bump bot can see the pin.
|
||||
|
||||
AUDIO_CPP_VERSION?=9c6a282337cc83f227cc10428867a478947706ad
|
||||
AUDIO_CPP_VERSION?=fa5aaac9266a98c68f8a5c9fcd1ba6ff65875416
|
||||
AUDIO_CPP_REPO?=https://github.com/0xShug0/audio.cpp
|
||||
|
||||
CURRENT_MAKEFILE_DIR := $(dir $(abspath $(lastword $(MAKEFILE_LIST))))
|
||||
|
||||
@@ -1,10 +1,10 @@
|
||||
# ds4 backend Makefile.
|
||||
#
|
||||
# Upstream pin lives below as DS4_VERSION?=f62ca29a308724cde5bc99134ede19104b2a3260
|
||||
# Upstream pin lives below as DS4_VERSION?=6289c516273979173abbc062209a81dd3706b804
|
||||
# (.github/bump_deps.sh) can find and update it - matches the
|
||||
# llama-cpp / ik-llama-cpp / turboquant convention.
|
||||
|
||||
DS4_VERSION?=f62ca29a308724cde5bc99134ede19104b2a3260
|
||||
DS4_VERSION?=6289c516273979173abbc062209a81dd3706b804
|
||||
DS4_REPO?=https://github.com/antirez/ds4
|
||||
|
||||
CURRENT_MAKEFILE_DIR := $(dir $(abspath $(lastword $(MAKEFILE_LIST))))
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
|
||||
IK_LLAMA_VERSION?=fe215a8ccdce6b844d2a3a3bbde08ae76a6284bf
|
||||
IK_LLAMA_VERSION?=1a2a8604a6c6c6413c06bf9adfc2f64329af4366
|
||||
LLAMA_REPO?=https://github.com/ikawrakow/ik_llama.cpp
|
||||
|
||||
CMAKE_ARGS?=
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
|
||||
LLAMA_VERSION?=67672dc5b76f8bc17785a19d3dc6d1463fc2902c
|
||||
LLAMA_VERSION?=f3f1a8f2760f28325a5ec20c05b171e5b7c83a29
|
||||
LLAMA_REPO?=https://github.com/ggerganov/llama.cpp
|
||||
|
||||
CMAKE_ARGS?=
|
||||
|
||||
@@ -615,10 +615,10 @@ func (w *CrispASR) TTSStream(req *pb.TTSRequest, results chan []byte) error {
|
||||
return fmt.Errorf("crispasr: tempfile: %w", err)
|
||||
}
|
||||
dst := tmp.Name()
|
||||
defer func() { _ = os.Remove(dst) }()
|
||||
if err := tmp.Close(); err != nil {
|
||||
return fmt.Errorf("crispasr: close tempfile: %w", err)
|
||||
}
|
||||
defer func() { _ = os.Remove(dst) }()
|
||||
|
||||
if err := writeWAV(dst, pcm, w.sampleRate); err != nil {
|
||||
return err
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"unsafe"
|
||||
|
||||
"github.com/mudler/LocalAI/pkg/grpc/base"
|
||||
@@ -109,30 +110,25 @@ func (r *LocateAnythingCpp) Detect(opts *pb.DetectOptions) (pb.DetectResponse, e
|
||||
return pb.DetectResponse{}, fmt.Errorf("locate-anything-cpp: a text prompt is required (open-vocabulary detection)")
|
||||
}
|
||||
|
||||
// Decode base64 image and write to temp file.
|
||||
imgData, err := base64.StdEncoding.DecodeString(opts.Src)
|
||||
if err != nil {
|
||||
return pb.DetectResponse{}, fmt.Errorf("locate-anything-cpp: failed to decode base64 image: %w", err)
|
||||
}
|
||||
|
||||
tmpFile, err := os.CreateTemp("", "locate-anything-*.img")
|
||||
if err != nil {
|
||||
return pb.DetectResponse{}, fmt.Errorf("locate-anything-cpp: failed to create temp file: %w", err)
|
||||
}
|
||||
defer func() { _ = os.Remove(tmpFile.Name()) }()
|
||||
|
||||
if _, err := tmpFile.Write(imgData); err != nil {
|
||||
_ = tmpFile.Close()
|
||||
return pb.DetectResponse{}, fmt.Errorf("locate-anything-cpp: failed to write temp file: %w", err)
|
||||
}
|
||||
if err := tmpFile.Close(); err != nil {
|
||||
return pb.DetectResponse{}, fmt.Errorf("locate-anything-cpp: failed to close temp file: %w", err)
|
||||
if len(imgData) == 0 {
|
||||
return pb.DetectResponse{}, fmt.Errorf("locate-anything-cpp: decoded image is empty")
|
||||
}
|
||||
|
||||
// mode 0 = hybrid (Parallel Box Decoding). The JSON return value is unused:
|
||||
// structured detections are read via the accessor functions. Still must
|
||||
// free the returned string.
|
||||
jsonPtr := CapiLocatePath(r.handle, tmpFile.Name(), prompt, 0)
|
||||
jsonPtr := CapiLocateBuffer(
|
||||
r.handle,
|
||||
uintptr(unsafe.Pointer(unsafe.SliceData(imgData))),
|
||||
uintptr(len(imgData)),
|
||||
prompt,
|
||||
0,
|
||||
)
|
||||
runtime.KeepAlive(imgData)
|
||||
if jsonPtr != 0 {
|
||||
CapiFreeString(jsonPtr)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,54 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"path/filepath"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("LocateAnythingCpp detection input", func() {
|
||||
It("detects from memory when the temporary directory is unavailable", func() {
|
||||
originalLocateBuffer := CapiLocateBuffer
|
||||
originalLocatePath := CapiLocatePath
|
||||
originalGetNDetections := CapiGetNDetections
|
||||
defer func() {
|
||||
CapiLocateBuffer = originalLocateBuffer
|
||||
CapiLocatePath = originalLocatePath
|
||||
CapiGetNDetections = originalGetNDetections
|
||||
}()
|
||||
|
||||
image := []byte("encoded-image")
|
||||
var receivedData uintptr
|
||||
var receivedLength uintptr
|
||||
CapiLocateBuffer = func(_ uintptr, data uintptr, length uintptr, _ string, _ int32) uintptr {
|
||||
receivedData = data
|
||||
receivedLength = length
|
||||
return 0
|
||||
}
|
||||
CapiLocatePath = func(_ uintptr, _ string, _ string, _ int32) uintptr {
|
||||
Fail("path-based detection must not be called")
|
||||
return 0
|
||||
}
|
||||
CapiGetNDetections = func(uintptr) int32 { return 0 }
|
||||
GinkgoT().Setenv("TMPDIR", filepath.Join(GinkgoT().TempDir(), "missing"))
|
||||
|
||||
result, err := (&LocateAnythingCpp{handle: 1}).Detect(&pb.DetectOptions{
|
||||
Src: base64.StdEncoding.EncodeToString(image),
|
||||
Prompt: "the object",
|
||||
})
|
||||
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(result.Detections).To(BeEmpty())
|
||||
Expect(receivedData).NotTo(BeZero())
|
||||
Expect(receivedLength).To(Equal(uintptr(len(image))))
|
||||
})
|
||||
|
||||
It("rejects an empty decoded image", func() {
|
||||
_, err := (&LocateAnythingCpp{handle: 1}).Detect(&pb.DetectOptions{Prompt: "the object"})
|
||||
|
||||
Expect(err).To(MatchError("locate-anything-cpp: decoded image is empty"))
|
||||
})
|
||||
})
|
||||
@@ -12,7 +12,7 @@
|
||||
# runs 'make -C backend/go/$(BACKEND) build' and then copies package/), so it
|
||||
# has to produce the binary and the package, not just the shared libraries.
|
||||
|
||||
NEMO_SPEECH_VERSION?=ffa38cb2408f1e832a36d46fef5e3e1e80d07e6c
|
||||
NEMO_SPEECH_VERSION?=a5b6953c4a579a2bbd1c0913ad8a85c2a4d99953
|
||||
NEMO_SPEECH_REPO?=https://github.com/NVIDIA/NeMo-Speech.cpp
|
||||
|
||||
GOCMD?=go
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"unsafe"
|
||||
|
||||
@@ -102,24 +103,12 @@ func (r *RFDetrCpp) Detect(opts *pb.DetectOptions) (pb.DetectResponse, error) {
|
||||
return pb.DetectResponse{}, fmt.Errorf("rfdetr-cpp: model not loaded")
|
||||
}
|
||||
|
||||
// Decode base64 image and write to temp file.
|
||||
imgData, err := base64.StdEncoding.DecodeString(opts.Src)
|
||||
if err != nil {
|
||||
return pb.DetectResponse{}, fmt.Errorf("rfdetr-cpp: failed to decode base64 image: %w", err)
|
||||
}
|
||||
|
||||
tmpFile, err := os.CreateTemp("", "rfdetr-*.img")
|
||||
if err != nil {
|
||||
return pb.DetectResponse{}, fmt.Errorf("rfdetr-cpp: failed to create temp file: %w", err)
|
||||
}
|
||||
defer func() { _ = os.Remove(tmpFile.Name()) }()
|
||||
|
||||
if _, err := tmpFile.Write(imgData); err != nil {
|
||||
_ = tmpFile.Close()
|
||||
return pb.DetectResponse{}, fmt.Errorf("rfdetr-cpp: failed to write temp file: %w", err)
|
||||
}
|
||||
if err := tmpFile.Close(); err != nil {
|
||||
return pb.DetectResponse{}, fmt.Errorf("rfdetr-cpp: failed to close temp file: %w", err)
|
||||
if len(imgData) == 0 {
|
||||
return pb.DetectResponse{}, fmt.Errorf("rfdetr-cpp: decoded image is empty")
|
||||
}
|
||||
|
||||
threshold := opts.Threshold
|
||||
@@ -127,10 +116,18 @@ func (r *RFDetrCpp) Detect(opts *pb.DetectOptions) (pb.DetectResponse, error) {
|
||||
threshold = 0.5
|
||||
}
|
||||
|
||||
// JSON output from detect_path is unused: we read structured detections via
|
||||
// JSON output from the detection ABI is unused: we read structured detections via
|
||||
// the accessor functions. Still must free the returned string.
|
||||
var jsonPtr uintptr
|
||||
rc := CapiDetectPath(r.handle, tmpFile.Name(), threshold, uint32(defaultTopK), &jsonPtr)
|
||||
rc := CapiDetectBuffer(
|
||||
r.handle,
|
||||
uintptr(unsafe.Pointer(unsafe.SliceData(imgData))),
|
||||
uintptr(len(imgData)),
|
||||
threshold,
|
||||
uint32(defaultTopK),
|
||||
&jsonPtr,
|
||||
)
|
||||
runtime.KeepAlive(imgData)
|
||||
if jsonPtr != 0 {
|
||||
CapiFreeString(jsonPtr)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"path/filepath"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("RFDetrCpp detection input", func() {
|
||||
It("detects from memory when the temporary directory is unavailable", func() {
|
||||
originalDetectBuffer := CapiDetectBuffer
|
||||
originalDetectPath := CapiDetectPath
|
||||
originalFreeString := CapiFreeString
|
||||
originalGetNDetections := CapiGetNDetections
|
||||
defer func() {
|
||||
CapiDetectBuffer = originalDetectBuffer
|
||||
CapiDetectPath = originalDetectPath
|
||||
CapiFreeString = originalFreeString
|
||||
CapiGetNDetections = originalGetNDetections
|
||||
}()
|
||||
|
||||
image := []byte("encoded-image")
|
||||
var receivedData uintptr
|
||||
var receivedLength uintptr
|
||||
CapiDetectBuffer = func(_ uintptr, data uintptr, length uintptr, _ float32, _ uint32, _ *uintptr) int32 {
|
||||
receivedData = data
|
||||
receivedLength = length
|
||||
return 0
|
||||
}
|
||||
CapiDetectPath = func(_ uintptr, _ string, _ float32, _ uint32, _ *uintptr) int32 {
|
||||
Fail("path-based detection must not be called")
|
||||
return -1
|
||||
}
|
||||
CapiFreeString = func(uintptr) {}
|
||||
CapiGetNDetections = func(uintptr) int32 { return 0 }
|
||||
GinkgoT().Setenv("TMPDIR", filepath.Join(GinkgoT().TempDir(), "missing"))
|
||||
|
||||
result, err := (&RFDetrCpp{handle: 1}).Detect(&pb.DetectOptions{
|
||||
Src: base64.StdEncoding.EncodeToString(image),
|
||||
})
|
||||
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(result.Detections).To(BeEmpty())
|
||||
Expect(receivedData).NotTo(BeZero())
|
||||
Expect(receivedLength).To(Equal(uintptr(len(image))))
|
||||
})
|
||||
|
||||
It("rejects an empty decoded image", func() {
|
||||
_, err := (&RFDetrCpp{handle: 1}).Detect(&pb.DetectOptions{})
|
||||
|
||||
Expect(err).To(MatchError("rfdetr-cpp: decoded image is empty"))
|
||||
})
|
||||
})
|
||||
@@ -1144,17 +1144,25 @@ static uint8_t* load_and_resize_image(const char* path, int target_width, int ta
|
||||
// Write sd.cpp's audio buffer to a temp WAV file (IEEE float, interleaved).
|
||||
// sd_audio_t.data is planar (all channel 0 samples, then channel 1, etc.) — we
|
||||
// interleave on the fly so ffmpeg's standard wav demuxer can read it directly.
|
||||
// Returns 0 on success and fills wav_path (must be at least 64 bytes).
|
||||
// Returns 0 on success and fills wav_path.
|
||||
static int write_planar_float_wav(const sd_audio_t* a, char* wav_path, size_t wav_path_sz) {
|
||||
if (!a || !a->data || a->sample_count == 0 || a->channels == 0 || a->sample_rate == 0) {
|
||||
return -1;
|
||||
}
|
||||
|
||||
snprintf(wav_path, wav_path_sz, "/tmp/gosd-audio-XXXXXX.wav");
|
||||
const char* temp_dir = getenv("TMPDIR");
|
||||
if (!temp_dir || temp_dir[0] == '\0') {
|
||||
temp_dir = "/tmp";
|
||||
}
|
||||
int path_len = snprintf(wav_path, wav_path_sz, "%s/gosd-audio-XXXXXX.wav", temp_dir);
|
||||
if (path_len < 0 || (size_t)path_len >= wav_path_sz) {
|
||||
fprintf(stderr, "temporary directory path is too long\n");
|
||||
return -1;
|
||||
}
|
||||
int fd = mkstemps(wav_path, 4);
|
||||
if (fd < 0) { perror("mkstemps wav"); return -1; }
|
||||
FILE* f = fdopen(fd, "wb");
|
||||
if (!f) { perror("fdopen wav"); close(fd); return -1; }
|
||||
if (!f) { perror("fdopen wav"); close(fd); unlink(wav_path); return -1; }
|
||||
|
||||
uint64_t frames = a->sample_count;
|
||||
uint32_t channels = a->channels;
|
||||
@@ -1221,7 +1229,7 @@ static int ffmpeg_mux_raw_to_mp4(sd_image_t* frames, int num_frames, int fps,
|
||||
snprintf(fps_str, sizeof(fps_str), "%d", fps);
|
||||
|
||||
// Optional audio: write a temp WAV file if the model produced audio.
|
||||
char wav_path[64] = {0};
|
||||
char wav_path[4096] = {0};
|
||||
bool have_audio = false;
|
||||
if (audio && audio->data && audio->sample_count > 0 && audio->channels > 0 && audio->sample_rate > 0) {
|
||||
if (write_planar_float_wav(audio, wav_path, sizeof(wav_path)) == 0) {
|
||||
|
||||
@@ -8,7 +8,7 @@ JOBS?=$(shell nproc --ignore=1)
|
||||
|
||||
# whisper.cpp version
|
||||
WHISPER_REPO?=https://github.com/ggml-org/whisper.cpp
|
||||
WHISPER_CPP_VERSION?=52a939a2a762224e255d366c1182b2af4dd1a032
|
||||
WHISPER_CPP_VERSION?=c44b60b8053bbf2a5c1e014f11323fb3f2485177
|
||||
SO_TARGET?=libgowhisper.so
|
||||
|
||||
CMAKE_ARGS+=-DBUILD_SHARED_LIBS=OFF
|
||||
|
||||
@@ -1194,7 +1194,6 @@
|
||||
tags:
|
||||
- image-generation
|
||||
- video-generation
|
||||
- sound-generation
|
||||
- diffusion-models
|
||||
license: apache-2.0
|
||||
alias: "diffusers"
|
||||
|
||||
@@ -19,6 +19,7 @@ import grpc
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..', 'common'))
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), 'common'))
|
||||
from grpc_auth import get_auth_interceptors
|
||||
from temp_utils import cleanup_paths
|
||||
|
||||
import tempfile
|
||||
|
||||
@@ -115,11 +116,6 @@ def merge_audio_files(audio_files, output_path, sample_rate):
|
||||
# Save the merged audio
|
||||
ta.save(output_path, merged_waveform, sample_rate)
|
||||
|
||||
# Clean up temporary files
|
||||
for audio_file in audio_files:
|
||||
if os.path.exists(audio_file):
|
||||
os.remove(audio_file)
|
||||
|
||||
_ONE_DAY_IN_SECONDS = 60 * 60 * 24
|
||||
|
||||
# If MAX_WORKERS are specified in the environment use it, otherwise default to 1
|
||||
@@ -226,19 +222,20 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
text_chunks = split_text_at_word_boundary(request.text, max_length=250)
|
||||
print(f"Splitting text into chunks of 250 characters: {len(text_chunks)}", file=sys.stderr)
|
||||
# Generate audio for each chunk
|
||||
temp_audio_files = []
|
||||
for i, chunk in enumerate(text_chunks):
|
||||
# Generate audio for this chunk
|
||||
wav = self.model.generate(chunk, **kwargs)
|
||||
|
||||
# Create temporary file for this chunk
|
||||
temp_file = tempfile.NamedTemporaryFile(delete=False, suffix='.wav')
|
||||
temp_file.close()
|
||||
ta.save(temp_file.name, wav, self.model.sr)
|
||||
temp_audio_files.append(temp_file.name)
|
||||
|
||||
# Merge all audio files
|
||||
merge_audio_files(temp_audio_files, request.dst, self.model.sr)
|
||||
with cleanup_paths() as temp_audio_files:
|
||||
for i, chunk in enumerate(text_chunks):
|
||||
# Generate audio for this chunk
|
||||
wav = self.model.generate(chunk, **kwargs)
|
||||
|
||||
# Register ownership before saving so a partial write is
|
||||
# removed too when generation or encoding fails.
|
||||
temp_file = tempfile.NamedTemporaryFile(delete=False, suffix='.wav')
|
||||
temp_file.close()
|
||||
temp_audio_files.append(temp_file.name)
|
||||
ta.save(temp_file.name, wav, self.model.sr)
|
||||
|
||||
# Merge all audio files
|
||||
merge_audio_files(temp_audio_files, request.dst, self.model.sr)
|
||||
else:
|
||||
# Generate audio using ChatterboxTTS for short text
|
||||
wav = self.model.generate(request.text, **kwargs)
|
||||
|
||||
@@ -0,0 +1,36 @@
|
||||
import base64
|
||||
import contextlib
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def materialize_base64(data, suffix=""):
|
||||
"""Materialize base64 data for a path-only library and always remove it."""
|
||||
descriptor, path = tempfile.mkstemp(prefix="localai-media-", suffix=suffix)
|
||||
try:
|
||||
with os.fdopen(descriptor, "wb") as output:
|
||||
descriptor = None
|
||||
output.write(base64.b64decode(data))
|
||||
yield path
|
||||
finally:
|
||||
if descriptor is not None:
|
||||
os.close(descriptor)
|
||||
try:
|
||||
os.remove(path)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def cleanup_paths():
|
||||
"""Collect temporary paths and remove them on success or failure."""
|
||||
paths = []
|
||||
try:
|
||||
yield paths
|
||||
finally:
|
||||
for path in paths:
|
||||
try:
|
||||
os.remove(path)
|
||||
except OSError:
|
||||
pass
|
||||
@@ -0,0 +1,41 @@
|
||||
import os
|
||||
import tempfile
|
||||
import unittest
|
||||
from unittest import mock
|
||||
|
||||
from temp_utils import cleanup_paths, materialize_base64
|
||||
|
||||
|
||||
class MaterializeBase64Test(unittest.TestCase):
|
||||
def test_removes_materialized_file_after_success(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
with mock.patch.object(tempfile, "tempdir", directory):
|
||||
with materialize_base64("aGVsbG8=", suffix=".data") as path:
|
||||
with open(path, "rb") as materialized:
|
||||
self.assertEqual(materialized.read(), b"hello")
|
||||
self.assertFalse(os.path.exists(path))
|
||||
|
||||
def test_removes_materialized_file_when_consumer_fails(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
with mock.patch.object(tempfile, "tempdir", directory):
|
||||
with self.assertRaisesRegex(RuntimeError, "decode failed"):
|
||||
with materialize_base64("aGVsbG8="):
|
||||
raise RuntimeError("decode failed")
|
||||
self.assertEqual(os.listdir(directory), [])
|
||||
|
||||
|
||||
class CleanupPathsTest(unittest.TestCase):
|
||||
def test_removes_every_registered_path_after_failure(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
paths = [os.path.join(directory, name) for name in ("one.wav", "two.wav")]
|
||||
with self.assertRaisesRegex(RuntimeError, "merge failed"):
|
||||
with cleanup_paths() as registered:
|
||||
for path in paths:
|
||||
open(path, "wb").close()
|
||||
registered.append(path)
|
||||
raise RuntimeError("merge failed")
|
||||
self.assertEqual(os.listdir(directory), [])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,24 +0,0 @@
|
||||
import array
|
||||
import sys
|
||||
import wave
|
||||
|
||||
|
||||
def write_pcm_wav(destination, samples, sampling_rate):
|
||||
"""Write normalized floating-point audio samples as mono 16-bit PCM."""
|
||||
pcm = array.array(
|
||||
"h",
|
||||
(
|
||||
max(-32768, min(32767, round(float(sample) * 32768)))
|
||||
for sample in samples
|
||||
),
|
||||
)
|
||||
if pcm.itemsize != 2:
|
||||
raise RuntimeError("16-bit PCM requires two-byte signed integers")
|
||||
if sys.byteorder != "little":
|
||||
pcm.byteswap()
|
||||
|
||||
with wave.open(destination, "wb") as output:
|
||||
output.setnchannels(1)
|
||||
output.setsampwidth(2)
|
||||
output.setframerate(sampling_rate)
|
||||
output.writeframes(pcm.tobytes())
|
||||
@@ -26,7 +26,6 @@ sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..', 'common'))
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), 'common'))
|
||||
from grpc_auth import get_auth_interceptors
|
||||
from model_utils import resolve_model_reference
|
||||
from audio_utils import write_pcm_wav
|
||||
|
||||
|
||||
# Import dynamic loader for pipeline discovery
|
||||
@@ -900,49 +899,6 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
|
||||
return backend_pb2.Result(message="Media generated", success=True)
|
||||
|
||||
def SoundGeneration(self, request, context):
|
||||
if not request.dst:
|
||||
return backend_pb2.Result(success=False, message="request.dst is required")
|
||||
|
||||
prompt = request.text or request.caption
|
||||
if not prompt:
|
||||
return backend_pb2.Result(success=False, message="request.text is required")
|
||||
|
||||
try:
|
||||
generation_options = dict(self.options)
|
||||
if "num_inference_steps" in generation_options:
|
||||
generation_options["num_inference_steps"] = int(
|
||||
generation_options["num_inference_steps"]
|
||||
)
|
||||
generation_options["prompt"] = prompt
|
||||
if request.HasField("duration"):
|
||||
generation_options["audio_length_in_s"] = request.duration
|
||||
if request.HasField("temperature"):
|
||||
generation_options["guidance_scale"] = request.temperature
|
||||
|
||||
generated = self.pipe(**generation_options)
|
||||
if not hasattr(generated, "audios") or len(generated.audios) == 0:
|
||||
return backend_pb2.Result(
|
||||
success=False,
|
||||
message="The diffusers pipeline returned no audio",
|
||||
)
|
||||
|
||||
samples = generated.audios[0]
|
||||
if hasattr(samples, "reshape"):
|
||||
samples = samples.reshape(-1)
|
||||
if hasattr(samples, "tolist"):
|
||||
samples = samples.tolist()
|
||||
|
||||
sampling_rate = getattr(
|
||||
getattr(getattr(self.pipe, "vae", None), "config", None),
|
||||
"sampling_rate",
|
||||
16000,
|
||||
)
|
||||
write_pcm_wav(request.dst, samples, sampling_rate)
|
||||
return backend_pb2.Result(success=True, message="Sound generated successfully")
|
||||
except Exception as err:
|
||||
return backend_pb2.Result(success=False, message=f"SoundGeneration error: {err}")
|
||||
|
||||
def UpscaleImage(self, request, context):
|
||||
try:
|
||||
if not request.src:
|
||||
|
||||
@@ -4,9 +4,6 @@ A test script to test the gRPC service and dynamic loader
|
||||
import unittest
|
||||
import subprocess
|
||||
import time
|
||||
import os
|
||||
import tempfile
|
||||
import wave
|
||||
from unittest.mock import patch, MagicMock
|
||||
|
||||
# Import dynamic loader for testing (these don't need gRPC)
|
||||
@@ -448,64 +445,3 @@ class TestDeviceSelection(unittest.TestCase):
|
||||
|
||||
def test_mps_overrides(self):
|
||||
self.assertEqual(backend.select_device(False, None, True, False, True), "mps")
|
||||
|
||||
|
||||
class TestWritePcmWav(unittest.TestCase):
|
||||
def test_writes_clipped_float_samples_as_mono_pcm(self):
|
||||
from audio_utils import write_pcm_wav
|
||||
|
||||
destination = tempfile.NamedTemporaryFile(suffix=".wav", delete=False)
|
||||
destination.close()
|
||||
|
||||
try:
|
||||
write_pcm_wav(destination.name, [0.0, 0.5, -0.5, 2.0], 16000)
|
||||
|
||||
with wave.open(destination.name, "rb") as generated:
|
||||
self.assertEqual(generated.getframerate(), 16000)
|
||||
self.assertEqual(generated.getnchannels(), 1)
|
||||
self.assertEqual(generated.getsampwidth(), 2)
|
||||
self.assertEqual(generated.getnframes(), 4)
|
||||
self.assertEqual(
|
||||
generated.readframes(4),
|
||||
b"\x00\x00\x00@\x00\xc0\xff\x7f",
|
||||
)
|
||||
finally:
|
||||
os.unlink(destination.name)
|
||||
|
||||
|
||||
@unittest.skipUnless(GRPC_AVAILABLE, "gRPC modules not available")
|
||||
class TestSoundGeneration(unittest.TestCase):
|
||||
def test_maps_request_options_and_writes_pipeline_audio(self):
|
||||
from backend import BackendServicer
|
||||
|
||||
service = BackendServicer.__new__(BackendServicer)
|
||||
service.options = {"num_inference_steps": 200.0}
|
||||
service.pipe = MagicMock()
|
||||
service.pipe.return_value.audios = [[0.0, 0.5, -0.5]]
|
||||
service.pipe.vae.config.sampling_rate = 16000
|
||||
|
||||
destination = tempfile.NamedTemporaryFile(suffix=".wav", delete=False)
|
||||
destination.close()
|
||||
|
||||
try:
|
||||
request = backend_pb2.SoundGenerationRequest(
|
||||
text="ocean waves",
|
||||
dst=destination.name,
|
||||
duration=2.5,
|
||||
temperature=0,
|
||||
)
|
||||
|
||||
result = service.SoundGeneration(request, context=None)
|
||||
|
||||
self.assertTrue(result.success, result.message)
|
||||
service.pipe.assert_called_once_with(
|
||||
num_inference_steps=200,
|
||||
prompt="ocean waves",
|
||||
audio_length_in_s=2.5,
|
||||
guidance_scale=0,
|
||||
)
|
||||
with wave.open(destination.name, "rb") as generated:
|
||||
self.assertEqual(generated.getframerate(), 16000)
|
||||
self.assertEqual(generated.getnframes(), 3)
|
||||
finally:
|
||||
os.unlink(destination.name)
|
||||
@@ -6,6 +6,7 @@ import datetime
|
||||
import gc
|
||||
import math
|
||||
import os
|
||||
import shutil
|
||||
import signal
|
||||
import subprocess
|
||||
import sys
|
||||
@@ -888,6 +889,13 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
def _release_model(self):
|
||||
self.pipeline = None
|
||||
self.model_kind = None
|
||||
try:
|
||||
if hasattr(self, "dist") and self.dist.is_initialized():
|
||||
self.dist.destroy_process_group()
|
||||
finally:
|
||||
if self._dist_store_dir is not None:
|
||||
shutil.rmtree(self._dist_store_dir, ignore_errors=True)
|
||||
self._dist_store_dir = None
|
||||
gc.collect()
|
||||
if hasattr(self, "torch") and self.torch.cuda.is_available():
|
||||
self.torch.cuda.empty_cache()
|
||||
|
||||
@@ -19,7 +19,6 @@ import base64
|
||||
import io
|
||||
import json
|
||||
import gc
|
||||
import tempfile
|
||||
|
||||
from PIL import Image
|
||||
import torch
|
||||
@@ -34,6 +33,7 @@ sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..', 'common'))
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), 'common'))
|
||||
from grpc_auth import get_auth_interceptors
|
||||
from model_utils import resolve_model_reference
|
||||
from temp_utils import materialize_base64
|
||||
from vllm_utils import parse_options, messages_to_dicts, setup_parsers
|
||||
|
||||
|
||||
@@ -118,13 +118,8 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
return video_to_ndarrays(video_path, num_frames=16)
|
||||
# Try base64 decode
|
||||
try:
|
||||
timestamp = str(int(time.time() * 1000))
|
||||
p = os.path.join(tempfile.gettempdir(), f"vl-{timestamp}.data")
|
||||
with open(p, "wb") as f:
|
||||
f.write(base64.b64decode(video_path))
|
||||
video = VideoAsset(name=p).np_ndarrays
|
||||
os.remove(p)
|
||||
return video
|
||||
with materialize_base64(video_path, suffix=".data") as path:
|
||||
return VideoAsset(name=path).np_ndarrays
|
||||
except:
|
||||
return None
|
||||
|
||||
@@ -136,15 +131,9 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
return (audio_signal.astype(np.float32), sr)
|
||||
# Try base64 decode
|
||||
try:
|
||||
audio_data = base64.b64decode(audio_path)
|
||||
# Save to temp file and load
|
||||
timestamp = str(int(time.time() * 1000))
|
||||
p = os.path.join(tempfile.gettempdir(), f"audio-{timestamp}.wav")
|
||||
with open(p, "wb") as f:
|
||||
f.write(audio_data)
|
||||
audio_signal, sr = librosa.load(p, sr=16000)
|
||||
os.remove(p)
|
||||
return (audio_signal.astype(np.float32), sr)
|
||||
with materialize_base64(audio_path, suffix=".wav") as path:
|
||||
audio_signal, sr = librosa.load(path, sr=16000)
|
||||
return (audio_signal.astype(np.float32), sr)
|
||||
except:
|
||||
return None
|
||||
|
||||
|
||||
@@ -10,7 +10,6 @@ import os
|
||||
import json
|
||||
import time
|
||||
import gc
|
||||
import tempfile
|
||||
from typing import List
|
||||
from PIL import Image
|
||||
|
||||
@@ -23,6 +22,7 @@ sys.path.insert(0, os.path.join(os.path.dirname(__file__), 'common'))
|
||||
from python_utils import attach_media_parts
|
||||
from grpc_auth import get_auth_interceptors
|
||||
from model_utils import resolve_model_reference
|
||||
from temp_utils import materialize_base64
|
||||
from vllm_utils import apply_options_to_engine_args, normalize_option_key
|
||||
|
||||
from vllm.engine.arg_utils import AsyncEngineArgs
|
||||
@@ -1005,13 +1005,8 @@ class BackendServicer(backend_pb2_grpc.BackendServicer):
|
||||
Video: The loaded video.
|
||||
"""
|
||||
try:
|
||||
timestamp = str(int(time.time() * 1000)) # Generate timestamp
|
||||
p = os.path.join(tempfile.gettempdir(), f"vl-{timestamp}.data")
|
||||
with open(p, "wb") as f:
|
||||
f.write(base64.b64decode(video_path))
|
||||
video = VideoAsset(name=p).np_ndarrays
|
||||
os.remove(p)
|
||||
return video
|
||||
with materialize_base64(video_path, suffix=".data") as path:
|
||||
return VideoAsset(name=path).np_ndarrays
|
||||
except Exception as e:
|
||||
print(f"Error loading video {video_path}: {e}", file=sys.stderr)
|
||||
return None
|
||||
|
||||
@@ -369,10 +369,10 @@ var BackendCapabilities = map[string]BackendCapability{
|
||||
|
||||
// --- Image/video generation backends ---
|
||||
"diffusers": {
|
||||
GRPCMethods: []GRPCMethod{MethodGenerateImage, MethodUpscaleImage, MethodGenerateVideo, MethodSoundGeneration},
|
||||
PossibleUsecases: []string{UsecaseImage, UsecaseVideo, UsecaseSoundGeneration},
|
||||
GRPCMethods: []GRPCMethod{MethodGenerateImage, MethodUpscaleImage, MethodGenerateVideo},
|
||||
PossibleUsecases: []string{UsecaseImage, UsecaseVideo},
|
||||
DefaultUsecases: []string{UsecaseImage},
|
||||
Description: "HuggingFace diffusers — image, video, and sound generation",
|
||||
Description: "HuggingFace diffusers — Stable Diffusion, Flux, video generation",
|
||||
},
|
||||
"longcat-video": {
|
||||
GRPCMethods: []GRPCMethod{MethodGenerateVideo},
|
||||
|
||||
@@ -57,12 +57,6 @@ var _ = Describe("BackendCapabilities", func() {
|
||||
})
|
||||
|
||||
var _ = Describe("GetBackendCapability", func() {
|
||||
It("advertises diffusers sound generation", func() {
|
||||
capability := GetBackendCapability("diffusers")
|
||||
Expect(capability.GRPCMethods).To(ContainElement(MethodSoundGeneration))
|
||||
Expect(capability.PossibleUsecases).To(ContainElement(UsecaseSoundGeneration))
|
||||
})
|
||||
|
||||
It("returns the capability for a known backend", func() {
|
||||
cap := GetBackendCapability("llama-cpp")
|
||||
Expect(cap).NotTo(BeNil())
|
||||
|
||||
@@ -99,3 +99,12 @@ var DiffusersSchedulerOptions = []FieldOption{
|
||||
{Value: "heun", Label: "Heun"},
|
||||
{Value: "unipc", Label: "UniPC"},
|
||||
}
|
||||
|
||||
// SystemMessagesAfterFirstOptions are the values of template.system_messages_after_first:
|
||||
// how system messages that appear after the first turn are handled before the chat
|
||||
// template runs (empty = pass through unchanged, which strict Jinja templates reject).
|
||||
var SystemMessagesAfterFirstOptions = []FieldOption{
|
||||
{Value: "", Label: "Pass through (default)"},
|
||||
{Value: "merge", Label: "Merge into the first system message"},
|
||||
{Value: "user", Label: "Forward as user messages"},
|
||||
}
|
||||
@@ -382,6 +382,14 @@ func DefaultRegistry() map[string]FieldMetaOverride {
|
||||
Description: "Use the chat template from the model's tokenizer config",
|
||||
Order: 44,
|
||||
},
|
||||
"template.system_messages_after_first": {
|
||||
Section: "templates",
|
||||
Label: "System Messages After First",
|
||||
Description: "How system messages that appear after the first turn are handled before templating: merge into the first system message, or forward as user messages. Empty passes them through unchanged, which strict Jinja templates reject.",
|
||||
Component: "select",
|
||||
Options: SystemMessagesAfterFirstOptions,
|
||||
Order: 45,
|
||||
},
|
||||
// Router section template — kept in the templates UI section
|
||||
// (rather than the router section under "other") so operators
|
||||
// editing prompt shapes find all template-typed fields in one
|
||||
|
||||
@@ -1351,6 +1351,16 @@ type TemplateConfig struct {
|
||||
// that can use the tokenizers specified in the JSON config files of the models
|
||||
UseTokenizerTemplate bool `yaml:"use_tokenizer_template,omitempty" json:"use_tokenizer_template,omitempty"`
|
||||
|
||||
// SystemMessagesAfterFirst controls what happens to system-role messages that
|
||||
// appear after the leading system block. Some tokenizer chat templates (e.g.
|
||||
// Qwen3.8 / Flash-Next) raise "System message must be at the beginning" for
|
||||
// them, while agent frameworks (cogito tool selection, adjustment prompts)
|
||||
// legitimately append system instructions mid-conversation.
|
||||
// ""/"error": pass through unchanged (template decides)
|
||||
// "merge": fold them into the leading system message
|
||||
// "user": forward them as user-role instructions (keeps their position)
|
||||
SystemMessagesAfterFirst string `yaml:"system_messages_after_first,omitempty" json:"system_messages_after_first,omitempty"`
|
||||
|
||||
// JoinChatMessagesByCharacter is a string that will be used to join chat messages together.
|
||||
// It defaults to \n
|
||||
JoinChatMessagesByCharacter *string `yaml:"join_chat_messages_by_character,omitempty" json:"join_chat_messages_by_character,omitempty"`
|
||||
|
||||
@@ -66,6 +66,65 @@ func stripEmptySystemMessages(messages []schema.Message) []schema.Message {
|
||||
return out
|
||||
}
|
||||
|
||||
// normalizeLateSystemMessages handles system-role messages that appear after the
|
||||
// leading system block, according to template.system_messages_after_first:
|
||||
// "merge" folds them into the first system message (created if absent), "user"
|
||||
// forwards them as user-role turns at their original position. Any other value
|
||||
// returns the messages unchanged. Needed for tokenizer templates that reject
|
||||
// late system turns (Qwen3.8: "System message must be at the beginning") while
|
||||
// agent frameworks append instructions mid-conversation.
|
||||
func normalizeLateSystemMessages(messages []schema.Message, mode string) []schema.Message {
|
||||
if mode != "merge" && mode != "user" {
|
||||
return messages
|
||||
}
|
||||
lead := 0
|
||||
for lead < len(messages) && messages[lead].Role == "system" {
|
||||
lead++
|
||||
}
|
||||
late := false
|
||||
for _, m := range messages[lead:] {
|
||||
if m.Role == "system" {
|
||||
late = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !late {
|
||||
return messages
|
||||
}
|
||||
out := make([]schema.Message, 0, len(messages)+1)
|
||||
out = append(out, messages[:lead]...)
|
||||
if mode == "merge" && lead == 0 {
|
||||
out = append(out, schema.Message{Role: "system"})
|
||||
}
|
||||
for _, m := range messages[lead:] {
|
||||
if m.Role != "system" {
|
||||
out = append(out, m)
|
||||
continue
|
||||
}
|
||||
text := strings.TrimSpace(messageText(m))
|
||||
if text == "" {
|
||||
continue
|
||||
}
|
||||
switch mode {
|
||||
case "merge":
|
||||
first := &out[0]
|
||||
joined := strings.TrimSpace(messageText(*first))
|
||||
if joined != "" {
|
||||
joined += "\n\n"
|
||||
}
|
||||
joined += text
|
||||
first.Content = joined
|
||||
first.StringContent = joined
|
||||
case "user":
|
||||
m.Role = "user"
|
||||
m.Content = text
|
||||
m.StringContent = text
|
||||
out = append(out, m)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// mergeToolCallDeltas merges streaming tool call deltas into complete tool calls.
|
||||
// In SSE streaming, a single tool call arrives as multiple chunks sharing the same Index:
|
||||
// the first chunk carries the ID, Type, and Name; subsequent chunks append to Arguments.
|
||||
@@ -182,6 +241,7 @@ func ChatEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator
|
||||
// Drop blank system turns from the web UI (and similar clients) so they
|
||||
// cannot suppress the model YAML system_prompt / tokenizer defaults.
|
||||
input.Messages = stripEmptySystemMessages(input.Messages)
|
||||
input.Messages = normalizeLateSystemMessages(input.Messages, config.TemplateConfig.SystemMessagesAfterFirst)
|
||||
|
||||
// Tokenizer-template models pass messages through to the backend as-is,
|
||||
// so apply the configured system_prompt when the request did not supply
|
||||
|
||||
@@ -378,6 +378,50 @@ var _ = Describe("system message helpers", func() {
|
||||
})
|
||||
})
|
||||
|
||||
Describe("normalizeLateSystemMessages", func() {
|
||||
msgs := func() []schema.Message {
|
||||
return []schema.Message{
|
||||
{Role: "system", Content: "lead", StringContent: "lead"},
|
||||
{Role: "user", Content: "q", StringContent: "q"},
|
||||
{Role: "assistant", Content: "a", StringContent: "a"},
|
||||
{Role: "system", Content: "late", StringContent: "late"},
|
||||
{Role: "user", Content: "q2", StringContent: "q2"},
|
||||
}
|
||||
}
|
||||
It("leaves messages untouched by default", func() {
|
||||
out := normalizeLateSystemMessages(msgs(), "")
|
||||
Expect(out).To(HaveLen(5))
|
||||
Expect(out[3].Role).To(Equal("system"))
|
||||
})
|
||||
It("merge folds late system turns into the leading one", func() {
|
||||
out := normalizeLateSystemMessages(msgs(), "merge")
|
||||
Expect(out).To(HaveLen(4))
|
||||
Expect(out[0].Role).To(Equal("system"))
|
||||
Expect(out[0].StringContent).To(Equal("lead\n\nlate"))
|
||||
for _, m := range out[1:] {
|
||||
Expect(m.Role).NotTo(Equal("system"))
|
||||
}
|
||||
})
|
||||
It("merge creates a leading system message when none exists", func() {
|
||||
in := msgs()[1:]
|
||||
out := normalizeLateSystemMessages(in, "merge")
|
||||
Expect(out[0].Role).To(Equal("system"))
|
||||
Expect(out[0].StringContent).To(Equal("late"))
|
||||
Expect(out).To(HaveLen(4))
|
||||
})
|
||||
It("user forwards late system turns as user turns in place", func() {
|
||||
out := normalizeLateSystemMessages(msgs(), "user")
|
||||
Expect(out).To(HaveLen(5))
|
||||
Expect(out[3].Role).To(Equal("user"))
|
||||
Expect(out[3].StringContent).To(Equal("late"))
|
||||
Expect(out[0].Role).To(Equal("system"))
|
||||
})
|
||||
It("does nothing when no late system turn exists", func() {
|
||||
in := msgs()[:3]
|
||||
Expect(normalizeLateSystemMessages(in, "user")).To(HaveLen(3))
|
||||
})
|
||||
})
|
||||
|
||||
Describe("stripEmptySystemMessages", func() {
|
||||
It("removes blank system turns and keeps the rest", func() {
|
||||
in := []schema.Message{
|
||||
|
||||
@@ -9,9 +9,11 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"maps"
|
||||
"math"
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
@@ -31,6 +33,7 @@ import (
|
||||
"github.com/mudler/LocalAI/core/http/endpoints/openai/types"
|
||||
"github.com/mudler/LocalAI/core/schema"
|
||||
"github.com/mudler/LocalAI/core/services/routing/router"
|
||||
"github.com/mudler/LocalAI/core/services/voiceprofile"
|
||||
"github.com/mudler/LocalAI/core/templates"
|
||||
laudio "github.com/mudler/LocalAI/pkg/audio"
|
||||
"github.com/mudler/LocalAI/pkg/functions"
|
||||
@@ -135,6 +138,8 @@ type Session struct {
|
||||
Instructions string
|
||||
DefaultConversationID string
|
||||
ModelInterface Model
|
||||
ttsParams map[string]string
|
||||
voiceRelease func()
|
||||
// The pipeline model config or the config for an any-to-any model
|
||||
ModelConfig *config.ModelConfig
|
||||
InputSampleRate int
|
||||
@@ -198,6 +203,22 @@ type Session struct {
|
||||
respSink *responseSink
|
||||
}
|
||||
|
||||
func (s *Session) installVoiceBinding(voice string, params map[string]string, release func()) {
|
||||
if release == nil {
|
||||
release = func() {}
|
||||
}
|
||||
var once sync.Once
|
||||
s.Voice = voice
|
||||
s.ttsParams = maps.Clone(params)
|
||||
s.voiceRelease = func() { once.Do(release) }
|
||||
}
|
||||
|
||||
func (s *Session) releaseVoiceBinding() {
|
||||
if s.voiceRelease != nil {
|
||||
s.voiceRelease()
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Session) FromClient(session *types.SessionUnion) {
|
||||
}
|
||||
|
||||
@@ -633,6 +654,17 @@ func runRealtimeSession(application *application.Application, t Transport, model
|
||||
sendError(t, "model_load_error", "Failed to load model", "", "")
|
||||
return
|
||||
}
|
||||
if wrapped, ok := m.(*wrappedModel); ok {
|
||||
resolvedVoice, params, release, resolveErr := resolveRealtimeVoice(context.Background(), wrapped.TTSConfig.TTSConfig.Voice, wrapped.TTSConfig, application.VoiceProfileStore())
|
||||
if resolveErr != nil {
|
||||
xlog.Error("failed to resolve realtime voice", "error", resolveErr)
|
||||
sendError(t, "voice_profile_error", resolveErr.Error(), "", "")
|
||||
return
|
||||
}
|
||||
session.installVoiceBinding(resolvedVoice, params, release)
|
||||
defer session.releaseVoiceBinding()
|
||||
wrapped.setTTSParams(params)
|
||||
}
|
||||
session.ModelInterface = m
|
||||
// A pipeline-seeded option list gets its scoring prompt prewarmed
|
||||
// alongside the model warm-up below, so the session's first turn
|
||||
@@ -826,6 +858,7 @@ func runRealtimeSession(application *application.Application, t Transport, model
|
||||
application.ApplicationConfig(),
|
||||
evaluator,
|
||||
buildRealtimeRoutingContext(application, session.ID),
|
||||
application.VoiceProfileStore(),
|
||||
); err != nil {
|
||||
xlog.Error("failed to update session", "error", err)
|
||||
sendError(t, "session_update_error", fmt.Sprintf("Failed to update session: %v", err), "", "")
|
||||
@@ -1164,7 +1197,7 @@ func updateTransSession(session *Session, update *types.SessionUnion, cl *config
|
||||
return nil
|
||||
}
|
||||
|
||||
func updateSession(session *Session, update *types.SessionUnion, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig, evaluator *templates.Evaluator, routing *RealtimeRoutingContext) error {
|
||||
func updateSession(session *Session, update *types.SessionUnion, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig, evaluator *templates.Evaluator, routing *RealtimeRoutingContext, profiles *voiceprofile.Store) error {
|
||||
sessionLock.Lock()
|
||||
defer sessionLock.Unlock()
|
||||
|
||||
@@ -1172,8 +1205,13 @@ func updateSession(session *Session, update *types.SessionUnion, cl *config.Mode
|
||||
return nil
|
||||
}
|
||||
|
||||
session.TranscriptionOnly = false
|
||||
rt := update.Realtime
|
||||
explicitVoice := rt.Audio != nil && rt.Audio.Output != nil && rt.Audio.Output.Voice != ""
|
||||
rebuild := rt.Model != "" || explicitVoice || (rt.Audio != nil && rt.Audio.Input != nil && rt.Audio.Input.Transcription != nil)
|
||||
|
||||
candidateModelName := session.Model
|
||||
candidateConfig := session.ModelConfig
|
||||
candidateTranscription := session.InputAudioTranscription
|
||||
|
||||
if rt.Model != "" {
|
||||
cfg, err := cl.LoadModelConfigFileByNameDefaultOptions(rt.Model, appConfig)
|
||||
@@ -1184,40 +1222,78 @@ func updateSession(session *Session, update *types.SessionUnion, cl *config.Mode
|
||||
return fmt.Errorf("model is not a valid pipeline model: %s", rt.Model)
|
||||
}
|
||||
|
||||
if session.InputAudioTranscription == nil {
|
||||
session.InputAudioTranscription = &types.AudioTranscription{}
|
||||
}
|
||||
session.InputAudioTranscription.Model = cfg.Pipeline.Transcription
|
||||
session.Voice = cfg.TTSConfig.Voice
|
||||
session.Model = rt.Model
|
||||
session.ModelConfig = cfg
|
||||
}
|
||||
|
||||
if rt.Audio != nil && rt.Audio.Output != nil && rt.Audio.Output.Voice != "" {
|
||||
session.Voice = string(rt.Audio.Output.Voice)
|
||||
candidateModelName = rt.Model
|
||||
candidateConfig = cfg
|
||||
candidateTranscription = &types.AudioTranscription{Model: cfg.Pipeline.Transcription}
|
||||
}
|
||||
|
||||
if rt.Audio != nil && rt.Audio.Input != nil && rt.Audio.Input.Transcription != nil {
|
||||
trUpd := rt.Audio.Input.Transcription
|
||||
trUpd := *rt.Audio.Input.Transcription
|
||||
// A language-only update (e.g. a client forcing the STT language) carries
|
||||
// an empty Model. Preserve the pipeline's configured transcription backend
|
||||
// instead of blanking it — otherwise the next utterance transcribes against
|
||||
// an empty model and the backend RPC fails with "unimplemented".
|
||||
if trUpd.Model == "" && session.InputAudioTranscription != nil {
|
||||
trUpd.Model = session.InputAudioTranscription.Model
|
||||
if trUpd.Model == "" && candidateTranscription != nil {
|
||||
trUpd.Model = candidateTranscription.Model
|
||||
}
|
||||
session.InputAudioTranscription = trUpd
|
||||
candidateTranscription = &trUpd
|
||||
if trUpd.Model != "" {
|
||||
session.ModelConfig.Pipeline.Transcription = trUpd.Model
|
||||
cfgCopy := *candidateConfig
|
||||
candidateConfig = &cfgCopy
|
||||
candidateConfig.Pipeline.Transcription = trUpd.Model
|
||||
}
|
||||
}
|
||||
|
||||
if rt.Model != "" || (rt.Audio != nil && rt.Audio.Output != nil && rt.Audio.Output.Voice != "") || (rt.Audio != nil && rt.Audio.Input != nil && rt.Audio.Input.Transcription != nil) {
|
||||
m, err := newModel(&session.ModelConfig.Pipeline, cl, ml, appConfig, evaluator, routing)
|
||||
candidateModel := session.ModelInterface
|
||||
candidateVoice := session.Voice
|
||||
candidateParams := maps.Clone(session.ttsParams)
|
||||
var candidateRelease func()
|
||||
selectVoice := rt.Model != "" || explicitVoice
|
||||
if rebuild {
|
||||
m, err := newModel(&candidateConfig.Pipeline, cl, ml, appConfig, evaluator, routing)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
session.ModelInterface = m
|
||||
candidateModel = m
|
||||
wrapped := m.(*wrappedModel)
|
||||
if selectVoice {
|
||||
configuredVoice := wrapped.TTSConfig.TTSConfig.Voice
|
||||
if explicitVoice {
|
||||
configuredVoice = string(rt.Audio.Output.Voice)
|
||||
}
|
||||
candidateVoice, candidateParams, candidateRelease, err = resolveRealtimeVoice(
|
||||
context.Background(), configuredVoice, wrapped.TTSConfig, profiles,
|
||||
)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
wrapped.setTTSParams(candidateParams)
|
||||
}
|
||||
|
||||
if rt.LocalAIClassifier != nil {
|
||||
if err := validateClassifierActivation(candidateModel, rt.LocalAIClassifier); err != nil {
|
||||
if candidateRelease != nil {
|
||||
candidateRelease()
|
||||
}
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
oldRelease := session.voiceRelease
|
||||
session.TranscriptionOnly = false
|
||||
session.Model = candidateModelName
|
||||
session.ModelConfig = candidateConfig
|
||||
session.ModelInterface = candidateModel
|
||||
session.InputAudioTranscription = candidateTranscription
|
||||
if selectVoice {
|
||||
session.installVoiceBinding(candidateVoice, candidateParams, candidateRelease)
|
||||
if oldRelease != nil {
|
||||
oldRelease()
|
||||
}
|
||||
}
|
||||
|
||||
if rebuild {
|
||||
// A session.update that swaps the model/voice rebuilds the pipeline, so
|
||||
// warm the new backends too (unless opted out) — otherwise the next turn
|
||||
// pays the cold-start load the original session warm-up already avoided.
|
||||
@@ -1226,9 +1302,9 @@ func updateSession(session *Session, update *types.SessionUnion, cl *config.Mode
|
||||
// stall every other session. Load errors are logged (and still surface on
|
||||
// first use); per-stage failures are already warned inside
|
||||
// backend.PreloadStages.
|
||||
if !session.ModelConfig.Pipeline.DisableWarmup {
|
||||
if !candidateConfig.Pipeline.DisableWarmup {
|
||||
go func() {
|
||||
if err := m.Warmup(context.Background()); err != nil {
|
||||
if err := candidateModel.Warmup(context.Background()); err != nil {
|
||||
xlog.Error("realtime warmup failed after session.update", "error", err)
|
||||
}
|
||||
}()
|
||||
@@ -1287,9 +1363,6 @@ func updateSession(session *Session, update *types.SessionUnion, cl *config.Mode
|
||||
// Replace-not-merge, like tools: the client owns the whole option
|
||||
// list. Invalid configs reject the update without touching the
|
||||
// session's current classifier.
|
||||
if err := validateClassifierActivation(session.ModelInterface, rt.LocalAIClassifier); err != nil {
|
||||
return err
|
||||
}
|
||||
session.Classifier = rt.LocalAIClassifier
|
||||
prewarmClassifier(session)
|
||||
}
|
||||
@@ -1923,7 +1996,7 @@ func commitUtteranceWithTranscript(ctx context.Context, utt []byte, live *liveUt
|
||||
// Generate an LLM response only when there is a transcript to feed it. A
|
||||
// sound-detection-only session (no transcription) has no LLM stage, so it
|
||||
// stops here after emitting the sound-detection event.
|
||||
if session.InputAudioTranscription != nil && !session.TranscriptionOnly {
|
||||
if session.InputAudioTranscription != nil && !session.TranscriptionOnly && strings.TrimSpace(transcript) != "" {
|
||||
generateResponse(ctx, session, utt, transcript, speaker, conv, t)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -6,7 +6,9 @@ import (
|
||||
"encoding/binary"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"maps"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -18,6 +20,7 @@ import (
|
||||
"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/core/services/voiceprofile"
|
||||
"github.com/mudler/LocalAI/core/templates"
|
||||
"github.com/mudler/LocalAI/pkg/functions"
|
||||
"github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
@@ -35,6 +38,7 @@ var (
|
||||
// which are for Any-To-Any models, but instead we will call a pipeline (for e.g STT->LLM->TTS)
|
||||
type wrappedModel struct {
|
||||
TTSConfig *config.ModelConfig
|
||||
ttsParams map[string]string
|
||||
TranscriptionConfig *config.ModelConfig
|
||||
LLMConfig *config.ModelConfig
|
||||
VADConfig *config.ModelConfig
|
||||
@@ -391,11 +395,39 @@ func newRealtimeDecisionID() string {
|
||||
}
|
||||
|
||||
func (m *wrappedModel) TTS(ctx context.Context, text, voice, language string) (string, *proto.Result, error) {
|
||||
return backend.ModelTTS(ctx, text, voice, language, "", nil, m.modelLoader, m.appConfig, *m.TTSConfig)
|
||||
return backend.ModelTTS(ctx, text, voice, language, "", maps.Clone(m.ttsParams), m.modelLoader, m.appConfig, *m.TTSConfig)
|
||||
}
|
||||
|
||||
func (m *wrappedModel) setTTSParams(params map[string]string) {
|
||||
m.ttsParams = maps.Clone(params)
|
||||
}
|
||||
|
||||
func (m *wrappedModel) TTSStream(ctx context.Context, text, voice, language string, onAudio func(pcm []byte, sampleRate int) error) error {
|
||||
return ttsStream(ctx, m.modelLoader, m.appConfig, *m.TTSConfig, text, voice, language, onAudio)
|
||||
return ttsStream(ctx, m.modelLoader, m.appConfig, *m.TTSConfig, text, voice, language, maps.Clone(m.ttsParams), onAudio)
|
||||
}
|
||||
|
||||
func resolveRealtimeVoice(ctx context.Context, configuredVoice string, ttsConfig *config.ModelConfig, profiles *voiceprofile.Store) (string, map[string]string, func(), error) {
|
||||
if !voiceprofile.IsReference(configuredVoice) {
|
||||
return configuredVoice, nil, func() {}, nil
|
||||
}
|
||||
profileID, valid := voiceprofile.ParseReference(configuredVoice)
|
||||
if !valid {
|
||||
return "", nil, nil, fmt.Errorf("invalid voice profile reference %q", configuredVoice)
|
||||
}
|
||||
if config.VoiceCloningForModel(ttsConfig) == nil {
|
||||
return "", nil, nil, fmt.Errorf("selected TTS model does not support reference-audio voice cloning")
|
||||
}
|
||||
if profiles == nil {
|
||||
return "", nil, nil, fmt.Errorf("voice profile store is unavailable")
|
||||
}
|
||||
profile, referencePath, release, err := profiles.LeaseAudio(ctx, profileID)
|
||||
if err != nil {
|
||||
if errors.Is(err, voiceprofile.ErrNotFound) {
|
||||
return "", nil, nil, fmt.Errorf("voice profile not found: %w", err)
|
||||
}
|
||||
return "", nil, nil, fmt.Errorf("resolve voice profile: %w", err)
|
||||
}
|
||||
return referencePath, map[string]string{"ref_text": profile.Transcript}, release, nil
|
||||
}
|
||||
|
||||
func (m *wrappedModel) TranscribeStream(ctx context.Context, audio, language string, translate, diarize bool, prompt string, onDelta func(text string)) (*schema.TranscriptionResult, error) {
|
||||
@@ -674,11 +706,11 @@ const wavStreamHeaderBytes = 44
|
||||
// callback, which wants raw PCM plus the sample rate. The header is buffered
|
||||
// until complete, the sample rate is read from it, and subsequent bytes are
|
||||
// forwarded as PCM.
|
||||
func ttsStream(ctx context.Context, ml *model.ModelLoader, appConfig *config.ApplicationConfig, ttsConfig config.ModelConfig, text, voice, language string, onAudio func(pcm []byte, sampleRate int) error) error {
|
||||
func ttsStream(ctx context.Context, ml *model.ModelLoader, appConfig *config.ApplicationConfig, ttsConfig config.ModelConfig, text, voice, language string, params map[string]string, onAudio func(pcm []byte, sampleRate int) error) error {
|
||||
var header []byte
|
||||
headerDone := false
|
||||
sampleRate := 0
|
||||
return backend.ModelTTSStream(ctx, text, voice, language, "", nil, ml, appConfig, ttsConfig, func(b []byte) error {
|
||||
return backend.ModelTTSStream(ctx, text, voice, language, "", params, ml, appConfig, ttsConfig, func(b []byte) error {
|
||||
if headerDone {
|
||||
if len(b) == 0 {
|
||||
return nil
|
||||
|
||||
@@ -355,6 +355,19 @@ var _ = Describe("commitUtteranceWithTranscript", func() {
|
||||
|
||||
Expect(tr.countEvents(types.ServerEventTypeConversationItemInputAudioTranscriptionCompleted)).To(Equal(1))
|
||||
})
|
||||
|
||||
It("does not generate a response for a blank transcript", func() {
|
||||
session, model := itSession(nil)
|
||||
model.transcribeFinal = &schema.TranscriptionResult{Text: " \t\n"}
|
||||
tr := &fakeTransport{}
|
||||
conv := &Conversation{}
|
||||
|
||||
commitUtterance(context.Background(), []byte{1, 2}, session, conv, tr)
|
||||
|
||||
Expect(tr.countEvents(types.ServerEventTypeConversationItemInputAudioTranscriptionCompleted)).To(Equal(1))
|
||||
Expect(conv.Items).To(BeEmpty())
|
||||
Expect(tr.countEvents(types.ServerEventTypeResponseCreated)).To(Equal(0))
|
||||
})
|
||||
})
|
||||
|
||||
// transcribeUtterance is the retranscribe gate's offline decode of the
|
||||
|
||||
@@ -0,0 +1,354 @@
|
||||
package openai
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
"github.com/mudler/LocalAI/core/http/endpoints/openai/types"
|
||||
grpcPkg "github.com/mudler/LocalAI/pkg/grpc"
|
||||
"github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
"github.com/mudler/LocalAI/pkg/model"
|
||||
"github.com/mudler/LocalAI/pkg/system"
|
||||
"google.golang.org/grpc"
|
||||
|
||||
"github.com/mudler/LocalAI/core/config"
|
||||
"github.com/mudler/LocalAI/core/services/voiceprofile"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
func realtimeProfileWAV(duration time.Duration) []byte {
|
||||
const (
|
||||
sampleRate = 16000
|
||||
channels = 1
|
||||
bitsPerSample = 16
|
||||
)
|
||||
dataSize := int(duration.Seconds() * sampleRate * channels * bitsPerSample / 8)
|
||||
buf := bytes.NewBuffer(nil)
|
||||
buf.WriteString("RIFF")
|
||||
_ = binary.Write(buf, binary.LittleEndian, uint32(36+dataSize))
|
||||
buf.WriteString("WAVEfmt ")
|
||||
_ = binary.Write(buf, binary.LittleEndian, uint32(16))
|
||||
_ = binary.Write(buf, binary.LittleEndian, uint16(1))
|
||||
_ = binary.Write(buf, binary.LittleEndian, uint16(channels))
|
||||
_ = binary.Write(buf, binary.LittleEndian, uint32(sampleRate))
|
||||
_ = binary.Write(buf, binary.LittleEndian, uint32(sampleRate*channels*bitsPerSample/8))
|
||||
_ = binary.Write(buf, binary.LittleEndian, uint16(channels*bitsPerSample/8))
|
||||
_ = binary.Write(buf, binary.LittleEndian, uint16(bitsPerSample))
|
||||
buf.WriteString("data")
|
||||
_ = binary.Write(buf, binary.LittleEndian, uint32(dataSize))
|
||||
buf.Write(make([]byte, dataSize))
|
||||
return buf.Bytes()
|
||||
}
|
||||
|
||||
var _ = Describe("realtime pipeline voice profiles", func() {
|
||||
It("resolves a saved profile to an immutable lease and transcript", func(ctx SpecContext) {
|
||||
store := voiceprofile.NewStore(GinkgoT().TempDir())
|
||||
DeferCleanup(func() { Expect(store.Close()).To(Succeed()) })
|
||||
profile, err := store.Create(ctx, voiceprofile.CreateInput{
|
||||
Name: "Narrator",
|
||||
Language: "en-US",
|
||||
Transcript: "The reference transcript.",
|
||||
ConsentConfirmed: true,
|
||||
}, bytes.NewReader(realtimeProfileWAV(time.Second)))
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
voice, params, release, err := resolveRealtimeVoice(ctx, profile.Voice, &config.ModelConfig{
|
||||
Name: "clone-base",
|
||||
Backend: "qwen3-tts-cpp",
|
||||
TTSConfig: config.TTSConfig{VoiceCloning: ptrTo(true)},
|
||||
}, store)
|
||||
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(voice).To(BeAnExistingFile())
|
||||
Expect(params).To(Equal(map[string]string{"ref_text": "The reference transcript."}))
|
||||
release()
|
||||
release()
|
||||
Expect(voice).NotTo(BeAnExistingFile())
|
||||
})
|
||||
|
||||
It("leaves an ordinary backend voice unchanged with no parameters", func() {
|
||||
voice, params, release, err := resolveRealtimeVoice(context.Background(), "speaker-7", &config.ModelConfig{}, nil)
|
||||
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(voice).To(Equal("speaker-7"))
|
||||
Expect(params).To(BeNil())
|
||||
Expect(release).NotTo(BeNil())
|
||||
Expect(func() { release(); release() }).NotTo(Panic())
|
||||
})
|
||||
|
||||
DescribeTable("returns actionable reference errors",
|
||||
func(configuredVoice string, cfg *config.ModelConfig, store *voiceprofile.Store, expected string) {
|
||||
_, _, release, err := resolveRealtimeVoice(context.Background(), configuredVoice, cfg, store)
|
||||
Expect(err).To(MatchError(ContainSubstring(expected)))
|
||||
Expect(release).To(BeNil())
|
||||
},
|
||||
Entry("malformed reference", "localai://voice-profiles/not-a-uuid", &config.ModelConfig{}, nil, "invalid voice profile reference"),
|
||||
Entry("unsupported model", "localai://voice-profiles/00000000-0000-0000-0000-000000000001", &config.ModelConfig{Backend: "piper"}, nil, "does not support reference-audio voice cloning"),
|
||||
Entry("unavailable store", "localai://voice-profiles/00000000-0000-0000-0000-000000000001", &config.ModelConfig{Name: "clone-base", Backend: "qwen3-tts-cpp", TTSConfig: config.TTSConfig{VoiceCloning: ptrTo(true)}}, nil, "voice profile store is unavailable"),
|
||||
)
|
||||
|
||||
It("reports a missing profile", func() {
|
||||
store := voiceprofile.NewStore(GinkgoT().TempDir())
|
||||
DeferCleanup(func() { Expect(store.Close()).To(Succeed()) })
|
||||
_, _, release, err := resolveRealtimeVoice(context.Background(), "localai://voice-profiles/00000000-0000-0000-0000-000000000001", &config.ModelConfig{
|
||||
Name: "clone-base", Backend: "qwen3-tts-cpp", TTSConfig: config.TTSConfig{VoiceCloning: ptrTo(true)},
|
||||
}, store)
|
||||
Expect(errors.Is(err, voiceprofile.ErrNotFound)).To(BeTrue())
|
||||
Expect(err.Error()).To(ContainSubstring("voice profile not found"))
|
||||
Expect(release).To(BeNil())
|
||||
})
|
||||
})
|
||||
|
||||
type recordingTTSBackend struct {
|
||||
grpcPkg.Backend
|
||||
requests []*proto.TTSRequest
|
||||
}
|
||||
|
||||
func (b *recordingTTSBackend) HealthCheck(context.Context) (bool, error) { return true, nil }
|
||||
func (b *recordingTTSBackend) IsBusy() bool { return false }
|
||||
|
||||
func (b *recordingTTSBackend) record(req *proto.TTSRequest) {
|
||||
b.requests = append(b.requests, req)
|
||||
req.Params["ref_text"] = "backend mutation"
|
||||
}
|
||||
|
||||
func (b *recordingTTSBackend) TTS(_ context.Context, req *proto.TTSRequest, _ ...grpc.CallOption) (*proto.Result, error) {
|
||||
b.record(req)
|
||||
return &proto.Result{Success: true}, nil
|
||||
}
|
||||
|
||||
func (b *recordingTTSBackend) TTSStream(_ context.Context, req *proto.TTSRequest, callback func(*proto.Reply), _ ...grpc.CallOption) error {
|
||||
b.record(req)
|
||||
header := make([]byte, wavStreamHeaderBytes)
|
||||
binary.LittleEndian.PutUint32(header[24:28], 24000)
|
||||
callback(&proto.Reply{Audio: header})
|
||||
return nil
|
||||
}
|
||||
|
||||
var _ = Describe("wrappedModel voice profile parameters", func() {
|
||||
var (
|
||||
wrapped *wrappedModel
|
||||
backendRecorder *recordingTTSBackend
|
||||
)
|
||||
|
||||
BeforeEach(func() {
|
||||
state, err := system.GetSystemState(system.WithModelPath(GinkgoT().TempDir()))
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
appConfig := config.NewApplicationConfig(config.WithSystemState(state))
|
||||
appConfig.GeneratedContentDir = GinkgoT().TempDir()
|
||||
loader := model.NewModelLoader(state)
|
||||
backendRecorder = &recordingTTSBackend{}
|
||||
cfg := &config.ModelConfig{Name: "tts-test", Backend: "test"}
|
||||
cfg.Model = "weights"
|
||||
loaded := model.NewModelWithClient(cfg.ModelID(), "in-process", backendRecorder)
|
||||
loaded.MarkHealthy()
|
||||
_, err = loader.LoadModel(cfg.ModelID(), cfg.Model, func(_, _, _ string) (*model.Model, error) { return loaded, nil })
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
wrapped = &wrappedModel{
|
||||
TTSConfig: cfg,
|
||||
ttsParams: map[string]string{"ref_text": "Original transcript"},
|
||||
modelLoader: loader,
|
||||
appConfig: appConfig,
|
||||
}
|
||||
})
|
||||
|
||||
It("forwards a fresh transcript parameter map to every unary request", func() {
|
||||
_, _, err := wrapped.TTS(context.Background(), "one", "voice.wav", "en")
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
_, _, err = wrapped.TTS(context.Background(), "two", "voice.wav", "en")
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
Expect(backendRecorder.requests).To(HaveLen(2))
|
||||
Expect(backendRecorder.requests[0].Params).To(HaveKeyWithValue("ref_text", "backend mutation"))
|
||||
Expect(backendRecorder.requests[1].Params).To(HaveKeyWithValue("ref_text", "backend mutation"))
|
||||
Expect(wrapped.ttsParams).To(HaveKeyWithValue("ref_text", "Original transcript"))
|
||||
})
|
||||
|
||||
It("forwards a copied transcript parameter map to streaming requests", func() {
|
||||
err := wrapped.TTSStream(context.Background(), "one", "voice.wav", "en", func([]byte, int) error { return nil })
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
Expect(backendRecorder.requests).To(HaveLen(1))
|
||||
Expect(backendRecorder.requests[0].Params).To(HaveKeyWithValue("ref_text", "backend mutation"))
|
||||
Expect(wrapped.ttsParams).To(HaveKeyWithValue("ref_text", "Original transcript"))
|
||||
})
|
||||
})
|
||||
|
||||
var _ = Describe("realtime session voice switching", func() {
|
||||
type fixture struct {
|
||||
store *voiceprofile.Store
|
||||
voiceDir string
|
||||
loader *config.ModelConfigLoader
|
||||
models *model.ModelLoader
|
||||
appConfig *config.ApplicationConfig
|
||||
profileA voiceprofile.Profile
|
||||
profileB voiceprofile.Profile
|
||||
}
|
||||
|
||||
newFixture := func(ctx SpecContext) *fixture {
|
||||
modelDir := GinkgoT().TempDir()
|
||||
voiceDir := GinkgoT().TempDir()
|
||||
store := voiceprofile.NewStore(voiceDir)
|
||||
DeferCleanup(func() { Expect(store.Close()).To(Succeed()) })
|
||||
profileA, err := store.Create(ctx, voiceprofile.CreateInput{
|
||||
Name: "Alpha", Language: "en", Transcript: "Alpha transcript", ConsentConfirmed: true,
|
||||
}, bytes.NewReader(realtimeProfileWAV(time.Second)))
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
profileB, err := store.Create(ctx, voiceprofile.CreateInput{
|
||||
Name: "Beta", Language: "it", Transcript: "Beta transcript", ConsentConfirmed: true,
|
||||
}, bytes.NewReader(realtimeProfileWAV(time.Second)))
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
configs := map[string]string{
|
||||
"vad": "name: vad\nbackend: test\nparameters:\n model: vad.bin\n",
|
||||
"stt": "name: stt\nbackend: test\nparameters:\n model: stt.bin\n",
|
||||
"llm": "name: llm\nbackend: test\nparameters:\n model: llm.bin\n",
|
||||
"tts-a": fmt.Sprintf("name: tts-a\nbackend: qwen3-tts-cpp\nparameters:\n model: tts-a.bin\ntts:\n voice: %s\n voice_cloning: true\n", profileA.Voice),
|
||||
"tts-b": fmt.Sprintf("name: tts-b\nbackend: qwen3-tts-cpp\nparameters:\n model: tts-b.bin\ntts:\n voice: %s\n voice_cloning: true\n", profileB.Voice),
|
||||
"pipe-a": "name: pipe-a\npipeline:\n vad: vad\n transcription: stt\n llm: llm\n tts: tts-a\n disable_warmup: true\n",
|
||||
"pipe-b": "name: pipe-b\npipeline:\n vad: vad\n transcription: stt\n llm: llm\n tts: tts-b\n disable_warmup: true\n",
|
||||
}
|
||||
for name, body := range configs {
|
||||
Expect(os.WriteFile(filepath.Join(modelDir, name+".yaml"), []byte(body), 0o644)).To(Succeed())
|
||||
}
|
||||
loader := config.NewModelConfigLoader(modelDir)
|
||||
Expect(loader.LoadModelConfigsFromPath(modelDir)).To(Succeed())
|
||||
state, err := system.GetSystemState(system.WithModelPath(modelDir))
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
return &fixture{
|
||||
store: store, voiceDir: voiceDir, loader: loader, models: model.NewModelLoader(state),
|
||||
appConfig: config.NewApplicationConfig(config.WithSystemState(state)),
|
||||
profileA: profileA, profileB: profileB,
|
||||
}
|
||||
}
|
||||
|
||||
newSession := func(f *fixture, voice string) *Session {
|
||||
cfg, err := f.loader.LoadModelConfigFileByNameDefaultOptions("pipe-a", f.appConfig)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
m, err := newModel(&cfg.Pipeline, f.loader, f.models, f.appConfig, nil, nil)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
session := &Session{
|
||||
Model: "pipe-a", Voice: voice, ModelConfig: cfg, ModelInterface: m,
|
||||
InputAudioTranscription: &types.AudioTranscription{Model: "stt"},
|
||||
}
|
||||
resolved, params, release, err := resolveRealtimeVoice(context.Background(), voice, m.(*wrappedModel).TTSConfig, f.store)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
session.installVoiceBinding(resolved, params, release)
|
||||
return session
|
||||
}
|
||||
|
||||
update := func(f *fixture, session *Session, rt *types.RealtimeSession) error {
|
||||
return updateSession(session, &types.SessionUnion{Realtime: rt}, f.loader, f.models, f.appConfig, nil, nil, f.store)
|
||||
}
|
||||
|
||||
It("switches ordinary voices to profiles and clears the lease at final cleanup", func(ctx SpecContext) {
|
||||
f := newFixture(ctx)
|
||||
session := newSession(f, "speaker-1")
|
||||
Expect(update(f, session, &types.RealtimeSession{Audio: &types.RealtimeSessionAudio{Output: &types.SessionAudioOutput{Voice: types.Voice(f.profileA.Voice)}}})).To(Succeed())
|
||||
|
||||
Expect(session.Voice).To(BeAnExistingFile())
|
||||
Expect(session.ttsParams).To(Equal(map[string]string{"ref_text": "Alpha transcript"}))
|
||||
leased := session.Voice
|
||||
session.releaseVoiceBinding()
|
||||
session.releaseVoiceBinding()
|
||||
Expect(leased).NotTo(BeAnExistingFile())
|
||||
})
|
||||
|
||||
It("replaces one profile lease with another and then an ordinary voice", func(ctx SpecContext) {
|
||||
f := newFixture(ctx)
|
||||
session := newSession(f, f.profileA.Voice)
|
||||
firstLease := session.Voice
|
||||
|
||||
Expect(update(f, session, &types.RealtimeSession{Audio: &types.RealtimeSessionAudio{Output: &types.SessionAudioOutput{Voice: types.Voice(f.profileB.Voice)}}})).To(Succeed())
|
||||
Expect(firstLease).NotTo(BeAnExistingFile())
|
||||
secondLease := session.Voice
|
||||
Expect(secondLease).To(BeAnExistingFile())
|
||||
Expect(session.ttsParams).To(HaveKeyWithValue("ref_text", "Beta transcript"))
|
||||
|
||||
Expect(update(f, session, &types.RealtimeSession{Audio: &types.RealtimeSessionAudio{Output: &types.SessionAudioOutput{Voice: "speaker-2"}}})).To(Succeed())
|
||||
Expect(secondLease).NotTo(BeAnExistingFile())
|
||||
Expect(session.Voice).To(Equal("speaker-2"))
|
||||
Expect(session.ttsParams).To(BeNil())
|
||||
Expect(session.ModelInterface.(*wrappedModel).ttsParams).To(BeNil())
|
||||
})
|
||||
|
||||
It("uses a new model default profile unless an explicit voice takes precedence", func(ctx SpecContext) {
|
||||
f := newFixture(ctx)
|
||||
session := newSession(f, "speaker-1")
|
||||
Expect(update(f, session, &types.RealtimeSession{Model: "pipe-b"})).To(Succeed())
|
||||
Expect(session.ttsParams).To(HaveKeyWithValue("ref_text", "Beta transcript"))
|
||||
defaultLease := session.Voice
|
||||
|
||||
Expect(update(f, session, &types.RealtimeSession{
|
||||
Model: "pipe-a",
|
||||
Audio: &types.RealtimeSessionAudio{Output: &types.SessionAudioOutput{Voice: "speaker-explicit"}},
|
||||
})).To(Succeed())
|
||||
Expect(defaultLease).NotTo(BeAnExistingFile())
|
||||
Expect(session.Voice).To(Equal("speaker-explicit"))
|
||||
Expect(session.ttsParams).To(BeNil())
|
||||
})
|
||||
|
||||
It("preserves a profile binding across a language-only rebuild", func(ctx SpecContext) {
|
||||
f := newFixture(ctx)
|
||||
session := newSession(f, f.profileA.Voice)
|
||||
lease := session.Voice
|
||||
Expect(update(f, session, &types.RealtimeSession{Audio: &types.RealtimeSessionAudio{Input: &types.SessionAudioInput{
|
||||
Transcription: &types.AudioTranscription{Language: "fr"},
|
||||
}}})).To(Succeed())
|
||||
|
||||
Expect(session.Voice).To(Equal(lease))
|
||||
Expect(session.InputAudioTranscription.Model).To(Equal("stt"))
|
||||
Expect(session.InputAudioTranscription.Language).To(Equal("fr"))
|
||||
wrapped := session.ModelInterface.(*wrappedModel)
|
||||
Expect(wrapped.ttsParams).To(Equal(map[string]string{"ref_text": "Alpha transcript"}))
|
||||
wrapped.ttsParams["ref_text"] = "wrapper mutation"
|
||||
Expect(session.ttsParams).To(HaveKeyWithValue("ref_text", "Alpha transcript"))
|
||||
})
|
||||
|
||||
It("rolls back the model, wrapper, voice, and lease when preparation fails", func(ctx SpecContext) {
|
||||
f := newFixture(ctx)
|
||||
session := newSession(f, f.profileA.Voice)
|
||||
oldModel, oldConfig, oldVoice := session.ModelInterface, session.ModelConfig, session.Voice
|
||||
err := update(f, session, &types.RealtimeSession{
|
||||
Model: "pipe-b",
|
||||
Audio: &types.RealtimeSessionAudio{Output: &types.SessionAudioOutput{Voice: "localai://voice-profiles/00000000-0000-0000-0000-000000000001"}},
|
||||
})
|
||||
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(session.Model).To(Equal("pipe-a"))
|
||||
Expect(session.ModelConfig).To(BeIdenticalTo(oldConfig))
|
||||
Expect(session.ModelInterface).To(BeIdenticalTo(oldModel))
|
||||
Expect(session.Voice).To(Equal(oldVoice))
|
||||
Expect(oldVoice).To(BeAnExistingFile())
|
||||
Expect(session.ttsParams).To(HaveKeyWithValue("ref_text", "Alpha transcript"))
|
||||
})
|
||||
|
||||
It("releases a candidate profile lease when later validation fails", func(ctx SpecContext) {
|
||||
f := newFixture(ctx)
|
||||
session := newSession(f, "speaker-1")
|
||||
leases := func() []string {
|
||||
matches, err := filepath.Glob(filepath.Join(f.voiceDir, voiceprofile.DirectoryName, ".leases", "*", "*.wav"))
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
return matches
|
||||
}
|
||||
Expect(leases()).To(BeEmpty())
|
||||
|
||||
err := update(f, session, &types.RealtimeSession{
|
||||
Audio: &types.RealtimeSessionAudio{Output: &types.SessionAudioOutput{Voice: types.Voice(f.profileA.Voice)}},
|
||||
LocalAIClassifier: classifierTestConfig(0, nil),
|
||||
})
|
||||
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(session.Voice).To(Equal("speaker-1"))
|
||||
Expect(leases()).To(BeEmpty())
|
||||
})
|
||||
})
|
||||
|
||||
func ptrTo[T any](value T) *T { return &value }
|
||||
@@ -103,3 +103,48 @@ test.describe("Models gallery - recommended panel prominence", () => {
|
||||
await expect(grid(page).locator(".lane__tag--evidence")).toHaveCount(1);
|
||||
});
|
||||
});
|
||||
|
||||
// Start with a fitting model so absence assertions cannot pass during loading.
|
||||
// Then change the polled hardware budget while keeping the same gallery.
|
||||
for (const view of ["models", "home"]) {
|
||||
test(`${view} removes GPU recommendations when no candidate fits`, async ({ page }) => {
|
||||
await mockGallery(page, 0);
|
||||
await page.route("**/v1/models", (route) =>
|
||||
route.fulfill({ json: { data: [] } }),
|
||||
);
|
||||
const gib = 1024 ** 3;
|
||||
let budget = 24 * gib;
|
||||
await page.route("**/api/resources", (route) =>
|
||||
route.fulfill({ json: {
|
||||
type: "gpu",
|
||||
aggregate: { total_memory: budget, gpu_count: 1 },
|
||||
gpus: [{ vendor: "nvidia", total_memory: budget }],
|
||||
} }),
|
||||
);
|
||||
await page.route("**/api/models/estimate/*", (route) =>
|
||||
route.fulfill({ json: {
|
||||
sizeBytes: 17.4 * gib,
|
||||
sizeDisplay: "17.4 GB",
|
||||
estimates: { 4096: { vramBytes: 18.4 * gib, vramDisplay: "18.4 GB" } },
|
||||
} }),
|
||||
);
|
||||
await page.goto(view === "models" ? "/app/models" : "/app/");
|
||||
const section = view === "models" ? panel(page) : page.locator(".home-starters");
|
||||
await expect(section).toBeVisible();
|
||||
await expect(section).toContainText("tiny-chat");
|
||||
|
||||
// Wait for BOTH recommendation estimates, not the hook's loading render
|
||||
// or the gallery rail's separate context-size requests.
|
||||
const estimatesFinished = REC_MODELS.map(model => page.waitForResponse(response => {
|
||||
const url = new URL(response.url());
|
||||
return url.pathname.endsWith('/api/models/estimate/' + model.name) &&
|
||||
url.searchParams.get('contexts') === '4096' && response.status() === 200;
|
||||
}).then(response => response.finished()));
|
||||
budget = 12 * gib;
|
||||
await Promise.all(estimatesFinished);
|
||||
await page.evaluate(() => new Promise(resolve =>
|
||||
requestAnimationFrame(() => requestAnimationFrame(resolve)),
|
||||
));
|
||||
await expect(section).toHaveCount(0, { timeout: 15_000 });
|
||||
});
|
||||
}
|
||||
Generated
+21
-10
@@ -24,7 +24,7 @@
|
||||
"@modelcontextprotocol/sdk": "^1.30.0",
|
||||
"dompurify": "^3.4.13",
|
||||
"highlight.js": "^11.11.1",
|
||||
"hono": "4.12.34",
|
||||
"hono": "4.13.5",
|
||||
"i18next": "^26.0.8",
|
||||
"i18next-browser-languagedetector": "^8.2.1",
|
||||
"i18next-http-backend": "^3.0.6",
|
||||
@@ -771,9 +771,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/@istanbuljs/load-nyc-config/node_modules/js-yaml": {
|
||||
"version": "3.14.2",
|
||||
"resolved": "https://registry.npmjs.org/js-yaml/-/js-yaml-3.14.2.tgz",
|
||||
"integrity": "sha512-PMSmkqxr106Xa156c2M265Z+FTrPl+oxd/rgOQy2tijQeK5TxQ43psO1ZCwhVOSdnn+RzkzlRz/eY4BgJBYVpg==",
|
||||
"version": "3.15.2",
|
||||
"resolved": "https://registry.npmjs.org/js-yaml/-/js-yaml-3.15.2.tgz",
|
||||
"integrity": "sha512-6EuL879VkRA+1Cz578mKMiKvjPNEuk6+r1JaFzoSWejZmtf7xWbIyw1e3KkxlkzTIt9Taw6JBhEppG7utc1P+w==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
@@ -3467,9 +3467,9 @@
|
||||
}
|
||||
},
|
||||
"node_modules/hono": {
|
||||
"version": "4.12.34",
|
||||
"resolved": "https://registry.npmjs.org/hono/-/hono-4.12.34.tgz",
|
||||
"integrity": "sha512-GqXJqY/xJkJmuloTrnV1ZEXG3fqte+VjkUqoRNZXcrUidiUOP4fMSIHHY4tsqZBK++kVyWmt/AAfSUuy57/eSA==",
|
||||
"version": "4.13.5",
|
||||
"resolved": "https://registry.npmjs.org/hono/-/hono-4.13.5.tgz",
|
||||
"integrity": "sha512-O6+/eCYRkzzzy0rPWwKLiGBR1nFuUPZynnwjxN1MBA62NNqbT0wQEzQyK2gSO5yDIDB336sXQleAhOHrzlYyKw==",
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">=16.9.0"
|
||||
@@ -4593,10 +4593,21 @@
|
||||
"license": "MIT"
|
||||
},
|
||||
"node_modules/js-yaml": {
|
||||
"version": "4.1.1",
|
||||
"resolved": "https://registry.npmjs.org/js-yaml/-/js-yaml-4.1.1.tgz",
|
||||
"integrity": "sha512-qQKT4zQxXl8lLwBtHMWwaTcGfFOZviOJet3Oy/xmGk2gZH677CJM9EvtfdSkgWcATZhj/55JZ0rmy3myCT5lsA==",
|
||||
"version": "4.3.2",
|
||||
"resolved": "https://registry.npmjs.org/js-yaml/-/js-yaml-4.3.2.tgz",
|
||||
"integrity": "sha512-SFNOvSJ+Dgf/9An904Yx+CgSlIPCkIpao4qo51lpee25TIRejdH3rhR4EZMGoNx3/TP3O+wzWuiTFl4sqbltzA==",
|
||||
"dev": true,
|
||||
"funding": [
|
||||
{
|
||||
"type": "github",
|
||||
"url": "https://github.com/sponsors/puzrin"
|
||||
},
|
||||
{
|
||||
"type": "github",
|
||||
"url": "https://github.com/sponsors/nodeca"
|
||||
}
|
||||
],
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"argparse": "^2.0.1"
|
||||
},
|
||||
|
||||
@@ -19,7 +19,7 @@
|
||||
"coverage:report": "nyc report"
|
||||
},
|
||||
"overrides": {
|
||||
"hono": "4.12.34",
|
||||
"hono": "4.13.5",
|
||||
"ip-address": "10.3.1",
|
||||
"path-to-regexp": "^8.4.0"
|
||||
},
|
||||
@@ -40,7 +40,7 @@
|
||||
"@modelcontextprotocol/sdk": "^1.30.0",
|
||||
"dompurify": "^3.4.13",
|
||||
"highlight.js": "^11.11.1",
|
||||
"hono": "4.12.34",
|
||||
"hono": "4.13.5",
|
||||
"i18next": "^26.0.8",
|
||||
"i18next-browser-languagedetector": "^8.2.1",
|
||||
"i18next-http-backend": "^3.0.6",
|
||||
|
||||
@@ -3,63 +3,17 @@ import { useTranslation } from 'react-i18next'
|
||||
import { modelsApi } from '../utils/api'
|
||||
import { useRecommendedModels, isNvfp4Name } from '../hooks/useRecommendedModels'
|
||||
|
||||
// Static fallback used only when the live gallery / estimates can't be reached
|
||||
// (offline, trimmed gallery). The hook is the primary, data-driven path; these
|
||||
// are real gallery names kept as a safety net so onboarding never shows nothing.
|
||||
// Gemma picks use the QAT (quantization-aware-trained) Q4 builds. NVIDIA boxes
|
||||
// get NVFP4 + MTP variants at the mid/large tiers (see NVIDIA below).
|
||||
const BASE = {
|
||||
cpu: [
|
||||
{ name: 'gemma-4-e2b-it-qat-q4_0', size: '~1.5 GB' },
|
||||
{ name: 'qwen3.5-4b-claude-4.6-opus-reasoning-distilled', size: '~2.5 GB' },
|
||||
{ name: 'gemma-4-e4b-it-qat-q4_0', size: '~3 GB' },
|
||||
{ name: 'lfm2.5-1.2b-instruct', size: '~0.8 GB' },
|
||||
],
|
||||
'gpu-small': [
|
||||
{ name: 'gemma-4-e4b-it-qat-q4_0', size: '~3 GB' },
|
||||
{ name: 'lfm2.5-8b-a1b', size: '~5 GB' },
|
||||
{ name: 'qwen3.5-9b', size: '~5.5 GB' },
|
||||
{ name: 'gemma-4-12b-it-qat-q4_0', size: '~7 GB' },
|
||||
],
|
||||
'gpu-mid': [
|
||||
{ name: 'qwen3.6-27b', size: '~16 GB' },
|
||||
{ name: 'qwen3.6-27b-mtp-pi-tune', size: '~16 GB' },
|
||||
{ name: 'gemma-4-26b-a4b-it-qat-q4_0', size: '~16 GB' },
|
||||
{ name: 'qwen3.5-27b', size: '~16 GB' },
|
||||
],
|
||||
'gpu-large': [
|
||||
{ name: 'qwen3.6-35b-a3b-apex', size: '~20 GB' },
|
||||
{ name: 'qwen3.6-35b-a3b-claude-4.6-opus-reasoning-distilled', size: '~20 GB' },
|
||||
{ name: 'gemma-4-31b-it-qat-q4_0', size: '~18 GB' },
|
||||
{ name: 'qwen3.5-35b-a3b-apex', size: '~20 GB' },
|
||||
],
|
||||
}
|
||||
|
||||
// NVIDIA-only overrides: NVFP4 is a Blackwell-optimised 4-bit format paired with
|
||||
// MTP (multi-token prediction) for speed. Only the mid/large tiers have these.
|
||||
const NVIDIA = {
|
||||
'gpu-mid': [
|
||||
{ name: 'qwen3.6-27b-nvfp4-mtp', size: '~14 GB' },
|
||||
{ name: 'qwen3.6-27b-mtp-pi-tune', size: '~16 GB' },
|
||||
{ name: 'gemma-4-26b-a4b-it-qat-q4_0', size: '~16 GB' },
|
||||
{ name: 'qwen3.6-27b', size: '~16 GB' },
|
||||
],
|
||||
'gpu-large': [
|
||||
{ name: 'qwen3.6-35b-a3b-nvfp4-mtp', size: '~18 GB' },
|
||||
{ name: 'qwen3.6-27b-nvfp4-mtp', size: '~14 GB' },
|
||||
{ name: 'qwen3.6-35b-a3b-apex', size: '~20 GB' },
|
||||
{ name: 'gemma-4-31b-it-qat-q4_0', size: '~18 GB' },
|
||||
],
|
||||
}
|
||||
|
||||
function fallbackFor(tierId, isNvidia) {
|
||||
if (isNvidia && NVIDIA[tierId]) return NVIDIA[tierId]
|
||||
return BASE[tierId] || BASE.cpu
|
||||
}
|
||||
// Offline CPU suggestions do not claim a measured GPU fit.
|
||||
const CPU_FALLBACK = [
|
||||
{ name: 'gemma-4-e2b-it-qat-q4_0', size: '~1.5 GB' },
|
||||
{ name: 'qwen3.5-4b-claude-4.6-opus-reasoning-distilled', size: '~2.5 GB' },
|
||||
{ name: 'gemma-4-e4b-it-qat-q4_0', size: '~3 GB' },
|
||||
{ name: 'lfm2.5-1.2b-instruct', size: '~0.8 GB' },
|
||||
]
|
||||
|
||||
export default function StarterModels({ addToast, onInstallStarted }) {
|
||||
const { t } = useTranslation('home')
|
||||
const { recommended, tier, isNvidia, loading } = useRecommendedModels({ count: 4 })
|
||||
const { recommended, tier, loading } = useRecommendedModels({ count: 4 })
|
||||
const [installing, setInstalling] = useState(() => new Set())
|
||||
|
||||
// While the hardware probe + gallery query are in flight, render nothing
|
||||
@@ -67,10 +21,11 @@ export default function StarterModels({ addToast, onInstallStarted }) {
|
||||
if (loading) return null
|
||||
|
||||
// Prefer live recommendations; fall back to the static list only when the
|
||||
// gallery yielded nothing.
|
||||
// gallery yielded nothing on a CPU host. Static GPU picks have no measured
|
||||
// fit and must not replace an empty set of fitting recommendations.
|
||||
const items = (recommended && recommended.length > 0)
|
||||
? recommended.map(r => ({ name: r.name, size: r.sizeDisplay }))
|
||||
: fallbackFor(tier.id, isNvidia)
|
||||
: tier.id === 'cpu' ? CPU_FALLBACK : []
|
||||
|
||||
if (items.length === 0) return null
|
||||
|
||||
|
||||
+3
-3
@@ -53,16 +53,16 @@ function rank(candidates, tier, count, isNvidia) {
|
||||
}
|
||||
const limit = tier.vram * 0.95
|
||||
const fits = pool.filter(c => c.vramBytes != null && c.vramBytes <= limit)
|
||||
const base = fits.length > 0 ? fits : pool // tiny GPU where nothing fits → fall through to smallest
|
||||
const byPreference = (a, b) => {
|
||||
// On NVIDIA, surface NVFP4 first; then largest-that-fits (best quality).
|
||||
if (isNvidia) {
|
||||
const an = isNvfp4Name(a.name), bn = isNvfp4Name(b.name)
|
||||
if (an !== bn) return an ? -1 : 1
|
||||
}
|
||||
return fits.length > 0 ? b.sizeBytes - a.sizeBytes : a.sizeBytes - b.sizeBytes
|
||||
return b.sizeBytes - a.sizeBytes
|
||||
}
|
||||
return [...base].sort(byPreference).slice(0, count)
|
||||
// An oversized or unestimated model cannot be labelled a hardware fit.
|
||||
return [...fits].sort(byPreference).slice(0, count)
|
||||
}
|
||||
|
||||
export function useRecommendedModels({ count = 4, candidatePool = 10 } = {}) {
|
||||
|
||||
@@ -439,6 +439,12 @@ func SubjectNodeFilesStage(nodeID string) string {
|
||||
return subjectNodePrefix + sanitizeSubjectToken(nodeID) + ".files.stage"
|
||||
}
|
||||
|
||||
// SubjectNodeFilesRelease tells a serve-backend node to evict one request's ephemeral cache keys.
|
||||
// Reply: {error}
|
||||
func SubjectNodeFilesRelease(nodeID string) string {
|
||||
return subjectNodePrefix + sanitizeSubjectToken(nodeID) + ".files.release"
|
||||
}
|
||||
|
||||
// SubjectNodeFilesTemp tells a serve-backend node to allocate a temp file.
|
||||
// Reply: {local_path, error}
|
||||
func SubjectNodeFilesTemp(nodeID string) string {
|
||||
|
||||
@@ -1,6 +1,11 @@
|
||||
package nodes
|
||||
|
||||
import "context"
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"path"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// FileStager abstracts file transfer between frontend and backend nodes
|
||||
// in distributed mode. Two implementations exist:
|
||||
@@ -29,7 +34,59 @@ type FileStager interface {
|
||||
// StageRemoteToStore uploads a remote file to shared storage.
|
||||
StageRemoteToStore(ctx context.Context, nodeID, remotePath, key string) error
|
||||
|
||||
// ReleaseRemote removes one ephemeral key from the remote node.
|
||||
ReleaseRemote(ctx context.Context, nodeID, key string) error
|
||||
|
||||
// ListRemoteDir returns relative file paths within a directory on the remote node.
|
||||
// keyPrefix is a storage-style key prefix (e.g. "models/mymodel").
|
||||
ListRemoteDir(ctx context.Context, nodeID, keyPrefix string) ([]string, error)
|
||||
}
|
||||
|
||||
// RequestFileReleaser removes all ephemeral keys staged for one inference in
|
||||
// one transport operation. FileStagingClient falls back to ReleaseRemote for
|
||||
// stagers that do not implement this optional rolling-upgrade extension.
|
||||
type RequestFileReleaser interface {
|
||||
ReleaseRemoteRequest(ctx context.Context, nodeID, requestID string, keys []string) error
|
||||
}
|
||||
|
||||
func validateEphemeralRequestRelease(requestID string, keys []string) error {
|
||||
if err := validateEphemeralRequestID(requestID); err != nil {
|
||||
return err
|
||||
}
|
||||
if len(keys) == 0 {
|
||||
return fmt.Errorf("release batch must contain at least one key")
|
||||
}
|
||||
for _, key := range keys {
|
||||
if err := validateEphemeralReleaseKey(key); err != nil {
|
||||
return err
|
||||
}
|
||||
parts := strings.Split(key, "/")
|
||||
if parts[2] != requestID {
|
||||
return fmt.Errorf("release batch mixes request IDs %q and %q", requestID, parts[2])
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateEphemeralRequestID(requestID string) error {
|
||||
if requestID == "" || strings.ContainsAny(requestID, "/\\") || path.Clean(requestID) != requestID || requestID == "." || requestID == ".." {
|
||||
return fmt.Errorf("invalid ephemeral request ID %q", requestID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateEphemeralReleaseKey(key string) error {
|
||||
if strings.Contains(key, "\\") || path.Clean(key) != key {
|
||||
return fmt.Errorf("invalid ephemeral key %q", key)
|
||||
}
|
||||
parts := strings.Split(key, "/")
|
||||
if len(parts) != 4 || parts[0] != "ephemeral" {
|
||||
return fmt.Errorf("release key %q must identify one file below ephemeral/", key)
|
||||
}
|
||||
for _, part := range parts[1:] {
|
||||
if part == "" || part == "." || part == ".." {
|
||||
return fmt.Errorf("invalid ephemeral key %q", key)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -1,6 +1,7 @@
|
||||
package nodes
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
@@ -10,6 +11,7 @@ import (
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
@@ -81,6 +83,86 @@ func NewHTTPFileStager(httpAddrFor func(nodeID string) (string, error), token st
|
||||
}
|
||||
}
|
||||
|
||||
// ReleaseRemote removes one exact ephemeral key from a backend node.
|
||||
func (h *HTTPFileStager) ReleaseRemote(ctx context.Context, nodeID, key string) error {
|
||||
if err := validateEphemeralReleaseKey(key); err != nil {
|
||||
return err
|
||||
}
|
||||
addr, err := h.httpAddrFor(nodeID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("resolving HTTP address for node %s: %w", nodeID, err)
|
||||
}
|
||||
releaseURL := (&url.URL{Scheme: "http", Host: addr, Path: "/v1/files/" + key}).String()
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodDelete, releaseURL, nil)
|
||||
if err != nil {
|
||||
return fmt.Errorf("creating release request for %q: %w", key, err)
|
||||
}
|
||||
if h.token != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+h.token)
|
||||
}
|
||||
resp, err := h.client.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("releasing %q from node %s: %w", key, nodeID, err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices {
|
||||
body, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
|
||||
return fmt.Errorf("releasing %q from node %s: status %d: %s", key, nodeID, resp.StatusCode, strings.TrimSpace(string(body)))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ReleaseRemoteRequest removes one inference's staged inputs with one HTTP
|
||||
// request. Older workers return 404 for the batch endpoint, so the client
|
||||
// retries through the exact-key API during rolling upgrades.
|
||||
func (h *HTTPFileStager) ReleaseRemoteRequest(ctx context.Context, nodeID, requestID string, keys []string) error {
|
||||
if err := validateEphemeralRequestRelease(requestID, keys); err != nil {
|
||||
return err
|
||||
}
|
||||
addr, err := h.httpAddrFor(nodeID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("resolving HTTP address for node %s: %w", nodeID, err)
|
||||
}
|
||||
payload, err := json.Marshal(struct {
|
||||
RequestID string `json:"request_id"`
|
||||
}{RequestID: requestID})
|
||||
if err != nil {
|
||||
return fmt.Errorf("encoding request release: %w", err)
|
||||
}
|
||||
releaseURL := (&url.URL{Scheme: "http", Host: addr, Path: "/v1/files-release"}).String()
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodDelete, releaseURL, bytes.NewReader(payload))
|
||||
if err != nil {
|
||||
return fmt.Errorf("creating request release: %w", err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
if h.token != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+h.token)
|
||||
}
|
||||
resp, err := h.client.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("releasing request inputs from node %s: %w", nodeID, err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode == http.StatusNotFound || resp.StatusCode == http.StatusMethodNotAllowed {
|
||||
return h.releaseRemoteKeys(ctx, nodeID, keys)
|
||||
}
|
||||
if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices {
|
||||
body, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
|
||||
return fmt.Errorf("releasing request inputs from node %s: status %d: %s", nodeID, resp.StatusCode, strings.TrimSpace(string(body)))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *HTTPFileStager) releaseRemoteKeys(ctx context.Context, nodeID string, keys []string) error {
|
||||
var releaseErrors []error
|
||||
for _, key := range keys {
|
||||
if err := h.ReleaseRemote(ctx, nodeID, key); err != nil {
|
||||
releaseErrors = append(releaseErrors, err)
|
||||
}
|
||||
}
|
||||
return errors.Join(releaseErrors...)
|
||||
}
|
||||
|
||||
func (h *HTTPFileStager) EnsureRemote(ctx context.Context, nodeID, localPath, key string) (string, error) {
|
||||
xlog.Debug("Staging file to remote node via HTTP", "node", nodeID, "localPath", localPath, "key", key)
|
||||
|
||||
@@ -90,7 +172,9 @@ func (h *HTTPFileStager) EnsureRemote(ctx context.Context, nodeID, localPath, ke
|
||||
}
|
||||
|
||||
// Probe: check if the remote already has the file with matching content hash.
|
||||
if remotePath, ok := h.probeExisting(ctx, addr, localPath, key); ok {
|
||||
if remotePath, ok, probeErr := h.probeExisting(ctx, addr, localPath, key); probeErr != nil {
|
||||
return "", fmt.Errorf("claiming existing file on node %s: %w", nodeID, probeErr)
|
||||
} else if ok {
|
||||
xlog.Info("Upload skipped (file already exists with matching hash)", "node", nodeID, "key", key, "remotePath", remotePath)
|
||||
return remotePath, nil
|
||||
}
|
||||
@@ -439,14 +523,15 @@ func isTransientError(err error) bool {
|
||||
|
||||
// probeExisting sends a HEAD request to check if the remote already has the
|
||||
// file with a matching SHA-256 hash. Returns the remote path and true if the
|
||||
// upload can be skipped. Any errors (including 405 from older servers) silently
|
||||
// fall through so the caller proceeds with a normal PUT.
|
||||
func (h *HTTPFileStager) probeExisting(ctx context.Context, addr, localPath, key string) (string, bool) {
|
||||
// upload can be skipped. HEAD and hash errors fall through to a normal PUT.
|
||||
// Matching ephemeral files are claimed first; a 404 or 405 claim response
|
||||
// identifies an older worker and also falls through to PUT.
|
||||
func (h *HTTPFileStager) probeExisting(ctx context.Context, addr, localPath, key string) (string, bool, error) {
|
||||
url := fmt.Sprintf("http://%s/v1/files/%s", addr, key)
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodHead, url, nil)
|
||||
if err != nil {
|
||||
return "", false
|
||||
return "", false, nil
|
||||
}
|
||||
if h.token != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+h.token)
|
||||
@@ -454,18 +539,18 @@ func (h *HTTPFileStager) probeExisting(ctx context.Context, addr, localPath, key
|
||||
|
||||
resp, err := h.client.Do(req)
|
||||
if err != nil {
|
||||
return "", false
|
||||
return "", false, nil
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", false
|
||||
return "", false, nil
|
||||
}
|
||||
|
||||
remotePath := resp.Header.Get(HeaderLocalPath)
|
||||
remoteHash := resp.Header.Get(HeaderContentSHA256)
|
||||
if remotePath == "" || remoteHash == "" {
|
||||
return "", false
|
||||
return "", false, nil
|
||||
}
|
||||
|
||||
// A 200 with a content hash is proof the worker is alive and serving right
|
||||
@@ -475,14 +560,53 @@ func (h *HTTPFileStager) probeExisting(ctx context.Context, addr, localPath, key
|
||||
|
||||
localHash, err := hashLocalCached(ctx, localPath)
|
||||
if err != nil {
|
||||
return "", false
|
||||
return "", false, nil
|
||||
}
|
||||
|
||||
if localHash != remoteHash {
|
||||
return "", false
|
||||
return "", false, nil
|
||||
}
|
||||
|
||||
return remotePath, true
|
||||
if strings.HasPrefix(key, "ephemeral/") {
|
||||
claimed, err := h.claimExisting(ctx, addr, key)
|
||||
if err != nil {
|
||||
return "", false, err
|
||||
}
|
||||
if !claimed {
|
||||
return "", false, nil
|
||||
}
|
||||
}
|
||||
|
||||
return remotePath, true, nil
|
||||
}
|
||||
|
||||
func (h *HTTPFileStager) claimExisting(ctx context.Context, addr, key string) (bool, error) {
|
||||
claimURL := (&url.URL{
|
||||
Scheme: "http",
|
||||
Host: addr,
|
||||
Path: "/v1/files/" + key,
|
||||
RawQuery: "claim=1",
|
||||
}).String()
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, claimURL, nil)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("creating claim request for %q: %w", key, err)
|
||||
}
|
||||
if h.token != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+h.token)
|
||||
}
|
||||
resp, err := h.client.Do(req)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("claiming %q: %w", key, err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode == http.StatusNotFound || resp.StatusCode == http.StatusMethodNotAllowed {
|
||||
return false, nil
|
||||
}
|
||||
if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices {
|
||||
body, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
|
||||
return false, fmt.Errorf("claiming %q: status %d: %s", key, resp.StatusCode, strings.TrimSpace(string(body)))
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// hashChunkSize is how much of a file is hashed between activity ticks and
|
||||
|
||||
@@ -0,0 +1,374 @@
|
||||
package nodes
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/mudler/LocalAI/core/services/messaging"
|
||||
"github.com/mudler/LocalAI/core/services/storage"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
type releaseTestSubscription struct{}
|
||||
|
||||
func (releaseTestSubscription) Unsubscribe() error { return nil }
|
||||
|
||||
type releaseTestMessaging struct {
|
||||
subject string
|
||||
payload []byte
|
||||
onRequest func()
|
||||
requestCalled bool
|
||||
requestCount int
|
||||
timeout time.Duration
|
||||
replies [][]byte
|
||||
}
|
||||
|
||||
func (m *releaseTestMessaging) Publish(string, any) error { return nil }
|
||||
func (m *releaseTestMessaging) Subscribe(string, func([]byte)) (messaging.Subscription, error) {
|
||||
return releaseTestSubscription{}, nil
|
||||
}
|
||||
func (m *releaseTestMessaging) QueueSubscribe(string, string, func([]byte)) (messaging.Subscription, error) {
|
||||
return releaseTestSubscription{}, nil
|
||||
}
|
||||
func (m *releaseTestMessaging) QueueSubscribeReply(string, string, func([]byte, func([]byte))) (messaging.Subscription, error) {
|
||||
return releaseTestSubscription{}, nil
|
||||
}
|
||||
func (m *releaseTestMessaging) SubscribeReply(string, func([]byte, func([]byte))) (messaging.Subscription, error) {
|
||||
return releaseTestSubscription{}, nil
|
||||
}
|
||||
func (m *releaseTestMessaging) Request(subject string, data []byte, timeout time.Duration) ([]byte, error) {
|
||||
m.subject = subject
|
||||
m.payload = append([]byte(nil), data...)
|
||||
m.requestCalled = true
|
||||
m.requestCount++
|
||||
m.timeout = timeout
|
||||
if m.onRequest != nil {
|
||||
m.onRequest()
|
||||
}
|
||||
if m.requestCount <= len(m.replies) {
|
||||
return append([]byte(nil), m.replies[m.requestCount-1]...), nil
|
||||
}
|
||||
return []byte(`{}`), nil
|
||||
}
|
||||
func (m *releaseTestMessaging) IsConnected() bool { return true }
|
||||
func (m *releaseTestMessaging) Close() {}
|
||||
|
||||
var _ = Describe("File stager exact-key release", func() {
|
||||
startReleaseServer := func(stagingDir, token string) (*HTTPFileStager, func()) {
|
||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
server, err := StartFileTransferServerWithListener(
|
||||
listener,
|
||||
stagingDir,
|
||||
GinkgoT().TempDir(),
|
||||
GinkgoT().TempDir(),
|
||||
token,
|
||||
0,
|
||||
)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
return NewHTTPFileStager(func(string) (string, error) {
|
||||
return listener.Addr().String(), nil
|
||||
}, token), func() {
|
||||
Expect(server.Shutdown(context.Background())).To(Succeed())
|
||||
}
|
||||
}
|
||||
|
||||
It("transmits URL metacharacters as the exact key", func() {
|
||||
stagingDir := GinkgoT().TempDir()
|
||||
categoryDir := filepath.Join(stagingDir, "ephemeral", "request-id", "audio")
|
||||
Expect(os.MkdirAll(categoryDir, 0750)).To(Succeed())
|
||||
|
||||
key := "ephemeral/request-id/audio/name ?#%2F.wav"
|
||||
exactPath := filepath.Join(categoryDir, "name ?#%2F.wav")
|
||||
wrongPath := filepath.Join(categoryDir, "name ")
|
||||
for _, path := range []string{exactPath, exactPath + hashSidecarSuffix, exactPath + targetSidecarSuffix, wrongPath} {
|
||||
Expect(os.WriteFile(path, []byte("data"), 0640)).To(Succeed())
|
||||
}
|
||||
|
||||
stager, stop := startReleaseServer(stagingDir, "release-token")
|
||||
DeferCleanup(stop)
|
||||
Expect(stager.ReleaseRemote(context.Background(), "node-1", key)).To(Succeed())
|
||||
|
||||
Expect(exactPath).NotTo(BeAnExistingFile())
|
||||
Expect(exactPath + hashSidecarSuffix).NotTo(BeAnExistingFile())
|
||||
Expect(exactPath + targetSidecarSuffix).NotTo(BeAnExistingFile())
|
||||
Expect(wrongPath).To(BeAnExistingFile())
|
||||
})
|
||||
|
||||
It("is idempotent and prunes empty category and request directories", func() {
|
||||
stagingDir := GinkgoT().TempDir()
|
||||
path := filepath.Join(stagingDir, "ephemeral", "request-id", "audio", "input.wav")
|
||||
Expect(os.MkdirAll(filepath.Dir(path), 0750)).To(Succeed())
|
||||
Expect(os.WriteFile(path, []byte("data"), 0640)).To(Succeed())
|
||||
|
||||
stager, stop := startReleaseServer(stagingDir, "release-token")
|
||||
DeferCleanup(stop)
|
||||
for range 2 {
|
||||
Expect(stager.ReleaseRemote(context.Background(), "node-1", "ephemeral/request-id/audio/input.wav")).To(Succeed())
|
||||
}
|
||||
|
||||
Expect(filepath.Join(stagingDir, "ephemeral", "request-id", "audio")).NotTo(BeADirectory())
|
||||
Expect(filepath.Join(stagingDir, "ephemeral", "request-id")).NotTo(BeADirectory())
|
||||
Expect(filepath.Join(stagingDir, "ephemeral")).To(BeADirectory())
|
||||
})
|
||||
|
||||
It("releases one request's HTTP inputs in one batch", func() {
|
||||
stagingDir := GinkgoT().TempDir()
|
||||
keys := []string{
|
||||
"ephemeral/audio/request-id/input.wav",
|
||||
"ephemeral/images/request-id/frame.jpg",
|
||||
}
|
||||
for _, key := range keys {
|
||||
path := filepath.Join(stagingDir, filepath.FromSlash(key))
|
||||
Expect(os.MkdirAll(filepath.Dir(path), 0750)).To(Succeed())
|
||||
Expect(os.WriteFile(path, []byte("data"), 0640)).To(Succeed())
|
||||
}
|
||||
|
||||
stager, stop := startReleaseServer(stagingDir, "release-token")
|
||||
DeferCleanup(stop)
|
||||
Expect(stager.ReleaseRemoteRequest(context.Background(), "node-1", "request-id", keys)).To(Succeed())
|
||||
|
||||
for _, key := range keys {
|
||||
Expect(filepath.Join(stagingDir, filepath.FromSlash(key))).NotTo(BeAnExistingFile())
|
||||
}
|
||||
})
|
||||
|
||||
It("falls back to exact HTTP releases for an older worker", func() {
|
||||
batchCalls := 0
|
||||
exactCalls := 0
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path == "/v1/files-release" {
|
||||
batchCalls++
|
||||
http.NotFound(w, r)
|
||||
return
|
||||
}
|
||||
if strings.HasPrefix(r.URL.Path, "/v1/files/") && r.Method == http.MethodDelete {
|
||||
exactCalls++
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
return
|
||||
}
|
||||
http.Error(w, "unexpected request", http.StatusInternalServerError)
|
||||
}))
|
||||
DeferCleanup(server.Close)
|
||||
stager := NewHTTPFileStager(func(string) (string, error) {
|
||||
return strings.TrimPrefix(server.URL, "http://"), nil
|
||||
}, "")
|
||||
keys := []string{
|
||||
"ephemeral/audio/request-id/input.wav",
|
||||
"ephemeral/images/request-id/frame.jpg",
|
||||
}
|
||||
|
||||
Expect(stager.ReleaseRemoteRequest(context.Background(), "node-1", "request-id", keys)).To(Succeed())
|
||||
|
||||
Expect(batchCalls).To(Equal(1))
|
||||
Expect(exactCalls).To(Equal(2))
|
||||
})
|
||||
|
||||
It("rejects non-ephemeral and traversing keys before making a request", func() {
|
||||
resolved := false
|
||||
stager := NewHTTPFileStager(func(string) (string, error) {
|
||||
resolved = true
|
||||
return "127.0.0.1:1", nil
|
||||
}, "token")
|
||||
|
||||
for _, key := range []string{
|
||||
"models/model.gguf",
|
||||
"ephemeral/../models/model.gguf",
|
||||
"ephemeral/request-id/../../model.gguf",
|
||||
"/ephemeral/request-id/audio/input.wav",
|
||||
} {
|
||||
Expect(stager.ReleaseRemote(context.Background(), "node-1", key)).NotTo(Succeed(), key)
|
||||
}
|
||||
Expect(resolved).To(BeFalse())
|
||||
})
|
||||
|
||||
It("rejects symlink escapes", func() {
|
||||
stagingDir := GinkgoT().TempDir()
|
||||
outsideDir := GinkgoT().TempDir()
|
||||
outsidePath := filepath.Join(outsideDir, "input.wav")
|
||||
Expect(os.WriteFile(outsidePath, []byte("keep"), 0640)).To(Succeed())
|
||||
requestDir := filepath.Join(stagingDir, "ephemeral", "request-id")
|
||||
Expect(os.MkdirAll(requestDir, 0750)).To(Succeed())
|
||||
Expect(os.Symlink(outsideDir, filepath.Join(requestDir, "audio"))).To(Succeed())
|
||||
|
||||
stager, stop := startReleaseServer(stagingDir, "release-token")
|
||||
DeferCleanup(stop)
|
||||
Expect(stager.ReleaseRemote(context.Background(), "node-1", "ephemeral/request-id/audio/input.wav")).NotTo(Succeed())
|
||||
Expect(outsidePath).To(BeAnExistingFile())
|
||||
})
|
||||
|
||||
It("rejects symlinked files and sidecars without deleting their targets", func() {
|
||||
for _, linkedName := range []string{"input.wav", "input.wav" + hashSidecarSuffix, "input.wav" + targetSidecarSuffix} {
|
||||
stagingDir := GinkgoT().TempDir()
|
||||
categoryDir := filepath.Join(stagingDir, "ephemeral", "request-id", "audio")
|
||||
Expect(os.MkdirAll(categoryDir, 0750)).To(Succeed())
|
||||
target := filepath.Join(categoryDir, "input.wav")
|
||||
if linkedName != "input.wav" {
|
||||
Expect(os.WriteFile(target, []byte("input"), 0640)).To(Succeed())
|
||||
}
|
||||
preserved := filepath.Join(stagingDir, "ephemeral", "preserved-"+linkedName)
|
||||
Expect(os.WriteFile(preserved, []byte("keep"), 0640)).To(Succeed())
|
||||
Expect(os.Symlink(preserved, filepath.Join(categoryDir, linkedName))).To(Succeed())
|
||||
|
||||
stager, stop := startReleaseServer(stagingDir, "release-token")
|
||||
Expect(stager.ReleaseRemote(context.Background(), "node-1", "ephemeral/request-id/audio/input.wav")).NotTo(Succeed(), linkedName)
|
||||
stop()
|
||||
Expect(preserved).To(BeAnExistingFile(), linkedName)
|
||||
}
|
||||
})
|
||||
|
||||
It("evicts the worker cache before deleting the shared object", func() {
|
||||
storeRoot := GinkgoT().TempDir()
|
||||
cacheRoot := GinkgoT().TempDir()
|
||||
store, err := storage.NewFilesystemStore(storeRoot)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
fm, err := storage.NewFileManager(store, cacheRoot)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
key := "ephemeral/request-id/audio/input.wav"
|
||||
Expect(store.Put(context.Background(), key, strings.NewReader("shared"))).To(Succeed())
|
||||
|
||||
client := &releaseTestMessaging{}
|
||||
client.onRequest = func() {
|
||||
exists, existsErr := store.Exists(context.Background(), key)
|
||||
Expect(existsErr).NotTo(HaveOccurred())
|
||||
Expect(exists).To(BeTrue())
|
||||
}
|
||||
stager := NewS3NATSFileStager(fm, client)
|
||||
Expect(stager.ReleaseRemote(context.Background(), "node.one", key)).To(Succeed())
|
||||
|
||||
Expect(client.requestCalled).To(BeTrue())
|
||||
Expect(client.subject).To(Equal(messaging.SubjectNodeFilesRelease("node.one")))
|
||||
var payload fileReleaseRequest
|
||||
Expect(json.Unmarshal(client.payload, &payload)).To(Succeed())
|
||||
Expect(payload.Key).To(Equal(key))
|
||||
exists, err := store.Exists(context.Background(), key)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(exists).To(BeFalse())
|
||||
})
|
||||
|
||||
It("evicts a request's S3 inputs with one NATS round trip", func() {
|
||||
storeRoot := GinkgoT().TempDir()
|
||||
store, err := storage.NewFilesystemStore(storeRoot)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
fm, err := storage.NewFileManager(store, GinkgoT().TempDir())
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
keys := []string{
|
||||
"ephemeral/audio/request-id/input.wav",
|
||||
"ephemeral/images/request-id/frame.jpg",
|
||||
}
|
||||
for _, key := range keys {
|
||||
Expect(store.Put(context.Background(), key, strings.NewReader("shared"))).To(Succeed())
|
||||
}
|
||||
client := &releaseTestMessaging{}
|
||||
stager := NewS3NATSFileStager(fm, client)
|
||||
|
||||
Expect(stager.ReleaseRemoteRequest(context.Background(), "node.one", "request-id", keys)).To(Succeed())
|
||||
|
||||
Expect(client.requestCount).To(Equal(1))
|
||||
var payload fileReleaseRequest
|
||||
Expect(json.Unmarshal(client.payload, &payload)).To(Succeed())
|
||||
Expect(payload.Key).To(BeEmpty())
|
||||
Expect(payload.RequestID).To(Equal("request-id"))
|
||||
for _, key := range keys {
|
||||
exists, existsErr := store.Exists(context.Background(), key)
|
||||
Expect(existsErr).NotTo(HaveOccurred())
|
||||
Expect(exists).To(BeFalse())
|
||||
}
|
||||
})
|
||||
|
||||
It("keeps worker coordination fixed-size for large requests", func() {
|
||||
store, err := storage.NewFilesystemStore(GinkgoT().TempDir())
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
fm, err := storage.NewFileManager(store, GinkgoT().TempDir())
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
keys := make([]string, 2048)
|
||||
for i := range keys {
|
||||
keys[i] = fmt.Sprintf("ephemeral/inputs/request-id/input-%d.bin", i)
|
||||
}
|
||||
client := &releaseTestMessaging{}
|
||||
stager := NewS3NATSFileStager(fm, client)
|
||||
|
||||
Expect(stager.ReleaseRemoteRequest(context.Background(), "node.one", "request-id", keys)).To(Succeed())
|
||||
|
||||
Expect(client.requestCount).To(Equal(1))
|
||||
Expect(len(client.payload)).To(BeNumerically("<", 128))
|
||||
var payload fileReleaseRequest
|
||||
Expect(json.Unmarshal(client.payload, &payload)).To(Succeed())
|
||||
Expect(payload.RequestID).To(Equal("request-id"))
|
||||
})
|
||||
|
||||
It("falls back to exact NATS releases for an older worker", func() {
|
||||
store, err := storage.NewFilesystemStore(GinkgoT().TempDir())
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
fm, err := storage.NewFileManager(store, GinkgoT().TempDir())
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
keys := []string{
|
||||
"ephemeral/audio/request-id/input.wav",
|
||||
"ephemeral/images/request-id/frame.jpg",
|
||||
}
|
||||
for _, key := range keys {
|
||||
Expect(store.Put(context.Background(), key, strings.NewReader("shared"))).To(Succeed())
|
||||
}
|
||||
client := &releaseTestMessaging{replies: [][]byte{
|
||||
[]byte(`{"error":"batch payload unsupported"}`),
|
||||
[]byte(`{}`),
|
||||
[]byte(`{}`),
|
||||
}}
|
||||
stager := NewS3NATSFileStager(fm, client)
|
||||
|
||||
Expect(stager.ReleaseRemoteRequest(context.Background(), "node.one", "request-id", keys)).To(Succeed())
|
||||
|
||||
Expect(client.requestCount).To(Equal(3))
|
||||
for _, key := range keys {
|
||||
exists, existsErr := store.Exists(context.Background(), key)
|
||||
Expect(existsErr).NotTo(HaveOccurred())
|
||||
Expect(exists).To(BeFalse())
|
||||
}
|
||||
})
|
||||
|
||||
It("rejects release batches that mix request IDs", func() {
|
||||
keys := []string{
|
||||
"ephemeral/audio/request-one/input.wav",
|
||||
"ephemeral/images/request-two/frame.jpg",
|
||||
}
|
||||
Expect(validateEphemeralRequestRelease("request-one", keys)).To(MatchError(ContainSubstring("mixes request IDs")))
|
||||
})
|
||||
|
||||
It("does not send a release request after cleanup is canceled", func() {
|
||||
store, err := storage.NewFilesystemStore(GinkgoT().TempDir())
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
fm, err := storage.NewFileManager(store, GinkgoT().TempDir())
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
client := &releaseTestMessaging{}
|
||||
stager := NewS3NATSFileStager(fm, client)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
Expect(stager.ReleaseRemote(ctx, "node.one", "ephemeral/request-id/audio/input.wav")).To(MatchError(context.Canceled))
|
||||
Expect(client.requestCalled).To(BeFalse())
|
||||
})
|
||||
|
||||
It("bounds the NATS release wait by the remaining cleanup deadline", func() {
|
||||
store, err := storage.NewFilesystemStore(GinkgoT().TempDir())
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
fm, err := storage.NewFileManager(store, GinkgoT().TempDir())
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
client := &releaseTestMessaging{}
|
||||
stager := NewS3NATSFileStager(fm, client)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel()
|
||||
|
||||
Expect(stager.ReleaseRemote(ctx, "node.one", "ephemeral/request-id/audio/input.wav")).To(Succeed())
|
||||
Expect(client.timeout).To(BeNumerically(">", time.Second))
|
||||
Expect(client.timeout).To(BeNumerically("<=", 2*time.Second))
|
||||
})
|
||||
})
|
||||
@@ -2,6 +2,7 @@ package nodes
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
@@ -48,6 +49,15 @@ type fileStageReply struct {
|
||||
Error string `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
type fileReleaseRequest struct {
|
||||
Key string `json:"key,omitempty"`
|
||||
RequestID string `json:"request_id,omitempty"`
|
||||
}
|
||||
|
||||
type fileReleaseReply struct {
|
||||
Error string `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
type fileTempRequest struct{}
|
||||
|
||||
type fileTempReply struct {
|
||||
@@ -181,3 +191,79 @@ func (s *S3NATSFileStager) StageRemoteToStore(ctx context.Context, nodeID, remot
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ReleaseRemote evicts one exact ephemeral key from the worker before deleting
|
||||
// the shared object.
|
||||
func (s *S3NATSFileStager) ReleaseRemote(ctx context.Context, nodeID, key string) error {
|
||||
if err := validateEphemeralReleaseKey(key); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.releaseWorkerKeys(ctx, nodeID, fileReleaseRequest{Key: key}); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.fm.Delete(ctx, key); err != nil {
|
||||
return fmt.Errorf("deleting shared object %q: %w", key, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ReleaseRemoteRequest evicts one inference's inputs with one NATS round trip.
|
||||
// A worker that only understands the exact-key payload returns an error, so the
|
||||
// frontend retries each key during a rolling upgrade.
|
||||
func (s *S3NATSFileStager) ReleaseRemoteRequest(ctx context.Context, nodeID, requestID string, keys []string) error {
|
||||
if err := validateEphemeralRequestRelease(requestID, keys); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.releaseWorkerKeys(ctx, nodeID, fileReleaseRequest{RequestID: requestID}); err != nil {
|
||||
var fallbackErrors []error
|
||||
for _, key := range keys {
|
||||
if fallbackErr := s.ReleaseRemote(ctx, nodeID, key); fallbackErr != nil {
|
||||
fallbackErrors = append(fallbackErrors, fallbackErr)
|
||||
}
|
||||
}
|
||||
if fallbackErr := errors.Join(fallbackErrors...); fallbackErr != nil {
|
||||
return errors.Join(err, fmt.Errorf("exact-key release fallback: %w", fallbackErr))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
var deleteErrors []error
|
||||
for _, key := range keys {
|
||||
if err := s.fm.Delete(ctx, key); err != nil {
|
||||
deleteErrors = append(deleteErrors, fmt.Errorf("deleting shared object %q: %w", key, err))
|
||||
}
|
||||
}
|
||||
return errors.Join(deleteErrors...)
|
||||
}
|
||||
|
||||
func (s *S3NATSFileStager) releaseWorkerKeys(ctx context.Context, nodeID string, request fileReleaseRequest) error {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
timeout := 30 * time.Second
|
||||
if deadline, ok := ctx.Deadline(); ok {
|
||||
remaining := time.Until(deadline)
|
||||
if remaining <= 0 {
|
||||
return context.DeadlineExceeded
|
||||
}
|
||||
timeout = min(timeout, remaining)
|
||||
}
|
||||
reply, err := messaging.RequestJSON[fileReleaseRequest, fileReleaseReply](
|
||||
s.nats,
|
||||
messaging.SubjectNodeFilesRelease(nodeID),
|
||||
request,
|
||||
timeout,
|
||||
)
|
||||
if err != nil {
|
||||
if contextErr := ctx.Err(); contextErr != nil {
|
||||
return contextErr
|
||||
}
|
||||
return err
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
if reply.Error != "" {
|
||||
return fmt.Errorf("backend release failed: %s", reply.Error)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -2,6 +2,7 @@ package nodes
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
@@ -19,6 +20,8 @@ import (
|
||||
"google.golang.org/protobuf/proto"
|
||||
)
|
||||
|
||||
const stagedInputReleaseTimeout = 30 * time.Second
|
||||
|
||||
// FileStagingClient wraps a grpc.Backend to transparently handle file transfer
|
||||
// for distributed mode. Input files are staged on the backend node before the
|
||||
// gRPC call. Output files are retrieved from the backend after the call.
|
||||
@@ -49,21 +52,70 @@ func NewFileStagingClient(inner grpc.Backend, stager FileStager, nodeID string)
|
||||
|
||||
// requestID generates a unique ID for ephemeral file keys.
|
||||
func requestID() string {
|
||||
return uuid.New().String()[:8]
|
||||
return uuid.NewString()
|
||||
}
|
||||
|
||||
type stagedInputLifecycle struct {
|
||||
client *FileStagingClient
|
||||
requestID string
|
||||
keys []string
|
||||
seen map[string]struct{}
|
||||
}
|
||||
|
||||
func (f *FileStagingClient) newStagedInputLifecycle() *stagedInputLifecycle {
|
||||
return &stagedInputLifecycle{
|
||||
client: f,
|
||||
requestID: requestID(),
|
||||
keys: []string{},
|
||||
seen: map[string]struct{}{},
|
||||
}
|
||||
}
|
||||
|
||||
func (l *stagedInputLifecycle) track(key string) {
|
||||
if _, ok := l.seen[key]; ok {
|
||||
return
|
||||
}
|
||||
l.seen[key] = struct{}{}
|
||||
l.keys = append(l.keys, key)
|
||||
}
|
||||
|
||||
func (l *stagedInputLifecycle) release() {
|
||||
if len(l.keys) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), stagedInputReleaseTimeout)
|
||||
defer cancel()
|
||||
if releaser, ok := l.client.stager.(RequestFileReleaser); ok {
|
||||
if err := releaser.ReleaseRemoteRequest(ctx, l.client.nodeID, l.requestID, l.keys); err != nil {
|
||||
xlog.Warn("Failed to release staged request inputs", "node", l.client.nodeID, "requestID", l.requestID, "keyCount", len(l.keys), "error", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
for _, key := range l.keys {
|
||||
if err := l.client.stager.ReleaseRemote(ctx, l.client.nodeID, key); err != nil {
|
||||
xlog.Warn("Failed to release staged input", "node", l.client.nodeID, "key", key, "error", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// stageInputFile uploads a local file to the remote node via the FileStager.
|
||||
// Returns the remote-local path and the ephemeral key.
|
||||
func (f *FileStagingClient) stageInputFile(ctx context.Context, reqID, localPath, category string) (string, string, error) {
|
||||
func (f *FileStagingClient) stageInputFile(
|
||||
ctx context.Context,
|
||||
lifecycle *stagedInputLifecycle,
|
||||
localPath,
|
||||
category string,
|
||||
) (string, error) {
|
||||
basename := filepath.Base(localPath)
|
||||
key := storage.EphemeralKey(reqID, category, basename)
|
||||
key := storage.EphemeralKey(lifecycle.requestID, category, basename)
|
||||
lifecycle.track(key)
|
||||
|
||||
remotePath, err := f.stager.EnsureRemote(ctx, f.nodeID, localPath, key)
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("staging input file: %w", err)
|
||||
return "", fmt.Errorf("staging input file: %w", err)
|
||||
}
|
||||
|
||||
return remotePath, key, nil
|
||||
return remotePath, nil
|
||||
}
|
||||
|
||||
// retrieveOutputFile retrieves an output file from the backend to a local path.
|
||||
@@ -101,23 +153,37 @@ func (f *FileStagingClient) translateModelPath(frontendPath string) string {
|
||||
}
|
||||
|
||||
func (f *FileStagingClient) Predict(ctx context.Context, in *pb.PredictOptions, opts ...ggrpc.CallOption) (*pb.Reply, error) {
|
||||
reqID := requestID()
|
||||
in, _ = f.stageMultimodalInputs(ctx, reqID, in)
|
||||
lifecycle := f.newStagedInputLifecycle()
|
||||
defer lifecycle.release()
|
||||
in = proto.Clone(in).(*pb.PredictOptions)
|
||||
var err error
|
||||
in, err = f.stageMultimodalInputs(ctx, lifecycle, in)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return f.Backend.Predict(ctx, in, opts...)
|
||||
}
|
||||
|
||||
func (f *FileStagingClient) PredictStream(ctx context.Context, in *pb.PredictOptions, fn func(reply *pb.Reply), opts ...ggrpc.CallOption) error {
|
||||
reqID := requestID()
|
||||
in, _ = f.stageMultimodalInputs(ctx, reqID, in)
|
||||
lifecycle := f.newStagedInputLifecycle()
|
||||
defer lifecycle.release()
|
||||
in = proto.Clone(in).(*pb.PredictOptions)
|
||||
var err error
|
||||
in, err = f.stageMultimodalInputs(ctx, lifecycle, in)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return f.Backend.PredictStream(ctx, in, fn, opts...)
|
||||
}
|
||||
|
||||
func (f *FileStagingClient) GenerateImage(ctx context.Context, in *pb.GenerateImageRequest, opts ...ggrpc.CallOption) (*pb.Result, error) {
|
||||
reqID := requestID()
|
||||
lifecycle := f.newStagedInputLifecycle()
|
||||
defer lifecycle.release()
|
||||
in = proto.Clone(in).(*pb.GenerateImageRequest)
|
||||
|
||||
// Stage input source image if present
|
||||
if in.Src != "" && isFilePath(in.Src) {
|
||||
backendPath, _, err := f.stageInputFile(ctx, reqID, in.Src, "inputs")
|
||||
backendPath, err := f.stageInputFile(ctx, lifecycle, in.Src, "inputs")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("staging image src: %w", err)
|
||||
}
|
||||
@@ -127,7 +193,7 @@ func (f *FileStagingClient) GenerateImage(ctx context.Context, in *pb.GenerateIm
|
||||
// Stage reference images
|
||||
for i, img := range in.RefImages {
|
||||
if isFilePath(img) {
|
||||
backendPath, _, err := f.stageInputFile(ctx, reqID, img, "inputs")
|
||||
backendPath, err := f.stageInputFile(ctx, lifecycle, img, "inputs")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("staging ref image: %w", err)
|
||||
}
|
||||
@@ -161,25 +227,27 @@ func (f *FileStagingClient) GenerateImage(ctx context.Context, in *pb.GenerateIm
|
||||
}
|
||||
|
||||
func (f *FileStagingClient) GenerateVideo(ctx context.Context, in *pb.GenerateVideoRequest, opts ...ggrpc.CallOption) (*pb.Result, error) {
|
||||
reqID := requestID()
|
||||
lifecycle := f.newStagedInputLifecycle()
|
||||
defer lifecycle.release()
|
||||
in = proto.Clone(in).(*pb.GenerateVideoRequest)
|
||||
|
||||
// Stage start/end images and optional audio conditioning.
|
||||
if in.StartImage != "" && isFilePath(in.StartImage) {
|
||||
backendPath, _, err := f.stageInputFile(ctx, reqID, in.StartImage, "inputs")
|
||||
backendPath, err := f.stageInputFile(ctx, lifecycle, in.StartImage, "inputs")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("staging start image: %w", err)
|
||||
}
|
||||
in.StartImage = backendPath
|
||||
}
|
||||
if in.EndImage != "" && isFilePath(in.EndImage) {
|
||||
backendPath, _, err := f.stageInputFile(ctx, reqID, in.EndImage, "inputs")
|
||||
backendPath, err := f.stageInputFile(ctx, lifecycle, in.EndImage, "inputs")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("staging end image: %w", err)
|
||||
}
|
||||
in.EndImage = backendPath
|
||||
}
|
||||
if in.Audio != "" && isFilePath(in.Audio) {
|
||||
backendPath, _, err := f.stageInputFile(ctx, reqID, in.Audio, "inputs")
|
||||
backendPath, err := f.stageInputFile(ctx, lifecycle, in.Audio, "inputs")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("staging video audio: %w", err)
|
||||
}
|
||||
@@ -211,11 +279,13 @@ func (f *FileStagingClient) GenerateVideo(ctx context.Context, in *pb.GenerateVi
|
||||
}
|
||||
|
||||
func (f *FileStagingClient) Generate3D(ctx context.Context, in *pb.Generate3DRequest, opts ...ggrpc.CallOption) (*pb.Result, error) {
|
||||
reqID := requestID()
|
||||
lifecycle := f.newStagedInputLifecycle()
|
||||
defer lifecycle.release()
|
||||
in = proto.Clone(in).(*pb.Generate3DRequest)
|
||||
|
||||
// Stage the conditioning image or existing GLB used by 3D post-processing.
|
||||
if in.Src != "" && isFilePath(in.Src) {
|
||||
backendPath, _, err := f.stageInputFile(ctx, reqID, in.Src, "inputs")
|
||||
backendPath, err := f.stageInputFile(ctx, lifecycle, in.Src, "inputs")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("staging 3D input asset: %w", err)
|
||||
}
|
||||
@@ -247,7 +317,9 @@ func (f *FileStagingClient) Generate3D(ctx context.Context, in *pb.Generate3DReq
|
||||
}
|
||||
|
||||
func (f *FileStagingClient) TTS(ctx context.Context, in *pb.TTSRequest, opts ...ggrpc.CallOption) (*pb.Result, error) {
|
||||
reqID := requestID()
|
||||
lifecycle := f.newStagedInputLifecycle()
|
||||
defer lifecycle.release()
|
||||
in = proto.Clone(in).(*pb.TTSRequest)
|
||||
|
||||
// Translate model path from frontend to remote worker path.
|
||||
// The model and its companion files (e.g. .onnx.json) were already staged
|
||||
@@ -258,7 +330,7 @@ func (f *FileStagingClient) TTS(ctx context.Context, in *pb.TTSRequest, opts ...
|
||||
// Voice may be a named backend speaker or a request-scoped reference WAV.
|
||||
// Only path-shaped values are staged; speaker IDs pass through unchanged.
|
||||
if in.Voice != "" && isFilePath(in.Voice) {
|
||||
backendPath, _, err := f.stageInputFile(ctx, reqID, in.Voice, "inputs")
|
||||
backendPath, err := f.stageInputFile(ctx, lifecycle, in.Voice, "inputs")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("staging TTS voice reference: %w", err)
|
||||
}
|
||||
@@ -290,14 +362,16 @@ func (f *FileStagingClient) TTS(ctx context.Context, in *pb.TTSRequest, opts ...
|
||||
}
|
||||
|
||||
func (f *FileStagingClient) TTSStream(ctx context.Context, in *pb.TTSRequest, fn func(*pb.Reply), opts ...ggrpc.CallOption) error {
|
||||
reqID := requestID()
|
||||
lifecycle := f.newStagedInputLifecycle()
|
||||
defer lifecycle.release()
|
||||
in = proto.Clone(in).(*pb.TTSRequest)
|
||||
|
||||
// Translate model path from frontend to remote worker path (same as TTS above)
|
||||
if in.Model != "" && isFilePath(in.Model) {
|
||||
in.Model = f.translateModelPath(in.Model)
|
||||
}
|
||||
if in.Voice != "" && isFilePath(in.Voice) {
|
||||
backendPath, _, err := f.stageInputFile(ctx, reqID, in.Voice, "inputs")
|
||||
backendPath, err := f.stageInputFile(ctx, lifecycle, in.Voice, "inputs")
|
||||
if err != nil {
|
||||
return fmt.Errorf("staging streaming TTS voice reference: %w", err)
|
||||
}
|
||||
@@ -308,11 +382,13 @@ func (f *FileStagingClient) TTSStream(ctx context.Context, in *pb.TTSRequest, fn
|
||||
}
|
||||
|
||||
func (f *FileStagingClient) SoundGeneration(ctx context.Context, in *pb.SoundGenerationRequest, opts ...ggrpc.CallOption) (*pb.Result, error) {
|
||||
reqID := requestID()
|
||||
lifecycle := f.newStagedInputLifecycle()
|
||||
defer lifecycle.release()
|
||||
in = proto.Clone(in).(*pb.SoundGenerationRequest)
|
||||
|
||||
// Stage input source
|
||||
if in.Src != nil && *in.Src != "" && isFilePath(*in.Src) {
|
||||
backendPath, _, err := f.stageInputFile(ctx, reqID, *in.Src, "inputs")
|
||||
backendPath, err := f.stageInputFile(ctx, lifecycle, *in.Src, "inputs")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("staging sound src: %w", err)
|
||||
}
|
||||
@@ -344,24 +420,27 @@ func (f *FileStagingClient) SoundGeneration(ctx context.Context, in *pb.SoundGen
|
||||
}
|
||||
|
||||
func (f *FileStagingClient) SoundDetection(ctx context.Context, in *pb.SoundDetectionRequest, opts ...ggrpc.CallOption) (*pb.SoundDetectionResponse, error) {
|
||||
lifecycle := f.newStagedInputLifecycle()
|
||||
defer lifecycle.release()
|
||||
in = proto.Clone(in).(*pb.SoundDetectionRequest)
|
||||
if in.Src != "" && isFilePath(in.Src) {
|
||||
backendPath, _, err := f.stageInputFile(ctx, requestID(), in.Src, "inputs")
|
||||
backendPath, err := f.stageInputFile(ctx, lifecycle, in.Src, "inputs")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("staging audio for sound detection: %w", err)
|
||||
}
|
||||
// Keep the frontend path available if the caller retries on another node.
|
||||
in = proto.Clone(in).(*pb.SoundDetectionRequest)
|
||||
in.Src = backendPath
|
||||
}
|
||||
return f.Backend.SoundDetection(ctx, in, opts...)
|
||||
}
|
||||
|
||||
func (f *FileStagingClient) AudioTranscription(ctx context.Context, in *pb.TranscriptRequest, opts ...ggrpc.CallOption) (*pb.TranscriptResult, error) {
|
||||
reqID := requestID()
|
||||
lifecycle := f.newStagedInputLifecycle()
|
||||
defer lifecycle.release()
|
||||
in = proto.Clone(in).(*pb.TranscriptRequest)
|
||||
|
||||
// Stage input audio file
|
||||
if in.Dst != "" && isFilePath(in.Dst) {
|
||||
backendPath, _, err := f.stageInputFile(ctx, reqID, in.Dst, "inputs")
|
||||
backendPath, err := f.stageInputFile(ctx, lifecycle, in.Dst, "inputs")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("staging audio for transcription: %w", err)
|
||||
}
|
||||
@@ -372,11 +451,13 @@ func (f *FileStagingClient) AudioTranscription(ctx context.Context, in *pb.Trans
|
||||
}
|
||||
|
||||
func (f *FileStagingClient) AudioTranscriptionStream(ctx context.Context, in *pb.TranscriptRequest, fn func(chunk *pb.TranscriptStreamResponse), opts ...ggrpc.CallOption) error {
|
||||
reqID := requestID()
|
||||
lifecycle := f.newStagedInputLifecycle()
|
||||
defer lifecycle.release()
|
||||
in = proto.Clone(in).(*pb.TranscriptRequest)
|
||||
|
||||
// Stage input audio file
|
||||
if in.Dst != "" && isFilePath(in.Dst) {
|
||||
backendPath, _, err := f.stageInputFile(ctx, reqID, in.Dst, "inputs")
|
||||
backendPath, err := f.stageInputFile(ctx, lifecycle, in.Dst, "inputs")
|
||||
if err != nil {
|
||||
return fmt.Errorf("staging audio for transcription stream: %w", err)
|
||||
}
|
||||
@@ -458,31 +539,46 @@ func (f *FileStagingClient) QuantizationProgress(ctx context.Context, in *pb.Qua
|
||||
|
||||
// stageMultimodalInputs stages Images, Videos, Audios fields in PredictOptions
|
||||
// if they are file paths (not base64 or URLs).
|
||||
func (f *FileStagingClient) stageMultimodalInputs(ctx context.Context, reqID string, in *pb.PredictOptions) (*pb.PredictOptions, []string) {
|
||||
var keys []string
|
||||
in.Images = f.stagePathSlice(ctx, reqID, in.Images, "inputs", &keys)
|
||||
in.Videos = f.stagePathSlice(ctx, reqID, in.Videos, "inputs", &keys)
|
||||
in.Audios = f.stagePathSlice(ctx, reqID, in.Audios, "inputs", &keys)
|
||||
return in, keys
|
||||
func (f *FileStagingClient) stageMultimodalInputs(
|
||||
ctx context.Context,
|
||||
lifecycle *stagedInputLifecycle,
|
||||
in *pb.PredictOptions,
|
||||
) (*pb.PredictOptions, error) {
|
||||
var err error
|
||||
in.Images, err = f.stagePathSlice(ctx, lifecycle, in.Images, "inputs")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("staging predict images: %w", err)
|
||||
}
|
||||
in.Videos, err = f.stagePathSlice(ctx, lifecycle, in.Videos, "inputs")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("staging predict videos: %w", err)
|
||||
}
|
||||
in.Audios, err = f.stagePathSlice(ctx, lifecycle, in.Audios, "inputs")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("staging predict audios: %w", err)
|
||||
}
|
||||
return in, nil
|
||||
}
|
||||
|
||||
func (f *FileStagingClient) stagePathSlice(ctx context.Context, reqID string, paths []string, category string, keys *[]string) []string {
|
||||
func (f *FileStagingClient) stagePathSlice(
|
||||
ctx context.Context,
|
||||
lifecycle *stagedInputLifecycle,
|
||||
paths []string,
|
||||
category string,
|
||||
) ([]string, error) {
|
||||
result := make([]string, len(paths))
|
||||
for i, p := range paths {
|
||||
if isFilePath(p) {
|
||||
backendPath, key, err := f.stageInputFile(ctx, reqID, p, category)
|
||||
backendPath, err := f.stageInputFile(ctx, lifecycle, p, category)
|
||||
if err != nil {
|
||||
xlog.Warn("Failed to stage multimodal file, passing through", "path", p, "error", err)
|
||||
result[i] = p
|
||||
continue
|
||||
return nil, fmt.Errorf("staging %q: %w", p, err)
|
||||
}
|
||||
result[i] = backendPath
|
||||
*keys = append(*keys, key)
|
||||
} else {
|
||||
result[i] = p
|
||||
}
|
||||
}
|
||||
return result
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// isFilePath checks if a string looks like a local file path (not base64 or URL).
|
||||
@@ -498,10 +594,27 @@ func isFilePath(s string) bool {
|
||||
if strings.HasPrefix(s, "http://") || strings.HasPrefix(s, "https://") {
|
||||
return false
|
||||
}
|
||||
// Raw JPEG base64 begins with /9j because every JPEG starts with the
|
||||
// FF D8 FF marker. Do not mistake that leading slash for an absolute path.
|
||||
if isRawJPEGBase64(s) {
|
||||
return false
|
||||
}
|
||||
// Starts with / (absolute path) or contains path separator
|
||||
return s[0] == '/' || filepath.IsAbs(s)
|
||||
}
|
||||
|
||||
func isRawJPEGBase64(s string) bool {
|
||||
if len(s) < 4 {
|
||||
return false
|
||||
}
|
||||
prefix, err := base64.StdEncoding.DecodeString(s[:4])
|
||||
if err != nil || len(prefix) != 3 || prefix[0] != 0xff || prefix[1] != 0xd8 || prefix[2] != 0xff {
|
||||
return false
|
||||
}
|
||||
_, err = io.Copy(io.Discard, base64.NewDecoder(base64.StdEncoding, strings.NewReader(s)))
|
||||
return err == nil
|
||||
}
|
||||
|
||||
// copyFile copies src to dst.
|
||||
func copyFile(src, dst string) error {
|
||||
if err := os.MkdirAll(filepath.Dir(dst), 0750); err != nil {
|
||||
|
||||
@@ -0,0 +1,358 @@
|
||||
package nodes
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
|
||||
grpc "github.com/mudler/LocalAI/pkg/grpc"
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
ggrpc "google.golang.org/grpc"
|
||||
"google.golang.org/protobuf/proto"
|
||||
)
|
||||
|
||||
const fullUUIDPattern = `[0-9a-f]{8}-[0-9a-f]{4}-4[0-9a-f]{3}-[89ab][0-9a-f]{3}-[0-9a-f]{12}`
|
||||
|
||||
type lifecycleStager struct {
|
||||
fakeFileStager
|
||||
ensureErr error
|
||||
ensureErrAt int
|
||||
releaseErr error
|
||||
releasedKeys []string
|
||||
releaseBatches [][]string
|
||||
releaseCtxErr []error
|
||||
releaseHasDeadline []bool
|
||||
releaseDeadlines []time.Time
|
||||
}
|
||||
|
||||
func (s *lifecycleStager) EnsureRemote(ctx context.Context, nodeID, localPath, key string) (string, error) {
|
||||
s.fakeFileStager.EnsureRemote(ctx, nodeID, localPath, key)
|
||||
if s.ensureErr != nil && (s.ensureErrAt == 0 || len(s.ensureCalls) == s.ensureErrAt) {
|
||||
return "", s.ensureErr
|
||||
}
|
||||
return "/remote/" + key, nil
|
||||
}
|
||||
|
||||
func (s *lifecycleStager) ReleaseRemote(ctx context.Context, _ string, key string) error {
|
||||
s.releasedKeys = append(s.releasedKeys, key)
|
||||
s.releaseCtxErr = append(s.releaseCtxErr, ctx.Err())
|
||||
deadline, ok := ctx.Deadline()
|
||||
s.releaseHasDeadline = append(s.releaseHasDeadline, ok)
|
||||
s.releaseDeadlines = append(s.releaseDeadlines, deadline)
|
||||
return s.releaseErr
|
||||
}
|
||||
|
||||
func (s *lifecycleStager) ReleaseRemoteRequest(ctx context.Context, _, _ string, keys []string) error {
|
||||
s.releaseBatches = append(s.releaseBatches, append([]string(nil), keys...))
|
||||
s.releasedKeys = append(s.releasedKeys, keys...)
|
||||
s.releaseCtxErr = append(s.releaseCtxErr, ctx.Err())
|
||||
deadline, ok := ctx.Deadline()
|
||||
s.releaseHasDeadline = append(s.releaseHasDeadline, ok)
|
||||
s.releaseDeadlines = append(s.releaseDeadlines, deadline)
|
||||
return s.releaseErr
|
||||
}
|
||||
|
||||
type lifecycleBackend struct {
|
||||
grpc.Backend
|
||||
predictResult *pb.Reply
|
||||
predictErr error
|
||||
predictCalls int
|
||||
predictInput *pb.PredictOptions
|
||||
streamCalls int
|
||||
streamBlock <-chan struct{}
|
||||
streamStarted chan<- struct{}
|
||||
}
|
||||
|
||||
func (b *lifecycleBackend) Predict(_ context.Context, in *pb.PredictOptions, _ ...ggrpc.CallOption) (*pb.Reply, error) {
|
||||
b.predictCalls++
|
||||
b.predictInput = proto.Clone(in).(*pb.PredictOptions)
|
||||
if b.predictResult == nil {
|
||||
b.predictResult = &pb.Reply{}
|
||||
}
|
||||
return b.predictResult, b.predictErr
|
||||
}
|
||||
|
||||
func (b *lifecycleBackend) PredictStream(_ context.Context, _ *pb.PredictOptions, _ func(*pb.Reply), _ ...ggrpc.CallOption) error {
|
||||
b.streamCalls++
|
||||
if b.streamStarted != nil {
|
||||
b.streamStarted <- struct{}{}
|
||||
}
|
||||
if b.streamBlock != nil {
|
||||
<-b.streamBlock
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (b *lifecycleBackend) GenerateImage(_ context.Context, _ *pb.GenerateImageRequest, _ ...ggrpc.CallOption) (*pb.Result, error) {
|
||||
return &pb.Result{Success: true}, nil
|
||||
}
|
||||
|
||||
func (b *lifecycleBackend) GenerateVideo(_ context.Context, _ *pb.GenerateVideoRequest, _ ...ggrpc.CallOption) (*pb.Result, error) {
|
||||
return &pb.Result{Success: true}, nil
|
||||
}
|
||||
|
||||
func (b *lifecycleBackend) Generate3D(_ context.Context, _ *pb.Generate3DRequest, _ ...ggrpc.CallOption) (*pb.Result, error) {
|
||||
return &pb.Result{Success: true}, nil
|
||||
}
|
||||
|
||||
func (b *lifecycleBackend) TTS(_ context.Context, _ *pb.TTSRequest, _ ...ggrpc.CallOption) (*pb.Result, error) {
|
||||
return &pb.Result{Success: true}, nil
|
||||
}
|
||||
|
||||
func (b *lifecycleBackend) TTSStream(_ context.Context, _ *pb.TTSRequest, _ func(*pb.Reply), _ ...ggrpc.CallOption) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (b *lifecycleBackend) SoundGeneration(_ context.Context, _ *pb.SoundGenerationRequest, _ ...ggrpc.CallOption) (*pb.Result, error) {
|
||||
return &pb.Result{Success: true}, nil
|
||||
}
|
||||
|
||||
func (b *lifecycleBackend) SoundDetection(_ context.Context, _ *pb.SoundDetectionRequest, _ ...ggrpc.CallOption) (*pb.SoundDetectionResponse, error) {
|
||||
return &pb.SoundDetectionResponse{}, nil
|
||||
}
|
||||
|
||||
func (b *lifecycleBackend) AudioTranscription(_ context.Context, _ *pb.TranscriptRequest, _ ...ggrpc.CallOption) (*pb.TranscriptResult, error) {
|
||||
return &pb.TranscriptResult{}, nil
|
||||
}
|
||||
|
||||
func (b *lifecycleBackend) AudioTranscriptionStream(_ context.Context, _ *pb.TranscriptRequest, _ func(*pb.TranscriptStreamResponse), _ ...ggrpc.CallOption) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
var _ = Describe("FileStagingClient request lifecycle", func() {
|
||||
It("uses a full UUID for ephemeral request keys", func() {
|
||||
Expect(requestID()).To(MatchRegexp(`^` + fullUUIDPattern + `$`))
|
||||
})
|
||||
|
||||
It("passes raw JPEG base64 to predict without staging it as a path", func(ctx SpecContext) {
|
||||
const jpegBase64 = "/9j/2Q=="
|
||||
backend := &lifecycleBackend{}
|
||||
stager := &lifecycleStager{}
|
||||
client := NewFileStagingClient(backend, stager, "worker-1")
|
||||
|
||||
_, err := client.Predict(ctx, &pb.PredictOptions{Images: []string{jpegBase64}})
|
||||
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(stager.ensureCalls).To(BeEmpty())
|
||||
Expect(backend.predictInput.Images).To(Equal([]string{jpegBase64}))
|
||||
})
|
||||
|
||||
It("still recognizes an invalid JPEG-like base64 string as a path", func() {
|
||||
Expect(isFilePath("/9j/not-a-jpeg")).To(BeTrue())
|
||||
})
|
||||
|
||||
It("releases every staged key and preserves caller requests", func(ctx SpecContext) {
|
||||
tests := []struct {
|
||||
name string
|
||||
keyCount int
|
||||
invoke func(*FileStagingClient) proto.Message
|
||||
}{
|
||||
{name: "predict", keyCount: 3, invoke: func(client *FileStagingClient) proto.Message {
|
||||
request := &pb.PredictOptions{Images: []string{"/tmp/image.png"}, Videos: []string{"/tmp/video.mp4"}, Audios: []string{"/tmp/audio.wav"}}
|
||||
original := proto.Clone(request)
|
||||
_, err := client.Predict(ctx, request)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(proto.Equal(request, original)).To(BeTrue())
|
||||
return request
|
||||
}},
|
||||
{name: "predict stream", keyCount: 1, invoke: func(client *FileStagingClient) proto.Message {
|
||||
request := &pb.PredictOptions{Images: []string{"/tmp/image.png"}}
|
||||
original := proto.Clone(request)
|
||||
Expect(client.PredictStream(ctx, request, func(*pb.Reply) {})).To(Succeed())
|
||||
Expect(proto.Equal(request, original)).To(BeTrue())
|
||||
return request
|
||||
}},
|
||||
{name: "image generation", keyCount: 2, invoke: func(client *FileStagingClient) proto.Message {
|
||||
request := &pb.GenerateImageRequest{Src: "/tmp/source.png", RefImages: []string{"/tmp/reference.png"}}
|
||||
original := proto.Clone(request)
|
||||
_, err := client.GenerateImage(ctx, request)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(proto.Equal(request, original)).To(BeTrue())
|
||||
return request
|
||||
}},
|
||||
{name: "video generation", keyCount: 3, invoke: func(client *FileStagingClient) proto.Message {
|
||||
request := &pb.GenerateVideoRequest{StartImage: "/tmp/start.png", EndImage: "/tmp/end.png", Audio: "/tmp/audio.wav"}
|
||||
original := proto.Clone(request)
|
||||
_, err := client.GenerateVideo(ctx, request)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(proto.Equal(request, original)).To(BeTrue())
|
||||
return request
|
||||
}},
|
||||
{name: "3D generation", keyCount: 1, invoke: func(client *FileStagingClient) proto.Message {
|
||||
request := &pb.Generate3DRequest{Src: "/tmp/source.glb"}
|
||||
original := proto.Clone(request)
|
||||
_, err := client.Generate3D(ctx, request)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(proto.Equal(request, original)).To(BeTrue())
|
||||
return request
|
||||
}},
|
||||
{name: "TTS", keyCount: 1, invoke: func(client *FileStagingClient) proto.Message {
|
||||
request := &pb.TTSRequest{Voice: "/tmp/voice.wav"}
|
||||
original := proto.Clone(request)
|
||||
_, err := client.TTS(ctx, request)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(proto.Equal(request, original)).To(BeTrue())
|
||||
return request
|
||||
}},
|
||||
{name: "streaming TTS", keyCount: 1, invoke: func(client *FileStagingClient) proto.Message {
|
||||
request := &pb.TTSRequest{Voice: "/tmp/voice.wav"}
|
||||
original := proto.Clone(request)
|
||||
Expect(client.TTSStream(ctx, request, func(*pb.Reply) {})).To(Succeed())
|
||||
Expect(proto.Equal(request, original)).To(BeTrue())
|
||||
return request
|
||||
}},
|
||||
{name: "sound generation", keyCount: 1, invoke: func(client *FileStagingClient) proto.Message {
|
||||
source := "/tmp/source.wav"
|
||||
request := &pb.SoundGenerationRequest{Src: &source}
|
||||
original := proto.Clone(request)
|
||||
_, err := client.SoundGeneration(ctx, request)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(proto.Equal(request, original)).To(BeTrue())
|
||||
return request
|
||||
}},
|
||||
{name: "sound detection", keyCount: 1, invoke: func(client *FileStagingClient) proto.Message {
|
||||
request := &pb.SoundDetectionRequest{Src: "/tmp/source.wav"}
|
||||
original := proto.Clone(request)
|
||||
_, err := client.SoundDetection(ctx, request)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(proto.Equal(request, original)).To(BeTrue())
|
||||
return request
|
||||
}},
|
||||
{name: "transcription", keyCount: 1, invoke: func(client *FileStagingClient) proto.Message {
|
||||
request := &pb.TranscriptRequest{Dst: "/tmp/source.wav"}
|
||||
original := proto.Clone(request)
|
||||
_, err := client.AudioTranscription(ctx, request)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(proto.Equal(request, original)).To(BeTrue())
|
||||
return request
|
||||
}},
|
||||
{name: "streaming transcription", keyCount: 1, invoke: func(client *FileStagingClient) proto.Message {
|
||||
request := &pb.TranscriptRequest{Dst: "/tmp/source.wav"}
|
||||
original := proto.Clone(request)
|
||||
Expect(client.AudioTranscriptionStream(ctx, request, func(*pb.TranscriptStreamResponse) {})).To(Succeed())
|
||||
Expect(proto.Equal(request, original)).To(BeTrue())
|
||||
return request
|
||||
}},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
By(test.name)
|
||||
stager := &lifecycleStager{}
|
||||
client := NewFileStagingClient(&lifecycleBackend{}, stager, "worker-1")
|
||||
test.invoke(client)
|
||||
Expect(stager.ensureCalls).To(HaveLen(test.keyCount))
|
||||
Expect(stager.releasedKeys).To(Equal(keysFromEnsureCalls(stager.ensureCalls)))
|
||||
Expect(stager.releaseBatches).To(Equal([][]string{keysFromEnsureCalls(stager.ensureCalls)}))
|
||||
Expect(stager.releaseCtxErr).To(Equal([]error{nil}))
|
||||
Expect(stager.releaseHasDeadline).To(HaveLen(1))
|
||||
for _, hasDeadline := range stager.releaseHasDeadline {
|
||||
Expect(hasDeadline).To(BeTrue())
|
||||
}
|
||||
for _, deadline := range stager.releaseDeadlines {
|
||||
Expect(time.Until(deadline)).To(BeNumerically(">", 0))
|
||||
Expect(time.Until(deadline)).To(BeNumerically("<=", time.Minute))
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
It("tracks a key before staging so a partial upload failure is released", func(ctx SpecContext) {
|
||||
uploadErr := errors.New("upload failed")
|
||||
stager := &lifecycleStager{ensureErr: uploadErr, releaseErr: errors.New("release failed")}
|
||||
client := NewFileStagingClient(&lifecycleBackend{}, stager, "worker-1")
|
||||
|
||||
_, err := client.GenerateImage(ctx, &pb.GenerateImageRequest{Src: "/tmp/source.png"})
|
||||
|
||||
Expect(err).To(MatchError(ContainSubstring("upload failed")))
|
||||
Expect(stager.ensureCalls).To(HaveLen(1))
|
||||
Expect(stager.releasedKeys).To(Equal(keysFromEnsureCalls(stager.ensureCalls)))
|
||||
})
|
||||
|
||||
It("does not invoke predict when multimodal staging fails", func(ctx SpecContext) {
|
||||
uploadErr := errors.New("ephemeral capacity exceeded")
|
||||
stager := &lifecycleStager{ensureErr: uploadErr, ensureErrAt: 2}
|
||||
backend := &lifecycleBackend{}
|
||||
client := NewFileStagingClient(backend, stager, "worker-1")
|
||||
request := &pb.PredictOptions{Images: []string{"/tmp/first.png", "/tmp/second.png"}}
|
||||
original := proto.Clone(request)
|
||||
|
||||
result, err := client.Predict(ctx, request)
|
||||
|
||||
Expect(result).To(BeNil())
|
||||
Expect(err).To(MatchError(ContainSubstring("ephemeral capacity exceeded")))
|
||||
Expect(backend.predictCalls).To(BeZero())
|
||||
Expect(proto.Equal(request, original)).To(BeTrue())
|
||||
Expect(stager.ensureCalls).To(HaveLen(2))
|
||||
Expect(stager.releasedKeys).To(Equal(keysFromEnsureCalls(stager.ensureCalls)))
|
||||
})
|
||||
|
||||
It("does not invoke streaming predict when multimodal staging fails", func(ctx SpecContext) {
|
||||
uploadErr := errors.New("ephemeral capacity exceeded")
|
||||
stager := &lifecycleStager{ensureErr: uploadErr}
|
||||
backend := &lifecycleBackend{}
|
||||
client := NewFileStagingClient(backend, stager, "worker-1")
|
||||
request := &pb.PredictOptions{Audios: []string{"/tmp/audio.wav"}}
|
||||
original := proto.Clone(request)
|
||||
|
||||
err := client.PredictStream(ctx, request, func(*pb.Reply) {})
|
||||
|
||||
Expect(err).To(MatchError(ContainSubstring("ephemeral capacity exceeded")))
|
||||
Expect(backend.streamCalls).To(BeZero())
|
||||
Expect(proto.Equal(request, original)).To(BeTrue())
|
||||
Expect(stager.releasedKeys).To(Equal(keysFromEnsureCalls(stager.ensureCalls)))
|
||||
})
|
||||
|
||||
It("uses an active bounded cleanup context after caller cancellation", func() {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
stager := &lifecycleStager{}
|
||||
client := NewFileStagingClient(&lifecycleBackend{}, stager, "worker-1")
|
||||
|
||||
_, err := client.Predict(ctx, &pb.PredictOptions{Images: []string{"/tmp/image.png"}})
|
||||
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(stager.releaseCtxErr).To(Equal([]error{nil}))
|
||||
})
|
||||
|
||||
It("does not release streaming inputs before the backend completes", func(ctx SpecContext) {
|
||||
block := make(chan struct{})
|
||||
started := make(chan struct{}, 1)
|
||||
backend := &lifecycleBackend{streamBlock: block, streamStarted: started}
|
||||
stager := &lifecycleStager{}
|
||||
client := NewFileStagingClient(backend, stager, "worker-1")
|
||||
done := make(chan error, 1)
|
||||
|
||||
go func() {
|
||||
done <- client.PredictStream(ctx, &pb.PredictOptions{Images: []string{"/tmp/image.png"}}, func(*pb.Reply) {})
|
||||
}()
|
||||
|
||||
Eventually(started).Should(Receive())
|
||||
Expect(stager.ensureCalls).To(HaveLen(1))
|
||||
Expect(stager.releasedKeys).To(BeEmpty())
|
||||
close(block)
|
||||
Eventually(done).Should(Receive(Succeed()))
|
||||
Expect(stager.releasedKeys).To(Equal(keysFromEnsureCalls(stager.ensureCalls)))
|
||||
})
|
||||
|
||||
It("does not replace a backend result when cleanup fails", func(ctx SpecContext) {
|
||||
reply := &pb.Reply{Message: []byte("ok")}
|
||||
backend := &lifecycleBackend{predictResult: reply}
|
||||
stager := &lifecycleStager{releaseErr: errors.New("release failed")}
|
||||
client := NewFileStagingClient(backend, stager, "worker-1")
|
||||
|
||||
result, err := client.Predict(ctx, &pb.PredictOptions{Images: []string{"/tmp/image.png"}})
|
||||
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(result).To(BeIdenticalTo(reply))
|
||||
})
|
||||
})
|
||||
|
||||
func keysFromEnsureCalls(calls []ensureCall) []string {
|
||||
keys := make([]string, len(calls))
|
||||
for i, call := range calls {
|
||||
keys[i] = call.key
|
||||
}
|
||||
return keys
|
||||
}
|
||||
@@ -28,6 +28,10 @@ func (s *soundStagingFailure) EnsureRemote(context.Context, string, string, stri
|
||||
return "", errors.New("upload failed")
|
||||
}
|
||||
|
||||
func (s *soundStagingFailure) ReleaseRemote(context.Context, string, string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
type soundRouteFactory struct{ client grpc.Backend }
|
||||
|
||||
func (f *soundRouteFactory) NewClient(string, bool) grpc.Backend { return f.client }
|
||||
|
||||
@@ -39,7 +39,8 @@ var _ = Describe("FileStagingClient TTS references", func() {
|
||||
Expect(stager.ensureCalls).To(HaveLen(1))
|
||||
Expect(stager.ensureCalls[0].localPath).To(Equal("/data/voice-profiles/profile/reference.wav"))
|
||||
Expect(backend.ttsRequest.Voice).To(HavePrefix("/remote/ephemeral/"))
|
||||
Expect(backend.ttsRequest.Voice).To(MatchRegexp(`/inputs/[0-9a-f]{8}/reference\.wav$`))
|
||||
voicePathPattern := `/inputs/` + fullUUIDPattern + `/reference\.wav$`
|
||||
Expect(backend.ttsRequest.Voice).To(MatchRegexp(voicePathPattern))
|
||||
})
|
||||
|
||||
It("stages a reference WAV before streaming synthesis", func(ctx SpecContext) {
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"crypto/subtle"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
@@ -22,6 +23,7 @@ import (
|
||||
"github.com/mudler/LocalAI/core/services/storage"
|
||||
"github.com/mudler/LocalAI/pkg/downloader"
|
||||
"github.com/mudler/LocalAI/pkg/model"
|
||||
"github.com/mudler/LocalAI/pkg/safefile"
|
||||
"github.com/mudler/xlog"
|
||||
)
|
||||
|
||||
@@ -49,17 +51,23 @@ const (
|
||||
// Auth is via Bearer token (registration token), using constant-time comparison.
|
||||
// A nil readiness fails open, keeping /readyz's historical always-200 answer.
|
||||
func StartFileTransferServer(addr, stagingDir, modelsDir, dataDir, token string, maxUploadSize int64, readiness *WorkerReadiness, logStore ...*model.BackendLogStore) (*http.Server, error) {
|
||||
return StartFileTransferServerWithCapacity(addr, stagingDir, modelsDir, dataDir, token, maxUploadSize, readiness, nil, logStore...)
|
||||
}
|
||||
|
||||
// StartFileTransferServerWithCapacity starts the file transfer server with a
|
||||
// worker-local guard for per-request ephemeral inputs.
|
||||
func StartFileTransferServerWithCapacity(addr, stagingDir, modelsDir, dataDir, token string, maxUploadSize int64, readiness *WorkerReadiness, capacity EphemeralCapacity, logStore ...*model.BackendLogStore) (*http.Server, error) {
|
||||
listener, err := net.Listen("tcp", addr)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("listen %s: %w", addr, err)
|
||||
}
|
||||
return StartFileTransferServerWithReadiness(listener, stagingDir, modelsDir, dataDir, token, maxUploadSize, readiness, logStore...)
|
||||
return startFileTransferServer(listener, stagingDir, modelsDir, dataDir, token, maxUploadSize, readiness, capacity, logStore...)
|
||||
}
|
||||
|
||||
// StartFileTransferServerWithListener starts the server on an existing listener.
|
||||
// This avoids the TOCTOU race of closing a listener and re-binding to the same port.
|
||||
func StartFileTransferServerWithListener(lis net.Listener, stagingDir, modelsDir, dataDir, token string, maxUploadSize int64, logStore ...*model.BackendLogStore) (*http.Server, error) {
|
||||
return StartFileTransferServerWithReadiness(lis, stagingDir, modelsDir, dataDir, token, maxUploadSize, nil, logStore...)
|
||||
return startFileTransferServer(lis, stagingDir, modelsDir, dataDir, token, maxUploadSize, nil, nil, logStore...)
|
||||
}
|
||||
|
||||
// StartFileTransferServerWithReadiness is StartFileTransferServerWithListener
|
||||
@@ -67,6 +75,10 @@ func StartFileTransferServerWithListener(lis net.Listener, stagingDir, modelsDir
|
||||
// the probe keeps its historical always-200 behaviour for callers that have no
|
||||
// meaningful readiness signal to report.
|
||||
func StartFileTransferServerWithReadiness(lis net.Listener, stagingDir, modelsDir, dataDir, token string, maxUploadSize int64, readiness *WorkerReadiness, logStore ...*model.BackendLogStore) (*http.Server, error) {
|
||||
return startFileTransferServer(lis, stagingDir, modelsDir, dataDir, token, maxUploadSize, readiness, nil, logStore...)
|
||||
}
|
||||
|
||||
func startFileTransferServer(lis net.Listener, stagingDir, modelsDir, dataDir, token string, maxUploadSize int64, readiness *WorkerReadiness, capacity EphemeralCapacity, logStore ...*model.BackendLogStore) (*http.Server, error) {
|
||||
if err := os.MkdirAll(stagingDir, 0750); err != nil {
|
||||
return nil, fmt.Errorf("creating staging dir %s: %w", stagingDir, err)
|
||||
}
|
||||
@@ -101,6 +113,18 @@ func StartFileTransferServerWithReadiness(lis net.Listener, stagingDir, modelsDi
|
||||
handleListDir(w, r, stagingDir, modelsDir, dataDir, key)
|
||||
})
|
||||
|
||||
mux.HandleFunc("/v1/files-release", func(w http.ResponseWriter, r *http.Request) {
|
||||
if !checkBearerToken(r, token) {
|
||||
http.Error(w, "unauthorized", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
if r.Method != http.MethodDelete {
|
||||
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
handleReleaseBatchWithCapacity(w, r, stagingDir, capacity)
|
||||
})
|
||||
|
||||
mux.HandleFunc("/v1/files/", func(w http.ResponseWriter, r *http.Request) {
|
||||
if !checkBearerToken(r, token) {
|
||||
xlog.Debug("HTTP file transfer: unauthorized request", "method", r.Method, "path", r.URL.Path, "remote", r.RemoteAddr)
|
||||
@@ -116,12 +140,16 @@ func StartFileTransferServerWithReadiness(lis net.Listener, stagingDir, modelsDi
|
||||
case http.MethodHead:
|
||||
handleHead(w, r, stagingDir, modelsDir, dataDir, key)
|
||||
case http.MethodPut:
|
||||
handleUpload(w, r, stagingDir, modelsDir, dataDir, key, maxUploadSize)
|
||||
handleUploadWithCapacity(w, r, stagingDir, modelsDir, dataDir, key, maxUploadSize, capacity)
|
||||
case http.MethodGet:
|
||||
handleDownload(w, r, stagingDir, modelsDir, dataDir, key)
|
||||
case http.MethodDelete:
|
||||
handleReleaseWithCapacity(w, r, stagingDir, key, capacity)
|
||||
case http.MethodPost:
|
||||
if key == "temp" {
|
||||
handleAllocTemp(w, r, stagingDir)
|
||||
} else if r.URL.Query().Get("claim") == "1" {
|
||||
handleClaimWithCapacity(w, r, stagingDir, key, capacity)
|
||||
} else {
|
||||
http.Error(w, "not found", http.StatusNotFound)
|
||||
}
|
||||
@@ -182,6 +210,163 @@ func StartFileTransferServerWithReadiness(lis net.Listener, stagingDir, modelsDi
|
||||
return server, nil
|
||||
}
|
||||
|
||||
// handleClaimWithCapacity marks an existing ephemeral file as owned by the
|
||||
// request that just verified its content.
|
||||
func handleClaimWithCapacity(w http.ResponseWriter, _ *http.Request, stagingDir, key string, capacity EphemeralCapacity) {
|
||||
if err := validateEphemeralReleaseKey(key); err != nil {
|
||||
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
if operations, ok := capacity.(ephemeralRequestOperationCapacity); ok {
|
||||
requestID := strings.Split(key, "/")[2]
|
||||
if err := operations.BeginRequestOperation(requestID); err != nil {
|
||||
http.Error(w, err.Error(), http.StatusConflict)
|
||||
return
|
||||
}
|
||||
defer operations.EndRequestOperation(requestID)
|
||||
}
|
||||
filePath := filepath.Join(stagingDir, filepath.FromSlash(key))
|
||||
if err := validatePathInDir(filePath, stagingDir); err != nil {
|
||||
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
info, err := os.Lstat(filePath)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
http.Error(w, "not found", http.StatusNotFound)
|
||||
} else {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
}
|
||||
return
|
||||
}
|
||||
if !info.Mode().IsRegular() {
|
||||
http.Error(w, "ephemeral path is not a regular file", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
if capacity != nil {
|
||||
if err := capacity.Claim(filePath); err != nil {
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
http.Error(w, "not found", http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
http.Error(w, err.Error(), http.StatusInsufficientStorage)
|
||||
return
|
||||
}
|
||||
}
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
func handleRelease(w http.ResponseWriter, _ *http.Request, stagingDir, key string) {
|
||||
handleReleaseWithCapacity(w, nil, stagingDir, key, nil)
|
||||
}
|
||||
|
||||
func handleReleaseWithCapacity(w http.ResponseWriter, _ *http.Request, stagingDir, key string, capacity EphemeralCapacity) {
|
||||
if err := releaseEphemeralStagingKey(stagingDir, key, capacity); err != nil {
|
||||
status := http.StatusInternalServerError
|
||||
if errors.Is(err, safefile.ErrUnsafePath) || validateEphemeralReleaseKey(key) != nil {
|
||||
status = http.StatusBadRequest
|
||||
}
|
||||
http.Error(w, err.Error(), status)
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
func handleReleaseBatchWithCapacity(w http.ResponseWriter, r *http.Request, stagingDir string, capacity EphemeralCapacity) {
|
||||
r.Body = http.MaxBytesReader(w, r.Body, 1<<20)
|
||||
var request struct {
|
||||
RequestID string `json:"request_id"`
|
||||
}
|
||||
if err := json.NewDecoder(r.Body).Decode(&request); err != nil {
|
||||
http.Error(w, fmt.Sprintf("decoding release batch: %v", err), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
if err := validateEphemeralRequestID(request.RequestID); err != nil {
|
||||
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
if err := releaseEphemeralStagingRequest(r.Context(), stagingDir, request.RequestID, capacity); err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
func releaseEphemeralStagingRequest(ctx context.Context, stagingDir, requestID string, capacity EphemeralCapacity) error {
|
||||
if err := validateEphemeralRequestID(requestID); err != nil {
|
||||
return err
|
||||
}
|
||||
if requestCapacity, ok := capacity.(ephemeralRequestCapacity); ok {
|
||||
if err := requestCapacity.BeginRequestRelease(ctx, requestID); err != nil {
|
||||
return fmt.Errorf("beginning release for request %q: %w", requestID, err)
|
||||
}
|
||||
defer requestCapacity.EndRequestRelease(requestID)
|
||||
}
|
||||
root := filepath.Join(stagingDir, "ephemeral")
|
||||
categories, err := os.ReadDir(root)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
var releaseErrors []error
|
||||
for _, category := range categories {
|
||||
if !category.IsDir() || category.Type()&os.ModeSymlink != 0 {
|
||||
continue
|
||||
}
|
||||
requestDir := filepath.Join(root, category.Name(), requestID)
|
||||
info, err := os.Lstat(requestDir)
|
||||
if err != nil {
|
||||
if !os.IsNotExist(err) {
|
||||
releaseErrors = append(releaseErrors, fmt.Errorf("stating request directory %q: %w", requestDir, err))
|
||||
}
|
||||
continue
|
||||
}
|
||||
if !info.IsDir() || info.Mode()&os.ModeSymlink != 0 {
|
||||
releaseErrors = append(releaseErrors, fmt.Errorf("ephemeral request path %q is not a real directory", requestDir))
|
||||
continue
|
||||
}
|
||||
entries, err := os.ReadDir(requestDir)
|
||||
if err != nil {
|
||||
if !os.IsNotExist(err) {
|
||||
releaseErrors = append(releaseErrors, fmt.Errorf("reading request directory %q: %w", requestDir, err))
|
||||
}
|
||||
continue
|
||||
}
|
||||
for _, entry := range entries {
|
||||
if entry.IsDir() {
|
||||
releaseErrors = append(releaseErrors, fmt.Errorf("unexpected directory in ephemeral request %q", filepath.Join(requestDir, entry.Name())))
|
||||
continue
|
||||
}
|
||||
key := filepath.ToSlash(filepath.Join("ephemeral", category.Name(), requestID, entry.Name()))
|
||||
if err := releaseEphemeralStagingKey(stagingDir, key, capacity); err != nil {
|
||||
releaseErrors = append(releaseErrors, fmt.Errorf("releasing %q: %w", key, err))
|
||||
}
|
||||
}
|
||||
}
|
||||
return errors.Join(releaseErrors...)
|
||||
}
|
||||
|
||||
func releaseEphemeralStagingKey(stagingDir, key string, capacity EphemeralCapacity) error {
|
||||
if err := validateEphemeralReleaseKey(key); err != nil {
|
||||
return err
|
||||
}
|
||||
relativePath := filepath.FromSlash(key)
|
||||
filePath := filepath.Join(stagingDir, relativePath)
|
||||
if err := safefile.RemoveExact(stagingDir, relativePath, []string{hashSidecarSuffix, targetSidecarSuffix}, 2); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, path := range []string{filePath, filePath + hashSidecarSuffix, filePath + targetSidecarSuffix} {
|
||||
if capacity != nil {
|
||||
if err := capacity.Release(path); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func handleHead(w http.ResponseWriter, r *http.Request, stagingDir, modelsDir, dataDir, key string) {
|
||||
if key == "" {
|
||||
http.Error(w, "key is required", http.StatusBadRequest)
|
||||
@@ -250,6 +435,74 @@ type contentRange struct {
|
||||
total int64
|
||||
}
|
||||
|
||||
// EphemeralCapacity bounds worker-local request input storage. Implementations
|
||||
// must reserve before bytes reach disk and may reconcile reservations with the
|
||||
// resulting file after a write ends.
|
||||
type EphemeralCapacity interface {
|
||||
Reserve(path string, size int64) error
|
||||
Commit(path string) error
|
||||
Claim(path string) error
|
||||
Release(path string) error
|
||||
CapacityWriter(path string, destination io.Writer) (io.WriteCloser, error)
|
||||
}
|
||||
|
||||
type ephemeralRequestCapacity interface {
|
||||
BeginRequestRelease(ctx context.Context, requestID string) error
|
||||
EndRequestRelease(requestID string)
|
||||
}
|
||||
|
||||
type ephemeralRequestOperationCapacity interface {
|
||||
BeginRequestOperation(requestID string) error
|
||||
EndRequestOperation(requestID string)
|
||||
}
|
||||
|
||||
type uploadStatusWriter struct {
|
||||
http.ResponseWriter
|
||||
status int
|
||||
}
|
||||
|
||||
func (w *uploadStatusWriter) WriteHeader(status int) {
|
||||
w.status = status
|
||||
w.ResponseWriter.WriteHeader(status)
|
||||
}
|
||||
|
||||
func (w *uploadStatusWriter) Write(payload []byte) (int, error) {
|
||||
if w.status == 0 {
|
||||
w.status = http.StatusOK
|
||||
}
|
||||
return w.ResponseWriter.Write(payload)
|
||||
}
|
||||
|
||||
type ephemeralCapacityWriteError struct{ err error }
|
||||
|
||||
func (e *ephemeralCapacityWriteError) Error() string { return e.err.Error() }
|
||||
func (e *ephemeralCapacityWriteError) Unwrap() error { return e.err }
|
||||
|
||||
type ephemeralCapacityWriteCloser struct{ io.WriteCloser }
|
||||
|
||||
func (w ephemeralCapacityWriteCloser) Write(payload []byte) (int, error) {
|
||||
written, err := w.WriteCloser.Write(payload)
|
||||
if err != nil {
|
||||
return written, &ephemeralCapacityWriteError{err: err}
|
||||
}
|
||||
return written, nil
|
||||
}
|
||||
|
||||
func (w ephemeralCapacityWriteCloser) Close() error {
|
||||
if err := w.WriteCloser.Close(); err != nil {
|
||||
return &ephemeralCapacityWriteError{err: err}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func uploadWriteStatus(err error) int {
|
||||
var capacityErr *ephemeralCapacityWriteError
|
||||
if errors.As(err, &capacityErr) {
|
||||
return http.StatusInsufficientStorage
|
||||
}
|
||||
return http.StatusInternalServerError
|
||||
}
|
||||
|
||||
// parseContentRange parses a Content-Range header value of the form
|
||||
// "bytes <start>-<end>/<total>". RFC 9110 §14.4.
|
||||
// Returns (nil, nil) when the header is empty (no range request).
|
||||
@@ -291,10 +544,29 @@ func parseContentRange(h string) (*contentRange, error) {
|
||||
}
|
||||
|
||||
func handleUpload(w http.ResponseWriter, r *http.Request, stagingDir, modelsDir, dataDir, key string, maxUploadSize int64) {
|
||||
handleUploadWithCapacity(w, r, stagingDir, modelsDir, dataDir, key, maxUploadSize, nil)
|
||||
}
|
||||
|
||||
func handleUploadWithCapacity(w http.ResponseWriter, r *http.Request, stagingDir, modelsDir, dataDir, key string, maxUploadSize int64, capacity EphemeralCapacity) {
|
||||
if key == "" {
|
||||
http.Error(w, "key is required", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
capacityEnabled := capacity != nil && strings.HasPrefix(key, "ephemeral/")
|
||||
if capacityEnabled {
|
||||
if err := validateEphemeralReleaseKey(key); err != nil {
|
||||
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
if operations, ok := capacity.(ephemeralRequestOperationCapacity); ok {
|
||||
requestID := strings.Split(key, "/")[2]
|
||||
if err := operations.BeginRequestOperation(requestID); err != nil {
|
||||
http.Error(w, err.Error(), http.StatusConflict)
|
||||
return
|
||||
}
|
||||
defer operations.EndRequestOperation(requestID)
|
||||
}
|
||||
}
|
||||
|
||||
if maxUploadSize > 0 {
|
||||
r.Body = http.MaxBytesReader(w, r.Body, maxUploadSize)
|
||||
@@ -329,18 +601,67 @@ func handleUpload(w http.ResponseWriter, r *http.Request, stagingDir, modelsDir,
|
||||
return
|
||||
}
|
||||
|
||||
if cr == nil {
|
||||
// Non-resumable (legacy) path: truncate-create, single fire-and-forget.
|
||||
handleFullUpload(w, r, dstPath, key, expectedFinalHash)
|
||||
return
|
||||
capacityEnabled = capacityEnabled && targetDir == stagingDir
|
||||
unknownLengthCapacity := capacityEnabled && r.ContentLength < 0
|
||||
capacityPaths := []string{dstPath, dstPath + hashSidecarSuffix, dstPath + targetSidecarSuffix}
|
||||
if capacityEnabled {
|
||||
if r.ContentLength >= 0 {
|
||||
if err := capacity.Reserve(dstPath, r.ContentLength); err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInsufficientStorage)
|
||||
return
|
||||
}
|
||||
}
|
||||
for _, sidecarPath := range capacityPaths[1:] {
|
||||
if err := capacity.Reserve(sidecarPath, sha256.Size*2); err != nil {
|
||||
for _, reservedPath := range capacityPaths {
|
||||
reconcileEphemeralCapacity(capacity, reservedPath, 0)
|
||||
}
|
||||
http.Error(w, err.Error(), http.StatusInsufficientStorage)
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
handleRangeUpload(w, r, dstPath, key, cr, expectedFinalHash)
|
||||
statusWriter := &uploadStatusWriter{ResponseWriter: w}
|
||||
var uploadCapacity EphemeralCapacity
|
||||
if unknownLengthCapacity {
|
||||
uploadCapacity = capacity
|
||||
}
|
||||
|
||||
if cr == nil {
|
||||
// Non-resumable (legacy) path: truncate-create, single fire-and-forget.
|
||||
handleFullUpload(statusWriter, r, dstPath, key, expectedFinalHash, uploadCapacity)
|
||||
} else {
|
||||
handleRangeUpload(statusWriter, r, dstPath, key, cr, expectedFinalHash, uploadCapacity)
|
||||
}
|
||||
|
||||
if !capacityEnabled {
|
||||
return
|
||||
}
|
||||
for _, capacityPath := range capacityPaths {
|
||||
reconcileEphemeralCapacity(capacity, capacityPath, statusWriter.status)
|
||||
}
|
||||
}
|
||||
|
||||
func reconcileEphemeralCapacity(capacity EphemeralCapacity, path string, status int) {
|
||||
if info, err := os.Lstat(path); err == nil && info.Mode().IsRegular() {
|
||||
if err := capacity.Commit(path); err != nil {
|
||||
xlog.Error("Committing ephemeral capacity failed", "path", path, "status", status, "error", err)
|
||||
if removeErr := os.Remove(path); removeErr != nil && !os.IsNotExist(removeErr) {
|
||||
xlog.Warn("Removing uncommitted ephemeral file failed", "path", path, "error", removeErr)
|
||||
}
|
||||
if releaseErr := capacity.Release(path); releaseErr != nil {
|
||||
xlog.Warn("Rolling back failed ephemeral commit", "path", path, "error", releaseErr)
|
||||
}
|
||||
}
|
||||
} else if err := capacity.Release(path); err != nil {
|
||||
xlog.Warn("Rolling back ephemeral capacity failed", "path", path, "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
// handleFullUpload writes the entire request body to dstPath, replacing any
|
||||
// existing content. This is the legacy happy-path with no Range header.
|
||||
func handleFullUpload(w http.ResponseWriter, r *http.Request, dstPath, key, expectedFinalHash string) {
|
||||
func handleFullUpload(w http.ResponseWriter, r *http.Request, dstPath, key, expectedFinalHash string, capacity EphemeralCapacity) {
|
||||
// Reset any in-progress resumable state.
|
||||
_ = os.Remove(dstPath + targetSidecarSuffix)
|
||||
|
||||
@@ -351,13 +672,32 @@ func handleFullUpload(w http.ResponseWriter, r *http.Request, dstPath, key, expe
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
var destination io.Writer = f
|
||||
var capacityWriter io.WriteCloser
|
||||
if capacity != nil {
|
||||
var writer io.WriteCloser
|
||||
writer, err = capacity.CapacityWriter(dstPath, f)
|
||||
if err != nil {
|
||||
_ = os.Remove(dstPath)
|
||||
http.Error(w, err.Error(), http.StatusInsufficientStorage)
|
||||
return
|
||||
}
|
||||
capacityWriter = ephemeralCapacityWriteCloser{WriteCloser: writer}
|
||||
destination = capacityWriter
|
||||
}
|
||||
|
||||
hasher := sha256.New()
|
||||
n, err := io.Copy(f, io.TeeReader(r.Body, hasher))
|
||||
n, err := io.Copy(destination, io.TeeReader(r.Body, hasher))
|
||||
if capacityWriter != nil {
|
||||
if closeErr := capacityWriter.Close(); err == nil {
|
||||
err = closeErr
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
os.Remove(dstPath)
|
||||
os.Remove(dstPath + hashSidecarSuffix)
|
||||
xlog.Error("File upload failed", "key", key, "bytesReceived", n, "contentLength", r.ContentLength, "remote", r.RemoteAddr, "error", err)
|
||||
http.Error(w, fmt.Sprintf("writing file: %v", err), http.StatusInternalServerError)
|
||||
http.Error(w, fmt.Sprintf("writing file: %v", err), uploadWriteStatus(err))
|
||||
return
|
||||
}
|
||||
|
||||
@@ -386,7 +726,7 @@ func handleFullUpload(w http.ResponseWriter, r *http.Request, dstPath, key, expe
|
||||
// the request starts at the current file size. When the slice completes the
|
||||
// transfer (end+1 == total), it validates the optional expected final hash and
|
||||
// writes the sidecar.
|
||||
func handleRangeUpload(w http.ResponseWriter, r *http.Request, dstPath, key string, cr *contentRange, expectedFinalHash string) {
|
||||
func handleRangeUpload(w http.ResponseWriter, r *http.Request, dstPath, key string, cr *contentRange, expectedFinalHash string, capacity EphemeralCapacity) {
|
||||
// Determine the current on-disk size (0 if missing).
|
||||
var currentSize int64
|
||||
if info, err := os.Stat(dstPath); err == nil {
|
||||
@@ -470,6 +810,18 @@ func handleRangeUpload(w http.ResponseWriter, r *http.Request, dstPath, key stri
|
||||
return
|
||||
}
|
||||
defer func() { _ = f.Close() }()
|
||||
var destination io.Writer = f
|
||||
var capacityWriter io.WriteCloser
|
||||
if capacity != nil {
|
||||
var writer io.WriteCloser
|
||||
writer, err = capacity.CapacityWriter(dstPath, f)
|
||||
if err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInsufficientStorage)
|
||||
return
|
||||
}
|
||||
capacityWriter = ephemeralCapacityWriteCloser{WriteCloser: writer}
|
||||
destination = capacityWriter
|
||||
}
|
||||
|
||||
// Persist the declared expected hash so subsequent chunks can be
|
||||
// cross-checked.
|
||||
@@ -481,10 +833,15 @@ func handleRangeUpload(w http.ResponseWriter, r *http.Request, dstPath, key stri
|
||||
|
||||
expectedChunkLen := cr.end - cr.start + 1
|
||||
limited := io.LimitReader(r.Body, expectedChunkLen)
|
||||
n, err := io.Copy(f, limited)
|
||||
n, err := io.Copy(destination, limited)
|
||||
if capacityWriter != nil {
|
||||
if closeErr := capacityWriter.Close(); err == nil {
|
||||
err = closeErr
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
xlog.Error("Range upload chunk failed", "key", key, "bytesReceived", n, "expected", expectedChunkLen, "remote", r.RemoteAddr, "error", err)
|
||||
http.Error(w, fmt.Sprintf("writing file: %v", err), http.StatusInternalServerError)
|
||||
http.Error(w, fmt.Sprintf("writing file: %v", err), uploadWriteStatus(err))
|
||||
return
|
||||
}
|
||||
if n != expectedChunkLen {
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
@@ -17,10 +18,99 @@ import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
type recordingEphemeralCapacity struct {
|
||||
reserved int64
|
||||
reserveErr error
|
||||
writerErr error
|
||||
claimCalls []string
|
||||
claimErr error
|
||||
commitErr error
|
||||
releases []string
|
||||
startedOps []string
|
||||
endedOps []string
|
||||
}
|
||||
|
||||
type nopWriteCloser struct{ io.Writer }
|
||||
|
||||
func (nopWriteCloser) Close() error { return nil }
|
||||
|
||||
type failingWriteCloser struct{ err error }
|
||||
|
||||
func (w failingWriteCloser) Write([]byte) (int, error) { return 0, w.err }
|
||||
func (failingWriteCloser) Close() error { return nil }
|
||||
|
||||
type blockingDestinationWriteCloser struct {
|
||||
destination io.Writer
|
||||
written chan struct{}
|
||||
release chan struct{}
|
||||
}
|
||||
|
||||
func (w *blockingDestinationWriteCloser) Write(payload []byte) (int, error) {
|
||||
n, err := w.destination.Write(payload)
|
||||
close(w.written)
|
||||
<-w.release
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (*blockingDestinationWriteCloser) Close() error { return nil }
|
||||
|
||||
type blockingDestinationCapacity struct {
|
||||
written chan struct{}
|
||||
release chan struct{}
|
||||
}
|
||||
|
||||
func (*blockingDestinationCapacity) Reserve(string, int64) error { return nil }
|
||||
func (*blockingDestinationCapacity) BeginRequestRelease(context.Context, string) error { return nil }
|
||||
func (*blockingDestinationCapacity) EndRequestRelease(string) {}
|
||||
func (*blockingDestinationCapacity) BeginRequestOperation(string) error { return nil }
|
||||
func (*blockingDestinationCapacity) EndRequestOperation(string) {}
|
||||
func (*blockingDestinationCapacity) Commit(string) error { return nil }
|
||||
func (*blockingDestinationCapacity) Claim(string) error { return nil }
|
||||
func (*blockingDestinationCapacity) Release(string) error { return nil }
|
||||
func (g *blockingDestinationCapacity) CapacityWriter(_ string, destination io.Writer) (io.WriteCloser, error) {
|
||||
return &blockingDestinationWriteCloser{
|
||||
destination: destination,
|
||||
written: g.written,
|
||||
release: g.release,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (g *recordingEphemeralCapacity) Reserve(_ string, size int64) error {
|
||||
g.reserved = size
|
||||
return g.reserveErr
|
||||
}
|
||||
|
||||
func (*recordingEphemeralCapacity) BeginRequestRelease(context.Context, string) error { return nil }
|
||||
func (*recordingEphemeralCapacity) EndRequestRelease(string) {}
|
||||
func (g *recordingEphemeralCapacity) BeginRequestOperation(requestID string) error {
|
||||
g.startedOps = append(g.startedOps, requestID)
|
||||
return nil
|
||||
}
|
||||
func (g *recordingEphemeralCapacity) EndRequestOperation(requestID string) {
|
||||
g.endedOps = append(g.endedOps, requestID)
|
||||
}
|
||||
|
||||
func (g *recordingEphemeralCapacity) Commit(string) error { return g.commitErr }
|
||||
func (g *recordingEphemeralCapacity) Release(path string) error {
|
||||
g.releases = append(g.releases, path)
|
||||
return nil
|
||||
}
|
||||
func (g *recordingEphemeralCapacity) Claim(path string) error {
|
||||
g.claimCalls = append(g.claimCalls, path)
|
||||
return g.claimErr
|
||||
}
|
||||
func (g *recordingEphemeralCapacity) CapacityWriter(_ string, destination io.Writer) (io.WriteCloser, error) {
|
||||
if g.writerErr != nil {
|
||||
return failingWriteCloser{err: g.writerErr}, nil
|
||||
}
|
||||
return nopWriteCloser{Writer: destination}, nil
|
||||
}
|
||||
|
||||
var _ = Describe("FileTransferServer", func() {
|
||||
setupTestServer := func(token string, maxUploadSize int64) (*httptest.Server, string, string, string) {
|
||||
stagingDir := GinkgoT().TempDir()
|
||||
@@ -54,6 +144,67 @@ var _ = Describe("FileTransferServer", func() {
|
||||
}
|
||||
|
||||
Describe("Upload and Download", func() {
|
||||
It("rejects a declared ephemeral upload before writing when capacity is exhausted", func() {
|
||||
stagingDir := GinkgoT().TempDir()
|
||||
modelsDir := GinkgoT().TempDir()
|
||||
dataDir := GinkgoT().TempDir()
|
||||
guard := &recordingEphemeralCapacity{reserveErr: fmt.Errorf("full")}
|
||||
recorder := httptest.NewRecorder()
|
||||
request := httptest.NewRequest(http.MethodPut, "/v1/files/ephemeral/audio/request/input.wav", strings.NewReader("payload"))
|
||||
|
||||
handleUploadWithCapacity(recorder, request, stagingDir, modelsDir, dataDir, "ephemeral/audio/request/input.wav", 0, guard)
|
||||
|
||||
Expect(recorder.Code).To(Equal(http.StatusInsufficientStorage))
|
||||
Expect(guard.reserved).To(Equal(int64(len("payload"))))
|
||||
Expect(filepath.Join(stagingDir, "ephemeral", "audio", "request", "input.wav")).NotTo(BeAnExistingFile())
|
||||
})
|
||||
|
||||
It("returns insufficient storage when a chunked upload reaches its bound", func() {
|
||||
stagingDir := GinkgoT().TempDir()
|
||||
guard := &recordingEphemeralCapacity{writerErr: fmt.Errorf("full")}
|
||||
recorder := httptest.NewRecorder()
|
||||
request := httptest.NewRequest(http.MethodPut, "/v1/files/ephemeral/audio/request/input.wav", strings.NewReader("payload"))
|
||||
request.ContentLength = -1
|
||||
|
||||
handleUploadWithCapacity(recorder, request, stagingDir, GinkgoT().TempDir(), GinkgoT().TempDir(), "ephemeral/audio/request/input.wav", 0, guard)
|
||||
|
||||
Expect(recorder.Code).To(Equal(http.StatusInsufficientStorage))
|
||||
})
|
||||
|
||||
It("keeps unknown-length bytes guarded until they reach the staged file", func() {
|
||||
stagingDir := GinkgoT().TempDir()
|
||||
modelsDir := GinkgoT().TempDir()
|
||||
dataDir := GinkgoT().TempDir()
|
||||
guard := &blockingDestinationCapacity{
|
||||
written: make(chan struct{}),
|
||||
release: make(chan struct{}),
|
||||
}
|
||||
released := false
|
||||
defer func() {
|
||||
if !released {
|
||||
close(guard.release)
|
||||
}
|
||||
}()
|
||||
recorder := httptest.NewRecorder()
|
||||
request := httptest.NewRequest(http.MethodPut, "/v1/files/ephemeral/audio/request/input.wav", strings.NewReader("payload"))
|
||||
request.ContentLength = -1
|
||||
done := make(chan struct{})
|
||||
|
||||
go func() {
|
||||
defer GinkgoRecover()
|
||||
handleUploadWithCapacity(recorder, request, stagingDir, modelsDir, dataDir, "ephemeral/audio/request/input.wav", 0, guard)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
Eventually(guard.written).Should(BeClosed())
|
||||
path := filepath.Join(stagingDir, "ephemeral", "audio", "request", "input.wav")
|
||||
Expect(os.ReadFile(path)).To(Equal([]byte("payload")))
|
||||
close(guard.release)
|
||||
released = true
|
||||
Eventually(done).Should(BeClosed())
|
||||
Expect(recorder.Code).To(Equal(http.StatusOK))
|
||||
})
|
||||
|
||||
It("round-trips file content correctly", func() {
|
||||
ts, _, _, _ := setupTestServer("secret-token", 0)
|
||||
|
||||
@@ -394,6 +545,17 @@ var _ = Describe("FileTransferServer", func() {
|
||||
})
|
||||
})
|
||||
|
||||
It("removes and releases a file whose capacity commit fails", func() {
|
||||
path := filepath.Join(GinkgoT().TempDir(), "input.wav")
|
||||
Expect(os.WriteFile(path, []byte("data"), 0o600)).To(Succeed())
|
||||
guard := &recordingEphemeralCapacity{commitErr: errors.New("request released")}
|
||||
|
||||
reconcileEphemeralCapacity(guard, path, http.StatusOK)
|
||||
|
||||
Expect(path).NotTo(BeAnExistingFile())
|
||||
Expect(guard.releases).To(Equal([]string{path}))
|
||||
})
|
||||
|
||||
// --- Upload sidecar tests ---
|
||||
|
||||
Describe("Upload hash sidecar", func() {
|
||||
@@ -438,6 +600,165 @@ var _ = Describe("FileTransferServer", func() {
|
||||
// --- EnsureRemote skip tests ---
|
||||
|
||||
Describe("EnsureRemote skip-if-exists", func() {
|
||||
It("reports a claim-time disappearance as a cache miss", func() {
|
||||
stagingDir := GinkgoT().TempDir()
|
||||
key := "ephemeral/audio/request/input.wav"
|
||||
remotePath := filepath.Join(stagingDir, filepath.FromSlash(key))
|
||||
Expect(os.MkdirAll(filepath.Dir(remotePath), 0o750)).To(Succeed())
|
||||
Expect(os.WriteFile(remotePath, []byte("stale"), 0o600)).To(Succeed())
|
||||
guard := &recordingEphemeralCapacity{claimErr: fmt.Errorf("claim raced recovery: %w", os.ErrNotExist)}
|
||||
recorder := httptest.NewRecorder()
|
||||
request := httptest.NewRequest(http.MethodPost, "/v1/files/"+key, nil)
|
||||
|
||||
handleClaimWithCapacity(recorder, request, stagingDir, key, guard)
|
||||
|
||||
Expect(recorder.Code).To(Equal(http.StatusNotFound))
|
||||
})
|
||||
|
||||
It("claims a matching ephemeral file before returning the worker path", func() {
|
||||
stagingDir := GinkgoT().TempDir()
|
||||
modelsDir := GinkgoT().TempDir()
|
||||
dataDir := GinkgoT().TempDir()
|
||||
guard := &recordingEphemeralCapacity{}
|
||||
key := "ephemeral/audio/request/input.wav"
|
||||
remotePath := filepath.Join(stagingDir, filepath.FromSlash(key))
|
||||
content := []byte("already on worker")
|
||||
Expect(os.MkdirAll(filepath.Dir(remotePath), 0o750)).To(Succeed())
|
||||
Expect(os.WriteFile(remotePath, content, 0o600)).To(Succeed())
|
||||
Expect(os.WriteFile(remotePath+hashSidecarSuffix, []byte(sha256Hex(content)), 0o600)).To(Succeed())
|
||||
localPath := filepath.Join(GinkgoT().TempDir(), "input.wav")
|
||||
Expect(os.WriteFile(localPath, content, 0o600)).To(Succeed())
|
||||
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("/v1/files/", func(w http.ResponseWriter, r *http.Request) {
|
||||
requestKey := strings.TrimPrefix(r.URL.Path, "/v1/files/")
|
||||
switch r.Method {
|
||||
case http.MethodHead:
|
||||
handleHead(w, r, stagingDir, modelsDir, dataDir, requestKey)
|
||||
case http.MethodPost:
|
||||
handleClaimWithCapacity(w, r, stagingDir, requestKey, guard)
|
||||
default:
|
||||
http.Error(w, "unexpected upload", http.StatusInternalServerError)
|
||||
}
|
||||
})
|
||||
ts := httptest.NewServer(mux)
|
||||
DeferCleanup(ts.Close)
|
||||
stager := NewHTTPFileStager(func(string) (string, error) {
|
||||
return strings.TrimPrefix(ts.URL, "http://"), nil
|
||||
}, "")
|
||||
|
||||
for range 2 {
|
||||
path, err := stager.EnsureRemote(context.Background(), "node-1", localPath, key)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(path).To(Equal(remotePath))
|
||||
}
|
||||
Expect(guard.claimCalls).To(Equal([]string{remotePath, remotePath}))
|
||||
Expect(guard.startedOps).To(Equal([]string{"request", "request"}))
|
||||
Expect(guard.endedOps).To(Equal([]string{"request", "request"}))
|
||||
})
|
||||
|
||||
It("propagates an ephemeral cache-hit claim failure", func() {
|
||||
content := []byte("already on worker")
|
||||
localPath := filepath.Join(GinkgoT().TempDir(), "input.wav")
|
||||
Expect(os.WriteFile(localPath, content, 0o600)).To(Succeed())
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("/v1/files/", func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method == http.MethodHead {
|
||||
w.Header().Set(HeaderLocalPath, "/remote/ephemeral/audio/request/input.wav")
|
||||
w.Header().Set(HeaderContentSHA256, sha256Hex(content))
|
||||
w.WriteHeader(http.StatusOK)
|
||||
return
|
||||
}
|
||||
http.Error(w, "ephemeral capacity exceeded", http.StatusInsufficientStorage)
|
||||
})
|
||||
ts := httptest.NewServer(mux)
|
||||
DeferCleanup(ts.Close)
|
||||
stager := NewHTTPFileStager(func(string) (string, error) {
|
||||
return strings.TrimPrefix(ts.URL, "http://"), nil
|
||||
}, "")
|
||||
|
||||
path, err := stager.EnsureRemote(context.Background(), "node-1", localPath, "ephemeral/audio/request/input.wav")
|
||||
|
||||
Expect(path).To(BeEmpty())
|
||||
Expect(err).To(MatchError(ContainSubstring("ephemeral capacity exceeded")))
|
||||
})
|
||||
|
||||
DescribeTable("uploads a matching ephemeral file when an old worker cannot claim it",
|
||||
func(claimStatus int) {
|
||||
stagingDir := GinkgoT().TempDir()
|
||||
content := []byte("compatible upload")
|
||||
localPath := filepath.Join(GinkgoT().TempDir(), "input.wav")
|
||||
Expect(os.WriteFile(localPath, content, 0o600)).To(Succeed())
|
||||
cacheHitPath := filepath.Join(stagingDir, "ephemeral", "audio", "stale", "input.wav")
|
||||
putCalls := 0
|
||||
putRemotePath := ""
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("/v1/files/", func(w http.ResponseWriter, r *http.Request) {
|
||||
switch r.Method {
|
||||
case http.MethodHead:
|
||||
w.Header().Set(HeaderLocalPath, cacheHitPath)
|
||||
w.Header().Set(HeaderContentSHA256, sha256Hex(content))
|
||||
w.WriteHeader(http.StatusOK)
|
||||
case http.MethodPost:
|
||||
http.Error(w, "claim unsupported", claimStatus)
|
||||
case http.MethodPut:
|
||||
putCalls++
|
||||
key := strings.TrimPrefix(r.URL.Path, "/v1/files/")
|
||||
putRemotePath = filepath.Join(stagingDir, filepath.FromSlash(key))
|
||||
handleUpload(w, r, stagingDir, "", "", key, 0)
|
||||
case http.MethodDelete:
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
})
|
||||
ts := httptest.NewServer(mux)
|
||||
DeferCleanup(ts.Close)
|
||||
stager := NewHTTPFileStager(func(string) (string, error) {
|
||||
return strings.TrimPrefix(ts.URL, "http://"), nil
|
||||
}, "")
|
||||
|
||||
backend := &lifecycleBackend{}
|
||||
client := NewFileStagingClient(backend, stager, "node-1")
|
||||
request := &pb.PredictOptions{Audios: []string{localPath}}
|
||||
_, err := client.Predict(context.Background(), request)
|
||||
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(putCalls).To(Equal(1))
|
||||
Expect(backend.predictInput).NotTo(BeNil())
|
||||
Expect(backend.predictInput.Audios).To(Equal([]string{putRemotePath}))
|
||||
},
|
||||
Entry("404", http.StatusNotFound),
|
||||
Entry("405", http.StatusMethodNotAllowed),
|
||||
)
|
||||
|
||||
It("keeps matching model probes read-only", func() {
|
||||
content := []byte("model")
|
||||
localPath := filepath.Join(GinkgoT().TempDir(), "model.bin")
|
||||
Expect(os.WriteFile(localPath, content, 0o600)).To(Succeed())
|
||||
unexpectedWrites := 0
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("/v1/files/", func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodHead {
|
||||
unexpectedWrites++
|
||||
http.Error(w, "unexpected write", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
w.Header().Set(HeaderLocalPath, "/models/tracking/model.bin")
|
||||
w.Header().Set(HeaderContentSHA256, sha256Hex(content))
|
||||
w.WriteHeader(http.StatusOK)
|
||||
})
|
||||
ts := httptest.NewServer(mux)
|
||||
DeferCleanup(ts.Close)
|
||||
stager := NewHTTPFileStager(func(string) (string, error) {
|
||||
return strings.TrimPrefix(ts.URL, "http://"), nil
|
||||
}, "")
|
||||
|
||||
path, err := stager.EnsureRemote(context.Background(), "node-1", localPath, "models/tracking/model.bin")
|
||||
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(path).To(Equal("/models/tracking/model.bin"))
|
||||
Expect(unexpectedWrites).To(BeZero())
|
||||
})
|
||||
|
||||
It("skips upload when file exists with matching hash", func() {
|
||||
ts, stagingDir, _, _ := setupTestServer("tok", 0)
|
||||
|
||||
|
||||
@@ -52,6 +52,8 @@ func (f *fakeFileStager) AllocRemoteTemp(_ context.Context, _ string) (string, e
|
||||
|
||||
func (f *fakeFileStager) StageRemoteToStore(_ context.Context, _, _, _ string) error { return nil }
|
||||
|
||||
func (f *fakeFileStager) ReleaseRemote(_ context.Context, _, _ string) error { return nil }
|
||||
|
||||
func (f *fakeFileStager) ListRemoteDir(_ context.Context, _, _ string) ([]string, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
@@ -47,8 +47,10 @@ type Config struct {
|
||||
PrefetchModels []string `env:"LOCALAI_PREFETCH_MODELS,PREFETCH_MODELS" help:"Comma-separated gallery model IDs to download from LOCALAI_GALLERIES at worker boot (e.g. 'llama-3.2-1b-instruct,phi-3-mini-4k'). Skipped if already on disk and SHA matches." group:"server"`
|
||||
|
||||
// HTTP file transfer
|
||||
HTTPAddr string `env:"LOCALAI_HTTP_ADDR" default:"" help:"HTTP file transfer server address (default: gRPC port + 1)" group:"server" hidden:""`
|
||||
AdvertiseHTTPAddr string `env:"LOCALAI_ADVERTISE_HTTP_ADDR" help:"HTTP address the frontend uses to reach this node for file transfer" group:"server" hidden:""`
|
||||
HTTPAddr string `env:"LOCALAI_HTTP_ADDR" default:"" help:"HTTP file transfer server address (default: gRPC port + 1)" group:"server" hidden:""`
|
||||
AdvertiseHTTPAddr string `env:"LOCALAI_ADVERTISE_HTTP_ADDR" help:"HTTP address the frontend uses to reach this node for file transfer" group:"server" hidden:""`
|
||||
EphemeralStagingByteLimit int64 `env:"LOCALAI_EPHEMERAL_STAGING_BYTE_LIMIT" default:"0" help:"Maximum bytes used by worker request-input staging across HTTP and S3 caches. Zero or negative uses min(10 GiB, 10% of filesystem capacity)." group:"server"`
|
||||
EphemeralStagingMinFreeBytes int64 `env:"LOCALAI_EPHEMERAL_STAGING_MIN_FREE_BYTES" default:"0" help:"Filesystem space kept free while staging request inputs. Zero or negative uses max(1 GiB, 5% of filesystem capacity)." group:"server"`
|
||||
|
||||
// Registration (required)
|
||||
AdvertiseAddr string `env:"LOCALAI_ADVERTISE_ADDR" help:"Address the frontend uses to reach this node (defaults to hostname:port from Addr)" group:"registration" hidden:""`
|
||||
|
||||
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,484 @@
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package worker
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
type capacityGatedFileWriter struct {
|
||||
file *os.File
|
||||
entered chan struct{}
|
||||
resume chan struct{}
|
||||
once sync.Once
|
||||
}
|
||||
|
||||
func (w *capacityGatedFileWriter) Write(p []byte) (int, error) {
|
||||
w.once.Do(func() {
|
||||
close(w.entered)
|
||||
<-w.resume
|
||||
})
|
||||
return w.file.Write(p)
|
||||
}
|
||||
|
||||
type capacityShortWriter struct{}
|
||||
|
||||
func (capacityShortWriter) Write(p []byte) (int, error) {
|
||||
if len(p) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
return 1, nil
|
||||
}
|
||||
|
||||
var _ = Describe("EphemeralCapacityGuard", func() {
|
||||
It("derives bounded defaults and preserves positive overrides", func() {
|
||||
root := canonicalWorkerTempDir()
|
||||
limit, headroom, err := effectiveEphemeralCapacity([]string{root}, 0, -1)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(limit).To(BeNumerically(">", 0))
|
||||
Expect(limit).To(BeNumerically("<=", defaultEphemeralByteLimitCeiling))
|
||||
Expect(headroom).To(BeNumerically(">=", defaultEphemeralMinFreeFloor))
|
||||
|
||||
limit, headroom, err = effectiveEphemeralCapacity([]string{root}, 123, 456)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(limit).To(Equal(int64(123)))
|
||||
Expect(headroom).To(Equal(int64(456)))
|
||||
})
|
||||
|
||||
It("accounts existing regular files without following symlinks", func() {
|
||||
root := canonicalWorkerTempDir()
|
||||
outside := filepath.Join(canonicalWorkerTempDir(), "outside.bin")
|
||||
Expect(os.WriteFile(filepath.Join(root, "existing.bin"), make([]byte, 6), 0o600)).To(Succeed())
|
||||
Expect(os.WriteFile(outside, make([]byte, 100), 0o600)).To(Succeed())
|
||||
Expect(os.Symlink(outside, filepath.Join(root, "outside-link"))).To(Succeed())
|
||||
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 10, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
err = guard.Reserve(filepath.Join(root, "next.bin"), 5)
|
||||
var capacityErr *EphemeralCapacityError
|
||||
Expect(errors.As(err, &capacityErr)).To(BeTrue())
|
||||
Expect(capacityErr.RequestedBytes).To(Equal(int64(5)))
|
||||
Expect(capacityErr.UsageBytes).To(Equal(int64(6)))
|
||||
Expect(capacityErr.LimitBytes).To(Equal(int64(10)))
|
||||
Expect(capacityErr.AvailableBytes).To(BeNumerically(">", 0))
|
||||
Expect(capacityErr.HeadroomBytes).To(BeZero())
|
||||
})
|
||||
|
||||
It("serializes competing reservations", func() {
|
||||
root := canonicalWorkerTempDir()
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 1, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
start := make(chan struct{})
|
||||
results := make(chan error, 32)
|
||||
var wait sync.WaitGroup
|
||||
for i := range 32 {
|
||||
wait.Add(1)
|
||||
go func(index int) {
|
||||
defer wait.Done()
|
||||
<-start
|
||||
results <- guard.Reserve(filepath.Join(root, string(rune('a'+index))), 1)
|
||||
}(i)
|
||||
}
|
||||
close(start)
|
||||
wait.Wait()
|
||||
close(results)
|
||||
|
||||
succeeded := 0
|
||||
for result := range results {
|
||||
if result == nil {
|
||||
succeeded++
|
||||
continue
|
||||
}
|
||||
var capacityErr *EphemeralCapacityError
|
||||
Expect(errors.As(result, &capacityErr)).To(BeTrue())
|
||||
}
|
||||
Expect(succeeded).To(Equal(1))
|
||||
})
|
||||
|
||||
It("makes only an equal active reservation idempotent", func() {
|
||||
root := canonicalWorkerTempDir()
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 10, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
path := filepath.Join(root, "nested", "payload.bin")
|
||||
|
||||
Expect(guard.Reserve(path, 4)).To(Succeed())
|
||||
Expect(guard.Reserve(filepath.Join(root, "nested", ".", "payload.bin"), 4)).To(Succeed())
|
||||
err = guard.Reserve(path, 5)
|
||||
var conflictErr *EphemeralReservationConflictError
|
||||
Expect(errors.As(err, &conflictErr)).To(BeTrue())
|
||||
Expect(conflictErr.ActiveBytes).To(Equal(int64(4)))
|
||||
Expect(conflictErr.RequestedBytes).To(Equal(int64(5)))
|
||||
Expect(guard.Reserve(filepath.Join(root, "other.bin"), 6)).To(Succeed())
|
||||
Expect(guard.Release(filepath.Join(root, "nested", ".", "payload.bin"))).To(Succeed())
|
||||
Expect(guard.Release(path)).To(Succeed())
|
||||
Expect(guard.Reserve(filepath.Join(root, "replacement.bin"), 4)).To(Succeed())
|
||||
})
|
||||
|
||||
It("retains committed bytes when the same path starts another reservation", func() {
|
||||
root := canonicalWorkerTempDir()
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 10, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
path := filepath.Join(root, "payload.bin")
|
||||
|
||||
Expect(guard.Reserve(path, 4)).To(Succeed())
|
||||
Expect(os.WriteFile(path, make([]byte, 4), 0o600)).To(Succeed())
|
||||
Expect(guard.Commit(path)).To(Succeed())
|
||||
Expect(guard.HasActiveReservation(path)).To(BeTrue())
|
||||
Expect(guard.Reserve(path, 6)).To(Succeed())
|
||||
Expect(guard.HasActiveReservation(path)).To(BeTrue())
|
||||
|
||||
err = guard.Reserve(filepath.Join(root, "overflow.bin"), 1)
|
||||
var capacityErr *EphemeralCapacityError
|
||||
Expect(errors.As(err, &capacityErr)).To(BeTrue())
|
||||
Expect(capacityErr.UsageBytes).To(Equal(int64(10)))
|
||||
})
|
||||
|
||||
It("retains startup-accounted bytes when the path is reserved", func() {
|
||||
root := canonicalWorkerTempDir()
|
||||
path := filepath.Join(root, "payload.bin")
|
||||
Expect(os.WriteFile(path, make([]byte, 4), 0o600)).To(Succeed())
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 10, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(guard.HasActiveReservation(path)).To(BeFalse())
|
||||
|
||||
Expect(guard.Reserve(path, 6)).To(Succeed())
|
||||
Expect(guard.HasActiveReservation(path)).To(BeTrue())
|
||||
err = guard.Reserve(filepath.Join(root, "overflow.bin"), 1)
|
||||
var capacityErr *EphemeralCapacityError
|
||||
Expect(errors.As(err, &capacityErr)).To(BeTrue())
|
||||
Expect(capacityErr.UsageBytes).To(Equal(int64(10)))
|
||||
})
|
||||
|
||||
It("commits the regular file's actual size", func() {
|
||||
root := canonicalWorkerTempDir()
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 10, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
path := filepath.Join(root, "payload.bin")
|
||||
|
||||
Expect(guard.Reserve(path, 10)).To(Succeed())
|
||||
Expect(os.WriteFile(path, []byte("four"), 0o600)).To(Succeed())
|
||||
Expect(guard.Commit(path)).To(Succeed())
|
||||
Expect(guard.Commit(filepath.Clean(path))).To(Succeed())
|
||||
Expect(guard.Reserve(filepath.Join(root, "six.bin"), 6)).To(Succeed())
|
||||
})
|
||||
|
||||
It("preserves configured filesystem headroom", func() {
|
||||
root := canonicalWorkerTempDir()
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 1<<30, 1<<62)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
err = guard.Reserve(filepath.Join(root, "payload.bin"), 1)
|
||||
var capacityErr *EphemeralCapacityError
|
||||
Expect(errors.As(err, &capacityErr)).To(BeTrue())
|
||||
Expect(capacityErr.RequestedBytes).To(Equal(int64(1)))
|
||||
Expect(capacityErr.AvailableBytes).To(BeNumerically(">", 0))
|
||||
Expect(capacityErr.HeadroomBytes).To(Equal(int64(1 << 62)))
|
||||
})
|
||||
|
||||
It("reserves bounded chunks before forwarding unknown-length input", func() {
|
||||
root := canonicalWorkerTempDir()
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, ephemeralCapacityWriteChunk+1, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
path := filepath.Join(root, "payload.bin")
|
||||
file, err := os.Create(path)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
DeferCleanup(file.Close)
|
||||
writer, err := guard.NewWriter(path, file)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
n, err := writer.Write(make([]byte, ephemeralCapacityWriteChunk+2))
|
||||
Expect(n).To(Equal(int(ephemeralCapacityWriteChunk)))
|
||||
var capacityErr *EphemeralCapacityError
|
||||
Expect(errors.As(err, &capacityErr)).To(BeTrue())
|
||||
Expect(writer.Close()).To(Succeed())
|
||||
info, err := file.Stat()
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(info.Size()).To(Equal(ephemeralCapacityWriteChunk))
|
||||
})
|
||||
|
||||
It("waits for an open bounded writer before committing", func() {
|
||||
root := canonicalWorkerTempDir()
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 10, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
path := filepath.Join(root, "payload.bin")
|
||||
file, err := os.Create(path)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
DeferCleanup(file.Close)
|
||||
gated := &capacityGatedFileWriter{
|
||||
file: file, entered: make(chan struct{}), resume: make(chan struct{}),
|
||||
}
|
||||
DeferCleanup(func() {
|
||||
select {
|
||||
case <-gated.resume:
|
||||
default:
|
||||
close(gated.resume)
|
||||
}
|
||||
})
|
||||
writer, err := guard.NewWriter(path, gated)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
writeDone := make(chan error, 1)
|
||||
go func() {
|
||||
_, writeErr := writer.Write([]byte("1234567"))
|
||||
writeDone <- writeErr
|
||||
}()
|
||||
Eventually(gated.entered).Should(BeClosed())
|
||||
commitDone := make(chan error, 1)
|
||||
go func() { commitDone <- guard.Commit(path) }()
|
||||
Eventually(func() int { return guard.commitWaiterCount(path) }).Should(Equal(1))
|
||||
Expect(commitDone).NotTo(Receive())
|
||||
|
||||
close(gated.resume)
|
||||
Expect(<-writeDone).To(Succeed())
|
||||
Expect(commitDone).NotTo(Receive())
|
||||
Expect(writer.Close()).To(Succeed())
|
||||
Eventually(commitDone).Should(Receive(Succeed()))
|
||||
|
||||
err = guard.Reserve(filepath.Join(root, "other.bin"), 4)
|
||||
var capacityErr *EphemeralCapacityError
|
||||
Expect(errors.As(err, &capacityErr)).To(BeTrue())
|
||||
Expect(capacityErr.UsageBytes).To(Equal(int64(7)))
|
||||
})
|
||||
|
||||
It("does not share pending capacity between concurrent writers", func() {
|
||||
root := canonicalWorkerTempDir()
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 10, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
path := filepath.Join(root, "payload.bin")
|
||||
file, err := os.Create(path)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
DeferCleanup(file.Close)
|
||||
gated := &capacityGatedFileWriter{
|
||||
file: file, entered: make(chan struct{}), resume: make(chan struct{}),
|
||||
}
|
||||
first, err := guard.NewWriter(path, gated)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
var secondDestination bytes.Buffer
|
||||
second, err := guard.NewWriter(path, &secondDestination)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
firstDone := make(chan error, 1)
|
||||
go func() {
|
||||
_, writeErr := first.Write([]byte("1234567"))
|
||||
firstDone <- writeErr
|
||||
}()
|
||||
Eventually(gated.entered).Should(BeClosed())
|
||||
|
||||
n, err := second.Write([]byte("7654321"))
|
||||
Expect(n).To(BeZero())
|
||||
var capacityErr *EphemeralCapacityError
|
||||
Expect(errors.As(err, &capacityErr)).To(BeTrue())
|
||||
Expect(secondDestination.Len()).To(BeZero())
|
||||
|
||||
close(gated.resume)
|
||||
Expect(<-firstDone).To(Succeed())
|
||||
Expect(first.Close()).To(Succeed())
|
||||
Expect(second.Close()).To(Succeed())
|
||||
})
|
||||
|
||||
It("rolls back bytes the destination writer does not accept", func() {
|
||||
root := canonicalWorkerTempDir()
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 5, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
writer, err := guard.NewWriter(filepath.Join(root, "payload.bin"), capacityShortWriter{})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
n, err := writer.Write([]byte("123"))
|
||||
Expect(n).To(Equal(1))
|
||||
Expect(err).To(MatchError(io.ErrShortWrite))
|
||||
Expect(writer.Close()).To(Succeed())
|
||||
Expect(guard.Reserve(filepath.Join(root, "other.bin"), 4)).To(Succeed())
|
||||
})
|
||||
|
||||
It("rejects paths outside roots and through symlinks", func() {
|
||||
root := canonicalWorkerTempDir()
|
||||
outside := canonicalWorkerTempDir()
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 100, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
Expect(guard.Reserve(filepath.Join(outside, "payload.bin"), 1)).To(
|
||||
MatchError(ContainSubstring("outside registered ephemeral roots")),
|
||||
)
|
||||
Expect(os.Symlink(outside, filepath.Join(root, "escape"))).To(Succeed())
|
||||
Expect(guard.Reserve(filepath.Join(root, "escape", "payload.bin"), 1)).To(
|
||||
MatchError(ContainSubstring("symlink")),
|
||||
)
|
||||
})
|
||||
|
||||
It("supports recovery tree accounting without dropping active reservations", func() {
|
||||
root := canonicalWorkerTempDir()
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 10, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
active := filepath.Join(root, "active", "payload.bin")
|
||||
stale := filepath.Join(root, "stale", "payload.bin")
|
||||
|
||||
Expect(guard.Reserve(active, 4)).To(Succeed())
|
||||
Expect(guard.Account(stale, 3)).To(Succeed())
|
||||
Expect(guard.HasActiveReservation(root)).To(BeTrue())
|
||||
Expect(guard.ReleaseTree(filepath.Join(root, "stale"))).To(Succeed())
|
||||
Expect(guard.Reserve(filepath.Join(root, "replacement.bin"), 6)).To(Succeed())
|
||||
Expect(guard.ReleaseTree(root)).To(Succeed())
|
||||
Expect(guard.HasActiveReservation(root)).To(BeFalse())
|
||||
})
|
||||
|
||||
It("waits for pre-release reservations before request cleanup scans", func() {
|
||||
root := canonicalWorkerTempDir()
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 10, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
path := filepath.Join(root, "audio", "request-1", "input.wav")
|
||||
Expect(os.MkdirAll(filepath.Dir(path), 0o750)).To(Succeed())
|
||||
Expect(guard.Reserve(path, 4)).To(Succeed())
|
||||
|
||||
released := make(chan error, 1)
|
||||
go func() {
|
||||
released <- guard.BeginRequestRelease(context.Background(), "request-1")
|
||||
}()
|
||||
Eventually(guard.releaseTombstoneCount).Should(Equal(1))
|
||||
Consistently(released, 50*time.Millisecond).ShouldNot(Receive())
|
||||
|
||||
Expect(os.WriteFile(path, []byte("data"), 0o600)).To(Succeed())
|
||||
Expect(guard.Commit(path)).To(Succeed())
|
||||
Eventually(released).Should(Receive(Succeed()))
|
||||
guard.EndRequestRelease("request-1")
|
||||
Expect(guard.HasActiveReservation(path)).To(BeFalse())
|
||||
})
|
||||
|
||||
It("rejects staging after request cleanup begins", func() {
|
||||
root := canonicalWorkerTempDir()
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 10, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(guard.BeginRequestRelease(context.Background(), "request-1")).To(Succeed())
|
||||
defer guard.EndRequestRelease("request-1")
|
||||
|
||||
err = guard.Reserve(filepath.Join(root, "audio", "request-1", "late.wav"), 1)
|
||||
var releasedErr *EphemeralRequestReleasedError
|
||||
Expect(errors.As(err, &releasedErr)).To(BeTrue())
|
||||
Expect(releasedErr.RequestID).To(Equal("request-1"))
|
||||
})
|
||||
|
||||
It("leaves a late commit recoverable when release times out", func() {
|
||||
root := canonicalWorkerTempDir()
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 10, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
path := filepath.Join(root, "audio", "request-1", "late.wav")
|
||||
Expect(os.MkdirAll(filepath.Dir(path), 0o750)).To(Succeed())
|
||||
Expect(guard.Reserve(path, 4)).To(Succeed())
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
Expect(guard.BeginRequestRelease(ctx, "request-1")).To(MatchError(context.Canceled))
|
||||
Expect(os.WriteFile(path, []byte("data"), 0o600)).To(Succeed())
|
||||
Expect(guard.Commit(path)).To(Succeed())
|
||||
Expect(guard.HasActiveReservation(path)).To(BeFalse())
|
||||
})
|
||||
|
||||
It("bounds release markers without reopening registered work", func() {
|
||||
root := canonicalWorkerTempDir()
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 10, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(guard.BeginRequestOperation("request-pinned")).To(Succeed())
|
||||
pinnedRelease := make(chan error, 1)
|
||||
go func() {
|
||||
pinnedRelease <- guard.BeginRequestRelease(context.Background(), "request-pinned")
|
||||
}()
|
||||
Eventually(guard.releaseTombstoneCount).Should(Equal(1))
|
||||
guard.EndRequestOperation("request-pinned")
|
||||
Eventually(pinnedRelease).Should(Receive(Succeed()))
|
||||
|
||||
for index := range maxEphemeralReleaseTombstones + 10 {
|
||||
requestID := fmt.Sprintf("request-%d", index)
|
||||
Expect(guard.BeginRequestRelease(context.Background(), requestID)).To(Succeed())
|
||||
guard.EndRequestRelease(requestID)
|
||||
}
|
||||
Expect(guard.releaseTombstoneCount()).To(Equal(maxEphemeralReleaseTombstones))
|
||||
err = guard.Reserve(filepath.Join(root, "audio", "request-pinned", "late.wav"), 1)
|
||||
var releasedErr *EphemeralRequestReleasedError
|
||||
Expect(errors.As(err, &releasedErr)).To(BeTrue())
|
||||
guard.EndRequestRelease("request-pinned")
|
||||
})
|
||||
|
||||
It("applies backpressure at the release-pin cap and clears ownership", func() {
|
||||
root := canonicalWorkerTempDir()
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 10, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
path := filepath.Join(root, "audio", "request-target", "input.wav")
|
||||
Expect(os.MkdirAll(filepath.Dir(path), 0o750)).To(Succeed())
|
||||
Expect(guard.Reserve(path, 4)).To(Succeed())
|
||||
Expect(os.WriteFile(path, []byte("data"), 0o600)).To(Succeed())
|
||||
Expect(guard.Commit(path)).To(Succeed())
|
||||
|
||||
guard.mu.Lock()
|
||||
for index := range maxEphemeralReleaseTombstones {
|
||||
guard.releasePins[fmt.Sprintf("pinned-%d", index)] = 1
|
||||
}
|
||||
guard.mu.Unlock()
|
||||
released := make(chan error, 1)
|
||||
go func() {
|
||||
released <- guard.BeginRequestRelease(context.Background(), "request-target")
|
||||
}()
|
||||
Consistently(released, 50*time.Millisecond).ShouldNot(Receive())
|
||||
|
||||
guard.EndRequestRelease("pinned-0")
|
||||
Eventually(released).Should(Receive(Succeed()))
|
||||
Expect(guard.HasActiveReservation(path)).To(BeFalse())
|
||||
guard.EndRequestRelease("request-target")
|
||||
})
|
||||
|
||||
It("makes committed files recoverable when pin backpressure expires", func() {
|
||||
root := canonicalWorkerTempDir()
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 10, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
path := filepath.Join(root, "audio", "request-target", "input.wav")
|
||||
Expect(os.MkdirAll(filepath.Dir(path), 0o750)).To(Succeed())
|
||||
Expect(guard.Reserve(path, 4)).To(Succeed())
|
||||
Expect(os.WriteFile(path, []byte("data"), 0o600)).To(Succeed())
|
||||
Expect(guard.Commit(path)).To(Succeed())
|
||||
guard.mu.Lock()
|
||||
for index := range maxEphemeralReleaseTombstones {
|
||||
guard.releasePins[fmt.Sprintf("pinned-%d", index)] = 1
|
||||
}
|
||||
guard.mu.Unlock()
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
Expect(guard.BeginRequestRelease(ctx, "request-target")).To(MatchError(context.Canceled))
|
||||
Expect(guard.HasActiveReservation(path)).To(BeFalse())
|
||||
})
|
||||
|
||||
It("rejects a registered cache-hit claim after pin backpressure expires", func() {
|
||||
root := canonicalWorkerTempDir()
|
||||
path := filepath.Join(root, "audio", "request-target", "input.wav")
|
||||
Expect(os.MkdirAll(filepath.Dir(path), 0o750)).To(Succeed())
|
||||
Expect(os.WriteFile(path, []byte("data"), 0o600)).To(Succeed())
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 10, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(guard.BeginRequestOperation("request-target")).To(Succeed())
|
||||
defer guard.EndRequestOperation("request-target")
|
||||
guard.mu.Lock()
|
||||
for index := range maxEphemeralReleaseTombstones {
|
||||
guard.releasePins[fmt.Sprintf("pinned-%d", index)] = 1
|
||||
}
|
||||
guard.mu.Unlock()
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
Expect(guard.BeginRequestRelease(ctx, "request-target")).To(MatchError(context.Canceled))
|
||||
err = guard.Claim(path)
|
||||
var releasedErr *EphemeralRequestReleasedError
|
||||
Expect(errors.As(err, &releasedErr)).To(BeTrue())
|
||||
Expect(guard.HasActiveReservation(path)).To(BeFalse())
|
||||
})
|
||||
})
|
||||
@@ -14,9 +14,9 @@ const (
|
||||
// outlive the request that needed it. Inference reads these files while the
|
||||
// request runs, so the window has to cover a slow multimodal request; it
|
||||
// does not have to cover anything longer.
|
||||
defaultEphemeralStagingTTL = 6 * time.Hour
|
||||
defaultEphemeralStagingTTL = time.Hour
|
||||
// defaultEphemeralStagingSweep is how often the worker sweeps.
|
||||
defaultEphemeralStagingSweep = 30 * time.Minute
|
||||
defaultEphemeralStagingSweep = 15 * time.Minute
|
||||
)
|
||||
|
||||
// StartEphemeralStagingCleanup sweeps the worker's own staging directory for
|
||||
@@ -32,6 +32,15 @@ func StartEphemeralStagingCleanup(ctx context.Context, stagingDir string, ttl, i
|
||||
if stagingDir == "" {
|
||||
return
|
||||
}
|
||||
StartEphemeralRootsCleanup(ctx, []string{filepath.Join(stagingDir, "ephemeral")}, nil, ttl, interval)
|
||||
}
|
||||
|
||||
// StartEphemeralRootsCleanup removes abandoned request inputs for every worker
|
||||
// transport while sharing accounting with live reservations.
|
||||
func StartEphemeralRootsCleanup(ctx context.Context, roots []string, guard *EphemeralCapacityGuard, ttl, interval time.Duration) {
|
||||
if len(roots) == 0 {
|
||||
return
|
||||
}
|
||||
if ttl <= 0 {
|
||||
ttl = defaultEphemeralStagingTTL
|
||||
}
|
||||
@@ -39,31 +48,41 @@ func StartEphemeralStagingCleanup(ctx context.Context, stagingDir string, ttl, i
|
||||
interval = defaultEphemeralStagingSweep
|
||||
}
|
||||
|
||||
// Reclaim crash leftovers before the caller starts accepting new work.
|
||||
CleanEphemeralRoots(roots, ttl, guard)
|
||||
|
||||
go func() {
|
||||
ticker := time.NewTicker(interval)
|
||||
defer ticker.Stop()
|
||||
// Sweep once at startup: a worker that crashed with staged files leaves
|
||||
// them behind, and waiting a full interval to reclaim that space is the
|
||||
// case that hurts on a volume that is already close to full.
|
||||
CleanEphemeralStaging(stagingDir, ttl)
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
CleanEphemeralStaging(stagingDir, ttl)
|
||||
CleanEphemeralRoots(roots, ttl, guard)
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
xlog.Info("Ephemeral staging cleanup started", "dir", stagingDir, "ttl", ttl, "interval", interval)
|
||||
xlog.Info("Ephemeral staging cleanup started", "roots", roots, "ttl", ttl, "interval", interval)
|
||||
}
|
||||
|
||||
// CleanEphemeralStaging removes staged per-request directories older than ttl.
|
||||
// It only ever descends into <stagingDir>/ephemeral, so staged model weights,
|
||||
// which live alongside it and are not scratch, are never considered.
|
||||
func CleanEphemeralStaging(stagingDir string, ttl time.Duration) {
|
||||
root := filepath.Join(stagingDir, "ephemeral")
|
||||
CleanEphemeralRoots([]string{filepath.Join(stagingDir, "ephemeral")}, ttl, nil)
|
||||
}
|
||||
|
||||
// CleanEphemeralRoots removes stale request directories from explicit
|
||||
// ephemeral roots. WalkDir never follows directory symlinks.
|
||||
func CleanEphemeralRoots(roots []string, ttl time.Duration, guard *EphemeralCapacityGuard) {
|
||||
for _, root := range roots {
|
||||
cleanEphemeralRoot(root, ttl, guard)
|
||||
}
|
||||
}
|
||||
|
||||
func cleanEphemeralRoot(root string, ttl time.Duration, guard *EphemeralCapacityGuard) {
|
||||
categories, err := os.ReadDir(root)
|
||||
if err != nil {
|
||||
// A worker that has never served a file-bearing request has no
|
||||
@@ -87,19 +106,31 @@ func CleanEphemeralStaging(stagingDir string, ttl time.Duration) {
|
||||
continue
|
||||
}
|
||||
for _, entry := range entries {
|
||||
path := filepath.Join(categoryDir, entry.Name())
|
||||
info, err := entry.Info()
|
||||
if !entry.IsDir() || entry.Type()&os.ModeSymlink != 0 {
|
||||
continue
|
||||
}
|
||||
requestPath := filepath.Join(categoryDir, entry.Name())
|
||||
newest, err := newestEphemeralModTime(requestPath)
|
||||
if err != nil {
|
||||
xlog.Warn("Ephemeral staging cleanup: cannot stat entry", "path", path, "error", err)
|
||||
xlog.Warn("Ephemeral staging cleanup: cannot inspect request", "path", requestPath, "error", err)
|
||||
continue
|
||||
}
|
||||
// A request rewrites nothing after staging, so the entry's own
|
||||
// modification time is when its request was served.
|
||||
if !info.ModTime().Before(cutoff) {
|
||||
if !newest.Before(cutoff) {
|
||||
continue
|
||||
}
|
||||
if err := os.RemoveAll(path); err != nil {
|
||||
xlog.Warn("Ephemeral staging cleanup: cannot remove", "path", path, "error", err)
|
||||
if guard != nil {
|
||||
removedTree, err := guard.RemoveTreeIfInactive(requestPath, func() error {
|
||||
return os.RemoveAll(requestPath)
|
||||
})
|
||||
if err != nil {
|
||||
xlog.Warn("Ephemeral staging cleanup: cannot remove", "path", requestPath, "error", err)
|
||||
continue
|
||||
}
|
||||
if !removedTree {
|
||||
continue
|
||||
}
|
||||
} else if err := os.RemoveAll(requestPath); err != nil {
|
||||
xlog.Warn("Ephemeral staging cleanup: cannot remove", "path", requestPath, "error", err)
|
||||
continue
|
||||
}
|
||||
removed++
|
||||
@@ -110,3 +141,24 @@ func CleanEphemeralStaging(stagingDir string, ttl time.Duration) {
|
||||
xlog.Info("Ephemeral staging cleanup removed stale request files", "count", removed, "dir", root)
|
||||
}
|
||||
}
|
||||
|
||||
func newestEphemeralModTime(root string) (time.Time, error) {
|
||||
var newest time.Time
|
||||
err := filepath.WalkDir(root, func(path string, entry os.DirEntry, walkErr error) error {
|
||||
if walkErr != nil {
|
||||
return walkErr
|
||||
}
|
||||
info, err := entry.Info()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if info.ModTime().After(newest) {
|
||||
newest = info.ModTime()
|
||||
}
|
||||
if entry.Type()&os.ModeSymlink != 0 && entry.IsDir() {
|
||||
return filepath.SkipDir
|
||||
}
|
||||
return nil
|
||||
})
|
||||
return newest, err
|
||||
}
|
||||
@@ -24,7 +24,7 @@ var _ = Describe("Worker ephemeral staging cleanup", func() {
|
||||
return dir
|
||||
}
|
||||
|
||||
BeforeEach(func() { stagingDir = GinkgoT().TempDir() })
|
||||
BeforeEach(func() { stagingDir = canonicalWorkerTempDir() })
|
||||
|
||||
It("removes staged request directories older than the TTL", func() {
|
||||
old := mkEphemeral("aaaa1111", 48*time.Hour)
|
||||
@@ -55,4 +55,60 @@ var _ = Describe("Worker ephemeral staging cleanup", func() {
|
||||
It("does nothing when no ephemeral directory exists", func() {
|
||||
Expect(func() { CleanEphemeralStaging(stagingDir, time.Hour) }).ToNot(Panic())
|
||||
})
|
||||
|
||||
It("sweeps both transport roots by newest descendant and skips active requests", func() {
|
||||
cacheDir := canonicalWorkerTempDir()
|
||||
httpRoot := filepath.Join(stagingDir, "ephemeral")
|
||||
s3Root := filepath.Join(cacheDir, "ephemeral")
|
||||
guard, err := NewEphemeralCapacityGuard([]string{httpRoot, s3Root}, 8, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
staleRequest := filepath.Join(httpRoot, "audio", "stale")
|
||||
activeRequest := filepath.Join(s3Root, "audio", "active")
|
||||
freshChildRequest := filepath.Join(s3Root, "audio", "fresh-child")
|
||||
for _, requestDir := range []string{staleRequest, activeRequest, freshChildRequest} {
|
||||
Expect(os.MkdirAll(requestDir, 0o750)).To(Succeed())
|
||||
Expect(os.WriteFile(filepath.Join(requestDir, "input.bin"), []byte("data"), 0o600)).To(Succeed())
|
||||
}
|
||||
old := time.Now().Add(-2 * time.Hour)
|
||||
fresh := time.Now().Add(-5 * time.Minute)
|
||||
for _, requestDir := range []string{staleRequest, activeRequest, freshChildRequest} {
|
||||
Expect(os.Chtimes(requestDir, old, old)).To(Succeed())
|
||||
}
|
||||
Expect(os.Chtimes(filepath.Join(staleRequest, "input.bin"), old, old)).To(Succeed())
|
||||
Expect(os.Chtimes(filepath.Join(activeRequest, "input.bin"), old, old)).To(Succeed())
|
||||
Expect(os.Chtimes(filepath.Join(freshChildRequest, "input.bin"), fresh, fresh)).To(Succeed())
|
||||
Expect(guard.Account(filepath.Join(staleRequest, "input.bin"), 4)).To(Succeed())
|
||||
Expect(guard.Reserve(filepath.Join(activeRequest, "input.bin"), 4)).To(Succeed())
|
||||
|
||||
CleanEphemeralRoots([]string{httpRoot, s3Root}, time.Hour, guard)
|
||||
|
||||
Expect(staleRequest).NotTo(BeADirectory())
|
||||
Expect(activeRequest).To(BeADirectory())
|
||||
Expect(freshChildRequest).To(BeADirectory())
|
||||
Expect(guard.Reserve(filepath.Join(httpRoot, "audio", "replacement", "input.bin"), 4)).To(Succeed())
|
||||
})
|
||||
|
||||
It("keeps committed request inputs until exact release ends ownership", func() {
|
||||
root := filepath.Join(stagingDir, "ephemeral")
|
||||
requestDir := filepath.Join(root, "audio", "owned")
|
||||
path := filepath.Join(requestDir, "input.bin")
|
||||
Expect(os.MkdirAll(requestDir, 0o750)).To(Succeed())
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 8, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
Expect(guard.Reserve(path, 4)).To(Succeed())
|
||||
Expect(os.WriteFile(path, []byte("data"), 0o600)).To(Succeed())
|
||||
Expect(guard.Commit(path)).To(Succeed())
|
||||
old := time.Now().Add(-2 * time.Hour)
|
||||
Expect(os.Chtimes(path, old, old)).To(Succeed())
|
||||
Expect(os.Chtimes(requestDir, old, old)).To(Succeed())
|
||||
|
||||
CleanEphemeralRoots([]string{root}, time.Hour, guard)
|
||||
Expect(requestDir).To(BeADirectory())
|
||||
|
||||
Expect(guard.Release(path)).To(Succeed())
|
||||
CleanEphemeralRoots([]string{root}, time.Hour, guard)
|
||||
Expect(requestDir).NotTo(BeADirectory())
|
||||
})
|
||||
})
|
||||
@@ -3,14 +3,19 @@ package worker
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/mudler/LocalAI/core/services/messaging"
|
||||
"github.com/mudler/LocalAI/core/services/storage"
|
||||
"github.com/mudler/LocalAI/pkg/safefile"
|
||||
"github.com/mudler/xlog"
|
||||
"golang.org/x/sync/singleflight"
|
||||
)
|
||||
|
||||
// isPathAllowed checks if path is within one of the allowed directories.
|
||||
@@ -37,7 +42,7 @@ func isPathAllowed(path string, allowedDirs []string) bool {
|
||||
}
|
||||
|
||||
// subscribeFileStaging subscribes to NATS file staging subjects for this node.
|
||||
func (cfg *Config) subscribeFileStaging(natsClient messaging.MessagingClient, nodeID string) error {
|
||||
func (cfg *Config) subscribeFileStaging(natsClient messaging.MessagingClient, nodeID string, capacity *EphemeralCapacityGuard) error {
|
||||
// Create FileManager with same S3 config as the frontend
|
||||
// TODO: propagate a caller-provided context once Config carries one
|
||||
s3Store, err := storage.NewS3Store(context.Background(), storage.S3Config{
|
||||
@@ -57,6 +62,10 @@ func (cfg *Config) subscribeFileStaging(natsClient messaging.MessagingClient, no
|
||||
if err != nil {
|
||||
return fmt.Errorf("initializing file manager: %w", err)
|
||||
}
|
||||
if err := subscribeFileReleaseWithCapacity(natsClient, nodeID, fm, cacheDir, capacity); err != nil {
|
||||
return err
|
||||
}
|
||||
var ensureGroup singleflight.Group
|
||||
|
||||
// Subscribe: files.ensure — download S3 key to local, reply with local path
|
||||
if _, err := natsClient.SubscribeReply(messaging.SubjectNodeFilesEnsure(nodeID), func(data []byte, reply func([]byte)) {
|
||||
@@ -68,12 +77,19 @@ func (cfg *Config) subscribeFileStaging(natsClient messaging.MessagingClient, no
|
||||
return
|
||||
}
|
||||
|
||||
localPath, err := fm.Download(context.Background(), req.Key)
|
||||
value, err, _ := ensureGroup.Do(req.Key, func() (any, error) {
|
||||
return ensureWorkerFile(context.Background(), fm, capacity, req.Key)
|
||||
})
|
||||
if err != nil {
|
||||
xlog.Error("File ensure failed", "key", req.Key, "error", err)
|
||||
replyJSON(reply, map[string]string{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
localPath, ok := value.(string)
|
||||
if !ok {
|
||||
replyJSON(reply, map[string]string{"error": fmt.Sprintf("unexpected file ensure result %T", value)})
|
||||
return
|
||||
}
|
||||
|
||||
xlog.Debug("File ensured locally", "key", req.Key, "path", localPath)
|
||||
replyJSON(reply, map[string]string{"local_path": localPath})
|
||||
@@ -199,3 +215,220 @@ func (cfg *Config) subscribeFileStaging(natsClient messaging.MessagingClient, no
|
||||
xlog.Info("Subscribed to file staging NATS subjects", "nodeID", nodeID)
|
||||
return nil
|
||||
}
|
||||
|
||||
func subscribeFileRelease(natsClient messaging.MessagingClient, nodeID string, fm *storage.FileManager, cacheDir string) error {
|
||||
return subscribeFileReleaseWithCapacity(natsClient, nodeID, fm, cacheDir, nil)
|
||||
}
|
||||
|
||||
func subscribeFileReleaseWithCapacity(natsClient messaging.MessagingClient, nodeID string, fm *storage.FileManager, cacheDir string, capacity *EphemeralCapacityGuard) error {
|
||||
if _, err := natsClient.SubscribeReply(messaging.SubjectNodeFilesRelease(nodeID), func(data []byte, reply func([]byte)) {
|
||||
var req struct {
|
||||
Key string `json:"key"`
|
||||
RequestID string `json:"request_id"`
|
||||
}
|
||||
if err := json.Unmarshal(data, &req); err != nil {
|
||||
replyJSON(reply, map[string]string{"error": "invalid request"})
|
||||
return
|
||||
}
|
||||
var err error
|
||||
if req.RequestID != "" {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
err = releaseEphemeralCacheRequest(ctx, cacheDir, req.RequestID, capacity)
|
||||
cancel()
|
||||
} else {
|
||||
cachePath, cacheErr := fm.CachePath(req.Key)
|
||||
err = cacheErr
|
||||
if err == nil {
|
||||
err = releaseEphemeralCachePathWithCapacity(cacheDir, req.Key, cachePath, capacity)
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
replyJSON(reply, map[string]string{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
replyJSON(reply, map[string]string{})
|
||||
}); err != nil {
|
||||
return fmt.Errorf("subscribing to files.release events: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func releaseEphemeralCacheKey(cacheDir, key string) error {
|
||||
return releaseEphemeralCachePath(cacheDir, key, filepath.Join(cacheDir, filepath.FromSlash(key)))
|
||||
}
|
||||
|
||||
func releaseEphemeralCachePath(cacheDir, key, filePath string) error {
|
||||
return releaseEphemeralCachePathWithCapacity(cacheDir, key, filePath, nil)
|
||||
}
|
||||
|
||||
func releaseEphemeralCachePathWithCapacity(cacheDir, key, filePath string, capacity *EphemeralCapacityGuard) error {
|
||||
if err := validateEphemeralCacheKey(key); err != nil {
|
||||
return err
|
||||
}
|
||||
relativePath := filepath.FromSlash(key)
|
||||
expectedPath := filepath.Join(cacheDir, relativePath)
|
||||
if filepath.Clean(filePath) != expectedPath {
|
||||
return fmt.Errorf("release path %q does not match key %q", filePath, key)
|
||||
}
|
||||
if err := safefile.RemoveExact(cacheDir, relativePath, []string{".sha256", ".sha256.target"}, 2); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, path := range []string{filePath, filePath + ".sha256", filePath + ".sha256.target"} {
|
||||
if capacity != nil {
|
||||
if err := capacity.Release(path); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func releaseEphemeralCacheRequest(ctx context.Context, cacheDir, requestID string, capacity *EphemeralCapacityGuard) error {
|
||||
if err := validateEphemeralCacheRequestID(requestID); err != nil {
|
||||
return err
|
||||
}
|
||||
if capacity != nil {
|
||||
if err := capacity.BeginRequestRelease(ctx, requestID); err != nil {
|
||||
return fmt.Errorf("beginning release for request %q: %w", requestID, err)
|
||||
}
|
||||
defer capacity.EndRequestRelease(requestID)
|
||||
}
|
||||
root := filepath.Join(cacheDir, "ephemeral")
|
||||
categories, err := os.ReadDir(root)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
var releaseErrors []error
|
||||
for _, category := range categories {
|
||||
if !category.IsDir() || category.Type()&os.ModeSymlink != 0 {
|
||||
continue
|
||||
}
|
||||
requestDir := filepath.Join(root, category.Name(), requestID)
|
||||
info, err := os.Lstat(requestDir)
|
||||
if err != nil {
|
||||
if !os.IsNotExist(err) {
|
||||
releaseErrors = append(releaseErrors, fmt.Errorf("stating request directory %q: %w", requestDir, err))
|
||||
}
|
||||
continue
|
||||
}
|
||||
if !info.IsDir() || info.Mode()&os.ModeSymlink != 0 {
|
||||
releaseErrors = append(releaseErrors, fmt.Errorf("ephemeral request path %q is not a real directory", requestDir))
|
||||
continue
|
||||
}
|
||||
entries, err := os.ReadDir(requestDir)
|
||||
if err != nil {
|
||||
if !os.IsNotExist(err) {
|
||||
releaseErrors = append(releaseErrors, fmt.Errorf("reading request directory %q: %w", requestDir, err))
|
||||
}
|
||||
continue
|
||||
}
|
||||
for _, entry := range entries {
|
||||
if entry.IsDir() {
|
||||
releaseErrors = append(releaseErrors, fmt.Errorf("unexpected directory in ephemeral request %q", filepath.Join(requestDir, entry.Name())))
|
||||
continue
|
||||
}
|
||||
key := filepath.ToSlash(filepath.Join("ephemeral", category.Name(), requestID, entry.Name()))
|
||||
filePath := filepath.Join(cacheDir, filepath.FromSlash(key))
|
||||
if err := releaseEphemeralCachePathWithCapacity(cacheDir, key, filePath, capacity); err != nil {
|
||||
releaseErrors = append(releaseErrors, fmt.Errorf("releasing %q: %w", key, err))
|
||||
}
|
||||
}
|
||||
}
|
||||
return errors.Join(releaseErrors...)
|
||||
}
|
||||
|
||||
type ephemeralStagingCapacity interface {
|
||||
Reserve(path string, size int64) error
|
||||
Commit(path string) error
|
||||
Claim(path string) error
|
||||
Release(path string) error
|
||||
}
|
||||
|
||||
func ensureWorkerFile(ctx context.Context, fm *storage.FileManager, capacity *EphemeralCapacityGuard, key string) (string, error) {
|
||||
if capacity == nil {
|
||||
return fm.Download(ctx, key)
|
||||
}
|
||||
if strings.HasPrefix(key, "ephemeral/") {
|
||||
if err := validateEphemeralCacheKey(key); err != nil {
|
||||
return "", err
|
||||
}
|
||||
requestID := strings.Split(key, "/")[2]
|
||||
if err := capacity.BeginRequestOperation(requestID); err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer capacity.EndRequestOperation(requestID)
|
||||
}
|
||||
return ensureWorkerFileWithCapacity(ctx, fm, capacity, key)
|
||||
}
|
||||
|
||||
func ensureWorkerFileWithCapacity(ctx context.Context, fm *storage.FileManager, capacity ephemeralStagingCapacity, key string) (string, error) {
|
||||
if capacity == nil || !strings.HasPrefix(key, "ephemeral/") {
|
||||
return fm.Download(ctx, key)
|
||||
}
|
||||
if err := validateEphemeralCacheKey(key); err != nil {
|
||||
return "", err
|
||||
}
|
||||
cachePath, err := fm.CachePath(key)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if info, statErr := os.Lstat(cachePath); statErr == nil {
|
||||
if !info.Mode().IsRegular() {
|
||||
return "", fmt.Errorf("ephemeral cache path %q is not a regular file", cachePath)
|
||||
}
|
||||
if err := capacity.Claim(cachePath); err != nil {
|
||||
if !errors.Is(err, os.ErrNotExist) {
|
||||
return "", err
|
||||
}
|
||||
} else {
|
||||
return cachePath, nil
|
||||
}
|
||||
} else if !os.IsNotExist(statErr) {
|
||||
return "", statErr
|
||||
}
|
||||
|
||||
meta, err := fm.Head(ctx, key)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("reading size for %s: %w", key, err)
|
||||
}
|
||||
if err := capacity.Reserve(cachePath, meta.Size); err != nil {
|
||||
return "", err
|
||||
}
|
||||
localPath, err := fm.Download(ctx, key)
|
||||
if err != nil {
|
||||
_ = capacity.Release(cachePath)
|
||||
return "", err
|
||||
}
|
||||
if err := capacity.Commit(cachePath); err != nil {
|
||||
_ = fm.EvictCache(key)
|
||||
_ = capacity.Release(cachePath)
|
||||
return "", err
|
||||
}
|
||||
return localPath, nil
|
||||
}
|
||||
|
||||
func validateEphemeralCacheKey(key string) error {
|
||||
if strings.Contains(key, "\\") || path.Clean(key) != key {
|
||||
return fmt.Errorf("invalid ephemeral key %q", key)
|
||||
}
|
||||
parts := strings.Split(key, "/")
|
||||
if len(parts) != 4 || parts[0] != "ephemeral" {
|
||||
return fmt.Errorf("release key %q must identify one file below ephemeral/", key)
|
||||
}
|
||||
for _, part := range parts[1:] {
|
||||
if part == "" || part == "." || part == ".." {
|
||||
return fmt.Errorf("invalid ephemeral key %q", key)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateEphemeralCacheRequestID(requestID string) error {
|
||||
if requestID == "" || strings.ContainsAny(requestID, "/\\") || path.Clean(requestID) != requestID || requestID == "." || requestID == ".." {
|
||||
return fmt.Errorf("invalid ephemeral request ID %q", requestID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,405 @@
|
||||
package worker
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/mudler/LocalAI/core/services/messaging"
|
||||
"github.com/mudler/LocalAI/core/services/nodes"
|
||||
"github.com/mudler/LocalAI/core/services/storage"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
type stagingObjectStore struct {
|
||||
payload []byte
|
||||
getCalls int
|
||||
getErr error
|
||||
}
|
||||
|
||||
type disappearingStagingCapacity struct{}
|
||||
|
||||
func (*disappearingStagingCapacity) Reserve(string, int64) error { return nil }
|
||||
func (*disappearingStagingCapacity) Commit(string) error { return nil }
|
||||
func (*disappearingStagingCapacity) Release(string) error { return nil }
|
||||
func (*disappearingStagingCapacity) Claim(path string) error {
|
||||
if err := os.Remove(path); err != nil {
|
||||
return err
|
||||
}
|
||||
return fmt.Errorf("claim raced recovery: %w", os.ErrNotExist)
|
||||
}
|
||||
|
||||
func (*stagingObjectStore) Put(context.Context, string, io.Reader) error { return nil }
|
||||
func (s *stagingObjectStore) Get(context.Context, string) (io.ReadCloser, error) {
|
||||
s.getCalls++
|
||||
if s.getErr != nil {
|
||||
return nil, s.getErr
|
||||
}
|
||||
return io.NopCloser(strings.NewReader(string(s.payload))), nil
|
||||
}
|
||||
func (s *stagingObjectStore) Head(_ context.Context, key string) (*storage.ObjectMeta, error) {
|
||||
return &storage.ObjectMeta{Key: key, Size: int64(len(s.payload))}, nil
|
||||
}
|
||||
func (*stagingObjectStore) Exists(context.Context, string) (bool, error) { return true, nil }
|
||||
func (*stagingObjectStore) Delete(context.Context, string) error { return nil }
|
||||
func (*stagingObjectStore) List(context.Context, string) ([]string, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
type releaseSubscription struct{}
|
||||
|
||||
func (releaseSubscription) Unsubscribe() error { return nil }
|
||||
|
||||
type releaseMessagingClient struct {
|
||||
subject string
|
||||
handler func([]byte, func([]byte))
|
||||
}
|
||||
|
||||
func (m *releaseMessagingClient) Publish(string, any) error { return nil }
|
||||
func (m *releaseMessagingClient) Subscribe(string, func([]byte)) (messaging.Subscription, error) {
|
||||
return releaseSubscription{}, nil
|
||||
}
|
||||
func (m *releaseMessagingClient) QueueSubscribe(string, string, func([]byte)) (messaging.Subscription, error) {
|
||||
return releaseSubscription{}, nil
|
||||
}
|
||||
func (m *releaseMessagingClient) QueueSubscribeReply(string, string, func([]byte, func([]byte))) (messaging.Subscription, error) {
|
||||
return releaseSubscription{}, nil
|
||||
}
|
||||
func (m *releaseMessagingClient) SubscribeReply(subject string, handler func([]byte, func([]byte))) (messaging.Subscription, error) {
|
||||
m.subject = subject
|
||||
m.handler = handler
|
||||
return releaseSubscription{}, nil
|
||||
}
|
||||
func (m *releaseMessagingClient) Request(string, []byte, time.Duration) ([]byte, error) {
|
||||
return nil, nil
|
||||
}
|
||||
func (m *releaseMessagingClient) IsConnected() bool { return true }
|
||||
func (m *releaseMessagingClient) Close() {}
|
||||
|
||||
var _ = Describe("Worker exact-key staging release", func() {
|
||||
It("protects a startup-accounted HTTP cache hit through authenticated repeated probes", func() {
|
||||
stagingDir := canonicalWorkerTempDir()
|
||||
root := filepath.Join(stagingDir, "ephemeral")
|
||||
key := "ephemeral/audio/request-id/input.wav"
|
||||
remotePath := filepath.Join(stagingDir, filepath.FromSlash(key))
|
||||
content := []byte("data")
|
||||
Expect(os.MkdirAll(filepath.Dir(remotePath), 0o750)).To(Succeed())
|
||||
Expect(os.WriteFile(remotePath, content, 0o600)).To(Succeed())
|
||||
hash := sha256.Sum256(content)
|
||||
Expect(os.WriteFile(remotePath+".sha256", []byte(fmt.Sprintf("%x", hash)), 0o600)).To(Succeed())
|
||||
old := time.Now().Add(-2 * time.Hour)
|
||||
for _, path := range []string{remotePath, remotePath + ".sha256", filepath.Dir(remotePath)} {
|
||||
Expect(os.Chtimes(path, old, old)).To(Succeed())
|
||||
}
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, int64(len(content)+sha256.Size*2), 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
addr := listener.Addr().String()
|
||||
Expect(listener.Close()).To(Succeed())
|
||||
server, err := nodes.StartFileTransferServerWithCapacity(addr, stagingDir, canonicalWorkerTempDir(), canonicalWorkerTempDir(), "secret", 0, nil, guard)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
DeferCleanup(nodes.ShutdownFileTransferServer, server)
|
||||
|
||||
localPath := filepath.Join(canonicalWorkerTempDir(), "input.wav")
|
||||
Expect(os.WriteFile(localPath, content, 0o600)).To(Succeed())
|
||||
stager := nodes.NewHTTPFileStager(func(string) (string, error) { return addr, nil }, "secret")
|
||||
for range 2 {
|
||||
path, ensureErr := stager.EnsureRemote(context.Background(), "worker", localPath, key)
|
||||
Expect(ensureErr).NotTo(HaveOccurred())
|
||||
Expect(path).To(Equal(remotePath))
|
||||
}
|
||||
|
||||
CleanEphemeralRoots([]string{root}, time.Hour, guard)
|
||||
Expect(remotePath).To(BeAnExistingFile())
|
||||
|
||||
Expect(stager.ReleaseRemote(context.Background(), "worker", key)).To(Succeed())
|
||||
CleanEphemeralRoots([]string{root}, time.Hour, guard)
|
||||
Expect(remotePath).NotTo(BeAnExistingFile())
|
||||
})
|
||||
|
||||
It("claims a startup-scanned cache hit against stale recovery until release", func() {
|
||||
cacheDir := canonicalWorkerTempDir()
|
||||
root := filepath.Join(cacheDir, "ephemeral")
|
||||
key := "ephemeral/audio/request-id/input.wav"
|
||||
cachePath := filepath.Join(cacheDir, filepath.FromSlash(key))
|
||||
Expect(os.MkdirAll(filepath.Dir(cachePath), 0o750)).To(Succeed())
|
||||
Expect(os.WriteFile(cachePath, []byte("data"), 0o600)).To(Succeed())
|
||||
old := time.Now().Add(-2 * time.Hour)
|
||||
Expect(os.Chtimes(cachePath, old, old)).To(Succeed())
|
||||
Expect(os.Chtimes(filepath.Dir(cachePath), old, old)).To(Succeed())
|
||||
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 4, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
store := &stagingObjectStore{payload: []byte("unused")}
|
||||
fm, err := storage.NewFileManager(store, cacheDir)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
localPath, err := ensureWorkerFile(context.Background(), fm, guard, key)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(localPath).To(Equal(cachePath))
|
||||
Expect(store.getCalls).To(BeZero())
|
||||
CleanEphemeralRoots([]string{root}, time.Hour, guard)
|
||||
Expect(cachePath).To(BeAnExistingFile())
|
||||
|
||||
Expect(guard.Release(cachePath)).To(Succeed())
|
||||
CleanEphemeralRoots([]string{root}, time.Hour, guard)
|
||||
Expect(cachePath).NotTo(BeAnExistingFile())
|
||||
})
|
||||
|
||||
It("downloads again when a cache file disappears while being claimed", func() {
|
||||
cacheDir := canonicalWorkerTempDir()
|
||||
key := "ephemeral/audio/request-id/input.wav"
|
||||
cachePath := filepath.Join(cacheDir, filepath.FromSlash(key))
|
||||
Expect(os.MkdirAll(filepath.Dir(cachePath), 0o750)).To(Succeed())
|
||||
Expect(os.WriteFile(cachePath, []byte("stale"), 0o600)).To(Succeed())
|
||||
store := &stagingObjectStore{payload: []byte("fresh")}
|
||||
fm, err := storage.NewFileManager(store, cacheDir)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
localPath, err := ensureWorkerFileWithCapacity(context.Background(), fm, &disappearingStagingCapacity{}, key)
|
||||
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(localPath).To(Equal(cachePath))
|
||||
Expect(os.ReadFile(localPath)).To(Equal([]byte("fresh")))
|
||||
Expect(store.getCalls).To(Equal(1))
|
||||
})
|
||||
|
||||
It("makes repeated cache-hit claims idempotent", func() {
|
||||
cacheDir := canonicalWorkerTempDir()
|
||||
root := filepath.Join(cacheDir, "ephemeral")
|
||||
key := "ephemeral/audio/request-id/input.wav"
|
||||
cachePath := filepath.Join(cacheDir, filepath.FromSlash(key))
|
||||
Expect(os.MkdirAll(filepath.Dir(cachePath), 0o750)).To(Succeed())
|
||||
Expect(os.WriteFile(cachePath, []byte("data"), 0o600)).To(Succeed())
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 4, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
fm, err := storage.NewFileManager(&stagingObjectStore{}, cacheDir)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
for range 2 {
|
||||
localPath, ensureErr := ensureWorkerFile(context.Background(), fm, guard, key)
|
||||
Expect(ensureErr).NotTo(HaveOccurred())
|
||||
Expect(localPath).To(Equal(cachePath))
|
||||
}
|
||||
err = guard.Reserve(filepath.Join(root, "other", "request-id", "input.wav"), 1)
|
||||
var capacityErr *EphemeralCapacityError
|
||||
Expect(errors.As(err, &capacityErr)).To(BeTrue())
|
||||
Expect(capacityErr.UsageBytes).To(Equal(int64(4)))
|
||||
})
|
||||
|
||||
It("capacity-checks growth of a startup-scanned cache file", func() {
|
||||
cacheDir := canonicalWorkerTempDir()
|
||||
root := filepath.Join(cacheDir, "ephemeral")
|
||||
key := "ephemeral/audio/request-id/input.wav"
|
||||
cachePath := filepath.Join(cacheDir, filepath.FromSlash(key))
|
||||
Expect(os.MkdirAll(filepath.Dir(cachePath), 0o750)).To(Succeed())
|
||||
Expect(os.WriteFile(cachePath, []byte("12"), 0o600)).To(Succeed())
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 4, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(os.WriteFile(cachePath, []byte("12345"), 0o600)).To(Succeed())
|
||||
fm, err := storage.NewFileManager(&stagingObjectStore{}, cacheDir)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
_, err = ensureWorkerFile(context.Background(), fm, guard, key)
|
||||
var capacityErr *EphemeralCapacityError
|
||||
Expect(errors.As(err, &capacityErr)).To(BeTrue())
|
||||
Expect(capacityErr.RequestedBytes).To(Equal(int64(3)))
|
||||
Expect(capacityErr.UsageBytes).To(Equal(int64(2)))
|
||||
Expect(guard.HasActiveReservation(cachePath)).To(BeFalse())
|
||||
})
|
||||
|
||||
It("reserves S3 object size before download and releases it with the exact key", func() {
|
||||
cacheDir := canonicalWorkerTempDir()
|
||||
root := filepath.Join(cacheDir, "ephemeral")
|
||||
store := &stagingObjectStore{payload: []byte("data")}
|
||||
fm, err := storage.NewFileManager(store, cacheDir)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 4, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
key := "ephemeral/audio/request-id/input.wav"
|
||||
|
||||
localPath, err := ensureWorkerFile(context.Background(), fm, guard, key)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(localPath).To(BeAnExistingFile())
|
||||
Expect(store.getCalls).To(Equal(1))
|
||||
Expect(guard.Reserve(filepath.Join(root, "audio", "other", "input.wav"), 1)).NotTo(Succeed())
|
||||
|
||||
Expect(releaseEphemeralCachePathWithCapacity(cacheDir, key, localPath, guard)).To(Succeed())
|
||||
Expect(guard.Reserve(filepath.Join(root, "audio", "other", "input.wav"), 4)).To(Succeed())
|
||||
})
|
||||
|
||||
It("rejects an oversized S3 object before starting its download", func() {
|
||||
cacheDir := canonicalWorkerTempDir()
|
||||
store := &stagingObjectStore{payload: []byte("oversized")}
|
||||
fm, err := storage.NewFileManager(store, cacheDir)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
guard, err := NewEphemeralCapacityGuard([]string{filepath.Join(cacheDir, "ephemeral")}, 4, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
_, err = ensureWorkerFile(context.Background(), fm, guard, "ephemeral/audio/request-id/input.wav")
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(store.getCalls).To(BeZero())
|
||||
})
|
||||
|
||||
It("rolls back an S3 reservation when the download fails", func() {
|
||||
cacheDir := canonicalWorkerTempDir()
|
||||
root := filepath.Join(cacheDir, "ephemeral")
|
||||
store := &stagingObjectStore{payload: []byte("data"), getErr: errors.New("download failed")}
|
||||
fm, err := storage.NewFileManager(store, cacheDir)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
guard, err := NewEphemeralCapacityGuard([]string{root}, 4, 0)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
|
||||
_, err = ensureWorkerFile(context.Background(), fm, guard, "ephemeral/audio/request-id/input.wav")
|
||||
Expect(err).To(MatchError(ContainSubstring("download failed")))
|
||||
Expect(guard.Reserve(filepath.Join(root, "audio", "replacement", "input.wav"), 4)).To(Succeed())
|
||||
})
|
||||
|
||||
It("removes only the exact cache file and upload sidecars", func() {
|
||||
cacheDir := canonicalWorkerTempDir()
|
||||
categoryDir := filepath.Join(cacheDir, "ephemeral", "request-id", "audio")
|
||||
Expect(os.MkdirAll(categoryDir, 0750)).To(Succeed())
|
||||
target := filepath.Join(categoryDir, "input.wav")
|
||||
sibling := filepath.Join(categoryDir, "keep.wav")
|
||||
for _, path := range []string{target, target + ".sha256", target + ".sha256.target", sibling} {
|
||||
Expect(os.WriteFile(path, []byte("data"), 0640)).To(Succeed())
|
||||
}
|
||||
|
||||
Expect(releaseEphemeralCacheKey(cacheDir, "ephemeral/request-id/audio/input.wav")).To(Succeed())
|
||||
Expect(target).NotTo(BeAnExistingFile())
|
||||
Expect(target + ".sha256").NotTo(BeAnExistingFile())
|
||||
Expect(target + ".sha256.target").NotTo(BeAnExistingFile())
|
||||
Expect(sibling).To(BeAnExistingFile())
|
||||
Expect(categoryDir).To(BeADirectory())
|
||||
})
|
||||
|
||||
It("succeeds for a missing file and prunes empty category and request directories", func() {
|
||||
cacheDir := canonicalWorkerTempDir()
|
||||
categoryDir := filepath.Join(cacheDir, "ephemeral", "request-id", "audio")
|
||||
Expect(os.MkdirAll(categoryDir, 0750)).To(Succeed())
|
||||
|
||||
for range 2 {
|
||||
Expect(releaseEphemeralCacheKey(cacheDir, "ephemeral/request-id/audio/missing.wav")).To(Succeed())
|
||||
}
|
||||
Expect(categoryDir).NotTo(BeADirectory())
|
||||
Expect(filepath.Dir(categoryDir)).NotTo(BeADirectory())
|
||||
Expect(filepath.Join(cacheDir, "ephemeral")).To(BeADirectory())
|
||||
})
|
||||
|
||||
It("rejects traversal and symlink escapes", func() {
|
||||
cacheDir := canonicalWorkerTempDir()
|
||||
outsideDir := canonicalWorkerTempDir()
|
||||
outsidePath := filepath.Join(outsideDir, "input.wav")
|
||||
Expect(os.WriteFile(outsidePath, []byte("keep"), 0640)).To(Succeed())
|
||||
requestDir := filepath.Join(cacheDir, "ephemeral", "request-id")
|
||||
Expect(os.MkdirAll(requestDir, 0750)).To(Succeed())
|
||||
Expect(os.Symlink(outsideDir, filepath.Join(requestDir, "audio"))).To(Succeed())
|
||||
|
||||
for _, key := range []string{
|
||||
"models/model.gguf",
|
||||
"ephemeral/../models/model.gguf",
|
||||
"ephemeral/request-id/audio/../../model.gguf",
|
||||
"ephemeral/request-id/audio/input.wav",
|
||||
} {
|
||||
Expect(releaseEphemeralCacheKey(cacheDir, key)).NotTo(Succeed(), key)
|
||||
}
|
||||
Expect(outsidePath).To(BeAnExistingFile())
|
||||
})
|
||||
|
||||
It("rejects symlinked files and sidecars without deleting their targets", func() {
|
||||
for _, linkedName := range []string{"input.wav", "input.wav.sha256", "input.wav.sha256.target"} {
|
||||
cacheDir := canonicalWorkerTempDir()
|
||||
categoryDir := filepath.Join(cacheDir, "ephemeral", "request-id", "audio")
|
||||
Expect(os.MkdirAll(categoryDir, 0750)).To(Succeed())
|
||||
target := filepath.Join(categoryDir, "input.wav")
|
||||
if linkedName != "input.wav" {
|
||||
Expect(os.WriteFile(target, []byte("input"), 0640)).To(Succeed())
|
||||
}
|
||||
preserved := filepath.Join(cacheDir, "ephemeral", "preserved-"+linkedName)
|
||||
Expect(os.WriteFile(preserved, []byte("keep"), 0640)).To(Succeed())
|
||||
Expect(os.Symlink(preserved, filepath.Join(categoryDir, linkedName))).To(Succeed())
|
||||
|
||||
Expect(releaseEphemeralCacheKey(cacheDir, "ephemeral/request-id/audio/input.wav")).NotTo(Succeed(), linkedName)
|
||||
Expect(preserved).To(BeAnExistingFile(), linkedName)
|
||||
}
|
||||
})
|
||||
|
||||
It("registers an exact release handler", func() {
|
||||
cacheDir := canonicalWorkerTempDir()
|
||||
path := filepath.Join(cacheDir, "ephemeral", "request-id", "audio", "input.wav")
|
||||
Expect(os.MkdirAll(filepath.Dir(path), 0750)).To(Succeed())
|
||||
Expect(os.WriteFile(path, []byte("data"), 0640)).To(Succeed())
|
||||
fm, err := storage.NewFileManager(nil, cacheDir)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
client := &releaseMessagingClient{}
|
||||
|
||||
Expect(subscribeFileRelease(client, "node.one", fm, cacheDir)).To(Succeed())
|
||||
Expect(client.subject).To(Equal(messaging.SubjectNodeFilesRelease("node.one")))
|
||||
request, err := json.Marshal(map[string]string{"key": "ephemeral/request-id/audio/input.wav"})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
var response []byte
|
||||
client.handler(request, func(data []byte) { response = append([]byte(nil), data...) })
|
||||
|
||||
var reply map[string]string
|
||||
Expect(json.Unmarshal(response, &reply)).To(Succeed())
|
||||
Expect(reply["error"]).To(BeEmpty())
|
||||
Expect(path).NotTo(BeAnExistingFile())
|
||||
})
|
||||
|
||||
It("releases a request batch through one worker message", func() {
|
||||
cacheDir := canonicalWorkerTempDir()
|
||||
keys := []string{
|
||||
"ephemeral/audio/request-id/input.wav",
|
||||
"ephemeral/images/request-id/frame.jpg",
|
||||
}
|
||||
for _, key := range keys {
|
||||
path := filepath.Join(cacheDir, filepath.FromSlash(key))
|
||||
Expect(os.MkdirAll(filepath.Dir(path), 0750)).To(Succeed())
|
||||
Expect(os.WriteFile(path, []byte("data"), 0640)).To(Succeed())
|
||||
}
|
||||
fm, err := storage.NewFileManager(nil, cacheDir)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
client := &releaseMessagingClient{}
|
||||
Expect(subscribeFileRelease(client, "node.one", fm, cacheDir)).To(Succeed())
|
||||
request, err := json.Marshal(map[string]any{"request_id": "request-id"})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
var response []byte
|
||||
|
||||
client.handler(request, func(data []byte) { response = append([]byte(nil), data...) })
|
||||
|
||||
var reply map[string]string
|
||||
Expect(json.Unmarshal(response, &reply)).To(Succeed())
|
||||
Expect(reply["error"]).To(BeEmpty())
|
||||
for _, key := range keys {
|
||||
Expect(filepath.Join(cacheDir, filepath.FromSlash(key))).NotTo(BeAnExistingFile())
|
||||
}
|
||||
})
|
||||
|
||||
It("returns validation errors through the release handler", func() {
|
||||
cacheDir := canonicalWorkerTempDir()
|
||||
fm, err := storage.NewFileManager(nil, cacheDir)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
client := &releaseMessagingClient{}
|
||||
Expect(subscribeFileRelease(client, "node-1", fm, cacheDir)).To(Succeed())
|
||||
|
||||
request, err := json.Marshal(map[string]string{"key": "models/model.gguf"})
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
var response []byte
|
||||
client.handler(request, func(data []byte) { response = append([]byte(nil), data...) })
|
||||
|
||||
var reply map[string]string
|
||||
Expect(json.Unmarshal(response, &reply)).To(Succeed())
|
||||
Expect(reply["error"]).NotTo(BeEmpty())
|
||||
})
|
||||
})
|
||||
@@ -597,6 +597,7 @@ func (s *backendSupervisor) reapDeadProcess(key string, bp *backendProcess) {
|
||||
if bp == nil {
|
||||
return
|
||||
}
|
||||
s.cleanupProcessRuntime(bp.proc)
|
||||
if bp.port <= 0 {
|
||||
xlog.Error("Cannot recycle backend port: dead process has invalid recorded port", "backend", key, "addr", bp.addr, "port", bp.port)
|
||||
return
|
||||
@@ -614,6 +615,7 @@ func (s *backendSupervisor) releaseBackendStart(key string, bp *backendProcess)
|
||||
return
|
||||
}
|
||||
delete(s.processes, key)
|
||||
s.cleanupProcessRuntime(bp.proc)
|
||||
if bp.port <= 0 {
|
||||
xlog.Error("Cannot recycle backend port: startup has invalid recorded port", "backend", key, "addr", bp.addr, "port", bp.port)
|
||||
return
|
||||
@@ -947,6 +949,7 @@ func (s *backendSupervisor) finishBackendStop(key string, bp *backendProcess, st
|
||||
return fmt.Errorf("stopping backend process %s: %w", key, stopErr)
|
||||
}
|
||||
delete(s.processes, key)
|
||||
s.cleanupProcessRuntime(bp.proc)
|
||||
if bp.port <= 0 {
|
||||
xlog.Error("Cannot recycle backend port: process has invalid recorded port", "backend", key, "addr", bp.addr, "port", bp.port)
|
||||
return nil
|
||||
@@ -955,6 +958,14 @@ func (s *backendSupervisor) finishBackendStop(key string, bp *backendProcess, st
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *backendSupervisor) cleanupProcessRuntime(proc *process.Process) {
|
||||
// Some focused supervisor tests provide synthetic process handles without a
|
||||
// ModelLoader. Production processes always come from s.ml.StartProcess.
|
||||
if s.ml != nil {
|
||||
s.ml.CleanupProcessRuntime(proc)
|
||||
}
|
||||
}
|
||||
|
||||
// stopAllBackends stops all running backend processes.
|
||||
func (s *backendSupervisor) stopAllBackends(force bool) {
|
||||
s.mu.Lock()
|
||||
|
||||
@@ -149,22 +149,37 @@ func Run(ctx *cliContext.Context, cfg *Config) error {
|
||||
// the top of Run so the worker fails before registering.)
|
||||
httpAddr := cfg.resolveHTTPAddr()
|
||||
stagingDir := filepath.Join(cfg.ModelsPath, "..", "staging")
|
||||
cacheDir := filepath.Join(cfg.ModelsPath, "..", "cache")
|
||||
dataDir := filepath.Join(cfg.ModelsPath, "..", "data")
|
||||
ephemeralRoots := []string{
|
||||
filepath.Join(stagingDir, "ephemeral"),
|
||||
filepath.Join(cacheDir, "ephemeral"),
|
||||
}
|
||||
byteLimit, minFreeBytes, err := effectiveEphemeralCapacity(
|
||||
ephemeralRoots,
|
||||
cfg.EphemeralStagingByteLimit,
|
||||
cfg.EphemeralStagingMinFreeBytes,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("resolving ephemeral staging capacity: %w", err)
|
||||
}
|
||||
ephemeralCapacity, err := NewEphemeralCapacityGuard(ephemeralRoots, byteLimit, minFreeBytes)
|
||||
if err != nil {
|
||||
return fmt.Errorf("initializing ephemeral staging capacity: %w", err)
|
||||
}
|
||||
xlog.Info("Ephemeral staging capacity configured", "roots", ephemeralRoots, "byteLimit", byteLimit, "minFreeBytes", minFreeBytes)
|
||||
StartEphemeralRootsCleanup(shutdownCtx, ephemeralRoots, ephemeralCapacity, 0, 0)
|
||||
// The readiness gate is created here but only armed once NATS is up and the
|
||||
// backend supervisor exists, below, because the gate probes both.
|
||||
// Until then /readyz reports ready, which is correct: reaching this line
|
||||
// means the worker has already registered with the frontend, so it is
|
||||
// mid-startup rather than broken.
|
||||
readiness := &nodes.WorkerReadiness{}
|
||||
httpServer, err := nodes.StartFileTransferServer(httpAddr, stagingDir, cfg.ModelsPath, dataDir, cfg.RegistrationToken, config.DefaultMaxUploadSize, readiness, ml.BackendLogs())
|
||||
httpServer, err := nodes.StartFileTransferServerWithCapacity(httpAddr, stagingDir, cfg.ModelsPath, dataDir, cfg.RegistrationToken, config.DefaultMaxUploadSize, readiness, ephemeralCapacity, ml.BackendLogs())
|
||||
if err != nil {
|
||||
return fmt.Errorf("starting HTTP file transfer server: %w", err)
|
||||
}
|
||||
|
||||
// Per-request input files land in stagingDir over that server and nothing
|
||||
// used to remove them, so a long-lived worker filled its own disk.
|
||||
StartEphemeralStagingCleanup(shutdownCtx, stagingDir, 0, 0)
|
||||
|
||||
// Connect to NATS
|
||||
xlog.Info("Connecting to NATS", "url", sanitize.URL(cfg.NatsURL))
|
||||
natsClient, err := connectNats()
|
||||
@@ -249,7 +264,7 @@ func Run(ctx *cliContext.Context, cfg *Config) error {
|
||||
|
||||
// Subscribe to file staging NATS subjects if S3 is configured
|
||||
if cfg.StorageURL != "" {
|
||||
if err := cfg.subscribeFileStaging(natsClient, nodeID); err != nil {
|
||||
if err := cfg.subscribeFileStaging(natsClient, nodeID, ephemeralCapacity); err != nil {
|
||||
nodes.ShutdownFileTransferServer(httpServer)
|
||||
return fmt.Errorf("subscribing to file staging subjects: %w", err)
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package worker
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
@@ -11,3 +12,11 @@ func TestWorker(t *testing.T) {
|
||||
RegisterFailHandler(Fail)
|
||||
RunSpecs(t, "Worker Suite")
|
||||
}
|
||||
|
||||
// Capacity guards reject symlink components, including macOS /var -> /private/var.
|
||||
func canonicalWorkerTempDir() string {
|
||||
GinkgoHelper()
|
||||
dir, err := filepath.EvalSymlinks(GinkgoT().TempDir())
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
return dir
|
||||
}
|
||||
@@ -674,6 +674,7 @@ Templates use Go templates with [Sprig functions](http://masterminds.github.io/s
|
||||
| `template.multimodal` | string | Template for multimodal interactions |
|
||||
| `template.reply_prefix` | string | Prefix to add to model replies |
|
||||
| `template.use_tokenizer_template` | bool | Use tokenizer's built-in template (vLLM/transformers) |
|
||||
| `template.system_messages_after_first` | string | What to do with `system`-role messages that appear after the leading system block: `merge` folds them into the first system message, `user` forwards them as user-role turns at their position. Unset keeps them as-is. Needed for tokenizer templates that reject late system turns (e.g. Qwen3.8) while agent frameworks append instructions mid-conversation. |
|
||||
| `template.join_chat_messages_by_character` | string | Character to join chat messages (default: `\n`) |
|
||||
|
||||
### Template Variables
|
||||
|
||||
@@ -297,6 +297,8 @@ local-ai worker \
|
||||
| `--advertise-addr` | `LOCALAI_ADVERTISE_ADDR` | *(auto)* | Address the frontend uses to reach this node (see below) |
|
||||
| `--http-addr` | `LOCALAI_HTTP_ADDR` | gRPC port - 1 | HTTP file transfer server bind address |
|
||||
| `--advertise-http-addr` | `LOCALAI_ADVERTISE_HTTP_ADDR` | *(auto)* | HTTP address the frontend uses for file transfer |
|
||||
| `--ephemeral-staging-byte-limit` | `LOCALAI_EPHEMERAL_STAGING_BYTE_LIMIT` | `0` (automatic) | Maximum bytes held by request-input staging across the worker's HTTP staging directory and S3 cache. Automatic mode uses the smaller of 10 GiB and 10% of filesystem capacity. |
|
||||
| `--ephemeral-staging-min-free-bytes` | `LOCALAI_EPHEMERAL_STAGING_MIN_FREE_BYTES` | `0` (automatic) | Free filesystem space preserved while staging request inputs. Automatic mode uses the larger of 1 GiB and 5% of filesystem capacity. |
|
||||
| `--register-to` | `LOCALAI_REGISTER_TO` | *(required)* | Frontend URL for self-registration |
|
||||
| `--node-name` | `LOCALAI_NODE_NAME` | hostname | Human-readable node name |
|
||||
| `--registration-token` | `LOCALAI_REGISTRATION_TOKEN` | *(empty)* | Token to authenticate with the frontend |
|
||||
@@ -320,6 +322,12 @@ local-ai worker \
|
||||
**HTTP file transfer:** Each worker also runs a small HTTP server for file transfer (model files, configs). By default it listens on the gRPC base port - 1 (e.g., if gRPC base is 50051, HTTP is on 50050). gRPC ports grow upward from the base port as additional models are loaded. Set `--advertise-http-addr` if the auto-detected address is not routable from the frontend.
|
||||
{{% /notice %}}
|
||||
|
||||
### Ephemeral request-input storage
|
||||
|
||||
Workers reserve local capacity before accepting per-request audio, image, and other ephemeral inputs. The limit covers both direct HTTP staging and the worker's S3 download cache. A request is rejected before inference when accepting its input would exceed the byte limit or the configured free-space headroom. One request-scoped cleanup operation releases all exact input keys and their reservations after inference, while a one-hour recovery sweep removes abandoned files after crashes. The sweep runs at startup and every 15 minutes, preserves active requests, and considers the newest file in each request directory.
|
||||
|
||||
Set both capacity variables to positive byte counts when a worker needs fixed limits. Leaving either value at zero selects its filesystem-based default. These settings apply only below the two `ephemeral` roots; model, data, and configuration files are excluded.
|
||||
|
||||
### Worker Health Probes
|
||||
|
||||
The worker's HTTP server (base port - 1, default 50050) exposes two unauthenticated probes:
|
||||
@@ -1223,8 +1231,9 @@ Notes:
|
||||
- Check the worker process is running and its NATS connection is up. `Scheduled node is not answering on the bus` in the frontend log names each node demoted this way.
|
||||
|
||||
**A worker fills its own disk over time:**
|
||||
- A request that carries a file (an image, an audio clip, a video) stages that file to the worker under `<models>/../staging/ephemeral/`. The worker deletes these 6 hours after the request that needed them, and sweeps every 30 minutes plus once at startup, so a worker that crashed mid-request still reclaims the space.
|
||||
- Releases before this sweep existed kept every staged input for the lifetime of the worker. Delete `<models>/../staging/ephemeral/` on an affected worker once, as the user the worker runs as; the sweep keeps it bounded from then on.
|
||||
- A request that carries a file (an image, an audio clip, a video) stages that file below the worker's HTTP staging or S3 cache `ephemeral/` directory. The frontend releases each request-owned input when inference finishes, and the worker reserves capacity before accepting it.
|
||||
- A one-hour recovery sweep runs at startup and every 15 minutes to reclaim inputs left by interrupted requests. It preserves active reservations and uses the newest file timestamp in each request directory.
|
||||
- Releases before request-owned cleanup existed can leave a legacy backlog. Delete the affected `ephemeral/` directory once, as the user the worker runs as; capacity admission and recovery cleanup keep new staging bounded.
|
||||
- Staged **model** files are not touched by this. They live beside the ephemeral directory and are not per-request scratch.
|
||||
- A worker whose volume is genuinely full reports `creating backend process state directory under ...: no space left on device` when a backend starts.
|
||||
|
||||
|
||||
@@ -124,6 +124,40 @@ Reference selection follows this order:
|
||||
|
||||
When a saved profile is selected, LocalAI supplies both its private WAV and exact transcript for that request. It does not rewrite the model YAML or copy the recording into the model directory.
|
||||
|
||||
### Realtime pipeline default
|
||||
|
||||
Set `tts.voice` on a realtime pipeline model to use a saved Voice Library profile as the session default:
|
||||
|
||||
```yaml
|
||||
name: gpt-realtime
|
||||
tts:
|
||||
voice: localai://voice-profiles/550e8400-e29b-41d4-a716-446655440000
|
||||
pipeline:
|
||||
vad: silero-vad-ggml
|
||||
transcription: whisper-large-turbo
|
||||
llm: qwen3-4b
|
||||
tts: qwen3-tts-base
|
||||
```
|
||||
|
||||
LocalAI resolves this profile when the realtime session starts. The selected TTS model must support Voice Library cloning.
|
||||
|
||||
You can also change the profile during a realtime session. Set `audio.output.voice` to a Voice Library URI in a `session.update` event:
|
||||
|
||||
```json
|
||||
{
|
||||
"type": "session.update",
|
||||
"session": {
|
||||
"audio": {
|
||||
"output": {
|
||||
"voice": "localai://voice-profiles/550e8400-e29b-41d4-a716-446655440000"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
If the same update changes `model`, the explicit `audio.output.voice` value takes precedence over the new model's `tts.voice` default. The selected model must support Voice Library cloning.
|
||||
|
||||
#### Supported backend and model variants
|
||||
|
||||
| Backend | Automatically compatible variants |
|
||||
@@ -302,31 +336,6 @@ The `/v1/sound-generation` endpoint is compatible with the [ElevenLabs sound gen
|
||||
|
||||
Error responses: `400` for a missing or invalid model or request parameters, and `500` for a backend error during sound generation.
|
||||
|
||||
### AudioLDM 2
|
||||
|
||||
[AudioLDM 2](https://github.com/haoheliu/AudioLDM2) generates sound effects,
|
||||
music, and speech from a text description. Install the gallery model:
|
||||
|
||||
```bash
|
||||
local-ai models install audioldm2
|
||||
```
|
||||
|
||||
Generate a WAV file through the sound-generation endpoint:
|
||||
|
||||
```bash
|
||||
curl http://localhost:8080/v1/sound-generation \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model_id": "audioldm2",
|
||||
"text": "Waves breaking on a rocky beach during a distant thunderstorm",
|
||||
"duration_seconds": 10
|
||||
}' --output storm.wav
|
||||
```
|
||||
|
||||
AudioLDM 2 uses the `AudioLDM2Pipeline` from the diffusers backend. The
|
||||
`duration_seconds` field maps to the pipeline's `audio_length_in_s` option, and
|
||||
`prompt_influence` maps to `guidance_scale`.
|
||||
|
||||
#### Configuration
|
||||
|
||||
You can configure ACE-Step models with various options:
|
||||
|
||||
@@ -19,6 +19,8 @@ This section covers everything you need to know about installing and configuring
|
||||
|
||||
The Model Gallery is the simplest way to install models. It provides pre-configured models ready to use.
|
||||
|
||||
GPU recommendations require a memory estimate within 95% of the detected model memory budget at a 4096-token context. If none of the sampled candidates fit, the recommendation section is hidden. You can still browse the gallery and check individual models at your intended context size. The Home page also omits static GPU suggestions when no fitting recommendation is available.
|
||||
|
||||
### Via WebUI
|
||||
|
||||
1. Open the LocalAI WebUI at `http://localhost:8080`
|
||||
|
||||
@@ -27,9 +27,16 @@ Complete reference for all LocalAI command-line interface (CLI) parameters and e
|
||||
| `--upload-path` | `TMPDIR/localai-UID/upload` | Path to store uploads from files API. Defaults under the OS temp dir (`$TMPDIR`, falling back to `/tmp`), scoped to the current user's UID. | `$LOCALAI_UPLOAD_PATH`, `$UPLOAD_PATH` |
|
||||
| `--localai-config-dir` | `BASEPATH/configuration` | Directory for dynamic loading of certain configuration files (currently runtime_settings.json, api_keys.json, and external_backends.json). See [Runtime Settings]({{%relref "features/runtime-settings" %}}) for web-based configuration. | `$LOCALAI_CONFIG_DIR` |
|
||||
| `--localai-config-dir-poll-interval` | | Time duration to poll the LocalAI Config Dir if your system has broken fsnotify events (example: `1m`) | `$LOCALAI_CONFIG_DIR_POLL_INTERVAL` |
|
||||
|
||||
| `--models-config-file` | | YAML file containing a list of model backend configs (alias: `--config-file`) | `$LOCALAI_MODELS_CONFIG_FILE`, `$CONFIG_FILE` |
|
||||
| `--artifact-download-concurrency` | `1` | How many files of a model artifact to download at once. `1` downloads sequentially. Raising it helps artifacts split into many files on a fast link, at the cost of more concurrent load on the models volume. Whole files only — a single file is never split, so resume and per-file checksum verification are unaffected | `$LOCALAI_ARTIFACT_DOWNLOAD_CONCURRENCY` |
|
||||
|
||||
Backend processes receive a private scratch directory through `TMPDIR`, `TMP`,
|
||||
and `TEMP`. LocalAI removes that directory when the backend exits and removes
|
||||
abandoned directories left by a LocalAI crash before starting another backend.
|
||||
Set `$LOCALAI_BACKEND_TEMP_DIR` to choose their base volume. LocalAI always
|
||||
appends `localai-UID/backend-runtime`; the default base is `TMPDIR`.
|
||||
|
||||
## Backend Flags
|
||||
|
||||
| Parameter | Default | Description | Environment Variable |
|
||||
|
||||
@@ -0,0 +1,158 @@
|
||||
# Request-owned ephemeral staging
|
||||
|
||||
## Problem
|
||||
|
||||
Distributed requests copy transient inputs below
|
||||
`<staging>/ephemeral/<category>/<request-id>`. The worker currently removes
|
||||
these files only when a periodic age sweep considers them stale. A Reachy Mini
|
||||
sending camera and sound data about once per second created more than 21,000
|
||||
request directories and filled its Mac worker before the six-hour retention
|
||||
window elapsed.
|
||||
|
||||
Reducing the retention window is insufficient. A time limit bounds residence
|
||||
time, but the retained bytes still scale with request rate and input size. A
|
||||
quota sweeper would also have to infer whether an old file is still in use.
|
||||
Neither rule prevents concurrent uploads from consuming the worker's last free
|
||||
space.
|
||||
|
||||
## Goals
|
||||
|
||||
- Give every ephemeral input an explicit owner and release it when that request
|
||||
finishes, fails, or is cancelled.
|
||||
- Keep cleanup transport-independent for HTTP and S3/NATS workers.
|
||||
- Reserve capacity before accepting ephemeral bytes so concurrent requests
|
||||
cannot consume configured disk headroom.
|
||||
- Reject a request cleanly when its input does not fit; never evict an input
|
||||
that a running request may still be reading.
|
||||
- Recover abandoned files after frontend or worker crashes.
|
||||
- Never inspect or remove models, data, configuration, or paths outside the
|
||||
worker's ephemeral staging tree.
|
||||
|
||||
## Non-goals
|
||||
|
||||
- Retaining request inputs as a cache.
|
||||
- Evicting persistent model or data files to make an inference request fit.
|
||||
- Treating modification timestamps as proof that a request is active.
|
||||
|
||||
## Request ownership
|
||||
|
||||
The `FileStagingClient` already creates one request ID before staging inputs and
|
||||
waits for synchronous and streaming backend calls to finish. It will track each
|
||||
ephemeral key before attempting to stage it and defer one request-scoped
|
||||
release identified by the request ID. Release runs after the backend call
|
||||
returns, including error and cancellation paths, using a short background
|
||||
timeout so cancellation of the request does not cancel its cleanup.
|
||||
|
||||
Request IDs will use the full UUID rather than the current eight-character
|
||||
prefix. The worker enumerates only category directories for that validated
|
||||
request ID and removes each entry with exact, symlink-safe deletion.
|
||||
|
||||
`FileStager` will expose an idempotent exact-key `ReleaseRemote` operation and
|
||||
an optional request-scoped operation. The client uses one fixed-size request
|
||||
message for the normal path and retains exact-key calls as a rolling-upgrade
|
||||
fallback:
|
||||
|
||||
- HTTP sends one authenticated request containing the fixed-size request ID.
|
||||
The worker derives and removes that request's exact files, then prunes empty
|
||||
request and category directories without following symlinks.
|
||||
- S3/NATS sends one request-reply containing the request ID so the selected
|
||||
worker evicts the request's local cached files. The frontend then deletes the
|
||||
matching objects from its tracked exact-key list.
|
||||
Either deletion may already have happened and still counts as success.
|
||||
|
||||
If staging fails partway through a request, the deferred release still includes
|
||||
the planned key, allowing it to remove a partial file when the transport can
|
||||
identify one. Cleanup errors are logged and do not replace the inference result.
|
||||
|
||||
HTTP and S3 ingress register request operations before any pre-reservation
|
||||
work. Before enumerating files, the capacity guard marks the request released
|
||||
and waits for registered operations and admitted writes to finish. Later
|
||||
operations, reservations, and cache claims for that request are rejected.
|
||||
Markers expire after one hour and are capped at 16,384 entries, but a marker is
|
||||
never evicted while its registered operation or cleanup scan is active.
|
||||
Concurrent operation and cleanup state have the same hard cap. Disk bytes
|
||||
remain independently bounded by capacity admission.
|
||||
Cleanup waits within its deadline when all cleanup-pin slots are occupied.
|
||||
If that deadline expires, existing entries lose active ownership so recovery
|
||||
can reclaim them; registered ingress for the request remains closed until it
|
||||
exits.
|
||||
|
||||
## Capacity admission
|
||||
|
||||
A worker-local ephemeral capacity guard is shared by its HTTP and S3/NATS input
|
||||
paths. It accounts for both `<staging>/ephemeral`, used by HTTP, and
|
||||
`<cache>/ephemeral`, used by S3 downloads. Before writing an ephemeral object,
|
||||
the transport reserves its declared size. HTTP obtains the size from the upload
|
||||
metadata; S3/NATS obtains it from object metadata. Reservations are serialized
|
||||
in memory, cover both committed ephemeral bytes and concurrent writes, and are
|
||||
returned on release or failed transfer.
|
||||
|
||||
Admission succeeds only when both conditions remain true after the reservation:
|
||||
|
||||
1. Total ephemeral bytes remain below the configured ephemeral staging limit.
|
||||
2. The filesystem retains the configured minimum free-space headroom.
|
||||
|
||||
The guard rejects the transfer before inference when either condition fails.
|
||||
An input with unknown size is written through a bounded accounting writer that
|
||||
reserves fixed-size chunks before writing each chunk and stops before crossing
|
||||
the limit. The existing maximum-upload-size check remains the per-file ceiling.
|
||||
|
||||
The limit and headroom are worker settings. By default, ephemeral data may use
|
||||
the smaller of 10 GiB or 10 percent of filesystem capacity, while the worker
|
||||
preserves the larger of 1 GiB or 5 percent as free-space headroom. The worker
|
||||
logs the effective values at startup. A zero or negative operator value selects
|
||||
the default rather than disabling protection. The guard scans the ephemeral
|
||||
tree at startup to account for abandoned committed bytes. Filesystem free-space
|
||||
checks are repeated at reservation time because other processes may share the
|
||||
volume.
|
||||
|
||||
## Crash recovery
|
||||
|
||||
The existing periodic cleanup remains as a fallback for ownership messages lost
|
||||
when a frontend or worker process dies. It uses a one-hour recovery TTL,
|
||||
performs one startup sweep, and repeats every 15 minutes. It skips every key
|
||||
held by an active reservation, considers the newest modification time in each
|
||||
remaining request tree, and does not follow directory symlinks. It removes only
|
||||
request directories below the registered `<staging>/ephemeral` and
|
||||
`<cache>/ephemeral` roots.
|
||||
|
||||
The recovery window does not control normal storage growth. Request completion
|
||||
and capacity reservations do. A recovery deletion updates the capacity guard's
|
||||
accounted bytes.
|
||||
|
||||
## Error handling and observability
|
||||
|
||||
Admission failures report the requested bytes, current ephemeral usage, limit,
|
||||
available bytes, and required headroom. Successful release and recovery update
|
||||
usage counters. Read, stat, and remove failures include the affected path and
|
||||
allow unrelated cleanup to continue. Missing ephemeral files and directories
|
||||
are normal for idempotent release.
|
||||
|
||||
## Testing
|
||||
|
||||
Regression tests will establish the following behavior:
|
||||
|
||||
1. Successful, failed, cancelled, and streaming calls issue one request-scoped
|
||||
worker cleanup only after the backend has returned.
|
||||
2. Partial staging failures release the planned key without changing the main
|
||||
error returned to the caller.
|
||||
3. HTTP and S3/NATS release remove local files; S3/NATS also removes the object.
|
||||
4. Release rejects persistent keys and path traversal, does not follow
|
||||
symlinks, and leaves paths outside `ephemeral` untouched.
|
||||
5. Concurrent reservations cannot exceed the byte limit or free-space
|
||||
headroom, and failed transfers return their reservations.
|
||||
6. Unknown-length writes stop at the capacity boundary.
|
||||
7. Startup accounting includes abandoned ephemeral files, and the recovery
|
||||
sweep removes only stale, inactive leftovers and updates accounting.
|
||||
|
||||
Focused package tests will run with race detection, followed by the relevant
|
||||
repository lint and vet checks.
|
||||
|
||||
## Rollout
|
||||
|
||||
The change requires a new LocalAI worker and frontend build because both sides
|
||||
participate in release. The Mac worker starts by accounting for its existing
|
||||
backlog and removing recovery-expired files. The deployment check will verify
|
||||
available space, admission and release logs, stable ephemeral usage under
|
||||
continuous camera and audio traffic, and successful vision, sound detection,
|
||||
and transcription requests.
|
||||
@@ -1,15 +0,0 @@
|
||||
---
|
||||
name: "audioldm2"
|
||||
|
||||
config_file: |
|
||||
backend: diffusers
|
||||
known_usecases:
|
||||
- sound_generation
|
||||
parameters:
|
||||
model: cvssp/audioldm2
|
||||
diffusers:
|
||||
pipeline_type: AudioLDM2Pipeline
|
||||
cuda: true
|
||||
options:
|
||||
- num_inference_steps:200
|
||||
- torch_dtype:fp16
|
||||
+4
-41
@@ -113,24 +113,7 @@
|
||||
url: "github:mudler/LocalAI/gallery/virtual.yaml@master"
|
||||
urls:
|
||||
- https://huggingface.co/unsloth/GLM-5.3-Flash-GGUF
|
||||
description: |
|
||||
# GLM-5.3-Flash
|
||||
|
||||
👋 Join our WeChat or Discord community.
|
||||
|
||||
📖 Check out the GLM-5.3-Flash blog and GLM-5 Technical report.
|
||||
|
||||
📍 Use GLM-5.3-Flash API services on Z.ai API Platform.
|
||||
|
||||
## Introduction
|
||||
|
||||
We introduce GLM-5.3-Flash, the first natively multimodal model in the GLM-5 series. With 320B total parameters and just 18B active parameters, it outperforms GLM-5.2 across benchmarks and real-world workloads at one-tenth the price, while approaching Claude Opus 4.8 on coding and agentic benchmarks.
|
||||
|
||||
GLM-5.3-Flash starts from a newly trained base model, with its architecture and training recipe redesigned around capability and efficiency. For the first time in the GLM series, we introduce a hybrid architecture combining sparse and linear attention, sharply reducing long-context serving costs while preserving precise long-context capabilities. The model also adopts Manifold-Constrained Hyper-Connections (mHC) to further improve scaling efficiency. Together with our latest 30T-token multimodal pre-training corpus, these changes enable GLM-5.3-Flash to deliver more intelligence with less compute.
|
||||
|
||||
## Serve GLM-5.3-Flash Locally
|
||||
|
||||
...
|
||||
description: "# GLM-5.3-Flash\n\n\U0001F44B Join our WeChat or Discord community.\n\n\U0001F4D6 Check out the GLM-5.3-Flash blog and GLM-5 Technical report.\n\n\U0001F4CD Use GLM-5.3-Flash API services on Z.ai API Platform.\n\n## Introduction\n\nWe introduce GLM-5.3-Flash, the first natively multimodal model in the GLM-5 series. With 320B total parameters and just 18B active parameters, it outperforms GLM-5.2 across benchmarks and real-world workloads at one-tenth the price, while approaching Claude Opus 4.8 on coding and agentic benchmarks.\n\nGLM-5.3-Flash starts from a newly trained base model, with its architecture and training recipe redesigned around capability and efficiency. For the first time in the GLM series, we introduce a hybrid architecture combining sparse and linear attention, sharply reducing long-context serving costs while preserving precise long-context capabilities. The model also adopts Manifold-Constrained Hyper-Connections (mHC) to further improve scaling efficiency. Together with our latest 30T-token multimodal pre-training corpus, these changes enable GLM-5.3-Flash to deliver more intelligence with less compute.\n\n## Serve GLM-5.3-Flash Locally\n\n...\n"
|
||||
license: "mit"
|
||||
tags:
|
||||
- llm
|
||||
@@ -768,9 +751,7 @@
|
||||
model: llama-cpp/models/s1-mini/s1-mini-q4_k_m.gguf
|
||||
temperature: 0
|
||||
system_prompt: >-
|
||||
You are a text normalizer for speech-to-text transcripts. The input begins
|
||||
with a control line specifying the styling, structure, and context settings;
|
||||
clean the transcript to match those settings and output only the cleaned text.
|
||||
You are a text normalizer for speech-to-text transcripts. The input begins with a control line specifying the styling, structure, and context settings; clean the transcript to match those settings and output only the cleaned text.
|
||||
template:
|
||||
use_tokenizer_template: true
|
||||
files:
|
||||
@@ -796,9 +777,7 @@
|
||||
model: llama-cpp/models/s1-mini/s1-mini-f16.gguf
|
||||
temperature: 0
|
||||
system_prompt: >-
|
||||
You are a text normalizer for speech-to-text transcripts. The input begins
|
||||
with a control line specifying the styling, structure, and context settings;
|
||||
clean the transcript to match those settings and output only the cleaned text.
|
||||
You are a text normalizer for speech-to-text transcripts. The input begins with a control line specifying the styling, structure, and context settings; clean the transcript to match those settings and output only the cleaned text.
|
||||
template:
|
||||
use_tokenizer_template: true
|
||||
files:
|
||||
@@ -1970,22 +1949,6 @@
|
||||
- filename: carbon-8b-q8_0.gguf
|
||||
uri: huggingface://HuggingFaceBio/Carbon-8B-GGUF/carbon-8b-q8_0.gguf
|
||||
sha256: ba5f7794d0768e639fcb0fe860ff7340b09bfaad5e0e2f2c021a4f49029352dd
|
||||
- name: audioldm2
|
||||
url: github:mudler/LocalAI/gallery/audioldm2.yaml@master
|
||||
urls:
|
||||
- https://huggingface.co/cvssp/audioldm2
|
||||
- https://github.com/haoheliu/AudioLDM2
|
||||
description: |
|
||||
AudioLDM 2 generates sound effects, music, and speech from natural-language
|
||||
descriptions through the diffusers backend and LocalAI sound-generation API.
|
||||
license: cc-by-nc-sa-4.0
|
||||
tags:
|
||||
- audio
|
||||
- sound-generation
|
||||
- text-to-audio
|
||||
- diffusers
|
||||
- gpu
|
||||
last_checked: "2026-08-13"
|
||||
- &ornith-1-0-9b
|
||||
name: "ornith-1.0-9b-q4"
|
||||
variants:
|
||||
@@ -14147,8 +14110,8 @@
|
||||
model: kokoro-int8-multi-lang-v1_0/model.int8.onnx
|
||||
files:
|
||||
- filename: kokoro-int8-multi-lang-v1_0.tar.bz2
|
||||
sha256: 75654a84864be26f345f020f4070c2c019e96dd1b7f9bf6e2ffd59efac6aa5a3
|
||||
uri: https://github.com/k2-fsa/sherpa-onnx/releases/download/tts-models/kokoro-int8-multi-lang-v1_0.tar.bz2
|
||||
sha256: 4c3052abaa60943a341f193888cf6abd68787dae6ab8ae5c925a706caa247e4e
|
||||
- name: supertonic-3
|
||||
url: github:mudler/LocalAI/gallery/supertonic.yaml@master
|
||||
urls:
|
||||
|
||||
@@ -173,7 +173,7 @@ func (ml *ModelLoader) spawnGRPCModel(backend, uri string, o *Options, modelID,
|
||||
if !ready {
|
||||
xlog.Debug("GRPC Service NOT ready")
|
||||
startupErr := grpcStartupError(client.Process())
|
||||
stopLoadProcess(client, modelID)
|
||||
ml.stopLoadProcess(client, modelID)
|
||||
return nil, startupErr
|
||||
}
|
||||
|
||||
@@ -189,11 +189,11 @@ func (ml *ModelLoader) spawnGRPCModel(backend, uri string, o *Options, modelID,
|
||||
|
||||
res, err := client.GRPC(o.parallelRequests, ml.wd).LoadModel(o.context, options)
|
||||
if err != nil {
|
||||
stopLoadProcess(client, modelID)
|
||||
ml.stopLoadProcess(client, modelID)
|
||||
return nil, fmt.Errorf("could not load model: %w", err)
|
||||
}
|
||||
if !res.Success {
|
||||
stopLoadProcess(client, modelID)
|
||||
ml.stopLoadProcess(client, modelID)
|
||||
return nil, fmt.Errorf("could not load model (no success): %s", res.Message)
|
||||
}
|
||||
|
||||
@@ -260,7 +260,7 @@ func lastNonEmptyLine(path string, maxBytes int64) string {
|
||||
|
||||
// stopLoadProcess tears down a backend process whose load did not complete.
|
||||
// The stop error is only logged: the load error is what the caller reports.
|
||||
func stopLoadProcess(client *Model, modelID string) {
|
||||
func (ml *ModelLoader) stopLoadProcess(client *Model, modelID string) {
|
||||
process := client.Process()
|
||||
if process == nil {
|
||||
return
|
||||
@@ -268,6 +268,7 @@ func stopLoadProcess(client *Model, modelID string) {
|
||||
if err := process.Stop(); err != nil {
|
||||
xlog.Warn("failed to stop backend process after failed load", "error", err, "modelID", modelID)
|
||||
}
|
||||
ml.cleanupProcessRuntime(process)
|
||||
}
|
||||
|
||||
// parallelSlotsFromOptions returns the effective n_parallel from the backend
|
||||
|
||||
@@ -106,6 +106,10 @@ type ModelLoader struct {
|
||||
// the exit code can't, since a child killed by our own SIGTERM/SIGKILL
|
||||
// reports -1, indistinguishable from a signal-induced crash.
|
||||
stoppingProcs sync.Map
|
||||
// processRuntimes keeps the owned state/scratch directory alive until the
|
||||
// loader has consumed any exit diagnostics. The exit watcher removes the
|
||||
// potentially large scratch contents immediately.
|
||||
processRuntimes sync.Map
|
||||
// loadFailures records, per modelID, the cooldown window applied after a
|
||||
// failed load so that a client repeatedly polling a broken model does not
|
||||
// spawn (and leak) a fresh backend process on every request. Guarded by mu.
|
||||
|
||||
+39
-18
@@ -177,12 +177,14 @@ func (ml *ModelLoader) deleteProcess(ctx context.Context, s string, force bool)
|
||||
// A concurrently crashed/already-reaped process can no longer own
|
||||
// resources even if Stop could not read or signal its PID.
|
||||
store.Delete(s)
|
||||
ml.cleanupProcessRuntime(process)
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
store.Delete(s)
|
||||
ml.cleanupProcessRuntime(process)
|
||||
return nil
|
||||
}
|
||||
func (ml *ModelLoader) StopGRPC(filter GRPCProcessFilter) error {
|
||||
@@ -231,16 +233,6 @@ func (ml *ModelLoader) StartProcess(grpcProcess, id string, serverAddress string
|
||||
return ml.startProcess(grpcProcess, id, serverAddress, args...)
|
||||
}
|
||||
|
||||
// newProcessStateDir creates the directory a backend process uses for its pid,
|
||||
// state and log files, and reports why when it cannot.
|
||||
func newProcessStateDir() (string, error) {
|
||||
dir, err := os.MkdirTemp(os.TempDir(), "go-processmanager")
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("creating backend process state directory under %s: %w", os.TempDir(), err)
|
||||
}
|
||||
return dir, nil
|
||||
}
|
||||
|
||||
func (ml *ModelLoader) startProcess(grpcProcess, id string, serverAddress string, args ...string) (*process.Process, error) {
|
||||
// Make sure the process is executable
|
||||
// Check first if it has executable permissions
|
||||
@@ -262,7 +254,12 @@ func (ml *ModelLoader) startProcess(grpcProcess, id string, serverAddress string
|
||||
return nil, err
|
||||
}
|
||||
|
||||
env := os.Environ()
|
||||
runtime, err := newBackendProcessRuntime()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
env := backendTempEnvironment(os.Environ(), runtime.tempDir)
|
||||
// Vulkan backends are self-contained: they bundle their own loader and
|
||||
// Mesa driver .so files in lib/ plus the matching ICD manifests in
|
||||
// vulkan/icd.d/. Point the loader at those manifests so it doesn't rely on
|
||||
@@ -271,16 +268,14 @@ func (ml *ModelLoader) startProcess(grpcProcess, id string, serverAddress string
|
||||
// and the GPU would silently fall back to CPU). No-op for other backends.
|
||||
env = append(env, vulkanICDEnv(workDir)...)
|
||||
|
||||
// Resolve the state directory here rather than through
|
||||
// Resolve and own the state directory here rather than through
|
||||
// process.WithTemporaryStateDir(). process.New applies its options but
|
||||
// discards the error they return, so a temp directory that cannot be
|
||||
// created leaves StateDir empty and every later option unapplied. Run()
|
||||
// then reported "mkdir : no such file or directory" with no path, hiding
|
||||
// the real cause (a full volume, or a TMPDIR that no longer resolves).
|
||||
stateDir, err := newProcessStateDir()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// then reports "mkdir : no such file or directory" with no useful path.
|
||||
// The same owned directory also contains backend scratch so an unexpected
|
||||
// exit cannot strand request files directly in the host's shared /tmp.
|
||||
stateDir := runtime.dir
|
||||
|
||||
grpcControlProcess := process.New(
|
||||
process.WithStateDir(stateDir),
|
||||
@@ -296,8 +291,10 @@ func (ml *ModelLoader) startProcess(grpcProcess, id string, serverAddress string
|
||||
}
|
||||
|
||||
if err := grpcControlProcess.Run(); err != nil {
|
||||
runtime.cleanup()
|
||||
return grpcControlProcess, err
|
||||
}
|
||||
ml.processRuntimes.Store(grpcControlProcess, runtime)
|
||||
|
||||
xlog.Debug("GRPC Service state dir", "dir", grpcControlProcess.StateDir())
|
||||
|
||||
@@ -376,11 +373,35 @@ func (ml *ModelLoader) startProcess(grpcProcess, id string, serverAddress string
|
||||
}
|
||||
xlog.Warn("Backend process exited unexpectedly", fields...)
|
||||
}
|
||||
runtime.cleanupScratch()
|
||||
close(runtime.diagnosticsDone)
|
||||
}()
|
||||
|
||||
return grpcControlProcess, nil
|
||||
}
|
||||
|
||||
func (ml *ModelLoader) cleanupProcessRuntime(process *process.Process) {
|
||||
if process == nil {
|
||||
return
|
||||
}
|
||||
value, ok := ml.processRuntimes.LoadAndDelete(process)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
runtime := value.(*backendProcessRuntime)
|
||||
go func() {
|
||||
<-runtime.diagnosticsDone
|
||||
runtime.cleanup()
|
||||
}()
|
||||
}
|
||||
|
||||
// CleanupProcessRuntime releases state and scratch owned by a process started
|
||||
// through StartProcess. Callers that supervise processes outside ModelLoader's
|
||||
// model store must invoke it after they have consumed exit diagnostics.
|
||||
func (ml *ModelLoader) CleanupProcessRuntime(process *process.Process) {
|
||||
ml.cleanupProcessRuntime(process)
|
||||
}
|
||||
|
||||
// vulkanICDEnv returns environment overrides that point the Vulkan loader at
|
||||
// the ICD manifests a backend bundles in <workDir>/vulkan/icd.d. Vulkan
|
||||
// backends ship a self-contained stack — their own loader and Mesa driver .so
|
||||
|
||||
@@ -15,8 +15,10 @@ import (
|
||||
var _ = Describe("backend process exit diagnostics", func() {
|
||||
It("includes the exit code and final stderr line for an unexpected exit", func() {
|
||||
tmpDir := GinkgoT().TempDir()
|
||||
backendTempRoot := filepath.Join(tmpDir, "backend-runtime")
|
||||
GinkgoT().Setenv(backendTempDirEnv, backendTempRoot)
|
||||
backendPath := filepath.Join(tmpDir, "failing-backend")
|
||||
Expect(os.WriteFile(backendPath, []byte("#!/bin/sh\necho 'first diagnostic' >&2\necho 'fatal metal pipeline error' >&2\nexit 42\n"), 0o700)).To(Succeed())
|
||||
Expect(os.WriteFile(backendPath, []byte("#!/bin/sh\nprintf '%s' \"$TMPDIR\" > \"$0.tmpdir\"\necho 'first diagnostic' >&2\necho 'fatal metal pipeline error' >&2\nexit 42\n"), 0o700)).To(Succeed())
|
||||
|
||||
captured := &bytes.Buffer{}
|
||||
handler := slog.NewTextHandler(captured, &slog.HandlerOptions{Level: slog.LevelWarn})
|
||||
@@ -29,10 +31,16 @@ var _ = Describe("backend process exit diagnostics", func() {
|
||||
process, err := loader.startProcess(backendPath, "test-model", "127.0.0.1:65535")
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Eventually(process.Done()).Should(BeClosed())
|
||||
backendTemp, err := os.ReadFile(backendPath + ".tmpdir")
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(string(backendTemp)).To(Equal(filepath.Join(process.StateDir(), "tmp")))
|
||||
Eventually(string(backendTemp)).ShouldNot(BeADirectory())
|
||||
Eventually(captured.String).Should(And(
|
||||
ContainSubstring("Backend process exited unexpectedly"),
|
||||
ContainSubstring("exitCode=42"),
|
||||
ContainSubstring(`stderr="fatal metal pipeline error"`),
|
||||
))
|
||||
loader.cleanupProcessRuntime(process)
|
||||
Eventually(process.StateDir()).ShouldNot(BeADirectory())
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,168 @@
|
||||
package model
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/gofrs/flock"
|
||||
"github.com/mudler/xlog"
|
||||
)
|
||||
|
||||
const (
|
||||
backendTempDirEnv = "LOCALAI_BACKEND_TEMP_DIR"
|
||||
backendRuntimeDirPrefix = "process-"
|
||||
backendRuntimeMarker = ".localai-backend-runtime"
|
||||
backendRuntimeMagic = "localai-backend-runtime-v1\n"
|
||||
)
|
||||
|
||||
// backendProcessRuntime owns both go-processmanager's state and all temporary
|
||||
// files created by one backend process. The held lock distinguishes a live
|
||||
// runtime from one abandoned when LocalAI was killed or crashed.
|
||||
type backendProcessRuntime struct {
|
||||
dir string
|
||||
tempDir string
|
||||
lock *flock.Flock
|
||||
scratch sync.Once
|
||||
once sync.Once
|
||||
// diagnosticsDone closes after the exit watcher has read the state files.
|
||||
diagnosticsDone chan struct{}
|
||||
}
|
||||
|
||||
func backendRuntimeRoot() string {
|
||||
base := os.TempDir()
|
||||
if configured := os.Getenv(backendTempDirEnv); configured != "" {
|
||||
base = configured
|
||||
}
|
||||
// Always append a LocalAI- and user-specific namespace. Even if an operator
|
||||
// points the configurable base at /tmp, the sweeper never inspects unrelated
|
||||
// process-* directories in that shared parent.
|
||||
return filepath.Join(base, fmt.Sprintf("localai-%d", os.Getuid()), "backend-runtime")
|
||||
}
|
||||
|
||||
func newBackendProcessRuntime() (*backendProcessRuntime, error) {
|
||||
root := backendRuntimeRoot()
|
||||
if err := os.MkdirAll(root, 0o700); err != nil {
|
||||
return nil, fmt.Errorf("creating backend runtime root %s: %w", root, err)
|
||||
}
|
||||
|
||||
// Serialize sweeping with creation. Otherwise a second LocalAI instance
|
||||
// could observe the new directory in the tiny window before its owner lock
|
||||
// is acquired and mistake it for an abandoned runtime.
|
||||
sweepLock := flock.New(filepath.Join(root, ".sweep.lock"))
|
||||
if err := sweepLock.Lock(); err != nil {
|
||||
return nil, fmt.Errorf("locking backend runtime root %s: %w", root, err)
|
||||
}
|
||||
defer func() {
|
||||
if err := sweepLock.Unlock(); err != nil {
|
||||
xlog.Warn("Failed to unlock backend runtime root", "root", root, "error", err)
|
||||
}
|
||||
}()
|
||||
|
||||
sweepAbandonedBackendRuntimes(root)
|
||||
|
||||
dir, err := os.MkdirTemp(root, backendRuntimeDirPrefix)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("creating backend process runtime under %s: %w", root, err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(dir, backendRuntimeMarker), []byte(backendRuntimeMagic), 0o600); err != nil {
|
||||
_ = os.RemoveAll(dir)
|
||||
return nil, fmt.Errorf("marking backend process runtime %s: %w", dir, err)
|
||||
}
|
||||
runtimeLock := flock.New(filepath.Join(dir, ".owner.lock"))
|
||||
if err := runtimeLock.Lock(); err != nil {
|
||||
_ = os.RemoveAll(dir)
|
||||
return nil, fmt.Errorf("locking backend process runtime %s: %w", dir, err)
|
||||
}
|
||||
tempDir := filepath.Join(dir, "tmp")
|
||||
if err := os.Mkdir(tempDir, 0o700); err != nil {
|
||||
_ = runtimeLock.Unlock()
|
||||
_ = os.RemoveAll(dir)
|
||||
return nil, fmt.Errorf("creating backend scratch directory %s: %w", tempDir, err)
|
||||
}
|
||||
|
||||
return &backendProcessRuntime{
|
||||
dir: dir,
|
||||
tempDir: tempDir,
|
||||
lock: runtimeLock,
|
||||
diagnosticsDone: make(chan struct{}),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func sweepAbandonedBackendRuntimes(root string) {
|
||||
entries, err := os.ReadDir(root)
|
||||
if err != nil {
|
||||
xlog.Warn("Failed to inspect backend runtime root", "root", root, "error", err)
|
||||
return
|
||||
}
|
||||
for _, entry := range entries {
|
||||
if !entry.IsDir() || !strings.HasPrefix(entry.Name(), backendRuntimeDirPrefix) {
|
||||
continue
|
||||
}
|
||||
dir := filepath.Join(root, entry.Name())
|
||||
marker, err := os.ReadFile(filepath.Join(dir, backendRuntimeMarker))
|
||||
if err != nil || string(marker) != backendRuntimeMagic {
|
||||
continue
|
||||
}
|
||||
ownerLock := flock.New(filepath.Join(dir, ".owner.lock"))
|
||||
available, err := ownerLock.TryLock()
|
||||
if err != nil {
|
||||
xlog.Warn("Failed to inspect backend runtime ownership", "dir", dir, "error", err)
|
||||
continue
|
||||
}
|
||||
if !available {
|
||||
continue
|
||||
}
|
||||
if err := ownerLock.Unlock(); err != nil {
|
||||
xlog.Warn("Failed to release abandoned backend runtime lock", "dir", dir, "error", err)
|
||||
continue
|
||||
}
|
||||
if err := os.RemoveAll(dir); err != nil {
|
||||
xlog.Warn("Failed to remove abandoned backend runtime", "dir", dir, "error", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (r *backendProcessRuntime) cleanup() {
|
||||
if r == nil {
|
||||
return
|
||||
}
|
||||
r.once.Do(func() {
|
||||
r.cleanupScratch()
|
||||
if err := r.lock.Unlock(); err != nil {
|
||||
xlog.Warn("Failed to unlock backend process runtime", "dir", r.dir, "error", err)
|
||||
}
|
||||
if err := os.RemoveAll(r.dir); err != nil {
|
||||
xlog.Warn("Failed to remove backend process runtime", "dir", r.dir, "error", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func (r *backendProcessRuntime) cleanupScratch() {
|
||||
if r == nil {
|
||||
return
|
||||
}
|
||||
r.scratch.Do(func() {
|
||||
if err := os.RemoveAll(r.tempDir); err != nil {
|
||||
xlog.Warn("Failed to remove backend scratch directory", "dir", r.tempDir, "error", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func backendTempEnvironment(env []string, tempDir string) []string {
|
||||
result := make([]string, 0, len(env)+3)
|
||||
for _, entry := range env {
|
||||
key, _, found := strings.Cut(entry, "=")
|
||||
if found && (key == "TMPDIR" || key == "TMP" || key == "TEMP") {
|
||||
continue
|
||||
}
|
||||
result = append(result, entry)
|
||||
}
|
||||
return append(result,
|
||||
"TMPDIR="+tempDir,
|
||||
"TMP="+tempDir,
|
||||
"TEMP="+tempDir,
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
package model
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("Backend process runtime directory", func() {
|
||||
It("keeps active runtimes while sweeping abandoned ones", func() {
|
||||
root := GinkgoT().TempDir()
|
||||
GinkgoT().Setenv(backendTempDirEnv, root)
|
||||
ownedRoot := backendRuntimeRoot()
|
||||
unrelated := filepath.Join(root, backendRuntimeDirPrefix+"unrelated")
|
||||
Expect(os.MkdirAll(unrelated, 0o700)).To(Succeed())
|
||||
Expect(os.WriteFile(filepath.Join(unrelated, "keep"), []byte("unrelated"), 0o600)).To(Succeed())
|
||||
|
||||
active, err := newBackendProcessRuntime()
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
DeferCleanup(active.cleanup)
|
||||
Expect(os.WriteFile(filepath.Join(active.tempDir, "active.img"), []byte("active"), 0o600)).To(Succeed())
|
||||
|
||||
abandoned := filepath.Join(ownedRoot, backendRuntimeDirPrefix+"abandoned")
|
||||
Expect(os.MkdirAll(abandoned, 0o700)).To(Succeed())
|
||||
Expect(os.WriteFile(filepath.Join(abandoned, backendRuntimeMarker), []byte(backendRuntimeMagic), 0o600)).To(Succeed())
|
||||
Expect(os.WriteFile(filepath.Join(abandoned, "orphan.img"), []byte("orphan"), 0o600)).To(Succeed())
|
||||
foreign := filepath.Join(ownedRoot, backendRuntimeDirPrefix+"foreign")
|
||||
Expect(os.MkdirAll(foreign, 0o700)).To(Succeed())
|
||||
Expect(os.WriteFile(filepath.Join(foreign, "keep"), []byte("foreign"), 0o600)).To(Succeed())
|
||||
|
||||
other, err := newBackendProcessRuntime()
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
DeferCleanup(other.cleanup)
|
||||
|
||||
Expect(active.dir).To(BeADirectory())
|
||||
Expect(abandoned).ToNot(BeAnExistingFile())
|
||||
Expect(foreign).To(BeADirectory())
|
||||
Expect(unrelated).To(BeADirectory())
|
||||
})
|
||||
|
||||
It("uses one private directory for process state and backend scratch", func() {
|
||||
root := GinkgoT().TempDir()
|
||||
GinkgoT().Setenv(backendTempDirEnv, root)
|
||||
|
||||
runtime, err := newBackendProcessRuntime()
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
DeferCleanup(runtime.cleanup)
|
||||
|
||||
Expect(filepath.Dir(runtime.dir)).To(Equal(backendRuntimeRoot()))
|
||||
Expect(runtime.tempDir).To(Equal(filepath.Join(runtime.dir, "tmp")))
|
||||
info, err := os.Stat(runtime.tempDir)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(info.IsDir()).To(BeTrue())
|
||||
Expect(info.Mode().Perm()).To(Equal(os.FileMode(0o700)))
|
||||
})
|
||||
|
||||
It("overrides inherited temp variables for the backend only", func() {
|
||||
env := backendTempEnvironment([]string{
|
||||
"PATH=/bin",
|
||||
"TMPDIR=/old/tmpdir",
|
||||
"TMP=/old/tmp",
|
||||
"TEMP=/old/temp",
|
||||
}, "/owned/scratch")
|
||||
|
||||
Expect(env).To(ConsistOf(
|
||||
"PATH=/bin",
|
||||
"TMPDIR=/owned/scratch",
|
||||
"TMP=/owned/scratch",
|
||||
"TEMP=/owned/scratch",
|
||||
))
|
||||
for _, key := range []string{"TMPDIR", "TMP", "TEMP"} {
|
||||
count := 0
|
||||
for _, entry := range env {
|
||||
if strings.HasPrefix(entry, key+"=") {
|
||||
count++
|
||||
}
|
||||
}
|
||||
Expect(count).To(Equal(1), key)
|
||||
}
|
||||
})
|
||||
|
||||
It("removes the runtime when its owner exits", func() {
|
||||
GinkgoT().Setenv(backendTempDirEnv, GinkgoT().TempDir())
|
||||
|
||||
runtime, err := newBackendProcessRuntime()
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
dir := runtime.dir
|
||||
runtime.cleanup()
|
||||
|
||||
Expect(dir).ToNot(BeAnExistingFile())
|
||||
})
|
||||
|
||||
It("reports which configured root cannot be used", func() {
|
||||
parent := GinkgoT().TempDir()
|
||||
file := filepath.Join(parent, "not-a-directory")
|
||||
Expect(os.WriteFile(file, []byte("x"), 0o600)).To(Succeed())
|
||||
base := filepath.Join(file, "backend-runtime")
|
||||
GinkgoT().Setenv(backendTempDirEnv, base)
|
||||
root := backendRuntimeRoot()
|
||||
|
||||
runtime, err := newBackendProcessRuntime()
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(runtime).To(BeNil())
|
||||
Expect(err.Error()).To(ContainSubstring(root))
|
||||
})
|
||||
})
|
||||
@@ -1,38 +0,0 @@
|
||||
package model
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("Backend process state directory", func() {
|
||||
It("reports why the state directory could not be created", func() {
|
||||
// A worker whose volume is full, or whose TMPDIR no longer resolves,
|
||||
// cannot get a state directory. go-processmanager's New() drops the
|
||||
// option error, leaving StateDir empty, and Run() then failed with
|
||||
// "mkdir : no such file or directory" naming no path at all. Resolving
|
||||
// the directory here keeps the real cause attached.
|
||||
GinkgoT().Setenv("TMPDIR", filepath.Join(GinkgoT().TempDir(), "does-not-exist"))
|
||||
|
||||
dir, err := newProcessStateDir()
|
||||
Expect(err).To(HaveOccurred())
|
||||
Expect(dir).To(BeEmpty())
|
||||
Expect(err.Error()).To(ContainSubstring("backend process state directory"))
|
||||
Expect(err.Error()).To(ContainSubstring("does-not-exist"),
|
||||
"the error must name the directory it could not create")
|
||||
})
|
||||
|
||||
It("returns a usable directory when the temp location works", func() {
|
||||
GinkgoT().Setenv("TMPDIR", GinkgoT().TempDir())
|
||||
|
||||
dir, err := newProcessStateDir()
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
Expect(dir).ToNot(BeEmpty())
|
||||
info, statErr := os.Stat(dir)
|
||||
Expect(statErr).ToNot(HaveOccurred())
|
||||
Expect(info.IsDir()).To(BeTrue())
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,17 @@
|
||||
//go:build !aix && !darwin && !dragonfly && !freebsd && !linux && !netbsd && !openbsd && !solaris
|
||||
|
||||
package safefile
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
// ErrUnsafePath reports a path shape or file type that exact removal refuses.
|
||||
var ErrUnsafePath = errors.New("unsafe removal path")
|
||||
|
||||
// RemoveExact fails closed on platforms without component-relative no-follow
|
||||
// filesystem operations.
|
||||
func RemoveExact(root, relativePath string, sidecarSuffixes []string, pruneParents int) error {
|
||||
return fmt.Errorf("secure exact removal is unsupported on this platform")
|
||||
}
|
||||
@@ -0,0 +1,144 @@
|
||||
//go:build aix || darwin || dragonfly || freebsd || linux || netbsd || openbsd || solaris
|
||||
|
||||
package safefile
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// ErrUnsafePath reports a path shape or file type that exact removal refuses.
|
||||
var ErrUnsafePath = errors.New("unsafe removal path")
|
||||
|
||||
// RemoveExact removes a file and its named sidecars below root without
|
||||
// following symbolic links. It prunes up to pruneParents empty parent
|
||||
// directories, but never removes root itself.
|
||||
func RemoveExact(root, relativePath string, sidecarSuffixes []string, pruneParents int) error {
|
||||
return removeExact(root, relativePath, sidecarSuffixes, pruneParents, nil)
|
||||
}
|
||||
|
||||
func removeExact(root, relativePath string, sidecarSuffixes []string, pruneParents int, parentsOpened func()) error {
|
||||
parts, err := cleanRelativeParts(relativePath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if pruneParents < 0 {
|
||||
return fmt.Errorf("%w: prune parent count must not be negative", ErrUnsafePath)
|
||||
}
|
||||
|
||||
rootFD, err := unix.Open(root, unix.O_RDONLY|unix.O_DIRECTORY|unix.O_CLOEXEC|unix.O_NOFOLLOW, 0)
|
||||
if err != nil {
|
||||
return fmt.Errorf("opening removal root %q: %w", root, err)
|
||||
}
|
||||
handles := []int{rootFD}
|
||||
defer func() {
|
||||
for i := len(handles) - 1; i >= 0; i-- {
|
||||
_ = unix.Close(handles[i])
|
||||
}
|
||||
}()
|
||||
|
||||
for _, component := range parts[:len(parts)-1] {
|
||||
fd, openErr := unix.Openat(handles[len(handles)-1], component, unix.O_RDONLY|unix.O_DIRECTORY|unix.O_CLOEXEC|unix.O_NOFOLLOW, 0)
|
||||
if errors.Is(openErr, unix.ENOENT) {
|
||||
return nil
|
||||
}
|
||||
if openErr != nil {
|
||||
if errors.Is(openErr, unix.ELOOP) || errors.Is(openErr, unix.ENOTDIR) {
|
||||
return fmt.Errorf("%w: path component %q is not a directory: %v", ErrUnsafePath, component, openErr)
|
||||
}
|
||||
return fmt.Errorf("opening removal path component %q: %w", component, openErr)
|
||||
}
|
||||
handles = append(handles, fd)
|
||||
}
|
||||
if parentsOpened != nil {
|
||||
parentsOpened()
|
||||
}
|
||||
|
||||
parentFD := handles[len(handles)-1]
|
||||
leaf := parts[len(parts)-1]
|
||||
names := make([]string, 0, len(sidecarSuffixes)+1)
|
||||
names = append(names, leaf)
|
||||
for _, suffix := range sidecarSuffixes {
|
||||
if suffix == "" || strings.ContainsAny(suffix, `/\\`) {
|
||||
return fmt.Errorf("%w: invalid sidecar suffix %q", ErrUnsafePath, suffix)
|
||||
}
|
||||
names = append(names, leaf+suffix)
|
||||
}
|
||||
|
||||
for _, name := range names {
|
||||
var stat unix.Stat_t
|
||||
statErr := unix.Fstatat(parentFD, name, &stat, unix.AT_SYMLINK_NOFOLLOW)
|
||||
if errors.Is(statErr, unix.ENOENT) {
|
||||
continue
|
||||
}
|
||||
if statErr != nil {
|
||||
return fmt.Errorf("stating removal entry %q: %w", name, statErr)
|
||||
}
|
||||
switch stat.Mode & unix.S_IFMT {
|
||||
case unix.S_IFLNK:
|
||||
return fmt.Errorf("%w: refusing to remove symbolic link %q", ErrUnsafePath, name)
|
||||
case unix.S_IFDIR:
|
||||
return fmt.Errorf("%w: release key identifies a directory", ErrUnsafePath)
|
||||
}
|
||||
}
|
||||
|
||||
for _, name := range names {
|
||||
if unlinkErr := unix.Unlinkat(parentFD, name, 0); unlinkErr != nil && !errors.Is(unlinkErr, unix.ENOENT) {
|
||||
return fmt.Errorf("removing entry %q: %w", name, unlinkErr)
|
||||
}
|
||||
}
|
||||
|
||||
maxPrune := min(pruneParents, len(handles)-1)
|
||||
for childIndex := len(handles) - 1; childIndex >= len(handles)-maxPrune; childIndex-- {
|
||||
parentIndex := childIndex - 1
|
||||
name := parts[childIndex-1]
|
||||
same, identityErr := sameDirectoryEntry(handles[parentIndex], name, handles[childIndex])
|
||||
if identityErr != nil {
|
||||
if errors.Is(identityErr, unix.ENOENT) {
|
||||
break
|
||||
}
|
||||
return fmt.Errorf("checking directory %q before pruning: %w", name, identityErr)
|
||||
}
|
||||
if !same {
|
||||
break
|
||||
}
|
||||
removeErr := unix.Unlinkat(handles[parentIndex], name, unix.AT_REMOVEDIR)
|
||||
if removeErr == nil {
|
||||
continue
|
||||
}
|
||||
if errors.Is(removeErr, unix.ENOENT) || errors.Is(removeErr, unix.ENOTEMPTY) || errors.Is(removeErr, unix.EEXIST) {
|
||||
break
|
||||
}
|
||||
return fmt.Errorf("pruning directory %q: %w", name, removeErr)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func cleanRelativeParts(relativePath string) ([]string, error) {
|
||||
if relativePath == "" || filepath.IsAbs(relativePath) || filepath.Clean(relativePath) != relativePath {
|
||||
return nil, fmt.Errorf("%w: %q is not a clean relative path", ErrUnsafePath, relativePath)
|
||||
}
|
||||
parts := strings.Split(relativePath, string(filepath.Separator))
|
||||
for _, part := range parts {
|
||||
if part == "" || part == "." || part == ".." {
|
||||
return nil, fmt.Errorf("%w: %q is not a clean relative path", ErrUnsafePath, relativePath)
|
||||
}
|
||||
}
|
||||
return parts, nil
|
||||
}
|
||||
|
||||
func sameDirectoryEntry(parentFD int, name string, openedFD int) (bool, error) {
|
||||
var opened unix.Stat_t
|
||||
if err := unix.Fstat(openedFD, &opened); err != nil {
|
||||
return false, err
|
||||
}
|
||||
var current unix.Stat_t
|
||||
if err := unix.Fstatat(parentFD, name, ¤t, unix.AT_SYMLINK_NOFOLLOW); err != nil {
|
||||
return false, err
|
||||
}
|
||||
return current.Mode&unix.S_IFMT == unix.S_IFDIR && current.Dev == opened.Dev && current.Ino == opened.Ino, nil
|
||||
}
|
||||
@@ -0,0 +1,57 @@
|
||||
//go:build aix || darwin || dragonfly || freebsd || linux || netbsd || openbsd || solaris
|
||||
|
||||
package safefile
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
|
||||
var _ = Describe("Race-safe exact removal", func() {
|
||||
DescribeTable("keeps replacement targets untouched after an opened parent is exchanged",
|
||||
func(targetInsideRoot bool) {
|
||||
root := GinkgoT().TempDir()
|
||||
originalRequest := filepath.Join(root, "ephemeral", "request-id")
|
||||
originalCategory := filepath.Join(originalRequest, "audio")
|
||||
Expect(os.MkdirAll(originalCategory, 0750)).To(Succeed())
|
||||
Expect(os.WriteFile(filepath.Join(originalCategory, "input.wav"), []byte("original"), 0640)).To(Succeed())
|
||||
|
||||
replacement := GinkgoT().TempDir()
|
||||
if targetInsideRoot {
|
||||
replacement = filepath.Join(root, "replacement")
|
||||
Expect(os.MkdirAll(replacement, 0750)).To(Succeed())
|
||||
}
|
||||
Expect(os.MkdirAll(filepath.Join(replacement, "audio"), 0750)).To(Succeed())
|
||||
replacementFile := filepath.Join(replacement, "audio", "input.wav")
|
||||
Expect(os.WriteFile(replacementFile, []byte("replacement"), 0640)).To(Succeed())
|
||||
|
||||
renamedRequest := filepath.Join(root, "ephemeral", "opened-request")
|
||||
opened := make(chan struct{})
|
||||
swapped := make(chan struct{})
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
done <- removeExact(root, filepath.Join("ephemeral", "request-id", "audio", "input.wav"), []string{".sha256", ".sha256.target"}, 2, func() {
|
||||
close(opened)
|
||||
<-swapped
|
||||
})
|
||||
}()
|
||||
|
||||
<-opened
|
||||
Expect(os.Rename(originalRequest, renamedRequest)).To(Succeed())
|
||||
Expect(os.Symlink(replacement, originalRequest)).To(Succeed())
|
||||
close(swapped)
|
||||
Expect(<-done).To(Succeed())
|
||||
|
||||
data, err := os.ReadFile(replacementFile)
|
||||
Expect(err).NotTo(HaveOccurred())
|
||||
Expect(data).To(Equal([]byte("replacement")))
|
||||
Expect(originalRequest).To(BeAnExistingFile())
|
||||
Expect(filepath.Join(renamedRequest, "audio", "input.wav")).NotTo(BeAnExistingFile())
|
||||
},
|
||||
Entry("for an internal replacement", true),
|
||||
Entry("for an external replacement", false),
|
||||
)
|
||||
})
|
||||
Reference in new issue
Block a user