From 1ff07a0c14ca220fec8185d27476c00dcfbf8c95 Mon Sep 17 00:00:00 2001 From: brxs Date: Sat, 8 Aug 2026 18:22:51 -0700 Subject: [PATCH 01/14] feat: add bounded MRT2 render worker protocol --- backend/lsdj/sidecar.py | 563 ++++++++++++++++++++++++++- backend/tests/test_render_sidecar.py | 529 +++++++++++++++++++++++++ 2 files changed, 1090 insertions(+), 2 deletions(-) create mode 100644 backend/tests/test_render_sidecar.py diff --git a/backend/lsdj/sidecar.py b/backend/lsdj/sidecar.py index fe16632..1747895 100644 --- a/backend/lsdj/sidecar.py +++ b/backend/lsdj/sidecar.py @@ -16,6 +16,12 @@ - CONTROL (engine → sidecar): a deck command (``play``/``stop``/``set_style``…) as UTF-8 JSON. +The dedicated MRT2 clip renderer uses the same authenticated connection but a +separate strict protocol: one JSON RENDER_REQUEST (or matching RENDER_CANCEL), +then RENDER_BEGIN metadata, bounded aligned RENDER_CHUNK frames, and a +RENDER_END carrying the exact byte count and SHA-256. It never accepts deck +control frames or unbounded PCM. + Shared-worker frames prefix payloads with a single deck byte (0 or 1). The Rust and Python transport tests cover both forms without loading either model stack. @@ -25,7 +31,9 @@ """ import argparse +import hashlib import json +import math import os import queue import re @@ -33,9 +41,12 @@ import struct import sys import threading +from dataclasses import dataclass +from typing import Any, BinaryIO, Callable, Mapping from .mrt2 import ( AUTO_RUNTIME, + PYTORCH_CUDA_RUNTIME, RUNTIME_CHOICES, create_engine, public_startup_error, @@ -53,18 +64,93 @@ # child's scrubbed environment and proves that the connector is the process Rust # just spawned, rather than another local process racing the loopback accept. FRAME_AUTH = 5 +# Native host -> dedicated MRT2 render worker. The JSON payload is a single +# strict request; render workers never accept deck-control frames. +FRAME_RENDER_REQUEST = 6 +# Render worker -> native host. BEGIN fixes the expected audio identity and byte +# count before any payload; CHUNK carries aligned f32le PCM; END authenticates +# the completed byte stream with its exact count and SHA-256. +FRAME_RENDER_BEGIN = 7 +FRAME_RENDER_CHUNK = 8 +FRAME_RENDER_END = 9 +# Native host -> render worker. In-flight model calls are not cooperatively +# cancellable, so this frame makes the disposable worker exit and release CUDA. +FRAME_RENDER_CANCEL = 10 +# Render worker -> native host. Diagnostics are bounded and path-free. +FRAME_RENDER_ERROR = 11 MAX_FRAME_BYTES = 16 * 1024 * 1024 MAX_EMBED_ID_BYTES = 4 * 1024 WORKER_TOKEN_ENV = "LSDJ_WORKER_LAUNCH_TOKEN" +RENDER_SCHEMA_VERSION = 1 +RENDER_SAMPLE_RATE = 48_000 +RENDER_CHANNELS = 2 +RENDER_SAMPLE_WIDTH = 4 +RENDER_BYTES_PER_FRAME = RENDER_CHANNELS * RENDER_SAMPLE_WIDTH +MIN_RENDER_SECONDS = 0.5 +MAX_RENDER_SECONDS = 180.0 +MAX_RENDER_PROMPT_CHARS = 32_000 +MAX_RENDER_REQUEST_BYTES = 64 * 1024 +MAX_RENDER_CONTROL_BYTES = 1024 +MAX_RENDER_PCM_BYTES = ( + round(MAX_RENDER_SECONDS * RENDER_SAMPLE_RATE) * RENDER_BYTES_PER_FRAME +) +RENDER_PCM_CHUNK_BYTES = 1024 * 1024 +MAX_RENDER_METADATA_BYTES = 8 * 1024 +_RENDER_JOB_ID = re.compile(r"^[A-Za-z0-9_-]{16,80}$") + # u8 frame type, u32 little-endian payload length. _HEADER = struct.Struct(" int: + return round(self.seconds * RENDER_SAMPLE_RATE) + + @property + def pcm_bytes(self) -> int: + return self.frames * RENDER_BYTES_PER_FRAME + + +@dataclass(frozen=True) +class RenderCancel: + job_id: str + + +@dataclass(frozen=True) +class _RenderReaderFailure: + message: str + + +_RenderCommand = RenderRequest | RenderCancel | _RenderReaderFailure | None + + def write_frame(sock: socket.socket, frame_type: int, payload: bytes) -> None: """Send one framed message. `sendall` is atomic enough here: the worker loop is the only writer, so frames never interleave.""" @@ -102,6 +188,272 @@ def authenticate_to_host( write_frame(sock, FRAME_AUTH, token.encode("ascii")) +def _strict_json_object(payload: bytes) -> dict[str, object]: + def pairs(values: list[tuple[str, object]]) -> dict[str, object]: + result: dict[str, object] = {} + for name, value in values: + if name in result: + raise RenderProtocolError("render JSON contains a duplicate field") + result[name] = value + return result + + try: + value = json.loads(payload.decode("utf-8"), object_pairs_hook=pairs) + except RenderProtocolError: + raise + except (UnicodeDecodeError, json.JSONDecodeError): + raise RenderProtocolError("render JSON is invalid") from None + if not isinstance(value, dict): + raise RenderProtocolError("render JSON must be an object") + return value + + +def _read_bounded_render_frame( + reader: BinaryIO, limits: Mapping[int, int] +) -> tuple[int, bytes] | None: + head = reader.read(_HEADER.size) + if not head: + return None + if len(head) != _HEADER.size: + raise RenderProtocolError("render frame header is truncated") + frame_type, length = _HEADER.unpack(head) + limit = limits.get(frame_type) + if limit is None: + raise RenderProtocolError(f"render frame type {frame_type} is out of order") + if length > limit: + raise RenderProtocolError( + f"render frame type {frame_type} exceeds its {limit}-byte cap" + ) + payload = reader.read(length) + if len(payload) != length: + raise RenderProtocolError("render frame payload is truncated") + return frame_type, payload + + +def _validate_job_id(value: object) -> str: + if not isinstance(value, str) or _RENDER_JOB_ID.fullmatch(value) is None: + raise RenderProtocolError("render jobId is invalid") + return value + + +def read_render_command(reader: BinaryIO) -> RenderRequest | RenderCancel | None: + """Read one strict host command, distinguishing clean EOF from truncation.""" + + frame = _read_bounded_render_frame( + reader, + { + FRAME_RENDER_REQUEST: MAX_RENDER_REQUEST_BYTES, + FRAME_RENDER_CANCEL: MAX_RENDER_CONTROL_BYTES, + }, + ) + if frame is None: + return None + frame_type, payload = frame + value = _strict_json_object(payload) + if value.get("schemaVersion") != RENDER_SCHEMA_VERSION: + raise RenderProtocolError("render command schema is unsupported") + job_id = _validate_job_id(value.get("jobId")) + if frame_type == FRAME_RENDER_CANCEL: + if set(value) != {"schemaVersion", "jobId"}: + raise RenderProtocolError("render cancel contains unknown fields") + return RenderCancel(job_id=job_id) + + if set(value) != {"schemaVersion", "jobId", "prompt", "seconds"}: + raise RenderProtocolError("render request contains missing or unknown fields") + prompt = value.get("prompt") + if not isinstance(prompt, str) or not prompt.strip(): + raise RenderProtocolError("render prompt must be a non-empty string") + prompt = prompt.strip() + if len(prompt) > MAX_RENDER_PROMPT_CHARS: + raise RenderProtocolError("render prompt exceeds its character cap") + seconds = value.get("seconds") + if ( + isinstance(seconds, bool) + or not isinstance(seconds, (int, float)) + or not math.isfinite(seconds) + or not MIN_RENDER_SECONDS <= float(seconds) <= MAX_RENDER_SECONDS + ): + raise RenderProtocolError( + f"render seconds must be {MIN_RENDER_SECONDS:g}-{MAX_RENDER_SECONDS:g}" + ) + request = RenderRequest(job_id=job_id, prompt=prompt, seconds=float(seconds)) + if request.pcm_bytes > MAX_RENDER_PCM_BYTES: + raise RenderProtocolError("render output exceeds its PCM byte cap") + return request + + +def _render_json(value: Mapping[str, object]) -> bytes: + payload = json.dumps(value, sort_keys=True, separators=(",", ":")).encode() + if len(payload) > MAX_RENDER_METADATA_BYTES: + raise RenderProtocolError("render metadata exceeds its cap") + return payload + + +def write_render_error( + sock: socket.socket, + *, + job_id: str | None, + code: str, + message: str, +) -> None: + write_frame( + sock, + FRAME_RENDER_ERROR, + _render_json( + { + "schemaVersion": RENDER_SCHEMA_VERSION, + "jobId": job_id, + "code": code[:64], + "message": message[:512], + } + ), + ) + + +def write_render_response( + sock: socket.socket, request: RenderRequest, pcm: bytes +) -> None: + """Write one complete, size-checked f32le render response.""" + + if not isinstance(pcm, bytes): + raise RenderProtocolError("render engine returned a non-bytes payload") + if len(pcm) != request.pcm_bytes or len(pcm) > MAX_RENDER_PCM_BYTES: + raise RenderProtocolError( + f"render engine returned {len(pcm)} PCM bytes; expected {request.pcm_bytes}" + ) + write_frame( + sock, + FRAME_RENDER_BEGIN, + _render_json( + { + "schemaVersion": RENDER_SCHEMA_VERSION, + "jobId": request.job_id, + "sampleRate": RENDER_SAMPLE_RATE, + "channels": RENDER_CHANNELS, + "sampleFormat": "f32le", + "frames": request.frames, + "pcmBytes": request.pcm_bytes, + } + ), + ) + digest = hashlib.sha256() + for start in range(0, len(pcm), RENDER_PCM_CHUNK_BYTES): + chunk = pcm[start : start + RENDER_PCM_CHUNK_BYTES] + digest.update(chunk) + write_frame(sock, FRAME_RENDER_CHUNK, chunk) + write_frame( + sock, + FRAME_RENDER_END, + _render_json( + { + "schemaVersion": RENDER_SCHEMA_VERSION, + "jobId": request.job_id, + "pcmBytes": request.pcm_bytes, + "sha256": digest.hexdigest(), + } + ), + ) + + +def read_render_response( + reader: BinaryIO, expected_job_id: str, *, require_eof: bool = False +) -> bytes: + """Validate and assemble one response; the native host mirrors this parser.""" + + expected_job_id = _validate_job_id(expected_job_id) + first = _read_bounded_render_frame( + reader, + { + FRAME_RENDER_BEGIN: MAX_RENDER_METADATA_BYTES, + FRAME_RENDER_ERROR: MAX_RENDER_METADATA_BYTES, + }, + ) + if first is None: + raise RenderProtocolError("render response ended before begin") + frame_type, payload = first + value = _strict_json_object(payload) + if frame_type == FRAME_RENDER_ERROR: + if ( + set(value) != {"schemaVersion", "jobId", "code", "message"} + or value.get("schemaVersion") != RENDER_SCHEMA_VERSION + or value.get("jobId") != expected_job_id + or not isinstance(value.get("code"), str) + or not isinstance(value.get("message"), str) + ): + raise RenderProtocolError("render error metadata is invalid") + code = value.get("code") + raise RenderProtocolError(f"render worker returned {code}") + expected_fields = { + "schemaVersion", + "jobId", + "sampleRate", + "channels", + "sampleFormat", + "frames", + "pcmBytes", + } + if ( + set(value) != expected_fields + or value.get("schemaVersion") != RENDER_SCHEMA_VERSION + ): + raise RenderProtocolError("render begin metadata is invalid") + if value.get("jobId") != expected_job_id: + raise RenderProtocolError("render response jobId is out of turn") + frames = value.get("frames") + pcm_bytes = value.get("pcmBytes") + if ( + value.get("sampleRate") != RENDER_SAMPLE_RATE + or value.get("channels") != RENDER_CHANNELS + or value.get("sampleFormat") != "f32le" + or isinstance(frames, bool) + or not isinstance(frames, int) + or frames < 1 + or isinstance(pcm_bytes, bool) + or not isinstance(pcm_bytes, int) + or pcm_bytes != frames * RENDER_BYTES_PER_FRAME + or pcm_bytes > MAX_RENDER_PCM_BYTES + ): + raise RenderProtocolError("render begin audio identity is invalid") + + output = bytearray() + digest = hashlib.sha256() + while True: + frame = _read_bounded_render_frame( + reader, + { + FRAME_RENDER_CHUNK: RENDER_PCM_CHUNK_BYTES, + FRAME_RENDER_END: MAX_RENDER_METADATA_BYTES, + }, + ) + if frame is None: + raise RenderProtocolError("render response is truncated") + frame_type, payload = frame + if frame_type == FRAME_RENDER_CHUNK: + if not payload or len(payload) % RENDER_BYTES_PER_FRAME: + raise RenderProtocolError("render PCM chunk is empty or misaligned") + if len(output) + len(payload) > pcm_bytes: + raise RenderProtocolError("render response contains extra PCM bytes") + output.extend(payload) + digest.update(payload) + continue + + end = _strict_json_object(payload) + if ( + set(end) != {"schemaVersion", "jobId", "pcmBytes", "sha256"} + or end.get("schemaVersion") != RENDER_SCHEMA_VERSION + or end.get("jobId") != expected_job_id + or end.get("pcmBytes") != pcm_bytes + or end.get("sha256") != digest.hexdigest() + or len(output) != pcm_bytes + ): + raise RenderProtocolError( + "render end metadata or exact byte total is invalid" + ) + if require_eof and reader.read(1): + raise RenderProtocolError("render response contains frames after end") + return bytes(output) + + class SocketOutQueue: """`run_deck_worker`'s `out_queue`, writing to the socket: ``('audio', bytes)`` → a PCM frame, ``('status', dict)`` → a status frame.""" @@ -349,6 +701,189 @@ def run_shared_sidecar( thread.join() +class _RenderCommandReader: + """Continuously read the control half so cancellation can preempt rendering.""" + + def __init__(self, reader: BinaryIO) -> None: + # Socket backpressure bounds requests that arrive faster than the single + # render slot can consume them; no unbounded JSON queue exists. + self.commands: queue.Queue[_RenderCommand] = queue.Queue(maxsize=4) + self._reader = reader + self._thread = threading.Thread( + target=self._pump, name="mrt2-render-control", daemon=True + ) + self._thread.start() + + def _pump(self) -> None: + while True: + try: + command = read_render_command(self._reader) + except (OSError, RenderProtocolError) as error: + message = ( + str(error) + if isinstance(error, RenderProtocolError) + else "render control connection failed" + ) + self.commands.put(_RenderReaderFailure(message[:512])) + return + self.commands.put(command) + if command is None: + return + + +def _render_startup_error(error: Exception) -> str: + # `public_startup_error` preserves deliberately bounded RuntimeUnavailable + # diagnostics and collapses unknown exceptions to their class only. + return public_startup_error(error)[:512] + + +def run_render_worker( + sock: socket.socket, + model: str, + *, + runtime: str = PYTORCH_CUDA_RUNTIME, + engine_factory=None, + terminate: Callable[[int], Any] = os._exit, +) -> None: + """Serve serial, authenticated MRT2 clip renders over one bounded socket. + + Authentication is emitted by :func:`main` before this function runs. The + loaded model remains warm across successful requests. A render call cannot + be interrupted safely inside upstream PyTorch, so cancellation or EOF while + it is active terminates this disposable process and releases its CUDA + context; the native supervisor may start a fresh worker for the next job. + """ + + try: + engine = ( + create_engine(model=model, runtime=runtime) + if engine_factory is None + else engine_factory(model=model) + ) + warm_up = getattr(engine, "warm_up", None) + if callable(warm_up): + warm_up() + except Exception as error: + write_render_error( + sock, + job_id=None, + code="startup_failed", + message=_render_startup_error(error), + ) + return + + write_frame( + sock, + FRAME_STATUS, + _render_json( + { + "schemaVersion": RENDER_SCHEMA_VERSION, + "event": "render_ready", + "model": model, + "runtime": runtime, + } + ), + ) + commands = _RenderCommandReader(sock.makefile("rb")) + + while True: + command = commands.commands.get() + if command is None: + return + if isinstance(command, _RenderReaderFailure): + write_render_error( + sock, + job_id=None, + code="protocol_error", + message=command.message, + ) + return + if isinstance(command, RenderCancel): + write_render_error( + sock, + job_id=command.job_id, + code="no_active_job", + message="render job is not active", + ) + continue + + result: queue.Queue[tuple[bool, bytes | Exception]] = queue.Queue(maxsize=1) + + def render() -> None: + try: + result.put((True, engine.render_clip(command.prompt, command.seconds))) + except Exception as error: # noqa: BLE001 - collapsed at the boundary + result.put((False, error)) + + threading.Thread( + target=render, + name=f"mrt2-render-{command.job_id[:16]}", + daemon=True, + ).start() + + while True: + try: + succeeded, value = result.get(timeout=0.025) + break + except queue.Empty: + pass + try: + pending = commands.commands.get_nowait() + except queue.Empty: + continue + + if pending is None: + terminate(0) + return + if isinstance(pending, RenderCancel) and pending.job_id == command.job_id: + write_render_error( + sock, + job_id=command.job_id, + code="cancelled", + message="render job was cancelled", + ) + terminate(2) + return + if isinstance(pending, _RenderReaderFailure): + message = pending.message + job_id = command.job_id + elif isinstance(pending, RenderRequest): + message = "render requests must not overlap" + job_id = pending.job_id + else: + message = "render cancellation is out of turn" + job_id = pending.job_id + write_render_error( + sock, + job_id=job_id, + code="protocol_error", + message=message, + ) + terminate(2) + return + + if not succeeded: + write_render_error( + sock, + job_id=command.job_id, + code="render_failed", + message="MRT2 render failed; the worker must be restarted", + ) + return + try: + if not isinstance(value, bytes): + raise RenderProtocolError("render engine returned a non-bytes payload") + write_render_response(sock, command, value) + except RenderProtocolError: + write_render_error( + sock, + job_id=command.job_id, + code="invalid_audio", + message="MRT2 render returned an invalid PCM payload", + ) + return + + # --- Model tooling (the in-app model manager, issue #43) ------------------- # # The Rust shell spawns this same binary to install Magenta assets without a @@ -445,9 +980,15 @@ def main(argv=None) -> None: # model-tooling modes below (issue #43) without a deck/port. parser.add_argument("--deck", help="deck id (e.g. a or b)") parser.add_argument("--model", help="model name (e.g. mrt2_small)") - parser.add_argument( + modes = parser.add_mutually_exclusive_group() + modes.add_argument( "--shared", action="store_true", help="run both decks in one worker" ) + modes.add_argument( + "--render-worker", + action="store_true", + help="run the dedicated authenticated MRT2 clip renderer", + ) parser.add_argument("--model-a", help="shared-worker model for deck a") parser.add_argument("--model-b", help="shared-worker model for deck b") parser.add_argument( @@ -489,6 +1030,24 @@ def main(argv=None) -> None: ) return + if args.render_worker: + missing = [name for name in ("model", "port") if getattr(args, name) is None] + if missing: + parser.error( + "the following arguments are required in render-worker mode: " + + ", ".join("--" + name for name in missing) + ) + if args.runtime != PYTORCH_CUDA_RUNTIME: + parser.error( + "render-worker mode requires the explicit pytorch-cuda runtime" + ) + sock = socket.create_connection(("127.0.0.1", args.port)) + sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) + authenticate_to_host(sock) + os.environ.pop(WORKER_TOKEN_ENV, None) + run_render_worker(sock, args.model, runtime=args.runtime) + return + if args.shared: missing = [ name diff --git a/backend/tests/test_render_sidecar.py b/backend/tests/test_render_sidecar.py new file mode 100644 index 0000000..22a8197 --- /dev/null +++ b/backend/tests/test_render_sidecar.py @@ -0,0 +1,529 @@ +"""Bounded, authenticated protocol for the dedicated MRT2 render worker.""" + +import io +import json +import socket +import struct +import threading + +import pytest + +from lsdj.sidecar import ( + FRAME_AUTH, + FRAME_RENDER_BEGIN, + FRAME_RENDER_CANCEL, + FRAME_RENDER_CHUNK, + FRAME_RENDER_END, + FRAME_RENDER_ERROR, + FRAME_RENDER_REQUEST, + FRAME_STATUS, + MAX_RENDER_PCM_BYTES, + MAX_RENDER_PROMPT_CHARS, + MAX_RENDER_REQUEST_BYTES, + PYTORCH_CUDA_RUNTIME, + RENDER_BYTES_PER_FRAME, + RENDER_CHANNELS, + RENDER_PCM_CHUNK_BYTES, + RENDER_SAMPLE_RATE, + RENDER_SCHEMA_VERSION, + RenderProtocolError, + RenderRequest, + read_frame, + read_render_command, + read_render_response, + run_render_worker, + write_frame, + write_render_response, +) + +JOB_ID = "render-job-0123456789abcdef" + + +class RecordingSock: + def __init__(self): + self.buffer = bytearray() + + def sendall(self, data): + self.buffer.extend(data) + + def setsockopt(self, *_args): + pass + + +class FakeRenderEngine: + def __init__(self): + self.warmups = 0 + self.requests = [] + + def warm_up(self): + self.warmups += 1 + + def render_clip(self, prompt, seconds): + self.requests.append((prompt, seconds)) + frames = round(seconds * RENDER_SAMPLE_RATE) + return b"\0" * (frames * RENDER_BYTES_PER_FRAME) + + +class BlockingRenderEngine(FakeRenderEngine): + def __init__(self): + super().__init__() + self.started = threading.Event() + self.release = threading.Event() + + def render_clip(self, prompt, seconds): + self.requests.append((prompt, seconds)) + self.started.set() + self.release.wait(timeout=5) + return b"\0" * (round(seconds * RENDER_SAMPLE_RATE) * RENDER_BYTES_PER_FRAME) + + +def request_payload(*, job_id=JOB_ID, prompt="bright piano", seconds=0.5, **extra): + return json.dumps( + { + "schemaVersion": RENDER_SCHEMA_VERSION, + "jobId": job_id, + "prompt": prompt, + "seconds": seconds, + **extra, + }, + separators=(",", ":"), + ).encode() + + +def cancel_payload(job_id=JOB_ID): + return json.dumps( + {"schemaVersion": RENDER_SCHEMA_VERSION, "jobId": job_id}, + separators=(",", ":"), + ).encode() + + +def begin_payload(*, job_id=JOB_ID, frames=1): + return json.dumps( + { + "schemaVersion": RENDER_SCHEMA_VERSION, + "jobId": job_id, + "sampleRate": RENDER_SAMPLE_RATE, + "channels": RENDER_CHANNELS, + "sampleFormat": "f32le", + "frames": frames, + "pcmBytes": frames * RENDER_BYTES_PER_FRAME, + }, + separators=(",", ":"), + ).encode() + + +def read_ready(reader): + frame_type, payload = read_frame(reader) + assert frame_type == FRAME_STATUS + assert json.loads(payload)["event"] == "render_ready" + + +def test_render_command_is_strict_and_bounded(): + valid = io.BytesIO( + struct.pack(" 1 + + +@pytest.mark.parametrize("delta", [-RENDER_BYTES_PER_FRAME, RENDER_BYTES_PER_FRAME]) +def test_render_response_writer_rejects_short_and_extra_pcm(delta): + request = RenderRequest(JOB_ID, "piano", 0.5) + with pytest.raises(RenderProtocolError, match="expected"): + write_render_response( + RecordingSock(), request, b"\0" * (request.pcm_bytes + delta) + ) + + +def test_render_response_reader_rejects_out_of_order_oversized_and_extra_pcm(): + out_of_order = RecordingSock() + write_frame(out_of_order, FRAME_RENDER_CHUNK, b"\0" * RENDER_BYTES_PER_FRAME) + with pytest.raises(RenderProtocolError, match="out of order"): + read_render_response(io.BytesIO(out_of_order.buffer), JOB_ID) + + oversized = bytearray() + begin = begin_payload() + oversized.extend(struct.pack(" Date: Sat, 8 Aug 2026 18:55:50 -0700 Subject: [PATCH 02/14] fix: harden MRT2 render protocol contract --- backend/lsdj/sidecar.py | 620 +++++++++++++++++++-------- backend/tests/test_render_sidecar.py | 312 ++++++++++++-- 2 files changed, 725 insertions(+), 207 deletions(-) diff --git a/backend/lsdj/sidecar.py b/backend/lsdj/sidecar.py index 1747895..b7def7b 100644 --- a/backend/lsdj/sidecar.py +++ b/backend/lsdj/sidecar.py @@ -17,10 +17,11 @@ as UTF-8 JSON. The dedicated MRT2 clip renderer uses the same authenticated connection but a -separate strict protocol: one JSON RENDER_REQUEST (or matching RENDER_CANCEL), -then RENDER_BEGIN metadata, bounded aligned RENDER_CHUNK frames, and a -RENDER_END carrying the exact byte count and SHA-256. It never accepts deck -control frames or unbounded PCM. +separate strict protocol: serial JSON RENDER_REQUEST messages with authoritative +integer frame counts and monotonically increasing sequence numbers (or a cancel +matching the active job/sequence), then exact RENDER_BEGIN metadata, bounded +aligned RENDER_CHUNK frames, and a RENDER_END carrying the exact identity, byte +count, and SHA-256. It never accepts deck control frames or unbounded PCM. Shared-worker frames prefix payloads with a single deck byte (0 or 1). The Rust and Python transport tests cover both forms without loading either model stack. @@ -41,8 +42,9 @@ import struct import sys import threading +import time from dataclasses import dataclass -from typing import Any, BinaryIO, Callable, Mapping +from typing import Any, BinaryIO, Callable, Mapping, MutableMapping from .mrt2 import ( AUTO_RUNTIME, @@ -90,14 +92,17 @@ RENDER_BYTES_PER_FRAME = RENDER_CHANNELS * RENDER_SAMPLE_WIDTH MIN_RENDER_SECONDS = 0.5 MAX_RENDER_SECONDS = 180.0 +MIN_RENDER_FRAMES = 24_000 +MAX_RENDER_FRAMES = 8_640_000 MAX_RENDER_PROMPT_CHARS = 32_000 MAX_RENDER_REQUEST_BYTES = 64 * 1024 MAX_RENDER_CONTROL_BYTES = 1024 -MAX_RENDER_PCM_BYTES = ( - round(MAX_RENDER_SECONDS * RENDER_SAMPLE_RATE) * RENDER_BYTES_PER_FRAME -) +MAX_RENDER_PCM_BYTES = MAX_RENDER_FRAMES * RENDER_BYTES_PER_FRAME RENDER_PCM_CHUNK_BYTES = 1024 * 1024 MAX_RENDER_METADATA_BYTES = 8 * 1024 +RENDER_WRITE_POLL_SECONDS = 0.01 +RENDER_WRITE_TIMEOUT_SECONDS = 5.0 +MAX_U64 = (1 << 64) - 1 _RENDER_JOB_ID = re.compile(r"^[A-Za-z0-9_-]{16,80}$") # u8 frame type, u32 little-endian payload length. @@ -126,12 +131,16 @@ class RenderProtocolError(ValueError): @dataclass(frozen=True) class RenderRequest: job_id: str + sequence: int prompt: str - seconds: float + frames: int @property - def frames(self) -> int: - return round(self.seconds * RENDER_SAMPLE_RATE) + def seconds(self) -> float: + # Frames are authoritative on the wire. Seconds exist only at the + # upstream engine boundary, so Python never independently rounds the + # user's duration. + return self.frames / RENDER_SAMPLE_RATE @property def pcm_bytes(self) -> int: @@ -141,6 +150,7 @@ def pcm_bytes(self) -> int: @dataclass(frozen=True) class RenderCancel: job_id: str + sequence: int @dataclass(frozen=True) @@ -151,6 +161,14 @@ class _RenderReaderFailure: _RenderCommand = RenderRequest | RenderCancel | _RenderReaderFailure | None +class _RenderWorkerStopped(Exception): + """Internal control-flow marker after the disposable worker is stopped.""" + + +class _RenderWriteError(Exception): + """A response frame could not be committed within the bounded deadline.""" + + def write_frame(sock: socket.socket, frame_type: int, payload: bytes) -> None: """Send one framed message. `sendall` is atomic enough here: the worker loop is the only writer, so frames never interleave.""" @@ -177,12 +195,18 @@ def read_frame(reader) -> tuple[int, bytes] | None: def authenticate_to_host( - sock: socket.socket, env: dict[str, str] | None = None + sock: socket.socket, env: MutableMapping[str, str] | None = None ) -> None: - """Send the in-memory launch capability before any worker traffic.""" + """Consume and send the per-child capability before any worker traffic. + + The token is removed before the write, including on failure, so a child can + never reconnect or retry with the same capability. The native listener must + accept it as the exact first frame, compare it in constant time, and consume + its expected token after that one connection attempt. + """ env = os.environ if env is None else env - token = env.get(WORKER_TOKEN_ENV, "") + token = env.pop(WORKER_TOKEN_ENV, "") if not 32 <= len(token) <= 256 or not token.isascii(): raise RuntimeError("the authenticated sidecar launch token is missing") write_frame(sock, FRAME_AUTH, token.encode("ascii")) @@ -201,7 +225,13 @@ def pairs(values: list[tuple[str, object]]) -> dict[str, object]: value = json.loads(payload.decode("utf-8"), object_pairs_hook=pairs) except RenderProtocolError: raise - except (UnicodeDecodeError, json.JSONDecodeError): + except ( + UnicodeDecodeError, + json.JSONDecodeError, + ValueError, + RecursionError, + OverflowError, + ): raise RenderProtocolError("render JSON is invalid") from None if not isinstance(value, dict): raise RenderProtocolError("render JSON must be an object") @@ -236,6 +266,38 @@ def _validate_job_id(value: object) -> str: return value +def _validate_exact_int(value: object, *, name: str, minimum: int, maximum: int) -> int: + if type(value) is not int or not minimum <= value <= maximum: + raise RenderProtocolError(f"render {name} is invalid") + return value + + +def render_frames_for_seconds(seconds: float) -> int: + """Reference conversion for the native gateway's user-facing duration. + + The gateway performs this once, before sending an integer frame count: + ``floor(seconds_f64 * 48000 + 0.5)``. This deliberately avoids Python's + ties-to-even ``round`` behavior at half-frame boundaries. + """ + + if ( + isinstance(seconds, bool) + or not isinstance(seconds, (int, float)) + or not math.isfinite(seconds) + or not MIN_RENDER_SECONDS <= float(seconds) <= MAX_RENDER_SECONDS + ): + raise RenderProtocolError( + f"render seconds must be {MIN_RENDER_SECONDS:g}-{MAX_RENDER_SECONDS:g}" + ) + frames = math.floor(float(seconds) * RENDER_SAMPLE_RATE + 0.5) + return _validate_exact_int( + frames, + name="frames", + minimum=MIN_RENDER_FRAMES, + maximum=MAX_RENDER_FRAMES, + ) + + def read_render_command(reader: BinaryIO) -> RenderRequest | RenderCancel | None: """Read one strict host command, distinguishing clean EOF from truncation.""" @@ -250,15 +312,19 @@ def read_render_command(reader: BinaryIO) -> RenderRequest | RenderCancel | None return None frame_type, payload = frame value = _strict_json_object(payload) - if value.get("schemaVersion") != RENDER_SCHEMA_VERSION: + schema_version = value.get("schemaVersion") + if type(schema_version) is not int or schema_version != RENDER_SCHEMA_VERSION: raise RenderProtocolError("render command schema is unsupported") job_id = _validate_job_id(value.get("jobId")) + sequence = _validate_exact_int( + value.get("sequence"), name="sequence", minimum=1, maximum=MAX_U64 + ) if frame_type == FRAME_RENDER_CANCEL: - if set(value) != {"schemaVersion", "jobId"}: + if set(value) != {"schemaVersion", "jobId", "sequence"}: raise RenderProtocolError("render cancel contains unknown fields") - return RenderCancel(job_id=job_id) + return RenderCancel(job_id=job_id, sequence=sequence) - if set(value) != {"schemaVersion", "jobId", "prompt", "seconds"}: + if set(value) != {"schemaVersion", "jobId", "sequence", "prompt", "frames"}: raise RenderProtocolError("render request contains missing or unknown fields") prompt = value.get("prompt") if not isinstance(prompt, str) or not prompt.strip(): @@ -266,20 +332,18 @@ def read_render_command(reader: BinaryIO) -> RenderRequest | RenderCancel | None prompt = prompt.strip() if len(prompt) > MAX_RENDER_PROMPT_CHARS: raise RenderProtocolError("render prompt exceeds its character cap") - seconds = value.get("seconds") - if ( - isinstance(seconds, bool) - or not isinstance(seconds, (int, float)) - or not math.isfinite(seconds) - or not MIN_RENDER_SECONDS <= float(seconds) <= MAX_RENDER_SECONDS - ): - raise RenderProtocolError( - f"render seconds must be {MIN_RENDER_SECONDS:g}-{MAX_RENDER_SECONDS:g}" - ) - request = RenderRequest(job_id=job_id, prompt=prompt, seconds=float(seconds)) - if request.pcm_bytes > MAX_RENDER_PCM_BYTES: - raise RenderProtocolError("render output exceeds its PCM byte cap") - return request + frames = _validate_exact_int( + value.get("frames"), + name="frames", + minimum=MIN_RENDER_FRAMES, + maximum=MAX_RENDER_FRAMES, + ) + return RenderRequest( + job_id=job_id, + sequence=sequence, + prompt=prompt, + frames=frames, + ) def _render_json(value: Mapping[str, object]) -> bytes: @@ -293,16 +357,28 @@ def write_render_error( sock: socket.socket, *, job_id: str | None, + sequence: int, code: str, message: str, + send_frame: Callable[[int, bytes], None] | None = None, ) -> None: - write_frame( - sock, + sequence = _validate_exact_int( + sequence, name="sequence", minimum=0, maximum=MAX_U64 + ) + if job_id is not None: + _validate_job_id(job_id) + sender = ( + (lambda frame_type, payload: write_frame(sock, frame_type, payload)) + if send_frame is None + else send_frame + ) + sender( FRAME_RENDER_ERROR, _render_json( { "schemaVersion": RENDER_SCHEMA_VERSION, "jobId": job_id, + "sequence": sequence, "code": code[:64], "message": message[:512], } @@ -311,23 +387,43 @@ def write_render_error( def write_render_response( - sock: socket.socket, request: RenderRequest, pcm: bytes + sock: socket, + request: RenderRequest, + pcm: bytes, + *, + send_frame: Callable[[int, bytes], None] | None = None, + before_frame: Callable[[], None] | None = None, ) -> None: """Write one complete, size-checked f32le render response.""" + _validate_job_id(request.job_id) + _validate_exact_int(request.sequence, name="sequence", minimum=1, maximum=MAX_U64) + _validate_exact_int( + request.frames, + name="frames", + minimum=MIN_RENDER_FRAMES, + maximum=MAX_RENDER_FRAMES, + ) if not isinstance(pcm, bytes): raise RenderProtocolError("render engine returned a non-bytes payload") if len(pcm) != request.pcm_bytes or len(pcm) > MAX_RENDER_PCM_BYTES: raise RenderProtocolError( f"render engine returned {len(pcm)} PCM bytes; expected {request.pcm_bytes}" ) - write_frame( - sock, + sender = ( + (lambda frame_type, payload: write_frame(sock, frame_type, payload)) + if send_frame is None + else send_frame + ) + check = (lambda: None) if before_frame is None else before_frame + check() + sender( FRAME_RENDER_BEGIN, _render_json( { "schemaVersion": RENDER_SCHEMA_VERSION, "jobId": request.job_id, + "sequence": request.sequence, "sampleRate": RENDER_SAMPLE_RATE, "channels": RENDER_CHANNELS, "sampleFormat": "f32le", @@ -338,16 +434,19 @@ def write_render_response( ) digest = hashlib.sha256() for start in range(0, len(pcm), RENDER_PCM_CHUNK_BYTES): + check() chunk = pcm[start : start + RENDER_PCM_CHUNK_BYTES] digest.update(chunk) - write_frame(sock, FRAME_RENDER_CHUNK, chunk) - write_frame( - sock, + sender(FRAME_RENDER_CHUNK, chunk) + check() + sender( FRAME_RENDER_END, _render_json( { "schemaVersion": RENDER_SCHEMA_VERSION, "jobId": request.job_id, + "sequence": request.sequence, + "frames": request.frames, "pcmBytes": request.pcm_bytes, "sha256": digest.hexdigest(), } @@ -355,12 +454,36 @@ def write_render_response( ) +def _validate_render_error(value: Mapping[str, object], request: RenderRequest) -> str: + if ( + set(value) != {"schemaVersion", "jobId", "sequence", "code", "message"} + or type(value.get("schemaVersion")) is not int + or value.get("schemaVersion") != RENDER_SCHEMA_VERSION + or value.get("jobId") != request.job_id + or type(value.get("sequence")) is not int + or value.get("sequence") != request.sequence + or type(value.get("code")) is not str + or not 1 <= len(value["code"]) <= 64 + or type(value.get("message")) is not str + or len(value["message"]) > 512 + ): + raise RenderProtocolError("render error metadata is invalid") + return value["code"] + + def read_render_response( - reader: BinaryIO, expected_job_id: str, *, require_eof: bool = False + reader: BinaryIO, request: RenderRequest, *, require_eof: bool = False ) -> bytes: """Validate and assemble one response; the native host mirrors this parser.""" - expected_job_id = _validate_job_id(expected_job_id) + _validate_job_id(request.job_id) + _validate_exact_int(request.sequence, name="sequence", minimum=1, maximum=MAX_U64) + _validate_exact_int( + request.frames, + name="frames", + minimum=MIN_RENDER_FRAMES, + maximum=MAX_RENDER_FRAMES, + ) first = _read_bounded_render_frame( reader, { @@ -373,19 +496,12 @@ def read_render_response( frame_type, payload = first value = _strict_json_object(payload) if frame_type == FRAME_RENDER_ERROR: - if ( - set(value) != {"schemaVersion", "jobId", "code", "message"} - or value.get("schemaVersion") != RENDER_SCHEMA_VERSION - or value.get("jobId") != expected_job_id - or not isinstance(value.get("code"), str) - or not isinstance(value.get("message"), str) - ): - raise RenderProtocolError("render error metadata is invalid") - code = value.get("code") + code = _validate_render_error(value, request) raise RenderProtocolError(f"render worker returned {code}") expected_fields = { "schemaVersion", "jobId", + "sequence", "sampleRate", "channels", "sampleFormat", @@ -394,24 +510,28 @@ def read_render_response( } if ( set(value) != expected_fields + or type(value.get("schemaVersion")) is not int or value.get("schemaVersion") != RENDER_SCHEMA_VERSION ): raise RenderProtocolError("render begin metadata is invalid") - if value.get("jobId") != expected_job_id: - raise RenderProtocolError("render response jobId is out of turn") + if ( + value.get("jobId") != request.job_id + or type(value.get("sequence")) is not int + or value.get("sequence") != request.sequence + ): + raise RenderProtocolError("render response identity is out of turn") frames = value.get("frames") pcm_bytes = value.get("pcmBytes") if ( - value.get("sampleRate") != RENDER_SAMPLE_RATE + type(value.get("sampleRate")) is not int + or value.get("sampleRate") != RENDER_SAMPLE_RATE + or type(value.get("channels")) is not int or value.get("channels") != RENDER_CHANNELS or value.get("sampleFormat") != "f32le" - or isinstance(frames, bool) - or not isinstance(frames, int) - or frames < 1 - or isinstance(pcm_bytes, bool) - or not isinstance(pcm_bytes, int) - or pcm_bytes != frames * RENDER_BYTES_PER_FRAME - or pcm_bytes > MAX_RENDER_PCM_BYTES + or type(frames) is not int + or frames != request.frames + or type(pcm_bytes) is not int + or pcm_bytes != request.pcm_bytes ): raise RenderProtocolError("render begin audio identity is invalid") @@ -423,6 +543,7 @@ def read_render_response( { FRAME_RENDER_CHUNK: RENDER_PCM_CHUNK_BYTES, FRAME_RENDER_END: MAX_RENDER_METADATA_BYTES, + FRAME_RENDER_ERROR: MAX_RENDER_METADATA_BYTES, }, ) if frame is None: @@ -437,12 +558,31 @@ def read_render_response( digest.update(payload) continue + if frame_type == FRAME_RENDER_ERROR: + code = _validate_render_error(_strict_json_object(payload), request) + raise RenderProtocolError(f"render worker returned {code}") + end = _strict_json_object(payload) if ( - set(end) != {"schemaVersion", "jobId", "pcmBytes", "sha256"} + set(end) + != { + "schemaVersion", + "jobId", + "sequence", + "frames", + "pcmBytes", + "sha256", + } + or type(end.get("schemaVersion")) is not int or end.get("schemaVersion") != RENDER_SCHEMA_VERSION - or end.get("jobId") != expected_job_id + or end.get("jobId") != request.job_id + or type(end.get("sequence")) is not int + or end.get("sequence") != request.sequence + or type(end.get("frames")) is not int + or end.get("frames") != request.frames + or type(end.get("pcmBytes")) is not int or end.get("pcmBytes") != pcm_bytes + or type(end.get("sha256")) is not str or end.get("sha256") != digest.hexdigest() or len(output) != pcm_bytes ): @@ -707,7 +847,7 @@ class _RenderCommandReader: def __init__(self, reader: BinaryIO) -> None: # Socket backpressure bounds requests that arrive faster than the single # render slot can consume them; no unbounded JSON queue exists. - self.commands: queue.Queue[_RenderCommand] = queue.Queue(maxsize=4) + self.commands: queue.Queue[_RenderCommand] = queue.Queue(maxsize=1) self._reader = reader self._thread = threading.Thread( target=self._pump, name="mrt2-render-control", daemon=True @@ -718,12 +858,12 @@ def _pump(self) -> None: while True: try: command = read_render_command(self._reader) - except (OSError, RenderProtocolError) as error: - message = ( - str(error) - if isinstance(error, RenderProtocolError) - else "render control connection failed" - ) + except RenderProtocolError as error: + message = str(error) + self.commands.put(_RenderReaderFailure(message[:512])) + return + except Exception: # noqa: BLE001 - never expose unexpected details + message = "render control connection failed" self.commands.put(_RenderReaderFailure(message[:512])) return self.commands.put(command) @@ -737,6 +877,64 @@ def _render_startup_error(error: Exception) -> str: return public_startup_error(error)[:512] +def _shutdown_render_socket(sock: socket.socket) -> None: + try: + sock.shutdown(socket.SHUT_RDWR) + except (AttributeError, OSError): + pass + + +def _write_render_frame_bounded( + sock: socket.socket, + frame_type: int, + payload: bytes, + *, + poll: Callable[[], None] | None = None, + timeout: float = RENDER_WRITE_TIMEOUT_SECONDS, +) -> None: + """Commit a frame without blocking the cancellation/control loop. + + ``sendall`` runs in a daemon because Python cannot portably make a socket's + send side nonblocking without also disturbing the concurrent buffered read + side. The caller remains live, polls control every 10 ms, and enforces one + absolute write deadline. Closing the socket before process termination + prevents a blocked writer from emitting late output in injected tests too. + """ + + result: queue.Queue[Exception | None] = queue.Queue(maxsize=1) + + def send() -> None: + try: + write_frame(sock, frame_type, payload) + except Exception as error: # noqa: BLE001 - sanitized below + result.put(error) + else: + result.put(None) + + if poll is not None: + poll() + threading.Thread( + target=send, + name=f"mrt2-render-write-{frame_type}", + daemon=True, + ).start() + deadline = time.monotonic() + timeout + while True: + if poll is not None: + poll() + remaining = deadline - time.monotonic() + if remaining <= 0: + _shutdown_render_socket(sock) + raise _RenderWriteError("render connection write timed out") + try: + error = result.get(timeout=min(RENDER_WRITE_POLL_SECONDS, remaining)) + except queue.Empty: + continue + if error is not None: + raise _RenderWriteError("render connection write failed") from None + return + + def run_render_worker( sock: socket.socket, model: str, @@ -752,6 +950,11 @@ def run_render_worker( be interrupted safely inside upstream PyTorch, so cancellation or EOF while it is active terminates this disposable process and releases its CUDA context; the native supervisor may start a fresh worker for the next job. + + The supervisor owns an outer startup/job deadline and must kill and reap the + full child tree on cancellation, disconnect, deadline, or owner drop, + including before ``render_ready``. A child has one connection and one + consumed auth token; it never reconnects or retries a sequence. """ try: @@ -764,124 +967,194 @@ def run_render_worker( if callable(warm_up): warm_up() except Exception as error: - write_render_error( + try: + write_render_error( + sock, + job_id=None, + sequence=0, + code="startup_failed", + message=_render_startup_error(error), + send_frame=lambda frame_type, payload: _write_render_frame_bounded( + sock, frame_type, payload, timeout=1.0 + ), + ) + except _RenderWriteError: + pass + return + + try: + _write_render_frame_bounded( sock, - job_id=None, - code="startup_failed", - message=_render_startup_error(error), + FRAME_STATUS, + _render_json( + { + "schemaVersion": RENDER_SCHEMA_VERSION, + "event": "render_ready", + "model": model, + "runtime": runtime, + "nextSequence": 1, + } + ), ) + except _RenderWriteError: return - - write_frame( - sock, - FRAME_STATUS, - _render_json( - { - "schemaVersion": RENDER_SCHEMA_VERSION, - "event": "render_ready", - "model": model, - "runtime": runtime, - } - ), - ) commands = _RenderCommandReader(sock.makefile("rb")) + expected_sequence: int | None = 1 - while True: - command = commands.commands.get() - if command is None: - return - if isinstance(command, _RenderReaderFailure): - write_render_error( - sock, - job_id=None, - code="protocol_error", - message=command.message, - ) - return - if isinstance(command, RenderCancel): + def stop(code: int) -> None: + _shutdown_render_socket(sock) + terminate(code) + raise _RenderWorkerStopped + + def send_error( + *, job_id: str | None, sequence: int, code: str, message: str + ) -> None: + try: write_render_error( sock, - job_id=command.job_id, - code="no_active_job", - message="render job is not active", + job_id=job_id, + sequence=sequence, + code=code, + message=message, + send_frame=lambda frame_type, payload: _write_render_frame_bounded( + sock, frame_type, payload, timeout=1.0 + ), ) - continue - - result: queue.Queue[tuple[bool, bytes | Exception]] = queue.Queue(maxsize=1) + except _RenderWriteError: + pass - def render() -> None: - try: - result.put((True, engine.render_clip(command.prompt, command.seconds))) - except Exception as error: # noqa: BLE001 - collapsed at the boundary - result.put((False, error)) + try: + while True: + command = commands.commands.get() + if command is None: + return + if isinstance(command, _RenderReaderFailure): + send_error( + job_id=None, + sequence=expected_sequence or MAX_U64, + code="protocol_error", + message=command.message, + ) + stop(2) + if isinstance(command, RenderCancel): + send_error( + job_id=command.job_id, + sequence=command.sequence, + code="no_active_job", + message="render job is not active", + ) + stop(2) + if expected_sequence is None or command.sequence != expected_sequence: + send_error( + job_id=command.job_id, + sequence=command.sequence, + code="sequence_error", + message="render sequence is duplicate or out of order", + ) + stop(2) + expected_sequence = ( + None if command.sequence == MAX_U64 else command.sequence + 1 + ) - threading.Thread( - target=render, - name=f"mrt2-render-{command.job_id[:16]}", - daemon=True, - ).start() + result: queue.Queue[tuple[bool, bytes | Exception]] = queue.Queue(maxsize=1) - while True: - try: - succeeded, value = result.get(timeout=0.025) - break - except queue.Empty: - pass - try: - pending = commands.commands.get_nowait() - except queue.Empty: - continue + def render() -> None: + try: + result.put( + (True, engine.render_clip(command.prompt, command.seconds)) + ) + except Exception as error: # noqa: BLE001 - boundary collapse + result.put((False, error)) + + threading.Thread( + target=render, + name=f"mrt2-render-{command.job_id[:16]}", + daemon=True, + ).start() + + def poll_active(*, can_reply: bool = True) -> None: + try: + pending = commands.commands.get_nowait() + except queue.Empty: + return + + if pending is None: + stop(0) + if ( + isinstance(pending, RenderCancel) + and pending.job_id == command.job_id + and pending.sequence == command.sequence + ): + if can_reply: + send_error( + job_id=command.job_id, + sequence=command.sequence, + code="cancelled", + message="render job was cancelled", + ) + stop(2) + message = ( + pending.message + if isinstance(pending, _RenderReaderFailure) + else "render command is overlapping or out of turn" + ) + if can_reply: + send_error( + job_id=command.job_id, + sequence=command.sequence, + code="protocol_error", + message=message, + ) + stop(2) + + # Control always gets first look. Once the result arrives, poll + # again before BEGIN so a cancel already queued behind it wins. + while True: + poll_active() + try: + succeeded, value = result.get(timeout=RENDER_WRITE_POLL_SECONDS) + break + except queue.Empty: + continue + poll_active() - if pending is None: - terminate(0) + if not succeeded: + send_error( + job_id=command.job_id, + sequence=command.sequence, + code="render_failed", + message="MRT2 render failed; the worker must be restarted", + ) return - if isinstance(pending, RenderCancel) and pending.job_id == command.job_id: - write_render_error( + try: + if not isinstance(value, bytes): + raise RenderProtocolError( + "render engine returned a non-bytes payload" + ) + write_render_response( sock, + command, + value, + before_frame=poll_active, + send_frame=lambda frame_type, payload: _write_render_frame_bounded( + sock, + frame_type, + payload, + poll=lambda: poll_active(can_reply=False), + ), + ) + except RenderProtocolError: + send_error( job_id=command.job_id, - code="cancelled", - message="render job was cancelled", + sequence=command.sequence, + code="invalid_audio", + message="MRT2 render returned an invalid PCM payload", ) - terminate(2) return - if isinstance(pending, _RenderReaderFailure): - message = pending.message - job_id = command.job_id - elif isinstance(pending, RenderRequest): - message = "render requests must not overlap" - job_id = pending.job_id - else: - message = "render cancellation is out of turn" - job_id = pending.job_id - write_render_error( - sock, - job_id=job_id, - code="protocol_error", - message=message, - ) - terminate(2) - return - - if not succeeded: - write_render_error( - sock, - job_id=command.job_id, - code="render_failed", - message="MRT2 render failed; the worker must be restarted", - ) - return - try: - if not isinstance(value, bytes): - raise RenderProtocolError("render engine returned a non-bytes payload") - write_render_response(sock, command, value) - except RenderProtocolError: - write_render_error( - sock, - job_id=command.job_id, - code="invalid_audio", - message="MRT2 render returned an invalid PCM payload", - ) - return + except _RenderWriteError: + stop(2) + except _RenderWorkerStopped: + return # --- Model tooling (the in-app model manager, issue #43) ------------------- @@ -1044,7 +1317,6 @@ def main(argv=None) -> None: sock = socket.create_connection(("127.0.0.1", args.port)) sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) authenticate_to_host(sock) - os.environ.pop(WORKER_TOKEN_ENV, None) run_render_worker(sock, args.model, runtime=args.runtime) return @@ -1062,7 +1334,6 @@ def main(argv=None) -> None: sock = socket.create_connection(("127.0.0.1", args.port)) sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) authenticate_to_host(sock) - os.environ.pop(WORKER_TOKEN_ENV, None) run_shared_sidecar( sock, (args.model_a, args.model_b), @@ -1082,7 +1353,6 @@ def main(argv=None) -> None: sock = socket.create_connection(("127.0.0.1", args.port)) sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) authenticate_to_host(sock) - os.environ.pop(WORKER_TOKEN_ENV, None) run_sidecar(sock, args.deck, args.model, runtime=args.runtime) diff --git a/backend/tests/test_render_sidecar.py b/backend/tests/test_render_sidecar.py index 22a8197..c3f4e79 100644 --- a/backend/tests/test_render_sidecar.py +++ b/backend/tests/test_render_sidecar.py @@ -17,6 +17,7 @@ FRAME_RENDER_ERROR, FRAME_RENDER_REQUEST, FRAME_STATUS, + MAX_RENDER_FRAMES, MAX_RENDER_PCM_BYTES, MAX_RENDER_PROMPT_CHARS, MAX_RENDER_REQUEST_BYTES, @@ -28,9 +29,11 @@ RENDER_SCHEMA_VERSION, RenderProtocolError, RenderRequest, + authenticate_to_host, read_frame, read_render_command, read_render_response, + render_frames_for_seconds, run_render_worker, write_frame, write_render_response, @@ -77,36 +80,72 @@ def render_clip(self, prompt, seconds): return b"\0" * (round(seconds * RENDER_SAMPLE_RATE) * RENDER_BYTES_PER_FRAME) -def request_payload(*, job_id=JOB_ID, prompt="bright piano", seconds=0.5, **extra): +class SignalAfterFirstChunkSock: + """Record frame writes and optionally signal after the first PCM chunk.""" + + def __init__(self, sock, signal=None): + self.sock = sock + self.signal = signal + self.signalled = False + self.frame_types = [] + + def sendall(self, data): + self.frame_types.append(data[0]) + self.sock.sendall(data) + if ( + data[0] == FRAME_RENDER_CHUNK + and not self.signalled + and self.signal is not None + ): + self.signalled = True + self.signal() + + def makefile(self, *args, **kwargs): + return self.sock.makefile(*args, **kwargs) + + def shutdown(self, how): + return self.sock.shutdown(how) + + +def request_payload( + *, job_id=JOB_ID, sequence=1, prompt="bright piano", frames=24_000, **extra +): return json.dumps( { "schemaVersion": RENDER_SCHEMA_VERSION, "jobId": job_id, + "sequence": sequence, "prompt": prompt, - "seconds": seconds, + "frames": frames, **extra, }, separators=(",", ":"), ).encode() -def cancel_payload(job_id=JOB_ID): +def cancel_payload(job_id=JOB_ID, sequence=1): return json.dumps( - {"schemaVersion": RENDER_SCHEMA_VERSION, "jobId": job_id}, + { + "schemaVersion": RENDER_SCHEMA_VERSION, + "jobId": job_id, + "sequence": sequence, + }, separators=(",", ":"), ).encode() -def begin_payload(*, job_id=JOB_ID, frames=1): +def begin_payload(*, job_id=JOB_ID, sequence=1, frames=24_000, **extra): return json.dumps( { "schemaVersion": RENDER_SCHEMA_VERSION, "jobId": job_id, + "sequence": sequence, "sampleRate": RENDER_SAMPLE_RATE, "channels": RENDER_CHANNELS, "sampleFormat": "f32le", "frames": frames, "pcmBytes": frames * RENDER_BYTES_PER_FRAME, + **extra, }, separators=(",", ":"), ).encode() @@ -124,22 +163,32 @@ def test_render_command_is_strict_and_bounded(): + request_payload() ) request = read_render_command(valid) - assert request == RenderRequest(JOB_ID, "bright piano", 0.5) - assert request.pcm_bytes == round(0.5 * RENDER_SAMPLE_RATE) * RENDER_BYTES_PER_FRAME + assert request == RenderRequest(JOB_ID, 1, "bright piano", 24_000) + assert request.seconds == 0.5 + assert request.pcm_bytes == 24_000 * RENDER_BYTES_PER_FRAME cancel = RecordingSock() write_frame(cancel, FRAME_RENDER_CANCEL, cancel_payload()) - assert read_render_command(io.BytesIO(cancel.buffer)).job_id == JOB_ID + assert read_render_command(io.BytesIO(cancel.buffer)).sequence == 1 invalid_payloads = [ request_payload(prompt=" "), request_payload(prompt="x" * (MAX_RENDER_PROMPT_CHARS + 1)), - request_payload(seconds=0.49), - request_payload(seconds=180.01), - request_payload(seconds=float("nan")), + request_payload(frames=23_999), + request_payload(frames=MAX_RENDER_FRAMES + 1), + request_payload(frames=True), + request_payload(frames=24_000.0), + request_payload(sequence=True), + request_payload(sequence=1.0), + request_payload(schemaVersion=True), request_payload(job_id="short"), request_payload(unexpected=True), - b'{"schemaVersion":1,"schemaVersion":1,"jobId":"render-job-0123456789abcdef","prompt":"p","seconds":1}', + b'{"schemaVersion":1,"schemaVersion":1,"jobId":"render-job-0123456789abcdef","sequence":1,"prompt":"p","frames":24000}', + ( + b'{"schemaVersion":1,"jobId":"render-job-0123456789abcdef",' + b'"sequence":1,"prompt":"p","frames":' + b"9" * 1000 + b"}" + ), + b"[" * 2000 + b"0" + b"]" * 2000, ] for payload in invalid_payloads: wire = RecordingSock() @@ -148,6 +197,12 @@ def test_render_command_is_strict_and_bounded(): read_render_command(io.BytesIO(wire.buffer)) +def test_render_frame_rounding_contract_uses_half_up_not_ties_to_even(): + half_frame = (24_000 + 0.5) / RENDER_SAMPLE_RATE + assert round(half_frame * RENDER_SAMPLE_RATE) == 24_000 + assert render_frames_for_seconds(half_frame) == 24_001 + + def test_render_command_rejects_truncation_oversize_and_out_of_order_frames(): with pytest.raises(RenderProtocolError, match="header is truncated"): read_render_command(io.BytesIO(b"\x06\x01")) @@ -174,14 +229,14 @@ def test_render_command_rejects_truncation_oversize_and_out_of_order_frames(): def test_render_response_round_trip_is_chunked_hashed_and_exact(): - request = RenderRequest(JOB_ID, "piano", 3.0) + request = RenderRequest(JOB_ID, 1, "piano", 3 * RENDER_SAMPLE_RATE) pcm = bytes(range(256)) * (request.pcm_bytes // 256) assert len(pcm) == request.pcm_bytes wire = RecordingSock() write_render_response(wire, request, pcm) reader = io.BytesIO(wire.buffer) - assert read_render_response(reader, JOB_ID, require_eof=True) == pcm + assert read_render_response(reader, request, require_eof=True) == pcm frame_reader = io.BytesIO(wire.buffer) frame_types = [] while frame := read_frame(frame_reader): @@ -196,7 +251,7 @@ def test_render_response_round_trip_is_chunked_hashed_and_exact(): @pytest.mark.parametrize("delta", [-RENDER_BYTES_PER_FRAME, RENDER_BYTES_PER_FRAME]) def test_render_response_writer_rejects_short_and_extra_pcm(delta): - request = RenderRequest(JOB_ID, "piano", 0.5) + request = RenderRequest(JOB_ID, 1, "piano", 24_000) with pytest.raises(RenderProtocolError, match="expected"): write_render_response( RecordingSock(), request, b"\0" * (request.pcm_bytes + delta) @@ -204,46 +259,75 @@ def test_render_response_writer_rejects_short_and_extra_pcm(delta): def test_render_response_reader_rejects_out_of_order_oversized_and_extra_pcm(): + request = RenderRequest(JOB_ID, 1, "piano", 24_000) out_of_order = RecordingSock() write_frame(out_of_order, FRAME_RENDER_CHUNK, b"\0" * RENDER_BYTES_PER_FRAME) with pytest.raises(RenderProtocolError, match="out of order"): - read_render_response(io.BytesIO(out_of_order.buffer), JOB_ID) + read_render_response(io.BytesIO(out_of_order.buffer), request) oversized = bytearray() begin = begin_payload() oversized.extend(struct.pack(" Date: Sat, 8 Aug 2026 19:00:22 -0700 Subject: [PATCH 03/14] test: lock render end identity contract --- backend/tests/test_render_sidecar.py | 35 ++++++++++++++++++++++++++++ 1 file changed, 35 insertions(+) diff --git a/backend/tests/test_render_sidecar.py b/backend/tests/test_render_sidecar.py index c3f4e79..553ffcd 100644 --- a/backend/tests/test_render_sidecar.py +++ b/backend/tests/test_render_sidecar.py @@ -1,5 +1,6 @@ """Bounded, authenticated protocol for the dedicated MRT2 render worker.""" +import hashlib import io import json import socket @@ -330,6 +331,40 @@ def test_render_response_rejects_coercible_scalar_types(override): read_render_response(io.BytesIO(wire.buffer), request) +@pytest.mark.parametrize( + "override", + [ + {"schemaVersion": 1.0}, + {"sequence": 2}, + {"frames": 24_001}, + {"pcmBytes": 192_000.0}, + {"sha256": "0" * 64}, + ], +) +def test_render_response_end_requires_exact_active_identity_total_and_hash(override): + request = RenderRequest(JOB_ID, 1, "piano", 24_000) + pcm = b"\0" * request.pcm_bytes + end = { + "schemaVersion": RENDER_SCHEMA_VERSION, + "jobId": request.job_id, + "sequence": request.sequence, + "frames": request.frames, + "pcmBytes": request.pcm_bytes, + "sha256": hashlib.sha256(pcm).hexdigest(), + **override, + } + wire = RecordingSock() + write_frame(wire, FRAME_RENDER_BEGIN, begin_payload()) + write_frame(wire, FRAME_RENDER_CHUNK, pcm) + write_frame( + wire, + FRAME_RENDER_END, + json.dumps(end, separators=(",", ":")).encode(), + ) + with pytest.raises(RenderProtocolError, match="end metadata"): + read_render_response(io.BytesIO(wire.buffer), request) + + def test_render_worker_reuses_one_warm_model_for_serial_requests(): shell, worker = socket.socketpair() engine = FakeRenderEngine() From af915fe50a6dc5c1f32736cd617114295d58140e Mon Sep 17 00:00:00 2001 From: brxs Date: Sat, 8 Aug 2026 19:24:13 -0700 Subject: [PATCH 04/14] feat: orchestrate managed native generation services --- src-tauri/src/analysis/live.rs | 2 +- src-tauri/src/generation.rs | 121 ++- src-tauri/src/lib.rs | 45 +- src-tauri/src/magenta_gateway.rs | 1405 ++++++++++++++++++++++++++++++ src-tauri/src/mcp.rs | 39 +- src-tauri/src/models.rs | 268 +++++- src-tauri/src/sidecar.rs | 187 +++- 7 files changed, 1963 insertions(+), 104 deletions(-) create mode 100644 src-tauri/src/magenta_gateway.rs diff --git a/src-tauri/src/analysis/live.rs b/src-tauri/src/analysis/live.rs index d14b79c..22f2ad9 100644 --- a/src-tauri/src/analysis/live.rs +++ b/src-tauri/src/analysis/live.rs @@ -133,7 +133,7 @@ pub struct AnalysisFeed { impl AnalysisFeed { /// A feed whose receivers are dropped — every send is a silent no-op. For /// tests that need the tee wiring without analysis threads (no `AppHandle`). - #[cfg(all(test, unix, not(feature = "managed-runtime")))] + #[cfg(test)] pub fn disconnected(deck_count: usize) -> Self { AnalysisFeed { senders: Arc::new((0..deck_count).map(|_| sync_channel(1).0).collect()), diff --git a/src-tauri/src/generation.rs b/src-tauri/src/generation.rs index 9f23053..a4d3b7b 100644 --- a/src-tauri/src/generation.rs +++ b/src-tauri/src/generation.rs @@ -2,10 +2,11 @@ //! //! The native shell hosts the realtime decks (the inference sidecars, [`crate::sidecar`]) //! and serves the frontend from the Tauri asset host, so FastAPI no longer serves -//! the UI. But the Stable Audio 3 / Magenta pad+track GENERATION still lives behind -//! HTTP (`/api/render`, `/api/generate`). This module spawns the FastAPI generation -//! server on a loopback port — the controller is generation-only: no deck workers, no -//! static mount — and the webview fetches it via `getApiBaseUrl()`. +//! the UI. This module supervises the Stable Audio 3 generation service on a +//! loopback port. In managed Linux/Windows builds Magenta rendering belongs to +//! the Rust gateway (`crate::magenta_gateway`), so this child receives no MRT2 +//! paths or dependencies. The bundled macOS backend retains its existing +//! combined `/api/render` + `/api/generate` behavior. //! //! Mirrors the sidecar's spawn/supervise/Drop-kill pattern. Started with the app; a //! failed spawn just leaves generation unreachable (the UI already surfaces those as @@ -25,9 +26,13 @@ use crate::child_process::{Readiness, SupervisedChild}; /// webview via `app_info`) and the child process. Held in Tauri managed state; /// dropping it kills the child. pub struct GenerationServer { + state: Mutex, +} + +struct GenerationState { port: Option, capability: Option, - child: Mutex>, + child: Option, } impl GenerationServer { @@ -35,25 +40,19 @@ impl GenerationServer { /// failed spawn yields `port() == None` and generation is simply unreachable (the /// webview surfaces that as fetch errors). pub fn start() -> GenerationServer { - let capability = crate::local_auth::generate_capability(); - match Self::spawn(&capability) { - Ok((port, child)) => { - println!("lsdj-app: generation server on 127.0.0.1:{port}"); - GenerationServer { - port: Some(port), - capability: Some(capability), - child: Mutex::new(Some(child)), - } - } - Err(e) => { - eprintln!("lsdj-app: generation server spawn failed: {e}"); - GenerationServer { - port: None, - capability: None, - child: Mutex::new(None), - } - } + let server = GenerationServer { + state: Mutex::new(GenerationState { + port: None, + capability: None, + child: None, + }), + }; + if let Err(error) = server.resume() { + // A fresh managed install intentionally has no runtime yet. The + // model manager calls `resume` immediately after first promotion. + eprintln!("lsdj-app: generation server unavailable: {error}"); } + server } fn spawn(capability: &str) -> io::Result<(u16, SupervisedChild)> { @@ -89,12 +88,62 @@ impl GenerationServer { /// The loopback port the generation server bound, or `None` if disabled / not /// running. The webview reads this through `app_info` to build the API base URL. pub fn port(&self) -> Option { - self.port + self.state + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + .port } /// The in-memory capability paired with [`port`](Self::port). Never persisted. pub fn capability(&self) -> Option { - self.capability.clone() + self.state + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + .capability + .clone() + } + + /// Start (or recover) the service from the currently promoted verified + /// generation. A running healthy child is left untouched. This is called on + /// startup and after every managed SA3 promotion/rollback. + pub fn resume(&self) -> io::Result<()> { + let mut state = self + .state + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + if let Some(child) = state.child.as_mut() { + if child.try_wait()?.is_none() { + return Ok(()); + } + state.child = None; + state.port = None; + state.capability = None; + } + let capability = crate::local_auth::generate_capability(); + let (port, child) = Self::spawn(&capability)?; + println!("lsdj-app: generation server on 127.0.0.1:{port}"); + state.port = Some(port); + state.capability = Some(capability); + state.child = Some(child); + Ok(()) + } + + /// Stop and reap the service before its managed generation is renamed. + /// Returns whether a live child was present so tests/lifecycle diagnostics + /// can distinguish first install from an update. + pub fn quiesce(&self) -> io::Result { + let mut state = self + .state + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + state.port = None; + state.capability = None; + let Some(mut child) = state.child.take() else { + return Ok(false); + }; + let report = child.shutdown(Duration::from_millis(500))?; + crate::child_process::log_shutdown("generation server", Ok(report)); + Ok(true) } /// Kill the generation server child. Called explicitly from the app's @@ -102,11 +151,8 @@ impl GenerationServer { /// macOS quit (`process::exit` skips destructors), so [`Drop`] alone would /// leak the process. pub fn shutdown(&self) { - if let Some(mut child) = self.child.lock().unwrap_or_else(|p| p.into_inner()).take() { - crate::child_process::log_shutdown( - "generation server", - child.shutdown(Duration::from_millis(500)), - ); + if let Err(error) = self.quiesce() { + crate::child_process::log_shutdown("generation server", Err(error)); } } } @@ -127,10 +173,17 @@ pub fn generation_command(port: u16, capability: &str) -> io::Result { use std::ffi::OsString; let paths = crate::platform_paths::get(); - let ephemeral = paths.backend_env().into_iter().chain(std::iter::once(( - OsString::from("LSDJ_API_CAPABILITY"), - OsString::from(capability), - ))); + let ephemeral = paths + .backend_env() + .into_iter() + // The managed SA3 interpreter has its own dependency closure and + // receives no MRT2 location. This makes accidental `/api/render` + // use fail closed instead of coupling the services again. + .filter(|(name, _)| name.to_str() != Some("MAGENTA_HOME")) + .chain(std::iter::once(( + OsString::from("LSDJ_API_CAPABILITY"), + OsString::from(capability), + ))); crate::managed_runtime::resolve( paths.assets(), crate::managed_runtime::Service::Sa3, diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 6e27734..850b9e2 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -46,6 +46,8 @@ mod generation; mod library; mod local_auth; mod loras; +#[cfg(feature = "managed-runtime")] +mod magenta_gateway; #[cfg_attr(not(feature = "managed-runtime"), allow(dead_code))] mod managed_runtime; mod mcp; @@ -251,22 +253,20 @@ fn start_sidecars( sidecar_status_sink(app.clone(), 0, feed.clone()), sidecar_status_sink(app.clone(), 1, feed.clone()), ]; - return match sidecar::SharedSidecar::spawn( + let mut shared = sidecar::SharedSidecar::parked( models, handles, sinks, taps.clone(), feed.clone(), - ) { - Ok(shared) => (sidecar::Sidecars::new_shared(shared), Vec::new()), - Err((error, handles)) => { - eprintln!("lsdj-app: shared MRT2 sidecar spawn failed: {error}"); - ( - sidecar::Sidecars::new((0..lsdj_engine::DECK_COUNT).map(|_| None).collect()), - handles.into_iter().collect(), - ) - } - }; + ); + if let Err(error) = shared.activate() { + // A fresh managed install has no runtime yet. Keep both permanent + // DeckHandles parked inside the supervisor; the first MRT2 install + // activates this same topology without restarting the app. + eprintln!("lsdj-app: shared MRT2 sidecar unavailable: {error}"); + } + return (sidecar::Sidecars::new_shared(shared), Vec::new()); } let mut decks = Vec::new(); @@ -305,6 +305,11 @@ struct AppInfo { /// Per-launch bearer capability for the generation service. It exists only in /// Rust state/the webview process and is never written to settings or logs. generation_capability: Option, + /// Managed Linux/Windows route Magenta through a separate Rust-owned + /// gateway. On bundled macOS these mirror the combined generation service, + /// preserving its established behavior. + magenta_port: Option, + magenta_capability: Option, /// The loopback port the MCP server bound (`None` only if the loopback bind /// failed — the server is otherwise always on), and the bearer token a client must /// present (ADR-0020 Phase 2). Surfaced so the client config can point at @@ -315,15 +320,27 @@ struct AppInfo { #[tauri::command] fn app_info( + app: tauri::AppHandle, state: tauri::State<'_, AudioState>, generation: tauri::State<'_, generation::GenerationServer>, mcp: tauri::State<'_, mcp::McpServer>, ) -> AppInfo { + #[cfg(feature = "managed-runtime")] + let (magenta_port, magenta_capability) = { + let gateway = app.state::(); + (gateway.port(), gateway.capability()) + }; + #[cfg(not(feature = "managed-runtime"))] + let (magenta_port, magenta_capability) = (generation.port(), generation.capability()); + #[cfg(not(feature = "managed-runtime"))] + let _ = app; AppInfo { version: env!("CARGO_PKG_VERSION").to_string(), audio_device_started: state.device_started, generation_port: generation.port(), generation_capability: generation.capability(), + magenta_port, + magenta_capability, mcp_port: mcp.port(), mcp_token: mcp.token(), } @@ -618,6 +635,8 @@ pub fn run() { // The sa3/Magenta generation server (gap 2): the gen-only FastAPI on a // loopback port the webview fetches; started with the app. let generation_server = generation::GenerationServer::start(); + #[cfg(feature = "managed-runtime")] + let magenta_gateway = magenta_gateway::MagentaGateway::start(); // The generated-songs library: the durable-data root from the platform // contract plus a JSON registry the take list restores from. On macOS // this remains Documents/LSDJ. Auto-save / list / load / delete all go @@ -795,6 +814,8 @@ pub fn run() { app.manage(analysis_feed); app.manage(analysis::track::TrackAnalysis::new(lsdj_engine::DECK_COUNT)); app.manage(generation_server); + #[cfg(feature = "managed-runtime")] + app.manage(magenta_gateway); // The native MCP server (ADR-0020 Phase 2): an external agent as a // co-DJ. Always on, loopback-only, token-guarded; its tools mutate the // same managed state the IPC commands do. Reaches that state through the @@ -918,6 +939,8 @@ pub fn run() { if let tauri::RunEvent::Exit = event { use tauri::Manager; app.state::().shutdown(); + #[cfg(feature = "managed-runtime")] + app.state::().shutdown(); app.state::().shutdown(); app.state::().shutdown(); app.state::().shutdown(); diff --git a/src-tauri/src/magenta_gateway.rs b/src-tauri/src/magenta_gateway.rs new file mode 100644 index 0000000..2e4beeb --- /dev/null +++ b/src-tauri/src/magenta_gateway.rs @@ -0,0 +1,1405 @@ +//! Native, authenticated Magenta render gateway for managed Linux/Windows. +//! +//! The public loopback HTTP service is deliberately separate from Stable Audio +//! 3. It owns one lazy, warm, disposable MRT2 render worker and translates the +//! user-facing `{prompt, seconds}` request into the reviewed binary protocol's +//! authoritative integer frame count and monotonic sequence. Every response is +//! bounded, sequence-bound, byte-counted, and SHA-256 checked before it becomes +//! a WAV. A cancellation, deadline, dropped HTTP request, or protocol mismatch +//! tears down and reaps the complete worker process tree; the next request starts +//! from a freshly revalidated managed generation. + +use std::collections::BTreeMap; +use std::io::{self, Read, Write}; +use std::net::{Shutdown, TcpListener, TcpStream}; +use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; +use std::sync::{Arc, Mutex}; +use std::time::{Duration, Instant}; + +use axum::body::Bytes; +use axum::extract::{DefaultBodyLimit, Request, State}; +use axum::http::{header, HeaderValue, Method, StatusCode}; +use axum::middleware::Next; +use axum::response::{IntoResponse, Response}; +use axum::routing::{get, post}; +use axum::Router; +use serde::{Deserialize, Serialize}; +use sha2::{Digest, Sha256}; +use tokio_util::sync::CancellationToken; + +use crate::child_process::SupervisedChild; + +const FRAME_STATUS: u8 = 2; +const FRAME_AUTH: u8 = 5; +const FRAME_RENDER_REQUEST: u8 = 6; +const FRAME_RENDER_BEGIN: u8 = 7; +const FRAME_RENDER_CHUNK: u8 = 8; +const FRAME_RENDER_END: u8 = 9; +const FRAME_RENDER_CANCEL: u8 = 10; +const FRAME_RENDER_ERROR: u8 = 11; + +const RENDER_SCHEMA_VERSION: u32 = 1; +const RENDER_SAMPLE_RATE: u64 = 48_000; +const RENDER_CHANNELS: u64 = 2; +const RENDER_BYTES_PER_FRAME: u64 = RENDER_CHANNELS * 4; +const MIN_RENDER_FRAMES: u64 = 24_000; +const MAX_RENDER_FRAMES: u64 = 8_640_000; +const MAX_RENDER_PCM_BYTES: usize = (MAX_RENDER_FRAMES * RENDER_BYTES_PER_FRAME) as usize; +const MAX_RENDER_REQUEST_BYTES: usize = 64 * 1024; +const MAX_RENDER_PROMPT_CHARS: usize = 32_000; +const MAX_RENDER_CONTROL_BYTES: usize = 1024; +const MAX_RENDER_METADATA_BYTES: usize = 8 * 1024; +const MAX_RENDER_CHUNK_BYTES: usize = 1024 * 1024; +const ACCEPT_TIMEOUT: Duration = Duration::from_secs(30); +const READY_TIMEOUT: Duration = Duration::from_secs(180); +const IO_POLL: Duration = Duration::from_millis(50); +const WRITE_TIMEOUT: Duration = Duration::from_secs(5); +const SAFE_ORIGINS: &[&str] = &[ + "tauri://localhost", + "http://tauri.localhost", + "https://tauri.localhost", +]; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum FailureKind { + Unavailable, + Protocol, + Deadline, + Cancelled, +} + +#[derive(Debug)] +struct RenderFailure { + kind: FailureKind, + detail: &'static str, +} + +impl RenderFailure { + fn unavailable() -> Self { + Self { + kind: FailureKind::Unavailable, + detail: "Magenta runtime is not installed or failed verification", + } + } + + fn protocol(detail: &'static str) -> Self { + Self { + kind: FailureKind::Protocol, + detail, + } + } + + fn deadline() -> Self { + Self { + kind: FailureKind::Deadline, + detail: "Magenta render timed out", + } + } + + fn cancelled() -> Self { + Self { + kind: FailureKind::Cancelled, + detail: "Magenta render was cancelled", + } + } +} + +impl From for RenderFailure { + fn from(_: io::Error) -> Self { + Self::protocol("Magenta render worker connection failed") + } +} + +#[derive(Clone)] +struct RequestCancellation { + request: Arc, + lifecycle: Arc, +} + +impl RequestCancellation { + fn cancelled(&self) -> bool { + self.request.load(Ordering::Acquire) || self.lifecycle.load(Ordering::Acquire) + } +} + +struct CancelOnDrop { + flag: Arc, + armed: bool, +} + +impl CancelOnDrop { + fn new(flag: Arc) -> Self { + Self { flag, armed: true } + } + + fn disarm(&mut self) { + self.armed = false; + } +} + +impl Drop for CancelOnDrop { + fn drop(&mut self) { + if self.armed { + self.flag.store(true, Ordering::Release); + } + } +} + +trait ProcessTree: Send { + fn shutdown(&mut self) -> io::Result<()>; +} + +struct ManagedProcess { + child: SupervisedChild, +} + +impl ProcessTree for ManagedProcess { + fn shutdown(&mut self) -> io::Result<()> { + let report = self.child.shutdown(Duration::from_millis(500))?; + crate::child_process::log_shutdown("MRT2 render worker", Ok(report)); + Ok(()) + } +} + +struct ManagedRenderWorker { + stream: TcpStream, + process: Box, +} + +impl ManagedRenderWorker { + fn render( + &mut self, + request: &WorkerRenderRequest, + cancellation: &RequestCancellation, + ) -> Result, RenderFailure> { + let payload = serde_json::to_vec(request) + .map_err(|_| RenderFailure::protocol("Magenta render request is invalid"))?; + if payload.len() > MAX_RENDER_REQUEST_BYTES { + return Err(RenderFailure::protocol( + "Magenta render request is too large", + )); + } + self.stream.set_write_timeout(Some(WRITE_TIMEOUT))?; + write_frame(&mut self.stream, FRAME_RENDER_REQUEST, &payload)?; + let duration = request.frames as f64 / RENDER_SAMPLE_RATE as f64; + let deadline = Instant::now() + Duration::from_secs_f64((duration * 2.0).max(90.0)); + match read_render_response(&mut self.stream, request, cancellation, deadline) { + Err(error) if error.kind == FailureKind::Cancelled => { + let cancel = WorkerRenderCancel { + schema_version: RENDER_SCHEMA_VERSION, + job_id: request.job_id.clone(), + sequence: request.sequence, + }; + if let Ok(payload) = serde_json::to_vec(&cancel) { + if payload.len() <= MAX_RENDER_CONTROL_BYTES { + let _ = write_frame(&mut self.stream, FRAME_RENDER_CANCEL, &payload); + } + } + Err(error) + } + result => result, + } + } + + fn shutdown(&mut self) -> io::Result<()> { + let _ = self.stream.shutdown(Shutdown::Both); + self.process.shutdown() + } +} + +trait WorkerFactory: Send + Sync { + fn spawn( + &self, + cancellation: &RequestCancellation, + ) -> Result; +} + +struct ManagedWorkerFactory; + +impl WorkerFactory for ManagedWorkerFactory { + fn spawn( + &self, + cancellation: &RequestCancellation, + ) -> Result { + spawn_managed_worker(cancellation) + } +} + +struct GatewayCore { + worker: Mutex>, + factory: Arc, + next_sequence: AtomicU64, + lifecycle: Mutex>, + quiescing: AtomicBool, +} + +impl GatewayCore { + fn new(factory: Arc) -> Self { + Self { + worker: Mutex::new(None), + factory, + next_sequence: AtomicU64::new(1), + lifecycle: Mutex::new(Arc::new(AtomicBool::new(false))), + quiescing: AtomicBool::new(false), + } + } + + fn cancellation(&self, request: Arc) -> RequestCancellation { + RequestCancellation { + request, + lifecycle: self + .lifecycle + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + .clone(), + } + } + + fn sequence(&self) -> Result { + self.next_sequence + .fetch_update(Ordering::AcqRel, Ordering::Acquire, |value| { + (value < u64::MAX).then_some(value + 1) + }) + .map_err(|_| RenderFailure::protocol("Magenta render sequence is exhausted")) + } + + fn render( + &self, + prompt: String, + frames: u64, + request_cancel: Arc, + ) -> Result, RenderFailure> { + if self.quiescing.load(Ordering::Acquire) { + return Err(RenderFailure::unavailable()); + } + let cancellation = self.cancellation(request_cancel); + if cancellation.cancelled() { + return Err(RenderFailure::cancelled()); + } + let sequence = self.sequence()?; + let request = WorkerRenderRequest { + schema_version: RENDER_SCHEMA_VERSION, + job_id: format!("render-{:032x}", rand::random::()), + sequence, + prompt, + frames, + }; + let mut worker = self + .worker + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + if cancellation.cancelled() || self.quiescing.load(Ordering::Acquire) { + return Err(RenderFailure::cancelled()); + } + if worker.is_none() { + *worker = Some(self.factory.spawn(&cancellation)?); + } + let result = worker + .as_mut() + .expect("worker was installed") + .render(&request, &cancellation); + if result.is_err() { + if let Some(mut failed) = worker.take() { + if failed.shutdown().is_err() { + // Keep ownership so a later quiesce can retry and, most + // importantly, an installer cannot mistake an uncertain + // process-tree state for "reaped" before a Windows rename. + *worker = Some(failed); + return Err(RenderFailure::protocol( + "Magenta render worker could not be reaped", + )); + } + } + } + result + } + + /// Cancel in-flight/queued renders, then kill and reap the warm worker. + /// Returns whether a worker was resident before the quiesce. + fn quiesce(&self) -> Result { + self.quiescing.store(true, Ordering::Release); + self.lifecycle + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + .store(true, Ordering::Release); + let mut worker = self + .worker + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + let Some(mut resident) = worker.take() else { + return Ok(false); + }; + match resident.shutdown() { + Ok(()) => Ok(true), + Err(_) => { + *worker = Some(resident); + Err("Magenta render worker could not be reaped".to_string()) + } + } + } + + /// Open a fresh request generation. If an update displaced a previously warm + /// renderer, eagerly restore it from the now-current verified generation. + fn resume(&self, restore_warm_worker: bool) -> Result<(), String> { + *self + .lifecycle + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) = Arc::new(AtomicBool::new(false)); + self.quiescing.store(false, Ordering::Release); + if !restore_warm_worker { + return Ok(()); + } + let cancellation = self.cancellation(Arc::new(AtomicBool::new(false))); + let mut worker = self + .worker + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + if worker.is_none() { + *worker = Some( + self.factory + .spawn(&cancellation) + .map_err(|error| error.detail.to_string())?, + ); + } + Ok(()) + } +} + +#[derive(Clone)] +struct HttpState { + core: Arc, +} + +#[derive(Clone)] +struct AuthState { + capability: Arc, +} + +/// The always-available public HTTP gateway. Absence of an MRT2 runtime affects +/// only render requests; it never prevents the window, model manager, or model +/// status endpoint from starting. +pub struct MagentaGateway { + port: Option, + capability: String, + cancel: CancellationToken, + core: Arc, +} + +impl MagentaGateway { + pub fn start() -> Self { + let capability = crate::local_auth::generate_capability(); + let core = Arc::new(GatewayCore::new(Arc::new(ManagedWorkerFactory))); + match bind_loopback() { + Ok((listener, port)) => { + let cancel = serve(listener, port, &capability, core.clone()); + Self { + port: Some(port), + capability, + cancel, + core, + } + } + Err(error) => { + eprintln!("lsdj-app: Magenta gateway bind failed: {error}"); + Self { + port: None, + capability, + cancel: CancellationToken::new(), + core, + } + } + } + } + + pub fn port(&self) -> Option { + self.port + } + + pub fn capability(&self) -> Option { + self.port.map(|_| self.capability.clone()) + } + + pub fn quiesce(&self) -> Result { + self.core.quiesce() + } + + pub fn resume(&self, restore_warm_worker: bool) -> Result<(), String> { + self.core.resume(restore_warm_worker) + } + + pub fn shutdown(&self) { + self.cancel.cancel(); + if let Err(error) = self.core.quiesce() { + eprintln!("lsdj-app: Magenta gateway shutdown failed: {error}"); + } + } +} + +impl Drop for MagentaGateway { + fn drop(&mut self) { + self.shutdown(); + } +} + +#[derive(Deserialize)] +#[serde(deny_unknown_fields)] +struct HttpRenderRequest { + prompt: String, + seconds: f64, +} + +#[derive(Serialize)] +#[serde(rename_all = "camelCase")] +struct WorkerRenderRequest { + schema_version: u32, + job_id: String, + sequence: u64, + prompt: String, + frames: u64, +} + +#[derive(Serialize)] +#[serde(rename_all = "camelCase")] +struct WorkerRenderCancel { + schema_version: u32, + job_id: String, + sequence: u64, +} + +#[derive(Deserialize)] +#[serde(rename_all = "camelCase", deny_unknown_fields)] +struct RenderReady { + schema_version: u32, + event: String, + model: String, + runtime: String, +} + +#[derive(Deserialize)] +#[serde(rename_all = "camelCase", deny_unknown_fields)] +struct RenderBegin { + schema_version: u32, + job_id: String, + sequence: u64, + sample_rate: u64, + channels: u64, + sample_format: String, + frames: u64, + pcm_bytes: u64, +} + +#[derive(Deserialize)] +#[serde(rename_all = "camelCase", deny_unknown_fields)] +struct RenderEnd { + schema_version: u32, + job_id: String, + sequence: u64, + frames: u64, + pcm_bytes: u64, + sha256: String, +} + +#[derive(Deserialize)] +#[serde(rename_all = "camelCase", deny_unknown_fields)] +struct RenderError { + schema_version: u32, + job_id: Option, + sequence: u64, + code: String, + message: String, +} + +async fn render_clip(State(state): State, body: Bytes) -> Response { + if body.len() > MAX_RENDER_REQUEST_BYTES { + return json_error(StatusCode::PAYLOAD_TOO_LARGE, "request body is too large"); + } + let parsed: HttpRenderRequest = match serde_json::from_slice(&body) { + Ok(parsed) => parsed, + Err(_) => return json_error(StatusCode::UNPROCESSABLE_ENTITY, "body must be JSON"), + }; + let prompt = parsed.prompt.trim().to_string(); + if prompt.is_empty() { + return json_error( + StatusCode::UNPROCESSABLE_ENTITY, + "'prompt' must be a non-empty string", + ); + } + if prompt.chars().count() > MAX_RENDER_PROMPT_CHARS { + return json_error( + StatusCode::UNPROCESSABLE_ENTITY, + "'prompt' must be at most 32000 characters", + ); + } + let frames = match frames_for_seconds(parsed.seconds) { + Some(frames) => frames, + None => { + return json_error( + StatusCode::UNPROCESSABLE_ENTITY, + "'seconds' must be 0.5-180", + ) + } + }; + + let request_cancel = Arc::new(AtomicBool::new(false)); + let mut drop_guard = CancelOnDrop::new(request_cancel.clone()); + let core = state.core.clone(); + let result = + tauri::async_runtime::spawn_blocking(move || core.render(prompt, frames, request_cancel)) + .await; + drop_guard.disarm(); + match result { + Ok(Ok(pcm)) => match float32_wav(&pcm) { + Ok(wav) => (StatusCode::OK, [(header::CONTENT_TYPE, "audio/wav")], wav).into_response(), + Err(_) => json_error(StatusCode::BAD_GATEWAY, "Magenta returned invalid audio"), + }, + Ok(Err(error)) => failure_response(error), + Err(_) => json_error(StatusCode::BAD_GATEWAY, "Magenta render task failed"), + } +} + +async fn model_info() -> Response { + let mut estimates = BTreeMap::new(); + estimates.insert("mrt2_small", 2.0); + estimates.insert("mrt2_base", 6.0); + axum::Json(serde_json::json!({ + "models": crate::models::magenta_models_for_gateway(), + "sample_rate": RENDER_SAMPLE_RATE, + "channels": RENDER_CHANNELS, + "chunk_seconds": 1.0, + "total_ram_gb": total_ram_gb(), + "model_ram_estimate_gb": estimates, + })) + .into_response() +} + +fn failure_response(error: RenderFailure) -> Response { + let status = match error.kind { + FailureKind::Unavailable => StatusCode::SERVICE_UNAVAILABLE, + FailureKind::Protocol => StatusCode::BAD_GATEWAY, + FailureKind::Deadline => StatusCode::GATEWAY_TIMEOUT, + FailureKind::Cancelled => StatusCode::from_u16(499).unwrap_or(StatusCode::BAD_GATEWAY), + }; + json_error(status, error.detail) +} + +fn json_error(status: StatusCode, detail: &str) -> Response { + (status, axum::Json(serde_json::json!({ "detail": detail }))).into_response() +} + +fn frames_for_seconds(seconds: f64) -> Option { + if !seconds.is_finite() || !(0.5..=180.0).contains(&seconds) { + return None; + } + let frames = (seconds * RENDER_SAMPLE_RATE as f64 + 0.5).floor() as u64; + (MIN_RENDER_FRAMES..=MAX_RENDER_FRAMES) + .contains(&frames) + .then_some(frames) +} + +fn float32_wav(pcm: &[u8]) -> Result, ()> { + if pcm.len() > MAX_RENDER_PCM_BYTES + || !pcm.len().is_multiple_of(RENDER_BYTES_PER_FRAME as usize) + { + return Err(()); + } + let data_len = u32::try_from(pcm.len()).map_err(|_| ())?; + let riff_len = 36u32.checked_add(data_len).ok_or(())?; + let mut wav = Vec::with_capacity(44 + pcm.len()); + wav.extend_from_slice(b"RIFF"); + wav.extend_from_slice(&riff_len.to_le_bytes()); + wav.extend_from_slice(b"WAVEfmt "); + wav.extend_from_slice(&16u32.to_le_bytes()); + wav.extend_from_slice(&3u16.to_le_bytes()); + wav.extend_from_slice(&(RENDER_CHANNELS as u16).to_le_bytes()); + wav.extend_from_slice(&(RENDER_SAMPLE_RATE as u32).to_le_bytes()); + wav.extend_from_slice(&((RENDER_SAMPLE_RATE * RENDER_BYTES_PER_FRAME) as u32).to_le_bytes()); + wav.extend_from_slice(&(RENDER_BYTES_PER_FRAME as u16).to_le_bytes()); + wav.extend_from_slice(&32u16.to_le_bytes()); + wav.extend_from_slice(b"data"); + wav.extend_from_slice(&data_len.to_le_bytes()); + wav.extend_from_slice(pcm); + Ok(wav) +} + +fn spawn_managed_worker( + cancellation: &RequestCancellation, +) -> Result { + let listener = TcpListener::bind("127.0.0.1:0").map_err(|_| RenderFailure::unavailable())?; + listener + .set_nonblocking(true) + .map_err(|_| RenderFailure::unavailable())?; + let port = listener + .local_addr() + .map_err(|_| RenderFailure::unavailable())? + .port(); + let token = crate::local_auth::generate_capability(); + let mut command = + crate::sidecar::authenticated_render_worker_command(crate::DEFAULT_MODEL, port, &token) + .map_err(|_| RenderFailure::unavailable())?; + let mut child = crate::child_process::spawn_grouped(&mut command) + .map_err(|_| RenderFailure::unavailable())?; + + let result = + accept_worker(&listener, &mut child, &token, cancellation).and_then(|mut stream| { + stream.set_nodelay(true).ok(); + let ready_deadline = Instant::now() + READY_TIMEOUT; + let (frame_type, payload) = read_bounded_frame( + &mut stream, + &[FRAME_STATUS, FRAME_RENDER_ERROR], + MAX_RENDER_METADATA_BYTES, + cancellation, + ready_deadline, + )?; + if frame_type == FRAME_RENDER_ERROR { + validate_startup_error(&payload)?; + return Err(RenderFailure::unavailable()); + } + let ready: RenderReady = serde_json::from_slice(&payload) + .map_err(|_| RenderFailure::protocol("Magenta worker readiness is invalid"))?; + if ready.schema_version != RENDER_SCHEMA_VERSION + || ready.event != "render_ready" + || ready.model != crate::DEFAULT_MODEL + || ready.runtime != "pytorch-cuda" + { + return Err(RenderFailure::protocol( + "Magenta worker readiness is invalid", + )); + } + Ok(stream) + }); + match result { + Ok(stream) => Ok(ManagedRenderWorker { + stream, + process: Box::new(ManagedProcess { child }), + }), + Err(error) => { + let _ = child.force_kill(); + Err(error) + } + } +} + +fn accept_worker( + listener: &TcpListener, + child: &mut SupervisedChild, + token: &str, + cancellation: &RequestCancellation, +) -> Result { + let deadline = Instant::now() + ACCEPT_TIMEOUT; + loop { + if cancellation.cancelled() { + return Err(RenderFailure::cancelled()); + } + if Instant::now() >= deadline { + return Err(RenderFailure::deadline()); + } + if child + .try_wait() + .map_err(|_| RenderFailure::unavailable())? + .is_some() + { + return Err(RenderFailure::unavailable()); + } + match listener.accept() { + Ok((mut stream, _)) => { + // The launch token is single-use: the first connection attempt + // consumes it, even when authentication fails. + stream + .set_nonblocking(false) + .map_err(|_| RenderFailure::protocol("Magenta worker connection failed"))?; + stream + .set_read_timeout(Some(IO_POLL)) + .map_err(|_| RenderFailure::protocol("Magenta worker connection failed"))?; + let (frame_type, payload) = read_bounded_frame( + &mut stream, + &[FRAME_AUTH], + 256, + cancellation, + deadline.min(Instant::now() + Duration::from_secs(1)), + )?; + if frame_type != FRAME_AUTH + || !(32..=256).contains(&payload.len()) + || !crate::local_auth::constant_time_eq(&payload, token.as_bytes()) + { + return Err(RenderFailure::protocol( + "Magenta worker authentication failed", + )); + } + return Ok(stream); + } + Err(error) if error.kind() == io::ErrorKind::WouldBlock => { + std::thread::sleep(IO_POLL); + } + Err(_) => return Err(RenderFailure::unavailable()), + } + } +} + +fn validate_startup_error(payload: &[u8]) -> Result<(), RenderFailure> { + let error: RenderError = serde_json::from_slice(payload) + .map_err(|_| RenderFailure::protocol("Magenta worker error is invalid"))?; + if error.schema_version != RENDER_SCHEMA_VERSION + || error.job_id.is_some() + || error.sequence != 0 + || error.code.is_empty() + || error.code.len() > 64 + || error.message.len() > 512 + { + return Err(RenderFailure::protocol("Magenta worker error is invalid")); + } + Ok(()) +} + +fn validate_render_error( + payload: &[u8], + request: &WorkerRenderRequest, +) -> Result<(), RenderFailure> { + let error: RenderError = serde_json::from_slice(payload) + .map_err(|_| RenderFailure::protocol("Magenta worker error is invalid"))?; + if error.schema_version != RENDER_SCHEMA_VERSION + || error.job_id.as_deref() != Some(&request.job_id) + || error.sequence != request.sequence + || error.code.is_empty() + || error.code.len() > 64 + || error.message.len() > 512 + { + return Err(RenderFailure::protocol("Magenta worker error is invalid")); + } + Err(RenderFailure::protocol("Magenta render worker failed")) +} + +fn read_render_response( + stream: &mut TcpStream, + request: &WorkerRenderRequest, + cancellation: &RequestCancellation, + deadline: Instant, +) -> Result, RenderFailure> { + let (frame_type, payload) = read_bounded_frame( + stream, + &[FRAME_RENDER_BEGIN, FRAME_RENDER_ERROR], + MAX_RENDER_METADATA_BYTES, + cancellation, + deadline, + )?; + if frame_type == FRAME_RENDER_ERROR { + validate_render_error(&payload, request)?; + unreachable!("a valid worker error is returned as a render failure"); + } + let begin: RenderBegin = serde_json::from_slice(&payload) + .map_err(|_| RenderFailure::protocol("Magenta render begin is invalid"))?; + let expected_bytes = request + .frames + .checked_mul(RENDER_BYTES_PER_FRAME) + .ok_or_else(|| RenderFailure::protocol("Magenta render size overflow"))?; + if begin.schema_version != RENDER_SCHEMA_VERSION + || begin.job_id != request.job_id + || begin.sequence != request.sequence + || begin.sample_rate != RENDER_SAMPLE_RATE + || begin.channels != RENDER_CHANNELS + || begin.sample_format != "f32le" + || begin.frames != request.frames + || begin.pcm_bytes != expected_bytes + || begin.pcm_bytes as usize > MAX_RENDER_PCM_BYTES + { + return Err(RenderFailure::protocol("Magenta render begin is invalid")); + } + + let mut pcm = Vec::with_capacity(expected_bytes as usize); + let mut digest = Sha256::new(); + loop { + let (frame_type, payload) = read_bounded_frame( + stream, + &[FRAME_RENDER_CHUNK, FRAME_RENDER_END, FRAME_RENDER_ERROR], + MAX_RENDER_CHUNK_BYTES, + cancellation, + deadline, + )?; + match frame_type { + FRAME_RENDER_CHUNK => { + if payload.is_empty() + || !payload + .len() + .is_multiple_of(RENDER_BYTES_PER_FRAME as usize) + || pcm.len().saturating_add(payload.len()) > expected_bytes as usize + { + return Err(RenderFailure::protocol("Magenta PCM chunk is invalid")); + } + digest.update(&payload); + pcm.extend_from_slice(&payload); + } + FRAME_RENDER_ERROR => { + if payload.len() > MAX_RENDER_METADATA_BYTES { + return Err(RenderFailure::protocol("Magenta worker error is too large")); + } + validate_render_error(&payload, request)?; + unreachable!("a valid worker error is returned as a render failure"); + } + FRAME_RENDER_END => { + if payload.len() > MAX_RENDER_METADATA_BYTES { + return Err(RenderFailure::protocol("Magenta render end is too large")); + } + let end: RenderEnd = serde_json::from_slice(&payload) + .map_err(|_| RenderFailure::protocol("Magenta render end is invalid"))?; + let actual_hash = hex::encode(digest.finalize()); + if end.schema_version != RENDER_SCHEMA_VERSION + || end.job_id != request.job_id + || end.sequence != request.sequence + || end.frames != request.frames + || end.pcm_bytes != expected_bytes + || pcm.len() != expected_bytes as usize + || end.sha256.len() != 64 + || !end.sha256.bytes().all(|byte| byte.is_ascii_hexdigit()) + || end.sha256 != actual_hash + { + return Err(RenderFailure::protocol("Magenta render end is invalid")); + } + return Ok(pcm); + } + _ => unreachable!("frame type was checked"), + } + } +} + +fn write_frame(writer: &mut impl Write, frame_type: u8, payload: &[u8]) -> io::Result<()> { + let length = u32::try_from(payload.len()) + .map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "frame is too large"))?; + writer.write_all(&[frame_type])?; + writer.write_all(&length.to_le_bytes())?; + writer.write_all(payload)?; + writer.flush() +} + +fn read_bounded_frame( + reader: &mut impl Read, + allowed_types: &[u8], + maximum: usize, + cancellation: &RequestCancellation, + deadline: Instant, +) -> Result<(u8, Vec), RenderFailure> { + let mut header = [0u8; 5]; + read_exact_cancellable(reader, &mut header, cancellation, deadline)?; + let frame_type = header[0]; + if !allowed_types.contains(&frame_type) { + return Err(RenderFailure::protocol("Magenta frame is out of order")); + } + let length = u32::from_le_bytes(header[1..5].try_into().expect("four bytes")) as usize; + // Metadata and PCM share this helper. A caller that accepts chunks passes + // the chunk cap, but control metadata remains capped before allocation or + // reading so an END/ERROR frame cannot consume a chunk-sized buffer. + let maximum = if frame_type == FRAME_RENDER_CHUNK { + maximum + } else { + maximum.min(MAX_RENDER_METADATA_BYTES) + }; + if length > maximum { + return Err(RenderFailure::protocol( + "Magenta frame exceeds its size cap", + )); + } + let mut payload = vec![0u8; length]; + read_exact_cancellable(reader, &mut payload, cancellation, deadline)?; + Ok((frame_type, payload)) +} + +fn read_exact_cancellable( + reader: &mut impl Read, + mut output: &mut [u8], + cancellation: &RequestCancellation, + deadline: Instant, +) -> Result<(), RenderFailure> { + while !output.is_empty() { + if cancellation.cancelled() { + return Err(RenderFailure::cancelled()); + } + if Instant::now() >= deadline { + return Err(RenderFailure::deadline()); + } + match reader.read(output) { + Ok(0) => { + return Err(RenderFailure::protocol( + "Magenta render response was truncated", + )) + } + Ok(read) => output = &mut output[read..], + Err(error) + if matches!( + error.kind(), + io::ErrorKind::WouldBlock + | io::ErrorKind::TimedOut + | io::ErrorKind::Interrupted + ) => {} + Err(_) => { + return Err(RenderFailure::protocol( + "Magenta render worker connection failed", + )) + } + } + } + Ok(()) +} + +fn bind_loopback() -> io::Result<(TcpListener, u16)> { + let listener = TcpListener::bind("127.0.0.1:0")?; + let port = listener.local_addr()?.port(); + listener.set_nonblocking(true)?; + Ok((listener, port)) +} + +fn serve( + listener: TcpListener, + port: u16, + capability: &str, + core: Arc, +) -> CancellationToken { + let auth = AuthState { + capability: Arc::from(capability), + }; + let router = Router::new() + .route("/api/render", post(render_clip).options(preflight)) + .route("/api/models", get(model_info).options(preflight)) + .layer(DefaultBodyLimit::max(MAX_RENDER_REQUEST_BYTES)) + .layer(axum::middleware::from_fn_with_state(auth, authenticate)) + .with_state(HttpState { core }); + let cancel = CancellationToken::new(); + let serve_cancel = cancel.clone(); + tauri::async_runtime::spawn(async move { + let listener = match tokio::net::TcpListener::from_std(listener) { + Ok(listener) => listener, + Err(error) => { + eprintln!("lsdj-app: Magenta gateway listener failed: {error}"); + return; + } + }; + println!("lsdj-app: Magenta gateway on http://127.0.0.1:{port}"); + if let Err(error) = axum::serve(listener, router) + .with_graceful_shutdown(async move { serve_cancel.cancelled().await }) + .await + { + eprintln!("lsdj-app: Magenta gateway stopped: {error}"); + } + }); + cancel +} + +async fn preflight() -> StatusCode { + StatusCode::NO_CONTENT +} + +async fn authenticate(State(auth): State, request: Request, next: Next) -> Response { + let origin = request + .headers() + .get(header::ORIGIN) + .and_then(|value| value.to_str().ok()) + .map(str::to_owned); + if origin + .as_deref() + .is_some_and(|origin| !SAFE_ORIGINS.contains(&origin)) + { + return json_error(StatusCode::FORBIDDEN, "origin is not allowed"); + } + if request.method() == Method::OPTIONS { + let requested_method = request + .headers() + .get(header::ACCESS_CONTROL_REQUEST_METHOD) + .and_then(|value| value.to_str().ok()); + let requested_headers = request + .headers() + .get(header::ACCESS_CONTROL_REQUEST_HEADERS) + .and_then(|value| value.to_str().ok()) + .unwrap_or_default() + .split(',') + .map(|value| value.trim().to_ascii_lowercase()) + .filter(|value| !value.is_empty()) + .collect::>(); + if origin.is_none() + || !matches!(requested_method, Some("GET" | "POST")) + || requested_headers + .iter() + .any(|value| !matches!(value.as_str(), "content-type" | "x-lsdj-capability")) + { + return json_error(StatusCode::FORBIDDEN, "preflight rejected"); + } + let mut response = StatusCode::NO_CONTENT.into_response(); + add_cors(&mut response, origin.as_deref().expect("origin checked")); + response.headers_mut().insert( + header::ACCESS_CONTROL_ALLOW_METHODS, + HeaderValue::from_static("GET, POST"), + ); + response.headers_mut().insert( + header::ACCESS_CONTROL_ALLOW_HEADERS, + HeaderValue::from_static("content-type, x-lsdj-capability"), + ); + response.headers_mut().insert( + header::ACCESS_CONTROL_MAX_AGE, + HeaderValue::from_static("600"), + ); + return response; + } + let supplied = request + .headers() + .get("x-lsdj-capability") + .map(|value| value.as_bytes()) + .unwrap_or_default(); + if !crate::local_auth::constant_time_eq(supplied, auth.capability.as_bytes()) { + return json_error(StatusCode::UNAUTHORIZED, "authentication required"); + } + let mut response = next.run(request).await; + if let Some(origin) = origin.as_deref() { + add_cors(&mut response, origin); + } + response +} + +fn add_cors(response: &mut Response, origin: &str) { + if let Ok(origin) = HeaderValue::from_str(origin) { + response + .headers_mut() + .insert(header::ACCESS_CONTROL_ALLOW_ORIGIN, origin); + response + .headers_mut() + .insert(header::VARY, HeaderValue::from_static("Origin")); + } +} + +#[cfg(target_os = "linux")] +fn total_ram_gb() -> Option { + std::fs::read_to_string("/proc/meminfo") + .ok()? + .lines() + .find_map(|line| line.strip_prefix("MemTotal:"))? + .split_whitespace() + .next()? + .parse::() + .ok() + .map(|kilobytes| kilobytes / 1024.0 / 1024.0) +} + +#[cfg(target_os = "windows")] +fn total_ram_gb() -> Option { + use windows_sys::Win32::System::SystemInformation::{GlobalMemoryStatusEx, MEMORYSTATUSEX}; + + let mut status: MEMORYSTATUSEX = unsafe { std::mem::zeroed() }; + status.dwLength = std::mem::size_of::() as u32; + // SAFETY: `status` is writable and advertises its exact structure size. + (unsafe { GlobalMemoryStatusEx(&mut status) } != 0) + .then_some(status.ullTotalPhys as f64 / 1024.0 / 1024.0 / 1024.0) +} + +#[cfg(not(any(target_os = "linux", target_os = "windows")))] +fn total_ram_gb() -> Option { + None +} + +#[cfg(test)] +mod tests { + use std::collections::VecDeque; + use std::sync::atomic::{AtomicUsize, Ordering}; + use std::thread; + + use super::*; + + #[derive(Clone, Copy)] + enum Scenario { + Valid, + WrongSequence, + WrongFrames, + WrongTotals, + WrongHash, + OutOfOrder, + OversizeEnd, + MisalignedChunk, + Stall, + } + + struct FakeProcess { + shutdowns: Arc, + shutdown_failures: Arc, + } + + impl ProcessTree for FakeProcess { + fn shutdown(&mut self) -> io::Result<()> { + self.shutdowns.fetch_add(1, Ordering::AcqRel); + if self + .shutdown_failures + .fetch_update(Ordering::AcqRel, Ordering::Acquire, |remaining| { + remaining.checked_sub(1) + }) + .is_ok() + { + return Err(io::Error::other("fake process is not reaped")); + } + Ok(()) + } + } + + struct FakeFactory { + scenarios: Mutex>, + spawns: Arc, + shutdowns: Arc, + shutdown_failures: Arc, + } + + impl FakeFactory { + fn new(scenarios: impl IntoIterator) -> Arc { + Arc::new(Self { + scenarios: Mutex::new(scenarios.into_iter().collect()), + spawns: Arc::new(AtomicUsize::new(0)), + shutdowns: Arc::new(AtomicUsize::new(0)), + shutdown_failures: Arc::new(AtomicUsize::new(0)), + }) + } + + fn fail_shutdowns(&self, count: usize) { + self.shutdown_failures.store(count, Ordering::Release); + } + } + + impl WorkerFactory for FakeFactory { + fn spawn( + &self, + _cancellation: &RequestCancellation, + ) -> Result { + let scenario = self + .scenarios + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + .pop_front() + .ok_or_else(RenderFailure::unavailable)?; + self.spawns.fetch_add(1, Ordering::AcqRel); + let listener = + TcpListener::bind("127.0.0.1:0").map_err(|_| RenderFailure::unavailable())?; + let address = listener + .local_addr() + .map_err(|_| RenderFailure::unavailable())?; + let client = TcpStream::connect(address).map_err(RenderFailure::from)?; + client + .set_read_timeout(Some(IO_POLL)) + .map_err(RenderFailure::from)?; + let (server, _) = listener.accept().map_err(RenderFailure::from)?; + thread::spawn(move || serve_scenario(server, scenario)); + Ok(ManagedRenderWorker { + stream: client, + process: Box::new(FakeProcess { + shutdowns: self.shutdowns.clone(), + shutdown_failures: self.shutdown_failures.clone(), + }), + }) + } + } + + fn read_test_frame(stream: &mut TcpStream) -> (u8, Vec) { + let mut header = [0u8; 5]; + stream.read_exact(&mut header).expect("request header"); + let length = u32::from_le_bytes(header[1..].try_into().unwrap()) as usize; + let mut payload = vec![0u8; length]; + stream.read_exact(&mut payload).expect("request payload"); + (header[0], payload) + } + + fn serve_scenario(mut stream: TcpStream, scenario: Scenario) { + let (frame_type, payload) = read_test_frame(&mut stream); + assert_eq!(frame_type, FRAME_RENDER_REQUEST); + let request: serde_json::Value = serde_json::from_slice(&payload).unwrap(); + let job_id = request["jobId"].as_str().unwrap(); + let sequence = request["sequence"].as_u64().unwrap(); + let frames = request["frames"].as_u64().unwrap(); + if matches!(scenario, Scenario::Stall) { + thread::sleep(Duration::from_millis(250)); + return; + } + if matches!(scenario, Scenario::OutOfOrder) { + write_frame(&mut stream, FRAME_RENDER_END, b"{}").ok(); + return; + } + + let pcm = vec![0x3fu8; frames as usize * RENDER_BYTES_PER_FRAME as usize]; + let begin_sequence = if matches!(scenario, Scenario::WrongSequence) { + sequence + 1 + } else { + sequence + }; + let begin_frames = if matches!(scenario, Scenario::WrongFrames) { + frames + 1 + } else { + frames + }; + let begin = serde_json::json!({ + "schemaVersion": RENDER_SCHEMA_VERSION, + "jobId": job_id, + "sequence": begin_sequence, + "sampleRate": RENDER_SAMPLE_RATE, + "channels": RENDER_CHANNELS, + "sampleFormat": "f32le", + "frames": begin_frames, + "pcmBytes": pcm.len(), + }); + if write_frame( + &mut stream, + FRAME_RENDER_BEGIN, + &serde_json::to_vec(&begin).unwrap(), + ) + .is_err() + { + return; + } + if matches!(scenario, Scenario::OversizeEnd) { + let _ = stream.write_all(&[FRAME_RENDER_END]); + let _ = stream.write_all(&((MAX_RENDER_METADATA_BYTES + 1) as u32).to_le_bytes()); + return; + } + let chunk = if matches!(scenario, Scenario::MisalignedChunk) { + &pcm[..3] + } else { + &pcm + }; + if write_frame(&mut stream, FRAME_RENDER_CHUNK, chunk).is_err() { + return; + } + let reported_bytes = if matches!(scenario, Scenario::WrongTotals) { + pcm.len() as u64 + RENDER_BYTES_PER_FRAME + } else { + pcm.len() as u64 + }; + let hash = if matches!(scenario, Scenario::WrongHash) { + "0".repeat(64) + } else { + hex::encode(Sha256::digest(&pcm)) + }; + let end = serde_json::json!({ + "schemaVersion": RENDER_SCHEMA_VERSION, + "jobId": job_id, + "sequence": sequence, + "frames": frames, + "pcmBytes": reported_bytes, + "sha256": hash, + }); + write_frame( + &mut stream, + FRAME_RENDER_END, + &serde_json::to_vec(&end).unwrap(), + ) + .ok(); + } + + fn render_with( + core: &GatewayCore, + cancellation: Arc, + ) -> Result, RenderFailure> { + core.render("test prompt".to_string(), 2, cancellation) + } + + #[test] + fn seconds_are_converted_to_authoritative_integer_frames() { + assert_eq!(frames_for_seconds(0.5), Some(24_000)); + assert_eq!( + frames_for_seconds((24_000.5) / RENDER_SAMPLE_RATE as f64), + Some(24_001) + ); + assert_eq!(frames_for_seconds(180.0), Some(MAX_RENDER_FRAMES)); + assert_eq!(frames_for_seconds(0.499), None); + assert_eq!(frames_for_seconds(f64::NAN), None); + } + + #[test] + fn valid_fake_worker_response_is_accepted_exactly() { + let factory = FakeFactory::new([Scenario::Valid]); + let core = GatewayCore::new(factory.clone()); + let pcm = render_with(&core, Arc::new(AtomicBool::new(false))).unwrap(); + assert_eq!(pcm, vec![0x3f; 2 * RENDER_BYTES_PER_FRAME as usize]); + assert_eq!(factory.spawns.load(Ordering::Acquire), 1); + assert_eq!(core.quiesce(), Ok(true)); + assert_eq!(factory.shutdowns.load(Ordering::Acquire), 1); + } + + #[test] + fn every_protocol_violation_discards_and_reaps_the_worker() { + for scenario in [ + Scenario::WrongSequence, + Scenario::WrongFrames, + Scenario::WrongTotals, + Scenario::WrongHash, + Scenario::OutOfOrder, + Scenario::OversizeEnd, + Scenario::MisalignedChunk, + ] { + let factory = FakeFactory::new([scenario]); + let core = GatewayCore::new(factory.clone()); + let error = render_with(&core, Arc::new(AtomicBool::new(false))).unwrap_err(); + assert_eq!(error.kind, FailureKind::Protocol); + assert_eq!(factory.shutdowns.load(Ordering::Acquire), 1); + assert!(core.worker.lock().unwrap().is_none()); + } + } + + #[test] + fn next_request_recovers_with_a_fresh_worker_after_failure() { + let factory = FakeFactory::new([Scenario::WrongHash, Scenario::Valid]); + let core = GatewayCore::new(factory.clone()); + assert!(render_with(&core, Arc::new(AtomicBool::new(false))).is_err()); + assert!(render_with(&core, Arc::new(AtomicBool::new(false))).is_ok()); + assert_eq!(factory.spawns.load(Ordering::Acquire), 2); + assert_eq!(factory.shutdowns.load(Ordering::Acquire), 1); + assert_eq!(core.quiesce(), Ok(true)); + assert_eq!(factory.shutdowns.load(Ordering::Acquire), 2); + } + + #[test] + fn uncertain_reap_state_is_retained_and_blocks_promotion() { + let factory = FakeFactory::new([Scenario::Valid]); + let core = GatewayCore::new(factory.clone()); + assert!(render_with(&core, Arc::new(AtomicBool::new(false))).is_ok()); + factory.fail_shutdowns(1); + + assert!(core.quiesce().is_err()); + assert!(core.worker.lock().unwrap().is_some()); + assert_eq!(factory.shutdowns.load(Ordering::Acquire), 1); + + assert_eq!(core.quiesce(), Ok(true)); + assert!(core.worker.lock().unwrap().is_none()); + assert_eq!(factory.shutdowns.load(Ordering::Acquire), 2); + } + + #[test] + fn cancellation_interrupts_a_stalled_worker_and_reaps_it() { + let factory = FakeFactory::new([Scenario::Stall]); + let core = Arc::new(GatewayCore::new(factory.clone())); + let cancellation = Arc::new(AtomicBool::new(false)); + let flag = cancellation.clone(); + thread::spawn(move || { + thread::sleep(Duration::from_millis(10)); + flag.store(true, Ordering::Release); + }); + let error = render_with(&core, cancellation).unwrap_err(); + assert_eq!(error.kind, FailureKind::Cancelled); + assert_eq!(factory.shutdowns.load(Ordering::Acquire), 1); + assert!(core.worker.lock().unwrap().is_none()); + } + + #[test] + fn deadline_and_drop_cancellation_are_observed_while_reading() { + let request = Arc::new(AtomicBool::new(false)); + { + let _guard = CancelOnDrop::new(request.clone()); + } + assert!(request.load(Ordering::Acquire)); + + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let address = listener.local_addr().unwrap(); + let mut client = TcpStream::connect(address).unwrap(); + client.set_read_timeout(Some(IO_POLL)).unwrap(); + let (_server, _) = listener.accept().unwrap(); + let cancellation = RequestCancellation { + request: Arc::new(AtomicBool::new(false)), + lifecycle: Arc::new(AtomicBool::new(false)), + }; + let error = read_bounded_frame( + &mut client, + &[FRAME_RENDER_BEGIN], + MAX_RENDER_METADATA_BYTES, + &cancellation, + Instant::now() + Duration::from_millis(20), + ) + .unwrap_err(); + assert_eq!(error.kind, FailureKind::Deadline); + } +} diff --git a/src-tauri/src/mcp.rs b/src-tauri/src/mcp.rs index aaabd54..be16dc3 100644 --- a/src-tauri/src/mcp.rs +++ b/src-tauri/src/mcp.rs @@ -36,6 +36,8 @@ use tokio_util::sync::CancellationToken; use crate::commands::{valid_deck, DrumModeArg, EqBandArg, FxKindArg}; use crate::generation::GenerationServer; +#[cfg(feature = "managed-runtime")] +use crate::magenta_gateway::MagentaGateway; use crate::samples::{NewSample, SampleLibrary}; use crate::sidecar::Sidecars; use crate::songs::{NewSong, SongLibrary}; @@ -915,13 +917,36 @@ impl McpHandler { /// validation. `magenta` routes to the Magenta renderer (`/api/render`, body /// `{prompt, seconds}`); the rest are Stable Audio 3 (`/api/generate`). async fn generate_clip(&self, prompt: &str, seconds: f32, kind: &str) -> Result, String> { - let generation = self.app.state::(); - let port = generation - .port() - .ok_or("the generation server is not running")?; - let capability = generation - .capability() - .ok_or("the generation server authentication capability is unavailable")?; + let (port, capability) = if kind == "magenta" { + #[cfg(feature = "managed-runtime")] + { + let gateway = self.app.state::(); + ( + gateway.port().ok_or("the Magenta gateway is not running")?, + gateway + .capability() + .ok_or("the Magenta gateway authentication capability is unavailable")?, + ) + } + #[cfg(not(feature = "managed-runtime"))] + { + let generation = self.app.state::(); + ( + generation.port().ok_or("the generation server is not running")?, + generation + .capability() + .ok_or("the generation server authentication capability is unavailable")?, + ) + } + } else { + let generation = self.app.state::(); + ( + generation.port().ok_or("the generation server is not running")?, + generation + .capability() + .ok_or("the generation server authentication capability is unavailable")?, + ) + }; // sa3 generation is serialised; a full track (medium model) can take minutes, // so allow generous headroom but never wait forever for a wedged worker. let client = reqwest::Client::builder() diff --git a/src-tauri/src/models.rs b/src-tauri/src/models.rs index 66396df..25b83be 100644 --- a/src-tauri/src/models.rs +++ b/src-tauri/src/models.rs @@ -25,6 +25,8 @@ use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::{Arc, Mutex}; use serde::{Deserialize, Serialize}; +#[cfg(feature = "managed-runtime")] +use tauri::Manager; use tauri::{AppHandle, Emitter}; use crate::child_process::{ @@ -127,17 +129,25 @@ const BACKEND_PATH_ENVIRONMENT: &[&str] = &[ "LSDJ_CONFIG_HOME", "LSDJ_DATA_HOME", "LSDJ_STAGING_HOME", - "MAGENTA_HOME", - "SA3_HOME", - "SA3_LORAS_HOME", - "SA3_MLX_HOME", ]; +const MRT2_PATH_ENVIRONMENT: &[&str] = &["MAGENTA_HOME"]; +const SA3_PATH_ENVIRONMENT: &[&str] = &["SA3_HOME", "SA3_LORAS_HOME", "SA3_MLX_HOME"]; const WINDOWS_CHILD_ENVIRONMENT: &[&str] = &["SYSTEMROOT", "WINDIR", "TEMP", "TMP"]; -fn service_ephemeral_environment(secret: &str) -> Vec { +fn service_ephemeral_environment( + service: crate::managed_runtime::Service, + secret: &str, +) -> Vec { + let service_paths = match service { + crate::managed_runtime::Service::Mrt2 => MRT2_PATH_ENVIRONMENT, + crate::managed_runtime::Service::Sa3 | crate::managed_runtime::Service::Sa3Cuda => { + SA3_PATH_ENVIRONMENT + } + }; std::iter::once(secret) .chain(BACKEND_PATH_ENVIRONMENT.iter().copied()) + .chain(service_paths.iter().copied()) .chain(WINDOWS_CHILD_ENVIRONMENT.iter().copied()) .map(str::to_string) .collect() @@ -595,6 +605,25 @@ fn status(active: Option<(Family, String)>) -> ModelStatus { } } +/// Installed MRT2 models for the Rust-owned managed render gateway's lightweight +/// `/api/models` compatibility endpoint. Only the authenticated, generation- +/// bound install identity is trusted; a candidate or hand-placed directory never +/// becomes loadable merely because files exist. +#[cfg(feature = "managed-runtime")] +pub(crate) fn magenta_models_for_gateway() -> Vec { + let models_dir = magenta_models_dir(); + let root = models_dir.parent().unwrap_or(&models_dir); + let pin = mrt2_pin(); + if !managed_mrt2_host() || validate_mrt2_identity(root, &pin).is_err() { + return Vec::new(); + } + pin.models + .iter() + .filter(|(name, snapshot)| mrt2_snapshot_present(root, name, snapshot)) + .map(|(name, _)| name.clone()) + .collect() +} + // --- Install / delete ------------------------------------------------------ /// Which family a command targets. `lowercase` serde is the single source of the @@ -1058,13 +1087,13 @@ impl InstallManager { return Err(format!("unknown model '{name}'")); } let model = name.clone(); - self.start(app, family, name, move |progress, shared| { - install_magenta(progress, shared, &model) + self.start(app, family, name, move |app, progress, shared| { + install_magenta(app, progress, shared, &model) }) } // `model://progress` carries the model name for Magenta, "" for SA3. - Family::Sa3 => self.start(app, family, String::new(), move |progress, shared| { - install_sa3(progress, shared, update) + Family::Sa3 => self.start(app, family, String::new(), move |app, progress, shared| { + install_sa3(app, progress, shared, update) }), Family::Lora => Err("adapters are installed via install_lora".into()), } @@ -1079,7 +1108,7 @@ impl InstallManager { spec: crate::loras::ImportSpec, ) -> Result<(), String> { let name = spec.display_name()?; - self.start(app, Family::Lora, name, move |progress, shared| { + self.start(app, Family::Lora, name, move |_app, progress, shared| { crate::loras::install(progress, shared, &spec) }) } @@ -1092,7 +1121,7 @@ impl InstallManager { app: AppHandle, family: Family, name: String, - job: impl FnOnce(&Progress, &InstallShared) -> Result<(), String> + Send + 'static, + job: impl FnOnce(&AppHandle, &Progress, &InstallShared) -> Result<(), String> + Send + 'static, ) -> Result<(), String> { if self.shared.busy.swap(true, Ordering::AcqRel) { return Err("an install is already running".into()); @@ -1108,7 +1137,7 @@ impl InstallManager { let progress = move |stage: &str, message: Option, file: Option| { emit(&progress_app, family, &name, stage, message, file); }; - let result = job(&progress, &shared); + let result = job(&app, &progress, &shared); *shared .current_child .lock() @@ -1283,14 +1312,92 @@ struct SidecarLine { /// emit; tests record the events while the install actually runs. pub(crate) type Progress = dyn Fn(&str, Option, Option); -fn install_magenta(progress: &Progress, shared: &InstallShared, name: &str) -> Result<(), String> { +#[cfg(feature = "managed-runtime")] +#[derive(Clone, Copy)] +struct Mrt2Lifecycle { + render_was_warm: bool, +} + +#[cfg(feature = "managed-runtime")] +fn quiesce_mrt2_services(app: &AppHandle) -> Result { + let gateway = app.state::(); + let render_was_warm = gateway + .quiesce() + .map_err(|error| format!("cannot quiesce Magenta renderer before promotion: {error}"))?; + if let Err(error) = app.state::().quiesce_shared() { + // No rename has happened. Restore the still-current verified render + // generation before returning the deck teardown error. + let _ = gateway.resume(render_was_warm); + return Err(format!( + "cannot quiesce realtime MRT2 decks before promotion: {error}" + )); + } + Ok(Mrt2Lifecycle { render_was_warm }) +} + +#[cfg(feature = "managed-runtime")] +fn resume_mrt2_services(app: &AppHandle, lifecycle: Mrt2Lifecycle) -> Result<(), String> { + // Restore a previously warm renderer completely before launching the deck + // process, avoiding concurrent cold model allocations during promotion. + app.state::() + .resume(lifecycle.render_was_warm) + .map_err(|error| format!("cannot resume Magenta renderer: {error}"))?; + app.state::() + .resume_shared() + .map_err(|error| format!("cannot resume realtime MRT2 decks: {error}")) +} + +/// Promotion owns the commit/rollback result, but service restart is part of +/// operational success. A failed promotion still attempts to restore the prior +/// verified generation, and reports both causes if that recovery also fails. +#[cfg(feature = "managed-runtime")] +fn finish_promotion( + label: &str, + promoted: Result<(), String>, + resumed: Result<(), String>, +) -> Result<(), String> { + match (promoted, resumed) { + (Ok(()), Ok(())) => Ok(()), + (Ok(()), Err(resume)) => Err(format!( + "{label} was promoted but its service could not resume: {resume}" + )), + (Err(promote), Ok(())) => Err(promote), + (Err(promote), Err(resume)) => Err(format!( + "{promote}; the previous {label} service also could not resume: {resume}" + )), + } +} + +/// Execute the rename window in its only safe order. Keeping the sequence in +/// one small, injectable helper makes the Windows "dead before rename" rule and +/// rollback restart behavior testable without Python, CUDA, or model files. +#[cfg(feature = "managed-runtime")] +fn run_promotion_lifecycle( + label: &str, + quiesce: impl FnOnce() -> Result, + promote: impl FnOnce() -> Result<(), String>, + resume: impl FnOnce(L) -> Result<(), String>, +) -> Result<(), String> { + let lifecycle = quiesce()?; + let promoted = promote(); + let resumed = resume(lifecycle); + finish_promotion(label, promoted, resumed) +} + +fn install_magenta( + app: &AppHandle, + progress: &Progress, + shared: &InstallShared, + name: &str, +) -> Result<(), String> { #[cfg(feature = "managed-runtime")] { - install_mrt2_managed(progress, shared, name) + install_mrt2_managed(app, progress, shared, name) } #[cfg(not(feature = "managed-runtime"))] { + let _ = app; progress("download", None, None); let mut cmd = crate::sidecar::sidecar_base_command().map_err(|e| e.to_string())?; if !resources_present() { @@ -1305,6 +1412,7 @@ fn install_magenta(progress: &Progress, shared: &InstallShared, name: &str) -> R #[cfg(feature = "managed-runtime")] fn install_mrt2_managed( + app: &AppHandle, progress: &Progress, shared: &InstallShared, name: &str, @@ -1402,9 +1510,16 @@ fn install_mrt2_managed( seal_mrt2_candidate(&candidate, &pin, python)?; validate_mrt2_candidate(&candidate, &pin, name, &cancelled_now)?; progress("promote", None, None); - promotion::promote(&candidate, &home, &backup, |root| { - validate_mrt2_candidate(root, &pin, name, &cancelled_now) - })?; + run_promotion_lifecycle( + "MRT2", + || quiesce_mrt2_services(app), + || { + promotion::promote(&candidate, &home, &backup, |root| { + validate_mrt2_candidate(root, &pin, name, &cancelled_now) + }) + }, + |lifecycle| resume_mrt2_services(app, lifecycle), + )?; let _ = std::fs::remove_dir_all(&work); Ok(()) } @@ -1713,7 +1828,12 @@ fn run_download(progress: &Progress, shared: &InstallShared, cmd: Command) -> Re result.map_err(|exit_err| sanitize_diagnostic(&last_error.unwrap_or(exit_err))) } -fn install_sa3(progress: &Progress, shared: &InstallShared, _update: bool) -> Result<(), String> { +fn install_sa3( + app: &AppHandle, + progress: &Progress, + shared: &InstallShared, + _update: bool, +) -> Result<(), String> { let pin = sa3_pin(); validate_sa3_pin(&pin)?; let backend = host_sa3_backend()?; @@ -1752,9 +1872,34 @@ fn install_sa3(progress: &Progress, shared: &InstallShared, _update: bool) -> Re )?; cancelled(shared)?; progress("promote", None, None); - promotion::promote(&candidate, &home, &backup, |path| { + #[cfg(feature = "managed-runtime")] + let service_was_running = app + .state::() + .quiesce() + .map_err(|error| format!("cannot quiesce SA3 before promotion: {error}"))?; + let promoted = promotion::promote(&candidate, &home, &backup, |path| { validate_sa3_install_cancellable(path, &pin, backend, &install_cancelled) - })?; + }); + #[cfg(feature = "managed-runtime")] + let resumed = app + .state::() + .resume() + .map_err(|error| format!("cannot resume SA3 after promotion: {error}")); + #[cfg(feature = "managed-runtime")] + let resumed = if promoted.is_err() && !service_was_running { + // First install had no prior service to restore; keep the promotion + // cause authoritative instead of appending the expected "not installed" + // resume failure. + Ok(()) + } else { + resumed + }; + #[cfg(feature = "managed-runtime")] + finish_promotion("SA3", promoted, resumed)?; + #[cfg(not(feature = "managed-runtime"))] + promoted?; + #[cfg(not(feature = "managed-runtime"))] + let _ = app; // Verified blobs are hard-linked into the promoted tree. Removing retry // state here reclaims only the staging directory entries, not model bytes. let _ = std::fs::remove_dir_all(&work); @@ -2619,7 +2764,10 @@ fn seal_sa3_candidate( .into_iter() .map(|(key, value)| (key.to_string(), value.to_string())) .collect(); - let ephemeral_environment = service_ephemeral_environment("LSDJ_API_CAPABILITY"); + let ephemeral_environment = service_ephemeral_environment( + crate::managed_runtime::Service::Sa3, + "LSDJ_API_CAPABILITY", + ); let spec = crate::managed_runtime::CommandSpec { program: relative_wire(candidate, &program)?, argv: vec!["launch.py".into(), "--generation-server".into()], @@ -2728,7 +2876,10 @@ fn seal_mrt2_candidate(candidate: &Path, pin: &Mrt2Pin, python: &PythonPin) -> R .into_iter() .map(|(key, value)| (key.to_string(), value.to_string())) .collect(); - let ephemeral_environment = service_ephemeral_environment("LSDJ_WORKER_LAUNCH_TOKEN"); + let ephemeral_environment = service_ephemeral_environment( + crate::managed_runtime::Service::Mrt2, + "LSDJ_WORKER_LAUNCH_TOKEN", + ); let spec = crate::managed_runtime::CommandSpec { program: relative_wire(candidate, &program)?, argv: vec!["launch.py".into()], @@ -3316,12 +3467,18 @@ mod tests { backend_sources_digest() ); - let sa3: BTreeSet<_> = service_ephemeral_environment("LSDJ_API_CAPABILITY") - .into_iter() - .collect(); - let mrt2: BTreeSet<_> = service_ephemeral_environment("LSDJ_WORKER_LAUNCH_TOKEN") - .into_iter() - .collect(); + let sa3: BTreeSet<_> = service_ephemeral_environment( + crate::managed_runtime::Service::Sa3, + "LSDJ_API_CAPABILITY", + ) + .into_iter() + .collect(); + let mrt2: BTreeSet<_> = service_ephemeral_environment( + crate::managed_runtime::Service::Mrt2, + "LSDJ_WORKER_LAUNCH_TOKEN", + ) + .into_iter() + .collect(); assert!(sa3.contains("LSDJ_API_CAPABILITY")); assert!(!sa3.contains("LSDJ_WORKER_LAUNCH_TOKEN")); assert!(mrt2.contains("LSDJ_WORKER_LAUNCH_TOKEN")); @@ -3333,6 +3490,61 @@ mod tests { assert!(sa3.contains(*name)); assert!(mrt2.contains(*name)); } + assert!(sa3.contains("SA3_HOME")); + assert!(!sa3.contains("MAGENTA_HOME")); + assert!(mrt2.contains("MAGENTA_HOME")); + assert!(!mrt2.contains("SA3_HOME")); + } + + #[cfg(feature = "managed-runtime")] + #[test] + fn managed_promotion_quiesces_before_rename_and_resumes_after_rollback() { + use std::cell::RefCell; + + let events = RefCell::new(Vec::new()); + let result = run_promotion_lifecycle( + "fake runtime", + || { + events.borrow_mut().push("quiesce"); + Ok("prior generation") + }, + || { + events.borrow_mut().push("promote"); + Err("fake promotion failed".to_string()) + }, + |generation| { + assert_eq!(generation, "prior generation"); + events.borrow_mut().push("resume"); + Ok(()) + }, + ); + assert_eq!(result.unwrap_err(), "fake promotion failed"); + assert_eq!(*events.borrow(), ["quiesce", "promote", "resume"]); + } + + #[cfg(feature = "managed-runtime")] + #[test] + fn failed_quiesce_never_enters_the_rename_window() { + use std::cell::RefCell; + + let events = RefCell::new(Vec::new()); + let result = run_promotion_lifecycle( + "fake runtime", + || { + events.borrow_mut().push("quiesce"); + Err::<(), _>("worker could not be reaped".to_string()) + }, + || { + events.borrow_mut().push("promote"); + Ok(()) + }, + |_| { + events.borrow_mut().push("resume"); + Ok(()) + }, + ); + assert_eq!(result.unwrap_err(), "worker could not be reaped"); + assert_eq!(*events.borrow(), ["quiesce"]); } #[test] diff --git a/src-tauri/src/sidecar.rs b/src-tauri/src/sidecar.rs index 95dd3d8..e018958 100644 --- a/src-tauri/src/sidecar.rs +++ b/src-tauri/src/sidecar.rs @@ -657,6 +657,30 @@ pub struct SharedSidecar { } impl SharedSidecar { + /// Construct the shared CUDA topology without launching Python. The native + /// engine's permanent ring producers remain parked here until a verified + /// managed MRT2 generation is installed (or a later retry succeeds). + pub fn parked( + models: [String; lsdj_engine::DECK_COUNT], + handles: [DeckHandle; lsdj_engine::DECK_COUNT], + on_status: DeckStatusSinks, + taps: PcmTaps, + feed: AnalysisFeed, + ) -> Self { + Self { + models, + taps, + feed, + on_status: on_status.map(|sink| Arc::new(Mutex::new(sink))), + control: Arc::new(Mutex::new(None)), + child: Arc::new(Mutex::new(None)), + stop: Arc::new(AtomicBool::new(true)), + reader: None, + parked: Some(SharedReaderExit { handles }), + } + } + + #[cfg(all(test, not(feature = "managed-runtime")))] pub fn spawn( models: [String; lsdj_engine::DECK_COUNT], handles: [DeckHandle; lsdj_engine::DECK_COUNT], @@ -664,27 +688,65 @@ impl SharedSidecar { taps: PcmTaps, feed: AnalysisFeed, ) -> Result { - let (listener, child, token) = match bind_and_launch_shared(&models) { - Ok(launch) => launch, - Err(error) => return Err((error, handles)), - }; - let on_status = on_status.map(|sink| Arc::new(Mutex::new(sink))); + let mut sidecar = Self::parked(models, handles, on_status, taps, feed); + if let Err(error) = sidecar.activate() { + let handles = sidecar + .parked + .take() + .expect("failed shared activation preserves deck handles") + .handles; + return Err((error, handles)); + } + Ok(sidecar) + } + + /// Start a worker from parked handles. Resolution and manifest validation + /// occur inside `bind_and_launch_shared` immediately before spawn, so an + /// install that completed after app startup becomes usable without restart. + pub fn activate(&mut self) -> io::Result<()> { + if self.reader.is_some() { + return Ok(()); + } + if self.parked.is_none() { + return Err(io::Error::other( + "shared sidecar has no parked deck handles", + )); + } + let (listener, child, token) = bind_and_launch_shared(&self.models)?; + let exit = self + .parked + .take() + .ok_or_else(|| io::Error::other("shared sidecar has no parked deck handles"))?; let on_pcm: DeckPcmSinks = [ - Box::new(pcm_tee(taps.clone(), feed.clone(), 0)), - Box::new(pcm_tee(taps.clone(), feed.clone(), 1)), + Box::new(pcm_tee(self.taps.clone(), self.feed.clone(), 0)), + Box::new(pcm_tee(self.taps.clone(), self.feed.clone(), 1)), ]; - let parts = start_shared_reader(listener, child, token, handles, on_status.clone(), on_pcm); - Ok(Self { - models, - taps, - feed, - on_status, - control: parts.control, - child: parts.child, - stop: parts.stop, - reader: Some(parts.reader), - parked: None, - }) + let parts = start_shared_reader( + listener, + child, + token, + exit.handles, + self.on_status.clone(), + on_pcm, + ); + self.control = parts.control; + self.child = parts.child; + self.stop = parts.stop; + self.reader = Some(parts.reader); + Ok(()) + } + + /// Stop and fully reap the worker while retaining the permanent deck ring + /// producers. Promotion may rename the managed generation only after this + /// succeeds (notably on Windows, where a live Python process holds DLLs). + #[cfg(feature = "managed-runtime")] + pub fn quiesce(&mut self) -> io::Result<()> { + if self.reader.is_none() { + return Ok(()); + } + let exit = self.stop_and_reclaim()?; + self.parked = Some(exit); + Ok(()) } fn send_control(&self, deck: usize, json: &str) { @@ -876,6 +938,28 @@ impl Sidecars { } } + /// Quiesce the managed shared worker before its verified generation is + /// renamed. A no-op for the macOS per-deck topology. + #[cfg(feature = "managed-runtime")] + pub fn quiesce_shared(&self) -> Result<(), String> { + let mut shared = self.shared.lock().unwrap_or_else(|p| p.into_inner()); + if let Some(shared) = shared.as_mut() { + shared.quiesce().map_err(|error| error.to_string())?; + } + Ok(()) + } + + /// Revalidate and activate a parked shared worker. This covers both first + /// install and post-promotion restart without reconstructing the audio host. + #[cfg(feature = "managed-runtime")] + pub fn resume_shared(&self) -> Result<(), String> { + let mut shared = self.shared.lock().unwrap_or_else(|p| p.into_inner()); + let shared = shared + .as_mut() + .ok_or_else(|| "shared MRT2 decks are unavailable on this platform".to_string())?; + shared.activate().map_err(|error| error.to_string()) + } + /// Forward a JSON deck command to the sidecar for `deck` (a no-op for a deck /// without a live sidecar). `deck` is validated by the IPC layer. pub fn send(&self, deck: usize, json: &str) { @@ -1074,10 +1158,19 @@ fn authenticated_sidecar_base_command(token: &str) -> io::Result { use std::ffi::OsString; let paths = crate::platform_paths::get(); - let ephemeral = paths.backend_env().into_iter().chain(std::iter::once(( - OsString::from("LSDJ_WORKER_LAUNCH_TOKEN"), - OsString::from(token), - ))); + let ephemeral = paths + .backend_env() + .into_iter() + .filter(|(name, _)| { + !matches!( + name.to_str(), + Some("SA3_HOME" | "SA3_MLX_HOME" | "SA3_LORAS_HOME") + ) + }) + .chain(std::iter::once(( + OsString::from("LSDJ_WORKER_LAUNCH_TOKEN"), + OsString::from(token), + ))); crate::managed_runtime::resolve( paths.assets(), crate::managed_runtime::Service::Mrt2, @@ -1086,6 +1179,29 @@ fn authenticated_sidecar_base_command(token: &str) -> io::Result { .map_err(io::Error::other) } +/// Build the dedicated managed MRT2 renderer command. The render worker is a +/// separate, disposable process from both realtime decks and the SA3 server; +/// every launch re-resolves and revalidates the promoted MRT2 manifest and gets +/// a fresh one-use loopback capability. +#[cfg(feature = "managed-runtime")] +pub(crate) fn authenticated_render_worker_command( + model: &str, + port: u16, + token: &str, +) -> io::Result { + let mut command = authenticated_sidecar_base_command(token)?; + command.args([ + "--render-worker", + "--model", + model, + "--runtime", + "pytorch-cuda", + "--port", + &port.to_string(), + ]); + Ok(command) +} + /// Runtime selected by the native platform. The value is always sent over the /// process boundary: Python never guesses and never falls back from CUDA to CPU. /// The override exists for model-free contract tests and qualification hosts; @@ -1219,6 +1335,31 @@ mod tests { } } + #[test] + fn shared_deck_handles_can_park_until_the_first_managed_install() { + let mut engine = Engine::new(); + let handles = [engine.create_deck(0), engine.create_deck(1)]; + let sinks: DeckStatusSinks = std::array::from_fn(|_| { + Box::new(|_message| {}) as StatusSink + }); + let shared = SharedSidecar::parked( + ["mrt2_small".into(), "mrt2_small".into()], + handles, + sinks, + PcmTaps::new(lsdj_engine::DECK_COUNT), + AnalysisFeed::disconnected(lsdj_engine::DECK_COUNT), + ); + + assert!(shared.reader.is_none()); + assert!(shared.parked.is_some()); + assert!(shared + .child + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + .is_none()); + assert!(shared.stop.load(Ordering::Acquire)); + } + #[test] fn transport_ended_matches_only_worker_end_events() { // The three events after which the worker is no longer generating. From 3a76a6f3ab1764ab4b9bf45c05627e26c49f9585 Mon Sep 17 00:00:00 2001 From: brxs Date: Sat, 8 Aug 2026 19:24:30 -0700 Subject: [PATCH 05/14] feat: route native generation by model family --- frontend/src/audio/nativeEngine.test.ts | 82 +++++++++++++++++++++++++ frontend/src/audio/nativeEngine.ts | 61 +++++++++++++++--- 2 files changed, 133 insertions(+), 10 deletions(-) diff --git a/frontend/src/audio/nativeEngine.test.ts b/frontend/src/audio/nativeEngine.test.ts index 15e9ffb..d3138e3 100644 --- a/frontend/src/audio/nativeEngine.test.ts +++ b/frontend/src/audio/nativeEngine.test.ts @@ -112,6 +112,88 @@ describe('createNativeEngine — control contract', () => { expect((init.headers as Headers).get('x-lsdj-capability')).toBe('b'.repeat(64)) }) + it('routes Magenta requests to the distinct native gateway', async () => { + const invoke = vi.fn((cmd: string) => + cmd === 'app_info' + ? Promise.resolve({ + generationPort: 4321, + generationCapability: 's'.repeat(64), + magentaPort: 9876, + magentaCapability: 'm'.repeat(64), + }) + : Promise.resolve(undefined), + ) + vi.stubGlobal('__TAURI__', { core: { invoke } }) + const fetchMock = vi.fn(async (_url: string, _init: RequestInit) => { + void _url + void _init + return { ok: true } + }) + vi.stubGlobal('fetch', fetchMock) + + await fetchGenerationApi('/api/render', { method: 'POST' }) + await fetchGenerationApi('/api/generate', { method: 'POST' }) + + const [renderUrl, renderInit] = fetchMock.mock.calls[0] + expect(renderUrl).toBe('http://127.0.0.1:9876/api/render') + expect((renderInit.headers as Headers).get('x-lsdj-capability')).toBe('m'.repeat(64)) + const [generateUrl, generateInit] = fetchMock.mock.calls[1] + expect(generateUrl).toBe('http://127.0.0.1:4321/api/generate') + expect((generateInit.headers as Headers).get('x-lsdj-capability')).toBe('s'.repeat(64)) + }) + + it('fails closed when a managed Magenta gateway could not bind', async () => { + const invoke = vi.fn((cmd: string) => + cmd === 'app_info' + ? Promise.resolve({ + generationPort: 4321, + generationCapability: 's'.repeat(64), + magentaPort: null, + magentaCapability: null, + }) + : Promise.resolve(undefined), + ) + vi.stubGlobal('__TAURI__', { core: { invoke } }) + const fetchMock = vi.fn(async (_url: string, _init: RequestInit) => { + void _url + void _init + return { ok: true } + }) + vi.stubGlobal('fetch', fetchMock) + + await expect(fetchGenerationApi('/api/render', { method: 'POST' })).rejects.toThrow( + 'authentication is unavailable', + ) + expect(fetchMock).not.toHaveBeenCalled() + }) + + it('refreshes app_info when a first SA3 install starts the service', async () => { + let appInfoCalls = 0 + const invoke = vi.fn((cmd: string) => { + if (cmd !== 'app_info') return Promise.resolve(undefined) + appInfoCalls += 1 + return Promise.resolve( + appInfoCalls === 1 + ? { generationPort: null, generationCapability: null } + : { generationPort: 2468, generationCapability: 'n'.repeat(64) }, + ) + }) + vi.stubGlobal('__TAURI__', { core: { invoke } }) + const fetchMock = vi.fn(async (_url: string, _init: RequestInit) => { + void _url + void _init + return { ok: true } + }) + vi.stubGlobal('fetch', fetchMock) + + await fetchGenerationApi('/api/generate', { method: 'POST' }) + + expect(appInfoCalls).toBe(2) + const [url, init] = fetchMock.mock.calls[0] + expect(url).toBe('http://127.0.0.1:2468/api/generate') + expect((init.headers as Headers).get('x-lsdj-capability')).toBe('n'.repeat(64)) + }) + it('createDeckChannel replays NO mixer config — the shell hydrates (phase C)', async () => { const engine = createNativeEngine() await engine.createDeckChannel( diff --git a/frontend/src/audio/nativeEngine.ts b/frontend/src/audio/nativeEngine.ts index bc25af8..15ed89f 100644 --- a/frontend/src/audio/nativeEngine.ts +++ b/frontend/src/audio/nativeEngine.ts @@ -61,12 +61,14 @@ export function isTauri(): boolean { } type ApiConnection = { baseUrl: string; capability: string | null } -let apiConnectionPromise: Promise | null = null +type ApiConnections = { sa3: ApiConnection; magenta: ApiConnection } +let apiConnectionPromise: Promise | null = null let apiConnectionOwner: TauriGlobal | null = null -function getApiConnection(): Promise { +function loadApiConnections(): Promise { const owner = tauriGlobal() - if (!owner) return Promise.resolve({ baseUrl: '', capability: null }) + const unavailable = { baseUrl: '', capability: null } + if (!owner) return Promise.resolve({ sa3: unavailable, magenta: unavailable }) // A webview has one bridge for its lifetime. Coupling the cache to that bridge // also avoids carrying a stale launch capability across test/dev hot reloads. if (apiConnectionOwner !== owner) { @@ -77,28 +79,67 @@ function getApiConnection(): Promise { apiConnectionPromise = invoke<{ generationPort: number | null generationCapability: string | null + magentaPort?: number | null + magentaCapability?: string | null }>('app_info') - .then((info) => ({ - baseUrl: info.generationPort ? `http://127.0.0.1:${info.generationPort}` : '', - capability: info.generationCapability ?? null, - })) - .catch(() => ({ baseUrl: '', capability: null })) + .then((info) => { + const sa3 = { + baseUrl: info.generationPort ? `http://127.0.0.1:${info.generationPort}` : '', + capability: info.generationCapability ?? null, + } + const hasDistinctMagentaGateway = + info.magentaPort !== undefined || info.magentaCapability !== undefined + return { + sa3, + // Bundled macOS reports no distinct gateway fields and deliberately + // retains the combined controller. Managed Linux/Windows supplies a + // Rust-owned Magenta endpoint here. Explicit nulls mean that gateway + // failed closed; they must never fall through to the SA3 controller. + magenta: hasDistinctMagentaGateway + ? { + baseUrl: info.magentaPort + ? `http://127.0.0.1:${info.magentaPort}` + : '', + capability: info.magentaCapability ?? null, + } + : sa3, + } + }) + .catch(() => ({ sa3: unavailable, magenta: unavailable })) } return apiConnectionPromise } +function isMagentaPath(path: string): boolean { + return path === '/api/render' || path === '/api/models' +} + +async function getApiConnection(path: string): Promise { + let connections = await loadApiConnections() + let connection = isMagentaPath(path) ? connections.magenta : connections.sa3 + // A fresh managed install legitimately starts without SA3. Do not pin that + // absence for the webview lifetime: the first request after promotion + // re-reads app_info and reaches the newly started generation server. + if (isTauri() && !connection.baseUrl) { + apiConnectionPromise = null + connections = await loadApiConnections() + connection = isMagentaPath(path) ? connections.magenta : connections.sa3 + } + return connection +} + /** Base URL for the backend `/api/*` generation endpoints (sa3/Magenta pad+track * render). FastAPI no longer serves the UI, so the Rust shell runs a generation * server on a loopback port it reports via `app_info`; the webview fetches * `http://127.0.0.1:/api/...`. Resolved once and cached; falls back to '' * (relative) if the port can't be resolved. */ export function getApiBaseUrl(): Promise { - return getApiConnection().then((connection) => connection.baseUrl) + return getApiConnection('/api/generate').then((connection) => connection.baseUrl) } /** Authenticated fetch to the app-owned loopback generation service. */ export async function fetchGenerationApi(path: string, init: RequestInit = {}): Promise { - const connection = await getApiConnection() + const connection = await getApiConnection(path) if (isTauri() && !connection.capability) { throw new Error('generation server authentication is unavailable') } From 340a71f4d1b4c3e63058dff61dbd9e27fade765d Mon Sep 17 00:00:00 2001 From: brxs Date: Sat, 8 Aug 2026 19:30:20 -0700 Subject: [PATCH 06/14] test: align gateway with finalized render protocol --- src-tauri/src/magenta_gateway.rs | 271 +++++++++++++++++++++++++------ 1 file changed, 217 insertions(+), 54 deletions(-) diff --git a/src-tauri/src/magenta_gateway.rs b/src-tauri/src/magenta_gateway.rs index 2e4beeb..0413206 100644 --- a/src-tauri/src/magenta_gateway.rs +++ b/src-tauri/src/magenta_gateway.rs @@ -12,7 +12,7 @@ use std::collections::BTreeMap; use std::io::{self, Read, Write}; use std::net::{Shutdown, TcpListener, TcpStream}; -use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; +use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::{Arc, Mutex}; use std::time::{Duration, Instant}; @@ -164,6 +164,7 @@ impl ProcessTree for ManagedProcess { struct ManagedRenderWorker { stream: TcpStream, process: Box, + next_sequence: u64, } impl ManagedRenderWorker { @@ -228,7 +229,6 @@ impl WorkerFactory for ManagedWorkerFactory { struct GatewayCore { worker: Mutex>, factory: Arc, - next_sequence: AtomicU64, lifecycle: Mutex>, quiescing: AtomicBool, } @@ -238,7 +238,6 @@ impl GatewayCore { Self { worker: Mutex::new(None), factory, - next_sequence: AtomicU64::new(1), lifecycle: Mutex::new(Arc::new(AtomicBool::new(false))), quiescing: AtomicBool::new(false), } @@ -255,14 +254,6 @@ impl GatewayCore { } } - fn sequence(&self) -> Result { - self.next_sequence - .fetch_update(Ordering::AcqRel, Ordering::Acquire, |value| { - (value < u64::MAX).then_some(value + 1) - }) - .map_err(|_| RenderFailure::protocol("Magenta render sequence is exhausted")) - } - fn render( &self, prompt: String, @@ -276,14 +267,6 @@ impl GatewayCore { if cancellation.cancelled() { return Err(RenderFailure::cancelled()); } - let sequence = self.sequence()?; - let request = WorkerRenderRequest { - schema_version: RENDER_SCHEMA_VERSION, - job_id: format!("render-{:032x}", rand::random::()), - sequence, - prompt, - frames, - }; let mut worker = self .worker .lock() @@ -294,21 +277,31 @@ impl GatewayCore { if worker.is_none() { *worker = Some(self.factory.spawn(&cancellation)?); } + let sequence = worker.as_ref().expect("worker was installed").next_sequence; + let request = WorkerRenderRequest { + schema_version: RENDER_SCHEMA_VERSION, + job_id: format!("render-{:032x}", rand::random::()), + sequence, + prompt, + frames, + }; let result = worker .as_mut() .expect("worker was installed") .render(&request, &cancellation); - if result.is_err() { - if let Some(mut failed) = worker.take() { - if failed.shutdown().is_err() { - // Keep ownership so a later quiesce can retry and, most - // importantly, an installer cannot mistake an uncertain - // process-tree state for "reaped" before a Windows rename. - *worker = Some(failed); - return Err(RenderFailure::protocol( - "Magenta render worker could not be reaped", - )); - } + if result.is_ok() && sequence < u64::MAX { + worker.as_mut().expect("worker was installed").next_sequence = sequence + 1; + return result; + } + if let Some(mut finished) = worker.take() { + if finished.shutdown().is_err() { + // Keep ownership so a later quiesce can retry and, most + // importantly, an installer cannot mistake an uncertain + // process-tree state for "reaped" before a Windows rename. + *worker = Some(finished); + return Err(RenderFailure::protocol( + "Magenta render worker could not be reaped", + )); } } result @@ -473,6 +466,7 @@ struct RenderReady { event: String, model: String, runtime: String, + next_sequence: u64, } #[derive(Deserialize)] @@ -642,35 +636,20 @@ fn spawn_managed_worker( let result = accept_worker(&listener, &mut child, &token, cancellation).and_then(|mut stream| { stream.set_nodelay(true).ok(); - let ready_deadline = Instant::now() + READY_TIMEOUT; - let (frame_type, payload) = read_bounded_frame( + let next_sequence = read_worker_ready( &mut stream, - &[FRAME_STATUS, FRAME_RENDER_ERROR], - MAX_RENDER_METADATA_BYTES, + crate::DEFAULT_MODEL, + "pytorch-cuda", cancellation, - ready_deadline, + Instant::now() + READY_TIMEOUT, )?; - if frame_type == FRAME_RENDER_ERROR { - validate_startup_error(&payload)?; - return Err(RenderFailure::unavailable()); - } - let ready: RenderReady = serde_json::from_slice(&payload) - .map_err(|_| RenderFailure::protocol("Magenta worker readiness is invalid"))?; - if ready.schema_version != RENDER_SCHEMA_VERSION - || ready.event != "render_ready" - || ready.model != crate::DEFAULT_MODEL - || ready.runtime != "pytorch-cuda" - { - return Err(RenderFailure::protocol( - "Magenta worker readiness is invalid", - )); - } - Ok(stream) + Ok((stream, next_sequence)) }); match result { - Ok(stream) => Ok(ManagedRenderWorker { + Ok((stream, next_sequence)) => Ok(ManagedRenderWorker { stream, process: Box::new(ManagedProcess { child }), + next_sequence, }), Err(error) => { let _ = child.force_kill(); @@ -679,6 +658,39 @@ fn spawn_managed_worker( } } +fn read_worker_ready( + stream: &mut TcpStream, + model: &str, + runtime: &str, + cancellation: &RequestCancellation, + deadline: Instant, +) -> Result { + let (frame_type, payload) = read_bounded_frame( + stream, + &[FRAME_STATUS, FRAME_RENDER_ERROR], + MAX_RENDER_METADATA_BYTES, + cancellation, + deadline, + )?; + if frame_type == FRAME_RENDER_ERROR { + validate_startup_error(&payload)?; + return Err(RenderFailure::unavailable()); + } + let ready: RenderReady = serde_json::from_slice(&payload) + .map_err(|_| RenderFailure::protocol("Magenta worker readiness is invalid"))?; + if ready.schema_version != RENDER_SCHEMA_VERSION + || ready.event != "render_ready" + || ready.model != model + || ready.runtime != runtime + || ready.next_sequence != 1 + { + return Err(RenderFailure::protocol( + "Magenta worker readiness is invalid", + )); + } + Ok(ready.next_sequence) +} + fn accept_worker( listener: &TcpListener, child: &mut SupervisedChild, @@ -1137,6 +1149,7 @@ mod tests { spawns: Arc, shutdowns: Arc, shutdown_failures: Arc, + sequences: Arc>>, } impl FakeFactory { @@ -1146,6 +1159,7 @@ mod tests { spawns: Arc::new(AtomicUsize::new(0)), shutdowns: Arc::new(AtomicUsize::new(0)), shutdown_failures: Arc::new(AtomicUsize::new(0)), + sequences: Arc::new(Mutex::new(Vec::new())), }) } @@ -1176,13 +1190,15 @@ mod tests { .set_read_timeout(Some(IO_POLL)) .map_err(RenderFailure::from)?; let (server, _) = listener.accept().map_err(RenderFailure::from)?; - thread::spawn(move || serve_scenario(server, scenario)); + let sequences = self.sequences.clone(); + thread::spawn(move || serve_scenario(server, scenario, sequences)); Ok(ManagedRenderWorker { stream: client, process: Box::new(FakeProcess { shutdowns: self.shutdowns.clone(), shutdown_failures: self.shutdown_failures.clone(), }), + next_sequence: 1, }) } } @@ -1196,12 +1212,16 @@ mod tests { (header[0], payload) } - fn serve_scenario(mut stream: TcpStream, scenario: Scenario) { + fn serve_scenario(mut stream: TcpStream, scenario: Scenario, sequences: Arc>>) { let (frame_type, payload) = read_test_frame(&mut stream); assert_eq!(frame_type, FRAME_RENDER_REQUEST); let request: serde_json::Value = serde_json::from_slice(&payload).unwrap(); let job_id = request["jobId"].as_str().unwrap(); let sequence = request["sequence"].as_u64().unwrap(); + sequences + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + .push(sequence); let frames = request["frames"].as_u64().unwrap(); if matches!(scenario, Scenario::Stall) { thread::sleep(Duration::from_millis(250)); @@ -1338,6 +1358,7 @@ mod tests { assert!(render_with(&core, Arc::new(AtomicBool::new(false))).is_err()); assert!(render_with(&core, Arc::new(AtomicBool::new(false))).is_ok()); assert_eq!(factory.spawns.load(Ordering::Acquire), 2); + assert_eq!(*factory.sequences.lock().unwrap(), [1, 1]); assert_eq!(factory.shutdowns.load(Ordering::Acquire), 1); assert_eq!(core.quiesce(), Ok(true)); assert_eq!(factory.shutdowns.load(Ordering::Acquire), 2); @@ -1402,4 +1423,146 @@ mod tests { .unwrap_err(); assert_eq!(error.kind, FailureKind::Deadline); } + + fn protocol_test_python() -> Option { + let mut candidates = Vec::new(); + if let Some(configured) = std::env::var_os("LSDJ_TEST_PYTHON") { + candidates.push(configured.into()); + } + candidates.push("/opt/homebrew/bin/python3".into()); + candidates.push("python3".into()); + candidates.push("python".into()); + candidates.into_iter().find(|candidate| { + let Ok(output) = std::process::Command::new(candidate) + .arg("--version") + .output() + else { + return false; + }; + let version = format!( + "{}{}", + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr) + ); + let Some(version) = version.split_whitespace().nth(1) else { + return false; + }; + let mut parts = version + .split('.') + .filter_map(|part| part.parse::().ok()); + matches!( + (parts.next(), parts.next()), + (Some(major), Some(minor)) if major > 3 || (major == 3 && minor >= 11) + ) + }) + } + + #[test] + fn rust_gateway_round_trips_two_requests_with_the_real_python_protocol() { + let Some(python) = protocol_test_python() else { + eprintln!("skipping Python protocol compatibility test: Python 3.11+ unavailable"); + return; + }; + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + listener.set_nonblocking(true).unwrap(); + let port = listener.local_addr().unwrap().port(); + let token = crate::local_auth::generate_capability(); + let sidecar = + std::path::Path::new(env!("CARGO_MANIFEST_DIR")).join("../backend/lsdj/sidecar.py"); + let harness = r#" +import importlib.util +import socket +import sys +import types + +package = types.ModuleType("lsdj") +package.__path__ = [] +sys.modules["lsdj"] = package +mrt2 = types.ModuleType("lsdj.mrt2") +mrt2.AUTO_RUNTIME = "auto" +mrt2.PYTORCH_CUDA_RUNTIME = "pytorch-cuda" +mrt2.RUNTIME_CHOICES = ("auto", "mlx", "pytorch-cuda") +mrt2.create_engine = lambda **kwargs: None +mrt2.public_startup_error = lambda error: str(error) +mrt2.runtime_manifest = lambda: {} +sys.modules["lsdj.mrt2"] = mrt2 +worker = types.ModuleType("lsdj.worker") +worker.run_deck_worker = lambda *args, **kwargs: None +sys.modules["lsdj.worker"] = worker +spec = importlib.util.spec_from_file_location("lsdj.sidecar", sys.argv[3]) +sidecar = importlib.util.module_from_spec(spec) +sys.modules["lsdj.sidecar"] = sidecar +spec.loader.exec_module(sidecar) + +class Engine: + def warm_up(self): + pass + def render_clip(self, prompt, seconds): + frames = int(seconds * sidecar.RENDER_SAMPLE_RATE + 0.5) + return b"\0" * (frames * sidecar.RENDER_BYTES_PER_FRAME) + +sock = socket.create_connection(("127.0.0.1", int(sys.argv[1]))) +sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) +sidecar.write_frame(sock, sidecar.FRAME_AUTH, sys.argv[2].encode()) +sidecar.run_render_worker( + sock, + "mrt2_small", + runtime="pytorch-cuda", + engine_factory=lambda model: Engine(), +) +"#; + let mut command = std::process::Command::new(python); + command + .arg("-c") + .arg(harness) + .arg(port.to_string()) + .arg(&token) + .arg(&sidecar); + let mut child = crate::child_process::spawn_grouped(&mut command).unwrap(); + let cancellation = RequestCancellation { + request: Arc::new(AtomicBool::new(false)), + lifecycle: Arc::new(AtomicBool::new(false)), + }; + let mut stream = accept_worker(&listener, &mut child, &token, &cancellation).unwrap(); + let next_sequence = read_worker_ready( + &mut stream, + "mrt2_small", + "pytorch-cuda", + &cancellation, + Instant::now() + Duration::from_secs(5), + ) + .unwrap(); + assert_eq!(next_sequence, 1); + + let core = GatewayCore::new(FakeFactory::new(std::iter::empty())); + *core.worker.lock().unwrap() = Some(ManagedRenderWorker { + stream, + process: Box::new(ManagedProcess { child }), + next_sequence, + }); + let first = core + .render( + "first compatibility render".to_string(), + MIN_RENDER_FRAMES, + Arc::new(AtomicBool::new(false)), + ) + .unwrap(); + let second = core + .render( + "second compatibility render".to_string(), + MIN_RENDER_FRAMES, + Arc::new(AtomicBool::new(false)), + ) + .unwrap(); + assert_eq!( + first.len(), + MIN_RENDER_FRAMES as usize * RENDER_BYTES_PER_FRAME as usize + ); + assert_eq!(second.len(), first.len()); + assert_eq!( + core.worker.lock().unwrap().as_ref().unwrap().next_sequence, + 3 + ); + assert_eq!(core.quiesce(), Ok(true)); + } } From 873002176a15e32ce1b836e561f931c40b677bb6 Mon Sep 17 00:00:00 2001 From: brxs Date: Sat, 8 Aug 2026 20:07:25 -0700 Subject: [PATCH 07/14] fix: retain native supervisors through promotion --- frontend/src/audio/nativeEngine.test.ts | 30 ++ frontend/src/audio/nativeEngine.ts | 13 +- src-tauri/src/generation.rs | 208 +++++++++--- src-tauri/src/lib.rs | 9 +- src-tauri/src/magenta_gateway.rs | 427 ++++++++++++++++++++---- src-tauri/src/mcp.rs | 14 +- src-tauri/src/models.rs | 67 +++- src-tauri/src/sidecar.rs | 161 ++++++--- 8 files changed, 747 insertions(+), 182 deletions(-) diff --git a/frontend/src/audio/nativeEngine.test.ts b/frontend/src/audio/nativeEngine.test.ts index d3138e3..0d93b83 100644 --- a/frontend/src/audio/nativeEngine.test.ts +++ b/frontend/src/audio/nativeEngine.test.ts @@ -194,6 +194,36 @@ describe('createNativeEngine — control contract', () => { expect((init.headers as Headers).get('x-lsdj-capability')).toBe('n'.repeat(64)) }) + it('uses the promoted SA3 connection on the next request without retrying either POST', async () => { + let appInfoCalls = 0 + const invoke = vi.fn((cmd: string) => { + if (cmd !== 'app_info') return Promise.resolve(undefined) + appInfoCalls += 1 + return Promise.resolve( + appInfoCalls === 1 + ? { generationPort: 1111, generationCapability: 'o'.repeat(64) } + : { generationPort: 2222, generationCapability: 'n'.repeat(64) }, + ) + }) + vi.stubGlobal('__TAURI__', { core: { invoke } }) + const fetchMock = vi.fn(async () => ({ ok: true })) + vi.stubGlobal('fetch', fetchMock) + + await fetchGenerationApi('/api/generate', { method: 'POST' }) + await fetchGenerationApi('/api/generate', { method: 'POST' }) + + expect(appInfoCalls).toBe(2) + expect(fetchMock).toHaveBeenCalledTimes(2) + expect(fetchMock.mock.calls[0][0]).toBe('http://127.0.0.1:1111/api/generate') + expect((fetchMock.mock.calls[0][1]?.headers as Headers).get('x-lsdj-capability')).toBe( + 'o'.repeat(64), + ) + expect(fetchMock.mock.calls[1][0]).toBe('http://127.0.0.1:2222/api/generate') + expect((fetchMock.mock.calls[1][1]?.headers as Headers).get('x-lsdj-capability')).toBe( + 'n'.repeat(64), + ) + }) + it('createDeckChannel replays NO mixer config — the shell hydrates (phase C)', async () => { const engine = createNativeEngine() await engine.createDeckChannel( diff --git a/frontend/src/audio/nativeEngine.ts b/frontend/src/audio/nativeEngine.ts index 15ed89f..2fe1078 100644 --- a/frontend/src/audio/nativeEngine.ts +++ b/frontend/src/audio/nativeEngine.ts @@ -115,15 +115,20 @@ function isMagentaPath(path: string): boolean { } async function getApiConnection(path: string): Promise { + const magenta = isMagentaPath(path) + // SA3 is deliberately replaced during managed promotion, including both its + // port and capability. Resolve one fresh atomic app_info snapshot before each + // request; a POST is never retried against either the old or new process. + if (isTauri() && !magenta) apiConnectionPromise = null let connections = await loadApiConnections() - let connection = isMagentaPath(path) ? connections.magenta : connections.sa3 + let connection = magenta ? connections.magenta : connections.sa3 // A fresh managed install legitimately starts without SA3. Do not pin that // absence for the webview lifetime: the first request after promotion // re-reads app_info and reaches the newly started generation server. if (isTauri() && !connection.baseUrl) { apiConnectionPromise = null connections = await loadApiConnections() - connection = isMagentaPath(path) ? connections.magenta : connections.sa3 + connection = magenta ? connections.magenta : connections.sa3 } return connection } @@ -131,8 +136,8 @@ async function getApiConnection(path: string): Promise { /** Base URL for the backend `/api/*` generation endpoints (sa3/Magenta pad+track * render). FastAPI no longer serves the UI, so the Rust shell runs a generation * server on a loopback port it reports via `app_info`; the webview fetches - * `http://127.0.0.1:/api/...`. Resolved once and cached; falls back to '' - * (relative) if the port can't be resolved. */ + * `http://127.0.0.1:/api/...`. SA3 is resolved fresh because promotion + * replaces both the port and capability; missing connections fall back to ''. */ export function getApiBaseUrl(): Promise { return getApiConnection('/api/generate').then((connection) => connection.baseUrl) } diff --git a/src-tauri/src/generation.rs b/src-tauri/src/generation.rs index a4d3b7b..e9b1886 100644 --- a/src-tauri/src/generation.rs +++ b/src-tauri/src/generation.rs @@ -16,7 +16,7 @@ use std::io; use std::net::{TcpListener, TcpStream}; #[cfg(not(feature = "managed-runtime"))] use std::path::Path; -use std::process::Command; +use std::process::{Command, ExitStatus}; use std::sync::Mutex; use std::time::Duration; @@ -26,13 +26,40 @@ use crate::child_process::{Readiness, SupervisedChild}; /// webview via `app_info`) and the child process. Held in Tauri managed state; /// dropping it kills the child. pub struct GenerationServer { - state: Mutex, + state: Mutex, } -struct GenerationState { - port: Option, - capability: Option, - child: Option, +trait GenerationProcess: Send { + fn try_wait(&mut self) -> io::Result>; + fn shutdown(&mut self) -> io::Result<()>; +} + +impl GenerationProcess for SupervisedChild { + fn try_wait(&mut self) -> io::Result> { + SupervisedChild::try_wait(self) + } + + fn shutdown(&mut self) -> io::Result<()> { + let report = SupervisedChild::shutdown(self, Duration::from_millis(500))?; + crate::child_process::log_shutdown("generation server", Ok(report)); + Ok(()) + } +} + +enum GenerationProcessState { + Stopped, + Running { + port: u16, + capability: String, + process: Box, + }, + /// Shutdown was requested, but the supervisor could not prove that the + /// complete process tree was reaped. The handle stays owned here so every + /// later quiesce/resume can retry; an installer must not rename through it. + Uncertain { + process: Box, + was_running: bool, + }, } impl GenerationServer { @@ -41,11 +68,7 @@ impl GenerationServer { /// webview surfaces that as fetch errors). pub fn start() -> GenerationServer { let server = GenerationServer { - state: Mutex::new(GenerationState { - port: None, - capability: None, - child: None, - }), + state: Mutex::new(GenerationProcessState::Stopped), }; if let Err(error) = server.resume() { // A fresh managed install intentionally has no runtime yet. The @@ -55,7 +78,7 @@ impl GenerationServer { server } - fn spawn(capability: &str) -> io::Result<(u16, SupervisedChild)> { + fn spawn(capability: &str) -> io::Result<(u16, Box)> { // Pick a free loopback port, then hand it to the child (uvicorn binds it). // The brief drop→rebind window on loopback is benign. let port = { @@ -77,7 +100,7 @@ impl GenerationServer { Readiness::Ready | Readiness::TimedOut => { // Preserve the existing macOS contract: a slow-but-running // service is advertised optimistically after the bounded wait. - Ok((port, child)) + Ok((port, Box::new(child))) } Readiness::Exited(status) => Err(io::Error::other(format!( "generation server exited before binding ({status})" @@ -85,22 +108,19 @@ impl GenerationServer { } } - /// The loopback port the generation server bound, or `None` if disabled / not - /// running. The webview reads this through `app_info` to build the API base URL. - pub fn port(&self) -> Option { - self.state - .lock() - .unwrap_or_else(|poisoned| poisoned.into_inner()) - .port - } - - /// The in-memory capability paired with [`port`](Self::port). Never persisted. - pub fn capability(&self) -> Option { - self.state + /// Return the port and capability from one lock acquisition. Neither half is + /// ever observable without the other across promotion/resume transitions. + pub fn connection(&self) -> Option<(u16, String)> { + match &*self + .state .lock() .unwrap_or_else(|poisoned| poisoned.into_inner()) - .capability - .clone() + { + GenerationProcessState::Running { + port, capability, .. + } => Some((*port, capability.clone())), + GenerationProcessState::Stopped | GenerationProcessState::Uncertain { .. } => None, + } } /// Start (or recover) the service from the currently promoted verified @@ -111,20 +131,52 @@ impl GenerationServer { .state .lock() .unwrap_or_else(|poisoned| poisoned.into_inner()); - if let Some(child) = state.child.as_mut() { - if child.try_wait()?.is_none() { - return Ok(()); + let previous = std::mem::replace(&mut *state, GenerationProcessState::Stopped); + match previous { + GenerationProcessState::Stopped => {} + GenerationProcessState::Running { + port, + capability, + mut process, + } => match process.try_wait() { + Ok(None) => { + *state = GenerationProcessState::Running { + port, + capability, + process, + }; + return Ok(()); + } + Ok(Some(_)) => {} + Err(error) => { + *state = GenerationProcessState::Uncertain { + process, + was_running: true, + }; + return Err(error); + } + }, + GenerationProcessState::Uncertain { + mut process, + was_running, + } => { + if let Err(error) = process.shutdown() { + *state = GenerationProcessState::Uncertain { + process, + was_running, + }; + return Err(error); + } } - state.child = None; - state.port = None; - state.capability = None; } let capability = crate::local_auth::generate_capability(); - let (port, child) = Self::spawn(&capability)?; + let (port, process) = Self::spawn(&capability)?; println!("lsdj-app: generation server on 127.0.0.1:{port}"); - state.port = Some(port); - state.capability = Some(capability); - state.child = Some(child); + *state = GenerationProcessState::Running { + port, + capability, + process, + }; Ok(()) } @@ -136,14 +188,25 @@ impl GenerationServer { .state .lock() .unwrap_or_else(|poisoned| poisoned.into_inner()); - state.port = None; - state.capability = None; - let Some(mut child) = state.child.take() else { - return Ok(false); + let previous = std::mem::replace(&mut *state, GenerationProcessState::Stopped); + let (mut process, was_running) = match previous { + GenerationProcessState::Stopped => return Ok(false), + GenerationProcessState::Running { process, .. } => (process, true), + GenerationProcessState::Uncertain { + process, + was_running, + } => (process, was_running), }; - let report = child.shutdown(Duration::from_millis(500))?; - crate::child_process::log_shutdown("generation server", Ok(report)); - Ok(true) + match process.shutdown() { + Ok(()) => Ok(was_running), + Err(error) => { + *state = GenerationProcessState::Uncertain { + process, + was_running, + }; + Err(error) + } + } } /// Kill the generation server child. Called explicitly from the app's @@ -254,8 +317,63 @@ mod tests { // Now-always-on `start()` never fails the app: a command that exits without // binding the port (echo) degrades to no advertised port. let server = GenerationServer::start(); - assert_eq!(server.port(), None); + assert_eq!(server.connection(), None); std::env::remove_var("LSDJ_GENERATION_CMD"); } } + +#[cfg(test)] +mod process_tests { + use std::sync::atomic::{AtomicUsize, Ordering}; + use std::sync::Arc; + + use super::*; + + struct RetryProcess { + shutdowns: Arc, + } + + impl GenerationProcess for RetryProcess { + fn try_wait(&mut self) -> io::Result> { + Ok(None) + } + + fn shutdown(&mut self) -> io::Result<()> { + if self.shutdowns.fetch_add(1, Ordering::AcqRel) == 0 { + Err(io::Error::other("first reap is uncertain")) + } else { + Ok(()) + } + } + } + + #[test] + fn failed_sa3_reap_retains_ownership_until_a_positive_retry() { + let shutdowns = Arc::new(AtomicUsize::new(0)); + let server = GenerationServer { + state: Mutex::new(GenerationProcessState::Running { + port: 4321, + capability: "capability".to_string(), + process: Box::new(RetryProcess { + shutdowns: shutdowns.clone(), + }), + }), + }; + + assert!(server.quiesce().is_err()); + assert_eq!(server.connection(), None); + assert!(matches!( + &*server.state.lock().unwrap(), + GenerationProcessState::Uncertain { .. } + )); + assert_eq!(shutdowns.load(Ordering::Acquire), 1); + + assert!(matches!(server.quiesce(), Ok(true))); + assert!(matches!( + &*server.state.lock().unwrap(), + GenerationProcessState::Stopped + )); + assert_eq!(shutdowns.load(Ordering::Acquire), 2); + } +} diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 850b9e2..3b956c9 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -325,20 +325,23 @@ fn app_info( generation: tauri::State<'_, generation::GenerationServer>, mcp: tauri::State<'_, mcp::McpServer>, ) -> AppInfo { + let generation_connection = generation.connection(); #[cfg(feature = "managed-runtime")] let (magenta_port, magenta_capability) = { let gateway = app.state::(); (gateway.port(), gateway.capability()) }; #[cfg(not(feature = "managed-runtime"))] - let (magenta_port, magenta_capability) = (generation.port(), generation.capability()); + let (magenta_port, magenta_capability) = generation_connection + .clone() + .map_or((None, None), |(port, capability)| (Some(port), Some(capability))); #[cfg(not(feature = "managed-runtime"))] let _ = app; AppInfo { version: env!("CARGO_PKG_VERSION").to_string(), audio_device_started: state.device_started, - generation_port: generation.port(), - generation_capability: generation.capability(), + generation_port: generation_connection.as_ref().map(|(port, _)| *port), + generation_capability: generation_connection.map(|(_, capability)| capability), magenta_port, magenta_capability, mcp_port: mcp.port(), diff --git a/src-tauri/src/magenta_gateway.rs b/src-tauri/src/magenta_gateway.rs index 0413206..92cd093 100644 --- a/src-tauri/src/magenta_gateway.rs +++ b/src-tauri/src/magenta_gateway.rs @@ -201,18 +201,33 @@ impl ManagedRenderWorker { result => result, } } - - fn shutdown(&mut self) -> io::Result<()> { - let _ = self.stream.shutdown(Shutdown::Both); - self.process.shutdown() - } } trait WorkerFactory: Send + Sync { fn spawn( &self, cancellation: &RequestCancellation, - ) -> Result; + ) -> Result; +} + +struct WorkerSpawnFailure { + failure: RenderFailure, + uncertain_process: Option>, +} + +impl WorkerSpawnFailure { + fn reaped(failure: RenderFailure) -> Self { + Self { + failure, + uncertain_process: None, + } + } +} + +impl From for WorkerSpawnFailure { + fn from(failure: RenderFailure) -> Self { + Self::reaped(failure) + } } struct ManagedWorkerFactory; @@ -221,13 +236,24 @@ impl WorkerFactory for ManagedWorkerFactory { fn spawn( &self, cancellation: &RequestCancellation, - ) -> Result { + ) -> Result { spawn_managed_worker(cancellation) } } +enum WorkerState { + Stopped, + Running(ManagedRenderWorker), + /// A failed teardown left process-tree ownership uncertain. No new worker + /// may launch, and promotion may not rename, until shutdown later succeeds. + Uncertain { + process: Box, + was_warm: bool, + }, +} + struct GatewayCore { - worker: Mutex>, + worker: Mutex, factory: Arc, lifecycle: Mutex>, quiescing: AtomicBool, @@ -236,7 +262,7 @@ struct GatewayCore { impl GatewayCore { fn new(factory: Arc) -> Self { Self { - worker: Mutex::new(None), + worker: Mutex::new(WorkerState::Stopped), factory, lifecycle: Mutex::new(Arc::new(AtomicBool::new(false))), quiescing: AtomicBool::new(false), @@ -274,10 +300,30 @@ impl GatewayCore { if cancellation.cancelled() || self.quiescing.load(Ordering::Acquire) { return Err(RenderFailure::cancelled()); } - if worker.is_none() { - *worker = Some(self.factory.spawn(&cancellation)?); + if matches!(&*worker, WorkerState::Stopped) { + match self.factory.spawn(&cancellation) { + Ok(spawned) => *worker = WorkerState::Running(spawned), + Err(spawn) => { + if let Some(process) = spawn.uncertain_process { + *worker = WorkerState::Uncertain { + process, + was_warm: false, + }; + } + return Err(spawn.failure); + } + } } - let sequence = worker.as_ref().expect("worker was installed").next_sequence; + let resident = match &mut *worker { + WorkerState::Running(resident) => resident, + WorkerState::Uncertain { .. } => { + return Err(RenderFailure::protocol( + "Magenta render worker could not be reaped", + )) + } + WorkerState::Stopped => unreachable!("worker spawn installed a running state"), + }; + let sequence = resident.next_sequence; let request = WorkerRenderRequest { schema_version: RENDER_SCHEMA_VERSION, job_id: format!("render-{:032x}", rand::random::()), @@ -285,24 +331,25 @@ impl GatewayCore { prompt, frames, }; - let result = worker - .as_mut() - .expect("worker was installed") - .render(&request, &cancellation); + let result = resident.render(&request, &cancellation); if result.is_ok() && sequence < u64::MAX { - worker.as_mut().expect("worker was installed").next_sequence = sequence + 1; + resident.next_sequence = sequence + 1; return result; } - if let Some(mut finished) = worker.take() { - if finished.shutdown().is_err() { - // Keep ownership so a later quiesce can retry and, most - // importantly, an installer cannot mistake an uncertain - // process-tree state for "reaped" before a Windows rename. - *worker = Some(finished); - return Err(RenderFailure::protocol( - "Magenta render worker could not be reaped", - )); - } + let WorkerState::Running(mut finished) = + std::mem::replace(&mut *worker, WorkerState::Stopped) + else { + unreachable!("render state stayed running while its lock was held") + }; + let _ = finished.stream.shutdown(Shutdown::Both); + if finished.process.shutdown().is_err() { + *worker = WorkerState::Uncertain { + process: finished.process, + was_warm: true, + }; + return Err(RenderFailure::protocol( + "Magenta render worker could not be reaped", + )); } result } @@ -319,13 +366,19 @@ impl GatewayCore { .worker .lock() .unwrap_or_else(|poisoned| poisoned.into_inner()); - let Some(mut resident) = worker.take() else { - return Ok(false); + let previous = std::mem::replace(&mut *worker, WorkerState::Stopped); + let (mut process, was_warm) = match previous { + WorkerState::Stopped => return Ok(false), + WorkerState::Running(resident) => { + let _ = resident.stream.shutdown(Shutdown::Both); + (resident.process, true) + } + WorkerState::Uncertain { process, was_warm } => (process, was_warm), }; - match resident.shutdown() { - Ok(()) => Ok(true), + match process.shutdown() { + Ok(()) => Ok(was_warm), Err(_) => { - *worker = Some(resident); + *worker = WorkerState::Uncertain { process, was_warm }; Err("Magenta render worker could not be reaped".to_string()) } } @@ -347,12 +400,24 @@ impl GatewayCore { .worker .lock() .unwrap_or_else(|poisoned| poisoned.into_inner()); - if worker.is_none() { - *worker = Some( - self.factory - .spawn(&cancellation) - .map_err(|error| error.detail.to_string())?, - ); + match &*worker { + WorkerState::Running(_) => return Ok(()), + WorkerState::Uncertain { .. } => { + return Err("Magenta render worker could not be reaped".to_string()) + } + WorkerState::Stopped => {} + } + match self.factory.spawn(&cancellation) { + Ok(spawned) => *worker = WorkerState::Running(spawned), + Err(spawn) => { + if let Some(process) = spawn.uncertain_process { + *worker = WorkerState::Uncertain { + process, + was_warm: false, + }; + } + return Err(spawn.failure.detail.to_string()); + } } Ok(()) } @@ -617,21 +682,22 @@ fn float32_wav(pcm: &[u8]) -> Result, ()> { fn spawn_managed_worker( cancellation: &RequestCancellation, -) -> Result { - let listener = TcpListener::bind("127.0.0.1:0").map_err(|_| RenderFailure::unavailable())?; +) -> Result { + let listener = TcpListener::bind("127.0.0.1:0") + .map_err(|_| WorkerSpawnFailure::reaped(RenderFailure::unavailable()))?; listener .set_nonblocking(true) - .map_err(|_| RenderFailure::unavailable())?; + .map_err(|_| WorkerSpawnFailure::reaped(RenderFailure::unavailable()))?; let port = listener .local_addr() - .map_err(|_| RenderFailure::unavailable())? + .map_err(|_| WorkerSpawnFailure::reaped(RenderFailure::unavailable()))? .port(); let token = crate::local_auth::generate_capability(); let mut command = crate::sidecar::authenticated_render_worker_command(crate::DEFAULT_MODEL, port, &token) - .map_err(|_| RenderFailure::unavailable())?; + .map_err(|_| WorkerSpawnFailure::reaped(RenderFailure::unavailable()))?; let mut child = crate::child_process::spawn_grouped(&mut command) - .map_err(|_| RenderFailure::unavailable())?; + .map_err(|_| WorkerSpawnFailure::reaped(RenderFailure::unavailable()))?; let result = accept_worker(&listener, &mut child, &token, cancellation).and_then(|mut stream| { @@ -652,8 +718,16 @@ fn spawn_managed_worker( next_sequence, }), Err(error) => { - let _ = child.force_kill(); - Err(error) + let mut process: Box = Box::new(ManagedProcess { child }); + match process.shutdown() { + Ok(()) => Err(WorkerSpawnFailure::reaped(error)), + Err(_) => Err(WorkerSpawnFailure { + failure: RenderFailure::protocol( + "Magenta render worker startup failed and could not be reaped", + ), + uncertain_process: Some(process), + }), + } } } } @@ -963,15 +1037,7 @@ fn serve( capability: &str, core: Arc, ) -> CancellationToken { - let auth = AuthState { - capability: Arc::from(capability), - }; - let router = Router::new() - .route("/api/render", post(render_clip).options(preflight)) - .route("/api/models", get(model_info).options(preflight)) - .layer(DefaultBodyLimit::max(MAX_RENDER_REQUEST_BYTES)) - .layer(axum::middleware::from_fn_with_state(auth, authenticate)) - .with_state(HttpState { core }); + let router = gateway_router(capability, core); let cancel = CancellationToken::new(); let serve_cancel = cancel.clone(); tauri::async_runtime::spawn(async move { @@ -993,6 +1059,18 @@ fn serve( cancel } +fn gateway_router(capability: &str, core: Arc) -> Router { + let auth = AuthState { + capability: Arc::from(capability), + }; + Router::new() + .route("/api/render", post(render_clip).options(preflight)) + .route("/api/models", get(model_info).options(preflight)) + .layer(DefaultBodyLimit::max(MAX_RENDER_REQUEST_BYTES)) + .layer(axum::middleware::from_fn_with_state(auth, authenticate)) + .with_state(HttpState { core }) +} + async fn preflight() -> StatusCode { StatusCode::NO_CONTENT } @@ -1121,6 +1199,7 @@ mod tests { OversizeEnd, MisalignedChunk, Stall, + LongStall, } struct FakeProcess { @@ -1172,7 +1251,7 @@ mod tests { fn spawn( &self, _cancellation: &RequestCancellation, - ) -> Result { + ) -> Result { let scenario = self .scenarios .lock() @@ -1223,8 +1302,13 @@ mod tests { .unwrap_or_else(|poisoned| poisoned.into_inner()) .push(sequence); let frames = request["frames"].as_u64().unwrap(); - if matches!(scenario, Scenario::Stall) { - thread::sleep(Duration::from_millis(250)); + if matches!(scenario, Scenario::Stall | Scenario::LongStall) { + let delay = if matches!(scenario, Scenario::LongStall) { + Duration::from_secs(2) + } else { + Duration::from_millis(250) + }; + thread::sleep(delay); return; } if matches!(scenario, Scenario::OutOfOrder) { @@ -1347,7 +1431,10 @@ mod tests { let error = render_with(&core, Arc::new(AtomicBool::new(false))).unwrap_err(); assert_eq!(error.kind, FailureKind::Protocol); assert_eq!(factory.shutdowns.load(Ordering::Acquire), 1); - assert!(core.worker.lock().unwrap().is_none()); + assert!(matches!( + &*core.worker.lock().unwrap(), + WorkerState::Stopped + )); } } @@ -1372,14 +1459,214 @@ mod tests { factory.fail_shutdowns(1); assert!(core.quiesce().is_err()); - assert!(core.worker.lock().unwrap().is_some()); + assert!(matches!( + &*core.worker.lock().unwrap(), + WorkerState::Uncertain { .. } + )); assert_eq!(factory.shutdowns.load(Ordering::Acquire), 1); assert_eq!(core.quiesce(), Ok(true)); - assert!(core.worker.lock().unwrap().is_none()); + assert!(matches!( + &*core.worker.lock().unwrap(), + WorkerState::Stopped + )); assert_eq!(factory.shutdowns.load(Ordering::Acquire), 2); } + struct FailedStartupFactory { + shutdowns: Arc, + } + + impl WorkerFactory for FailedStartupFactory { + fn spawn( + &self, + _cancellation: &RequestCancellation, + ) -> Result { + // The production factory reaches this shape only after failed + // accept/readiness cleanup. Count that first failed reap here; the + // retained process succeeds when quiesce retries it. + self.shutdowns.store(1, Ordering::Release); + Err(WorkerSpawnFailure { + failure: RenderFailure::protocol( + "Magenta render worker startup failed and could not be reaped", + ), + uncertain_process: Some(Box::new(FakeProcess { + shutdowns: self.shutdowns.clone(), + shutdown_failures: Arc::new(AtomicUsize::new(0)), + })), + }) + } + } + + #[test] + fn failed_startup_cleanup_retains_process_until_positive_reap() { + let shutdowns = Arc::new(AtomicUsize::new(0)); + let core = GatewayCore::new(Arc::new(FailedStartupFactory { + shutdowns: shutdowns.clone(), + })); + + assert!(render_with(&core, Arc::new(AtomicBool::new(false))).is_err()); + assert!(matches!( + &*core.worker.lock().unwrap(), + WorkerState::Uncertain { .. } + )); + assert_eq!(shutdowns.load(Ordering::Acquire), 1); + + assert_eq!(core.quiesce(), Ok(false)); + assert!(matches!( + &*core.worker.lock().unwrap(), + WorkerState::Stopped + )); + assert_eq!(shutdowns.load(Ordering::Acquire), 2); + } + + async fn host_test_router( + core: Arc, + capability: &str, + ) -> (std::net::SocketAddr, CancellationToken) { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let cancel = CancellationToken::new(); + let serve_cancel = cancel.clone(); + let router = gateway_router(capability, core); + tokio::spawn(async move { + axum::serve(listener, router) + .with_graceful_shutdown(async move { serve_cancel.cancelled().await }) + .await + .unwrap(); + }); + (address, cancel) + } + + #[tokio::test(flavor = "multi_thread")] + async fn router_enforces_auth_cors_body_and_request_bounds() { + let capability = "c".repeat(64); + let core = Arc::new(GatewayCore::new(FakeFactory::new([]))); + let (address, cancel) = host_test_router(core, &capability).await; + let client = reqwest::Client::new(); + let render = format!("http://{address}/api/render"); + + assert_eq!( + client.get(&render).send().await.unwrap().status(), + StatusCode::UNAUTHORIZED + ); + assert_eq!( + client + .get(&render) + .header("x-lsdj-capability", "wrong") + .send() + .await + .unwrap() + .status(), + StatusCode::UNAUTHORIZED + ); + assert_eq!( + client + .get(&render) + .header("x-lsdj-capability", &capability) + .header(header::ORIGIN, "https://hostile.example") + .send() + .await + .unwrap() + .status(), + StatusCode::FORBIDDEN + ); + + let allowed = client + .get(&render) + .header("x-lsdj-capability", &capability) + .header(header::ORIGIN, SAFE_ORIGINS[0]) + .send() + .await + .unwrap(); + assert_eq!(allowed.status(), StatusCode::METHOD_NOT_ALLOWED); + assert_eq!( + allowed.headers().get(header::ACCESS_CONTROL_ALLOW_ORIGIN), + Some(&HeaderValue::from_static("tauri://localhost")) + ); + + let preflight = client + .request(Method::OPTIONS, &render) + .header(header::ORIGIN, SAFE_ORIGINS[0]) + .header(header::ACCESS_CONTROL_REQUEST_METHOD, "POST") + .header( + header::ACCESS_CONTROL_REQUEST_HEADERS, + "content-type, x-lsdj-capability", + ) + .send() + .await + .unwrap(); + assert_eq!(preflight.status(), StatusCode::NO_CONTENT); + assert_eq!( + preflight.headers().get(header::ACCESS_CONTROL_ALLOW_ORIGIN), + Some(&HeaderValue::from_static("tauri://localhost")) + ); + + let oversized = client + .post(&render) + .header("x-lsdj-capability", &capability) + .header(header::CONTENT_TYPE, "application/json") + .body(vec![b'x'; MAX_RENDER_REQUEST_BYTES + 1]) + .send() + .await + .unwrap(); + assert_eq!(oversized.status(), StatusCode::PAYLOAD_TOO_LARGE); + + for invalid in [ + serde_json::json!({"prompt": "ok", "seconds": 2.0, "extra": true}), + serde_json::json!({"prompt": " ", "seconds": 2.0}), + serde_json::json!({"prompt": "x".repeat(MAX_RENDER_PROMPT_CHARS + 1), "seconds": 2.0}), + serde_json::json!({"prompt": "ok", "seconds": 0.49}), + serde_json::json!({"prompt": "ok", "seconds": 180.01}), + ] { + let response = client + .post(&render) + .header("x-lsdj-capability", &capability) + .json(&invalid) + .send() + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::UNPROCESSABLE_ENTITY); + } + cancel.cancel(); + } + + #[tokio::test(flavor = "multi_thread")] + async fn real_http_disconnect_cancels_kills_and_reaps_worker() { + let capability = "d".repeat(64); + let factory = FakeFactory::new([Scenario::LongStall]); + let core = Arc::new(GatewayCore::new(factory.clone())); + let (address, cancel) = host_test_router(core.clone(), &capability).await; + let body = br#"{"prompt":"disconnect me","seconds":2.0}"#; + let mut stream = TcpStream::connect(address).unwrap(); + write!( + stream, + "POST /api/render HTTP/1.1\r\nHost: {address}\r\nContent-Type: application/json\r\nx-lsdj-capability: {capability}\r\nContent-Length: {}\r\n\r\n", + body.len() + ) + .unwrap(); + stream.write_all(body).unwrap(); + stream.flush().unwrap(); + + let spawn_deadline = Instant::now() + Duration::from_secs(1); + while factory.spawns.load(Ordering::Acquire) == 0 && Instant::now() < spawn_deadline { + tokio::time::sleep(Duration::from_millis(10)).await; + } + assert_eq!(factory.spawns.load(Ordering::Acquire), 1); + drop(stream); + + let reap_deadline = Instant::now() + Duration::from_secs(1); + while factory.shutdowns.load(Ordering::Acquire) == 0 && Instant::now() < reap_deadline { + tokio::time::sleep(Duration::from_millis(10)).await; + } + assert_eq!(factory.shutdowns.load(Ordering::Acquire), 1); + assert!(matches!( + &*core.worker.lock().unwrap(), + WorkerState::Stopped + )); + cancel.cancel(); + } + #[test] fn cancellation_interrupts_a_stalled_worker_and_reaps_it() { let factory = FakeFactory::new([Scenario::Stall]); @@ -1393,7 +1680,10 @@ mod tests { let error = render_with(&core, cancellation).unwrap_err(); assert_eq!(error.kind, FailureKind::Cancelled); assert_eq!(factory.shutdowns.load(Ordering::Acquire), 1); - assert!(core.worker.lock().unwrap().is_none()); + assert!(matches!( + &*core.worker.lock().unwrap(), + WorkerState::Stopped + )); } #[test] @@ -1535,7 +1825,7 @@ sidecar.run_render_worker( assert_eq!(next_sequence, 1); let core = GatewayCore::new(FakeFactory::new(std::iter::empty())); - *core.worker.lock().unwrap() = Some(ManagedRenderWorker { + *core.worker.lock().unwrap() = WorkerState::Running(ManagedRenderWorker { stream, process: Box::new(ManagedProcess { child }), next_sequence, @@ -1559,10 +1849,13 @@ sidecar.run_render_worker( MIN_RENDER_FRAMES as usize * RENDER_BYTES_PER_FRAME as usize ); assert_eq!(second.len(), first.len()); - assert_eq!( - core.worker.lock().unwrap().as_ref().unwrap().next_sequence, - 3 - ); + assert!(matches!( + &*core.worker.lock().unwrap(), + WorkerState::Running(ManagedRenderWorker { + next_sequence: 3, + .. + }) + )); assert_eq!(core.quiesce(), Ok(true)); } } diff --git a/src-tauri/src/mcp.rs b/src-tauri/src/mcp.rs index be16dc3..99e2b08 100644 --- a/src-tauri/src/mcp.rs +++ b/src-tauri/src/mcp.rs @@ -931,21 +931,11 @@ impl McpHandler { #[cfg(not(feature = "managed-runtime"))] { let generation = self.app.state::(); - ( - generation.port().ok_or("the generation server is not running")?, - generation - .capability() - .ok_or("the generation server authentication capability is unavailable")?, - ) + generation.connection().ok_or("the generation server is not running")? } } else { let generation = self.app.state::(); - ( - generation.port().ok_or("the generation server is not running")?, - generation - .capability() - .ok_or("the generation server authentication capability is unavailable")?, - ) + generation.connection().ok_or("the generation server is not running")? }; // sa3 generation is serialised; a full track (medium model) can take minutes, // so allow generous headroom but never wait forever for a wedged worker. diff --git a/src-tauri/src/models.rs b/src-tauri/src/models.rs index 25b83be..57640a4 100644 --- a/src-tauri/src/models.rs +++ b/src-tauri/src/models.rs @@ -1327,14 +1327,27 @@ fn quiesce_mrt2_services(app: &AppHandle) -> Result { if let Err(error) = app.state::().quiesce_shared() { // No rename has happened. Restore the still-current verified render // generation before returning the deck teardown error. - let _ = gateway.resume(render_was_warm); - return Err(format!( - "cannot quiesce realtime MRT2 decks before promotion: {error}" + return Err(mrt2_quiesce_recovery_error( + error, + gateway.resume(render_was_warm), )); } Ok(Mrt2Lifecycle { render_was_warm }) } +#[cfg(feature = "managed-runtime")] +fn mrt2_quiesce_recovery_error(shared_error: String, gateway_resume: Result<(), String>) -> String { + let primary = format!( + "cannot quiesce realtime MRT2 decks before promotion: {shared_error}" + ); + match gateway_resume { + Ok(()) => primary, + Err(resume_error) => format!( + "{primary}; the Magenta renderer also could not resume: {resume_error}" + ), + } +} + #[cfg(feature = "managed-runtime")] fn resume_mrt2_services(app: &AppHandle, lifecycle: Mrt2Lifecycle) -> Result<(), String> { // Restore a previously warm renderer completely before launching the deck @@ -3547,6 +3560,54 @@ mod tests { assert_eq!(*events.borrow(), ["quiesce"]); } + #[cfg(feature = "managed-runtime")] + #[test] + fn promotion_waits_for_a_second_positive_reap_before_rename() { + use std::cell::{Cell, RefCell}; + + let attempts = Cell::new(0usize); + let events = RefCell::new(Vec::new()); + let run = || { + run_promotion_lifecycle( + "fake runtime", + || { + events.borrow_mut().push("quiesce"); + let attempt = attempts.get(); + attempts.set(attempt + 1); + if attempt == 0 { + Err("process reap is uncertain".to_string()) + } else { + Ok(()) + } + }, + || { + events.borrow_mut().push("promote"); + Ok(()) + }, + |_| { + events.borrow_mut().push("resume"); + Ok(()) + }, + ) + }; + + assert!(run().is_err()); + assert_eq!(*events.borrow(), ["quiesce"]); + assert!(run().is_ok()); + assert_eq!(*events.borrow(), ["quiesce", "quiesce", "promote", "resume"]); + } + + #[cfg(feature = "managed-runtime")] + #[test] + fn shared_quiesce_and_gateway_recovery_errors_are_both_reported() { + let error = mrt2_quiesce_recovery_error( + "shared process is not reaped".to_string(), + Err("gateway spawn failed".to_string()), + ); + assert!(error.contains("shared process is not reaped")); + assert!(error.contains("gateway spawn failed")); + } + #[test] fn materialized_backend_imports_are_isolated_and_missing_modules_fail_closed() { let root = std::env::temp_dir().join(format!( diff --git a/src-tauri/src/sidecar.rs b/src-tauri/src/sidecar.rs index e018958..5a7ba52 100644 --- a/src-tauri/src/sidecar.rs +++ b/src-tauri/src/sidecar.rs @@ -278,11 +278,31 @@ struct SharedReaderExit { struct SharedReaderParts { control: Arc>>, - child: Arc>>, + process: Arc>, stop: Arc, reader: JoinHandle, } +trait SharedProcess: Send { + fn shutdown(&mut self) -> io::Result<()>; +} + +impl SharedProcess for SupervisedChild { + fn shutdown(&mut self) -> io::Result<()> { + let report = SupervisedChild::shutdown(self, Duration::from_millis(500))?; + crate::child_process::log_shutdown("shared sidecar", Ok(report)); + Ok(()) + } +} + +enum SharedProcessState { + Stopped, + Running(Box), + /// Teardown did not positively reap the process tree. Ownership is retained + /// and promotion remains blocked until a later retry reaches `Stopped`. + Uncertain(Box), +} + /// One supervised deck sidecar: the spawned Python process, the control writer /// (engine → sidecar), and the reader thread (sidecar → engine). Dropping it /// stops the reader, closes the socket, and kills the child. @@ -487,7 +507,7 @@ fn start_shared_reader( .expect("failed to spawn shared LSDJ sidecar reader thread"); SharedReaderParts { control, - child: Arc::new(Mutex::new(Some(child))), + process: Arc::new(Mutex::new(SharedProcessState::Running(Box::new(child)))), stop, reader, } @@ -648,7 +668,7 @@ pub struct SharedSidecar { feed: AnalysisFeed, on_status: SharedStatusSinks, control: Arc>>, - child: Arc>>, + process: Arc>, stop: Arc, reader: Option>, /// Reclaimed ring producers parked after a replacement launch failure. A @@ -673,7 +693,7 @@ impl SharedSidecar { feed, on_status: on_status.map(|sink| Arc::new(Mutex::new(sink))), control: Arc::new(Mutex::new(None)), - child: Arc::new(Mutex::new(None)), + process: Arc::new(Mutex::new(SharedProcessState::Stopped)), stop: Arc::new(AtomicBool::new(true)), reader: None, parked: Some(SharedReaderExit { handles }), @@ -704,8 +724,14 @@ impl SharedSidecar { /// occur inside `bind_and_launch_shared` immediately before spawn, so an /// install that completed after app startup becomes usable without restart. pub fn activate(&mut self) -> io::Result<()> { - if self.reader.is_some() { - return Ok(()); + match &*self.process.lock().unwrap_or_else(|p| p.into_inner()) { + SharedProcessState::Running(_) => return Ok(()), + SharedProcessState::Uncertain(_) => { + return Err(io::Error::other( + "shared sidecar process reap is still uncertain", + )) + } + SharedProcessState::Stopped => {} } if self.parked.is_none() { return Err(io::Error::other( @@ -730,7 +756,7 @@ impl SharedSidecar { on_pcm, ); self.control = parts.control; - self.child = parts.child; + self.process = parts.process; self.stop = parts.stop; self.reader = Some(parts.reader); Ok(()) @@ -741,9 +767,6 @@ impl SharedSidecar { /// succeeds (notably on Windows, where a live Python process holds DLLs). #[cfg(feature = "managed-runtime")] pub fn quiesce(&mut self) -> io::Result<()> { - if self.reader.is_none() { - return Ok(()); - } let exit = self.stop_and_reclaim()?; self.parked = Some(exit); Ok(()) @@ -839,17 +862,13 @@ impl SharedSidecar { ); self.models = models; self.control = parts.control; - self.child = parts.child; + self.process = parts.process; self.stop = parts.stop; self.reader = Some(parts.reader); Ok(()) } fn stop_and_reclaim(&mut self) -> io::Result { - if let Some(exit) = self.parked.take() { - return Ok(exit); - } - self.stop.store(true, Ordering::Release); if let Some(writer) = self .control @@ -859,28 +878,18 @@ impl SharedSidecar { { let _ = writer.shutdown(std::net::Shutdown::Both); } - let mut shutdown_error = None; - if let Some(mut old) = self.child.lock().unwrap_or_else(|p| p.into_inner()).take() { - match old.shutdown(Duration::from_millis(500)) { - Ok(report) => { - crate::child_process::log_shutdown("shared sidecar restart", Ok(report)) - } - Err(error) => { - if let Err(force_error) = old.force_kill() { - shutdown_error = Some(io::Error::other(format!( - "cannot reap old shared CUDA worker ({error}); forced teardown also failed ({force_error})" - ))); - } - } - } - } - let exit = self - .reader - .take() - .ok_or_else(|| io::Error::other("shared sidecar has no reader to reclaim"))? - .join() - .map_err(|_| io::Error::other("shared sidecar reader thread panicked"))?; - if let Some(error) = shutdown_error { + let shutdown_result = + stop_shared_process(&mut self.process.lock().unwrap_or_else(|p| p.into_inner())); + let exit = if let Some(reader) = self.reader.take() { + reader + .join() + .map_err(|_| io::Error::other("shared sidecar reader thread panicked"))? + } else { + self.parked + .take() + .ok_or_else(|| io::Error::other("shared sidecar has no deck handles to reclaim"))? + }; + if let Err(error) = shutdown_result { self.parked = Some(exit); return Err(error); } @@ -888,6 +897,21 @@ impl SharedSidecar { } } +fn stop_shared_process(state: &mut SharedProcessState) -> io::Result<()> { + let previous = std::mem::replace(state, SharedProcessState::Stopped); + let mut process = match previous { + SharedProcessState::Stopped => return Ok(()), + SharedProcessState::Running(process) | SharedProcessState::Uncertain(process) => process, + }; + match process.shutdown() { + Ok(()) => Ok(()), + Err(error) => { + *state = SharedProcessState::Uncertain(process); + Err(error) + } + } +} + impl Drop for SharedSidecar { fn drop(&mut self) { self.stop.store(true, Ordering::Release); @@ -899,11 +923,10 @@ impl Drop for SharedSidecar { { let _ = writer.shutdown(std::net::Shutdown::Both); } - if let Some(mut child) = self.child.lock().unwrap_or_else(|p| p.into_inner()).take() { - crate::child_process::log_shutdown( - "shared sidecar", - child.shutdown(Duration::from_millis(500)), - ); + if let Err(error) = + stop_shared_process(&mut self.process.lock().unwrap_or_else(|p| p.into_inner())) + { + crate::child_process::log_shutdown("shared sidecar", Err(error)); } if let Some(reader) = self.reader.take() { let _ = reader.join(); @@ -1321,6 +1344,7 @@ mod tests { use std::net::TcpStream; #[cfg(all(unix, not(feature = "managed-runtime")))] use std::os::unix::fs::PermissionsExt; + use std::sync::atomic::AtomicUsize; #[cfg(all(unix, not(feature = "managed-runtime")))] static SIDECAR_ENV_LOCK: Mutex<()> = Mutex::new(()); @@ -1342,7 +1366,7 @@ mod tests { let sinks: DeckStatusSinks = std::array::from_fn(|_| { Box::new(|_message| {}) as StatusSink }); - let shared = SharedSidecar::parked( + let mut shared = SharedSidecar::parked( ["mrt2_small".into(), "mrt2_small".into()], handles, sinks, @@ -1352,12 +1376,53 @@ mod tests { assert!(shared.reader.is_none()); assert!(shared.parked.is_some()); - assert!(shared - .child - .lock() - .unwrap_or_else(|poisoned| poisoned.into_inner()) - .is_none()); + assert!(matches!( + &*shared + .process + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()), + SharedProcessState::Stopped + )); assert!(shared.stop.load(Ordering::Acquire)); + + let shutdowns = Arc::new(AtomicUsize::new(1)); + *shared.process.lock().unwrap() = + SharedProcessState::Uncertain(Box::new(RetrySharedProcess { + shutdowns: shutdowns.clone(), + })); + assert!(shared.reader.is_none()); + assert!(shared.activate().is_err()); + assert_eq!(shutdowns.load(Ordering::Acquire), 1); + } + + struct RetrySharedProcess { + shutdowns: Arc, + } + + impl SharedProcess for RetrySharedProcess { + fn shutdown(&mut self) -> io::Result<()> { + if self.shutdowns.fetch_add(1, Ordering::AcqRel) == 0 { + Err(io::Error::other("first reap is uncertain")) + } else { + Ok(()) + } + } + } + + #[test] + fn failed_shared_reap_retains_supervisor_until_a_positive_retry() { + let shutdowns = Arc::new(AtomicUsize::new(0)); + let mut state = SharedProcessState::Running(Box::new(RetrySharedProcess { + shutdowns: shutdowns.clone(), + })); + + assert!(stop_shared_process(&mut state).is_err()); + assert!(matches!(state, SharedProcessState::Uncertain(_))); + assert_eq!(shutdowns.load(Ordering::Acquire), 1); + + assert!(stop_shared_process(&mut state).is_ok()); + assert!(matches!(state, SharedProcessState::Stopped)); + assert_eq!(shutdowns.load(Ordering::Acquire), 2); } #[test] From 1aedb8f49445e74f3610a2724857515e842e66a6 Mon Sep 17 00:00:00 2001 From: brxs Date: Sat, 8 Aug 2026 20:19:14 -0700 Subject: [PATCH 08/14] fix: retain SA3 supervisor on readiness error --- src-tauri/src/generation.rs | 172 +++++++++++++++++++++++++++++------- 1 file changed, 140 insertions(+), 32 deletions(-) diff --git a/src-tauri/src/generation.rs b/src-tauri/src/generation.rs index e9b1886..ad44bca 100644 --- a/src-tauri/src/generation.rs +++ b/src-tauri/src/generation.rs @@ -62,6 +62,26 @@ enum GenerationProcessState { }, } +struct GenerationSpawnFailure { + error: io::Error, + uncertain_process: Option>, +} + +impl GenerationSpawnFailure { + fn reaped(error: io::Error) -> Self { + Self { + error, + uncertain_process: None, + } + } +} + +impl From for GenerationSpawnFailure { + fn from(error: io::Error) -> Self { + Self::reaped(error) + } +} + impl GenerationServer { /// Spawn the generation server — started with the app. Never fails the app: a /// failed spawn yields `port() == None` and generation is simply unreachable (the @@ -78,7 +98,9 @@ impl GenerationServer { server } - fn spawn(capability: &str) -> io::Result<(u16, Box)> { + fn spawn( + capability: &str, + ) -> Result<(u16, Box), GenerationSpawnFailure> { // Pick a free loopback port, then hand it to the child (uvicorn binds it). // The brief drop→rebind window on loopback is benign. let port = { @@ -94,39 +116,18 @@ impl GenerationServer { // a slow-but-working server is reported optimistically rather than // blocking the window; a child that EXITS is reported as a failure. let addr = ("127.0.0.1", port); - match child.wait_for_readiness(Duration::from_millis(1500), || { + let readiness = child.wait_for_readiness(Duration::from_millis(1500), || { Ok(TcpStream::connect(addr).is_ok()) - })? { - Readiness::Ready | Readiness::TimedOut => { - // Preserve the existing macOS contract: a slow-but-running - // service is advertised optimistically after the bounded wait. - Ok((port, Box::new(child))) - } - Readiness::Exited(status) => Err(io::Error::other(format!( - "generation server exited before binding ({status})" - ))), - } + }); + finish_generation_startup(port, Box::new(child), readiness) } - /// Return the port and capability from one lock acquisition. Neither half is - /// ever observable without the other across promotion/resume transitions. - pub fn connection(&self) -> Option<(u16, String)> { - match &*self - .state - .lock() - .unwrap_or_else(|poisoned| poisoned.into_inner()) - { - GenerationProcessState::Running { - port, capability, .. - } => Some((*port, capability.clone())), - GenerationProcessState::Stopped | GenerationProcessState::Uncertain { .. } => None, - } - } - - /// Start (or recover) the service from the currently promoted verified - /// generation. A running healthy child is left untouched. This is called on - /// startup and after every managed SA3 promotion/rollback. - pub fn resume(&self) -> io::Result<()> { + fn resume_with_spawn( + &self, + spawn: impl FnOnce( + &str, + ) -> Result<(u16, Box), GenerationSpawnFailure>, + ) -> io::Result<()> { let mut state = self .state .lock() @@ -170,7 +171,18 @@ impl GenerationServer { } } let capability = crate::local_auth::generate_capability(); - let (port, process) = Self::spawn(&capability)?; + let (port, process) = match spawn(&capability) { + Ok(spawned) => spawned, + Err(spawn) => { + if let Some(process) = spawn.uncertain_process { + *state = GenerationProcessState::Uncertain { + process, + was_running: false, + }; + } + return Err(spawn.error); + } + }; println!("lsdj-app: generation server on 127.0.0.1:{port}"); *state = GenerationProcessState::Running { port, @@ -180,6 +192,28 @@ impl GenerationServer { Ok(()) } + /// Return the port and capability from one lock acquisition. Neither half is + /// ever observable without the other across promotion/resume transitions. + pub fn connection(&self) -> Option<(u16, String)> { + match &*self + .state + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + { + GenerationProcessState::Running { + port, capability, .. + } => Some((*port, capability.clone())), + GenerationProcessState::Stopped | GenerationProcessState::Uncertain { .. } => None, + } + } + + /// Start (or recover) the service from the currently promoted verified + /// generation. A running healthy child is left untouched. This is called on + /// startup and after every managed SA3 promotion/rollback. + pub fn resume(&self) -> io::Result<()> { + self.resume_with_spawn(Self::spawn) + } + /// Stop and reap the service before its managed generation is renamed. /// Returns whether a live child was present so tests/lifecycle diagnostics /// can distinguish first install from an update. @@ -220,6 +254,35 @@ impl GenerationServer { } } +fn finish_generation_startup( + port: u16, + mut process: Box, + readiness: io::Result, +) -> Result<(u16, Box), GenerationSpawnFailure> { + match readiness { + Ok(Readiness::Ready | Readiness::TimedOut) => { + // Preserve the existing macOS contract: a slow-but-running + // service is advertised optimistically after the bounded wait. + Ok((port, process)) + } + Ok(Readiness::Exited(status)) => Err(GenerationSpawnFailure::reaped(io::Error::other( + format!("generation server exited before binding ({status})"), + ))), + Err(readiness_error) => match process.shutdown() { + Ok(()) => Err(GenerationSpawnFailure::reaped(readiness_error)), + Err(cleanup_error) => Err(GenerationSpawnFailure { + error: io::Error::new( + readiness_error.kind(), + format!( + "{readiness_error}; generation startup cleanup also failed: {cleanup_error}" + ), + ), + uncertain_process: Some(process), + }), + } + } +} + impl Drop for GenerationServer { fn drop(&mut self) { self.shutdown(); @@ -376,4 +439,49 @@ mod process_tests { )); assert_eq!(shutdowns.load(Ordering::Acquire), 2); } + + #[test] + fn readiness_error_with_failed_cleanup_retains_ownership_until_second_shutdown() { + let shutdowns = Arc::new(AtomicUsize::new(0)); + let startup = finish_generation_startup( + 4321, + Box::new(RetryProcess { + shutdowns: shutdowns.clone(), + }), + Err(io::Error::other("readiness OS error")), + ); + let failure = match startup { + Err(failure) => failure, + Ok(_) => panic!("readiness error must fail startup"), + }; + assert!(failure.error.to_string().contains("readiness OS error")); + assert!(failure + .error + .to_string() + .contains("startup cleanup also failed")); + assert!(failure.uncertain_process.is_some()); + assert_eq!(shutdowns.load(Ordering::Acquire), 1); + + let server = GenerationServer { + state: Mutex::new(GenerationProcessState::Stopped), + }; + let error = server + .resume_with_spawn(move |_| Err(failure)) + .unwrap_err(); + assert!(error.to_string().contains("readiness OS error")); + assert_eq!(server.connection(), None); + assert!(matches!( + &*server.state.lock().unwrap(), + GenerationProcessState::Uncertain { .. } + )); + + // The startup cleanup was the first shutdown attempt. Quiesce owns the + // second attempt and cannot expose Stopped until that positive reap. + assert!(matches!(server.quiesce(), Ok(false))); + assert!(matches!( + &*server.state.lock().unwrap(), + GenerationProcessState::Stopped + )); + assert_eq!(shutdowns.load(Ordering::Acquire), 2); + } } From bd10f8efdcef297d9a6d3c540997c55edc1987dd Mon Sep 17 00:00:00 2001 From: brxs Date: Sat, 8 Aug 2026 20:33:51 -0700 Subject: [PATCH 09/14] test: type promoted SA3 fetch mock --- frontend/src/audio/nativeEngine.test.ts | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/frontend/src/audio/nativeEngine.test.ts b/frontend/src/audio/nativeEngine.test.ts index 0d93b83..f24f69b 100644 --- a/frontend/src/audio/nativeEngine.test.ts +++ b/frontend/src/audio/nativeEngine.test.ts @@ -206,7 +206,11 @@ describe('createNativeEngine — control contract', () => { ) }) vi.stubGlobal('__TAURI__', { core: { invoke } }) - const fetchMock = vi.fn(async () => ({ ok: true })) + const fetchMock = vi.fn(async (_url: string, _init: RequestInit) => { + void _url + void _init + return { ok: true } + }) vi.stubGlobal('fetch', fetchMock) await fetchGenerationApi('/api/generate', { method: 'POST' }) From e7c4c3dfbc095482d08cc063954106cdd7c8cabd Mon Sep 17 00:00:00 2001 From: brxs Date: Sat, 8 Aug 2026 20:33:07 -0700 Subject: [PATCH 10/14] test: qualify managed runtime on Windows --- .github/workflows/ci.yml | 43 ++++++ src-tauri/src/managed_runtime.rs | 228 +++++++++++++++++++++++++++++++ src-tauri/src/models.rs | 31 +++++ src-tauri/src/sidecar.rs | 112 +++++++++++++++ 4 files changed, 414 insertions(+) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 9bf1456..1ec2192 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -15,6 +15,49 @@ concurrency: cancel-in-progress: true jobs: + windows-managed-runtime: + name: Managed runtime qualification (Windows) + runs-on: windows-2025 + timeout-minutes: 90 + + steps: + - name: Check out source and test corpus + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 + with: + lfs: true + persist-credentials: false + + - name: Set up Python + uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 + with: + python-version: "3.13" + + - name: Set up Node.js + uses: actions/setup-node@820762786026740c76f36085b0efc47a31fe5020 # v7.0.0 + with: + node-version: 24 + cache: npm + cache-dependency-path: frontend/package-lock.json + + - name: Build frontend assets for native shell + working-directory: frontend + run: npm ci && npm run build + + - name: Set up Rust + run: rustup toolchain install stable --profile minimal --no-self-update + + - name: Test managed runtime workspace + working-directory: src-tauri + run: cargo test --locked --workspace --features managed-runtime + + - name: Lint managed runtime workspace + working-directory: src-tauri + run: cargo clippy --locked --workspace --all-targets --features managed-runtime -- -D warnings + + - name: Build managed runtime release + working-directory: src-tauri + run: cargo build --locked --workspace --release --features managed-runtime + shared: name: Shared checks (${{ matrix.name }}) runs-on: ${{ matrix.runner }} diff --git a/src-tauri/src/managed_runtime.rs b/src-tauri/src/managed_runtime.rs index bbb4ba8..2970017 100644 --- a/src-tauri/src/managed_runtime.rs +++ b/src-tauri/src/managed_runtime.rs @@ -676,6 +676,234 @@ mod tests { (root, home) } + #[cfg(windows)] + const WINDOWS_HELPER_ROLE: &str = "LSDJ_API_CAPABILITY"; + #[cfg(windows)] + const WINDOWS_HELPER_PID_FILE: &str = "LSDJ_STAGING_HOME"; + + #[cfg(windows)] + fn windows_helper_command(role: &str, pid_file: &Path) -> Command { + let mut command = Command::new(std::env::current_exe().expect("current test executable")); + command + .args([ + "--ignored", + "--exact", + "managed_runtime::tests::windows_managed_runtime_process_helper", + "--nocapture", + ]) + .env(WINDOWS_HELPER_ROLE, role) + .env(WINDOWS_HELPER_PID_FILE, pid_file); + command + } + + #[cfg(windows)] + fn install_windows_spawnable(root: &Path, revision: &str) { + let program = root.join("runtime/bin/managed helper.exe"); + fs::create_dir_all(program.parent().unwrap()).unwrap(); + fs::copy(std::env::current_exe().unwrap(), &program).unwrap(); + fs::write( + root.join("runtime").join(format!("{revision}.marker")), + revision, + ) + .unwrap(); + let spec = CommandSpec { + program: "runtime/bin/managed helper.exe".into(), + argv: vec![ + "--ignored".into(), + "--exact".into(), + "managed_runtime::tests::windows_managed_runtime_process_helper".into(), + "--nocapture".into(), + ], + cwd: "runtime".into(), + environment: BTreeMap::new(), + ephemeral_environment: [ + WINDOWS_HELPER_ROLE, + WINDOWS_HELPER_PID_FILE, + "SYSTEMROOT", + "WINDIR", + "TEMP", + "TMP", + ] + .into_iter() + .map(str::to_string) + .collect(), + }; + seal_candidate( + root, + &host_target(), + BTreeMap::from([("sourceRevision".into(), revision.into())]), + BTreeMap::from([("mrt2".into(), spec)]), + ) + .unwrap(); + } + + #[cfg(windows)] + fn wait_for_windows_pids(path: &Path) -> (u32, u32) { + let deadline = std::time::Instant::now() + std::time::Duration::from_secs(10); + loop { + if let Ok(contents) = fs::read_to_string(path) { + let pids = contents + .split_whitespace() + .filter_map(|value| value.parse::().ok()) + .collect::>(); + if pids.len() == 2 { + return (pids[0], pids[1]); + } + } + assert!( + std::time::Instant::now() < deadline, + "managed runtime helper did not report its process tree" + ); + std::thread::sleep(std::time::Duration::from_millis(20)); + } + } + + #[cfg(windows)] + fn windows_process_is_alive(pid: u32) -> bool { + use windows_sys::Win32::Foundation::{CloseHandle, WAIT_TIMEOUT}; + use windows_sys::Win32::System::Threading::{ + OpenProcess, WaitForSingleObject, PROCESS_QUERY_LIMITED_INFORMATION, + }; + const SYNCHRONIZE_ACCESS: u32 = 0x0010_0000; + // SAFETY: this opens a read-only liveness handle for a test-owned pid. + let process = unsafe { + OpenProcess( + SYNCHRONIZE_ACCESS | PROCESS_QUERY_LIMITED_INFORMATION, + 0, + pid, + ) + }; + if process.is_null() { + return false; + } + // SAFETY: `process` is a live handle and the zero timeout cannot block. + let result = unsafe { WaitForSingleObject(process, 0) }; + // SAFETY: close exactly the handle opened above. + unsafe { CloseHandle(process) }; + result == WAIT_TIMEOUT + } + + #[cfg(windows)] + fn wait_until_windows_process_is_gone(pid: u32) -> bool { + let deadline = std::time::Instant::now() + std::time::Duration::from_secs(10); + while std::time::Instant::now() < deadline { + if !windows_process_is_alive(pid) { + return true; + } + std::thread::sleep(std::time::Duration::from_millis(20)); + } + false + } + + /// Process-tree stand-in copied into a sealed generation. The service role + /// launches a descendant from the same locked executable so promotion is + /// qualified against the real Windows Job Object and filesystem semantics. + #[cfg(windows)] + #[test] + #[ignore] + #[allow(clippy::zombie_processes)] + fn windows_managed_runtime_process_helper() { + let role = std::env::var(WINDOWS_HELPER_ROLE).expect("helper role"); + let pid_file = + PathBuf::from(std::env::var_os(WINDOWS_HELPER_PID_FILE).expect("helper pid file")); + match role.as_str() { + "service" => { + let grandchild = windows_helper_command("grandchild", &pid_file) + .spawn() + .expect("spawn managed runtime grandchild"); + fs::write( + &pid_file, + format!("{} {}", std::process::id(), grandchild.id()), + ) + .expect("write managed runtime pid file"); + loop { + std::thread::sleep(std::time::Duration::from_secs(60)); + } + } + "grandchild" => loop { + std::thread::sleep(std::time::Duration::from_secs(60)); + }, + other => panic!("unknown managed runtime helper role {other}"), + } + } + + #[cfg(windows)] + #[test] + fn windows_reaps_active_managed_tree_before_promoting_and_removing_old_generation() { + let root = root("windows locked promotion 资产"); + let first_candidate = root.join("first candidate"); + let next_candidate = root.join("next candidate"); + let home = root.join("active generation"); + let backup = root.join("old generation"); + let pid_file = root.join("managed process ids"); + + install_windows_spawnable(&first_candidate, "old-revision"); + crate::runtime_installer::promotion::promote(&first_candidate, &home, &backup, |path| { + resolve_at(path, "mrt2", &host_target()).map(|_| ()) + }) + .unwrap(); + let old_generation = resolve_at(&home, "mrt2", &host_target()) + .unwrap() + .generation() + .to_string(); + install_windows_spawnable(&next_candidate, "new-revision"); + + let mut ephemeral = vec![ + ( + OsString::from(WINDOWS_HELPER_ROLE), + OsString::from("service"), + ), + ( + OsString::from(WINDOWS_HELPER_PID_FILE), + pid_file.as_os_str().to_owned(), + ), + ]; + ephemeral.extend( + ["SYSTEMROOT", "WINDIR", "TEMP", "TMP"] + .into_iter() + .filter_map(|name| { + std::env::var_os(name).map(|value| (OsString::from(name), value)) + }), + ); + let mut command = resolve_at(&home, "mrt2", &host_target()) + .unwrap() + .into_command([], ephemeral) + .unwrap(); + let mut process = crate::child_process::spawn_grouped(&mut command).unwrap(); + let (service_pid, grandchild_pid) = wait_for_windows_pids(&pid_file); + assert!(windows_process_is_alive(service_pid)); + assert!(windows_process_is_alive(grandchild_pid)); + + let report = process + .shutdown(std::time::Duration::from_millis(100)) + .unwrap(); + assert!( + report.forced, + "live managed tree should require Job teardown" + ); + assert!(report.status.is_some(), "service leader was not reaped"); + assert!( + wait_until_windows_process_is_gone(service_pid), + "managed service survived quiesce" + ); + assert!( + wait_until_windows_process_is_gone(grandchild_pid), + "managed descendant survived quiesce" + ); + + crate::runtime_installer::promotion::promote(&next_candidate, &home, &backup, |path| { + resolve_at(path, "mrt2", &host_target()).map(|_| ()) + }) + .unwrap(); + let promoted = resolve_at(&home, "mrt2", &host_target()).unwrap(); + assert_ne!(promoted.generation(), old_generation); + assert!(home.join("runtime/new-revision.marker").is_file()); + assert!(!home.join("runtime/old-revision.marker").exists()); + assert!(!next_candidate.exists(), "candidate should be promoted"); + assert!(!backup.exists(), "old generation should be removed"); + fs::remove_dir_all(root).unwrap(); + } + #[test] fn clean_host_fails_closed_and_install_produces_structured_commands() { let root = root("clean host with spaces 资产"); diff --git a/src-tauri/src/models.rs b/src-tauri/src/models.rs index 57640a4..c40afb8 100644 --- a/src-tauri/src/models.rs +++ b/src-tauri/src/models.rs @@ -3560,6 +3560,37 @@ mod tests { assert_eq!(*events.borrow(), ["quiesce"]); } + #[cfg(all(feature = "managed-runtime", windows))] + #[test] + fn windows_uncertain_reap_leaves_active_and_candidate_generations_unrenamed() { + let root = std::env::temp_dir().join(format!( + "lsdj-windows-uncertain-reap-{}-{:?}", + std::process::id(), + std::thread::current().id() + )); + let _ = std::fs::remove_dir_all(&root); + let home = root.join("active generation"); + let candidate = root.join("candidate generation"); + let backup = root.join("old generation"); + std::fs::create_dir_all(&home).unwrap(); + std::fs::create_dir_all(&candidate).unwrap(); + std::fs::write(home.join("active.marker"), b"active").unwrap(); + std::fs::write(candidate.join("candidate.marker"), b"candidate").unwrap(); + + let result = run_promotion_lifecycle( + "Windows managed runtime", + || Err::<(), _>("process-tree reap is uncertain".to_string()), + || crate::runtime_installer::promotion::promote(&candidate, &home, &backup, |_| Ok(())), + |_| Ok(()), + ); + + assert_eq!(result.unwrap_err(), "process-tree reap is uncertain"); + assert!(home.join("active.marker").is_file()); + assert!(candidate.join("candidate.marker").is_file()); + assert!(!backup.exists(), "rename window must remain unopened"); + std::fs::remove_dir_all(root).unwrap(); + } + #[cfg(feature = "managed-runtime")] #[test] fn promotion_waits_for_a_second_positive_reap_before_rename() { diff --git a/src-tauri/src/sidecar.rs b/src-tauri/src/sidecar.rs index 5a7ba52..b6e783e 100644 --- a/src-tauri/src/sidecar.rs +++ b/src-tauri/src/sidecar.rs @@ -1622,6 +1622,118 @@ mod tests { ); } + #[cfg(all(windows, feature = "managed-runtime"))] + fn windows_python() -> std::path::PathBuf { + let search = std::env::var_os("PATH").expect("Python is available on CI PATH"); + for directory in std::env::split_paths(&search) { + for name in ["python.exe", "python3.exe"] { + let candidate = directory.join(name); + if candidate.is_file() { + return candidate; + } + } + } + panic!("Python executable is unavailable for Windows protocol qualification"); + } + + /// Native Windows qualification for the real Rust/Python wire boundary. + /// The stand-in is stdlib-only and model-free, but the socket, authenticated + /// handshake, bidirectional framing, structured argv, and supervised Python + /// process are the same primitives used by the packaged sidecar. + #[cfg(all(windows, feature = "managed-runtime"))] + #[test] + fn windows_python_round_trips_authenticated_control_status_and_pcm() { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let port = listener.local_addr().unwrap().port(); + let token = "windows-protocol-token-0123456789abcdef"; + let root = std::env::temp_dir().join(format!( + "lsdj-windows-python-protocol-{}-{port}-资产", + std::process::id() + )); + let _ = std::fs::remove_dir_all(&root); + std::fs::create_dir_all(&root).unwrap(); + let script = root.join("sidecar stand-in.py"); + std::fs::write( + &script, + r#"import json +import os +import socket +import struct +import sys + +def receive_exact(sock, length): + payload = b"" + while len(payload) < length: + chunk = sock.recv(length - len(payload)) + if not chunk: + raise RuntimeError("truncated frame") + payload += chunk + return payload + +def receive_frame(sock): + header = receive_exact(sock, 5) + frame_type, length = struct.unpack(">(); + assert_eq!(read_frame(&mut stream).unwrap(), Some((FRAME_PCM, pcm))); + drop(stream); + let status = child.wait().unwrap(); + assert!( + status.success(), + "Python protocol stand-in failed: {status}" + ); + std::fs::remove_dir_all(root).unwrap(); + } + /// In-process model switch: `restart` respawns the sidecar with a new model, /// reusing the deck's permanent ring producer, and suppresses a false /// `worker_died` across the deliberate switch. Wires a minimal stdlib-only From 925a3dfcfe4d8ebdf0fc458e22796b322b5ba163 Mon Sep 17 00:00:00 2001 From: brxs Date: Sat, 8 Aug 2026 21:12:23 -0700 Subject: [PATCH 11/14] fix: scope Unix-only downloader tests --- src-tauri/src/models.rs | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src-tauri/src/models.rs b/src-tauri/src/models.rs index c40afb8..654c54f 100644 --- a/src-tauri/src/models.rs +++ b/src-tauri/src/models.rs @@ -1299,7 +1299,7 @@ pub(crate) fn stream_child( /// One parsed line of the sidecar's JSON progress contract. #[derive(Deserialize)] -#[cfg(any(not(feature = "managed-runtime"), test))] +#[cfg(any(not(feature = "managed-runtime"), all(test, unix)))] struct SidecarLine { event: String, file: Option, @@ -1817,7 +1817,7 @@ fn validate_mrt2_candidate( /// Spawn the download tooling and map its JSON progress contract onto the sink. /// Takes the fully-built command so the spawn+parse path is testable against a /// stub without mutating the process environment. -#[cfg(any(not(feature = "managed-runtime"), test))] +#[cfg(any(not(feature = "managed-runtime"), all(test, unix)))] fn run_download(progress: &Progress, shared: &InstallShared, cmd: Command) -> Result<(), String> { let mut last_error: Option = None; let result = stream_child(shared, "download-model", cmd, |line| { From 2d483eaf581888e09c64603875f96e35f8e6bc78 Mon Sep 17 00:00:00 2001 From: brxs Date: Sat, 8 Aug 2026 21:21:15 -0700 Subject: [PATCH 12/14] fix: synchronize media browser controls with rows --- frontend/src/media/MediaExplorer.tsx | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/frontend/src/media/MediaExplorer.tsx b/frontend/src/media/MediaExplorer.tsx index 701b99d..7c47688 100644 --- a/frontend/src/media/MediaExplorer.tsx +++ b/frontend/src/media/MediaExplorer.tsx @@ -725,7 +725,11 @@ export function MediaExplorer({ ) const bus = useControlBus() - useEffect(() => + // Refresh the hardware handler in the same commit that exposes new rows. + // A passive effect leaves a frame where the DOM shows the new list but the + // bus still holds the previous render's empty/stale list closure, so a rotary + // tick in that window is silently ignored. + useLayoutEffect(() => bus.subscribe((intent) => { if (intent.kind === 'browse_tab') { // Rotary press: cycle the visible tab from the hardware. From dbb54b1f70618025827dc4aaa1d988903ef48a9f Mon Sep 17 00:00:00 2001 From: brxs Date: Sat, 8 Aug 2026 21:55:36 -0700 Subject: [PATCH 13/14] fix: scope Unix-only shared sidecar test helper --- src-tauri/src/sidecar.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src-tauri/src/sidecar.rs b/src-tauri/src/sidecar.rs index b6e783e..899138f 100644 --- a/src-tauri/src/sidecar.rs +++ b/src-tauri/src/sidecar.rs @@ -700,7 +700,7 @@ impl SharedSidecar { } } - #[cfg(all(test, not(feature = "managed-runtime")))] + #[cfg(all(test, unix, not(feature = "managed-runtime")))] pub fn spawn( models: [String; lsdj_engine::DECK_COUNT], handles: [DeckHandle; lsdj_engine::DECK_COUNT], From 298ec51032bf1e470aa06d8e7f955b326989cd10 Mon Sep 17 00:00:00 2001 From: brxs Date: Sat, 8 Aug 2026 22:04:50 -0700 Subject: [PATCH 14/14] fix: order overlapping library refreshes --- frontend/src/media/MediaExplorer.test.tsx | 50 +++++++++++++ frontend/src/media/MediaExplorer.tsx | 89 +++++++++++++++-------- 2 files changed, 108 insertions(+), 31 deletions(-) diff --git a/frontend/src/media/MediaExplorer.test.tsx b/frontend/src/media/MediaExplorer.test.tsx index 0bc69da..59861f1 100644 --- a/frontend/src/media/MediaExplorer.test.tsx +++ b/frontend/src/media/MediaExplorer.test.tsx @@ -701,6 +701,56 @@ describe('MediaExplorer', () => { expect(screen.getByText('#2')).toBeInTheDocument() }) + it('keeps the newer sample scan when overlapping refreshes finish out of order', async () => { + type ResolveSamples = (rows: { + file: string + title: string + prompt: string + model: string + oneShot: boolean + }[]) => void + const scans: ResolveSamples[] = [] + let onChange: ((e: { payload: unknown }) => void) | null = null + const invoke = vi.fn((cmd: string) => { + if (cmd === 'list_generated_samples') { + return new Promise((resolve: ResolveSamples) => scans.push(resolve)) + } + return Promise.resolve([]) + }) + const listen = vi.fn( + async (event: string, handler: (e: { payload: unknown }) => void) => { + if (event === 'library://changed') onChange = handler + return () => {} + }, + ) + vi.stubGlobal('__TAURI__', { core: { invoke }, event: { listen } }) + renderExplorer() + fireEvent.click(screen.getByRole('tab', { name: 'Samples' })) + expect(scans).toHaveLength(1) + + // A watcher scan starts while startup's scan is still pending. Finish the newer + // scan first, then the stale startup scan in the same batch: the old implementation + // would replace the two current rows with `one #3` because both completions read + // the same stale passive-effect ref and minted fresh ids. + act(() => onChange?.({ payload: { library: 'samples' } })) + expect(scans).toHaveLength(2) + await act(async () => { + scans[1]([ + { file: 'one.wav', title: 'one', prompt: 'one', model: 'sfx', oneShot: false }, + { file: 'two.wav', title: 'two', prompt: 'two', model: 'music', oneShot: false }, + ]) + scans[0]([ + { file: 'one.wav', title: 'one', prompt: 'one', model: 'sfx', oneShot: false }, + ]) + await Promise.resolve() + }) + + expect(screen.getByText('one', { selector: '.media__name-text' })).toBeInTheDocument() + expect(screen.getByText('two', { selector: '.media__name-text' })).toBeInTheDocument() + expect(screen.getByText('#1')).toBeInTheDocument() + expect(screen.getByText('#2')).toBeInTheDocument() + }) + it('restores samples, tagging a freeze and a hand-added file', async () => { const invoke = vi.fn(async (cmd: string) => { if (cmd === 'list_generated_samples') { diff --git a/frontend/src/media/MediaExplorer.tsx b/frontend/src/media/MediaExplorer.tsx index 7c47688..47001b1 100644 --- a/frontend/src/media/MediaExplorer.tsx +++ b/frontend/src/media/MediaExplorer.tsx @@ -221,23 +221,33 @@ function hasVersionedRecipe(value: unknown): boolean { ) } +type LibraryRefreshState = { + issued: number + applied: number + idsByFile: Map +} + /** Re-list one library (songs or samples) from its on-disk registry, reconciled * against the folder by the Rust shell (hand-added files appear; deleted files drop * out). A row already held for a file keeps its id + in-memory wav (reuse by * filename), so a live re-list never churns; a row whose file vanished is dropped; an - * in-session take not yet on disk is kept. `ref` is read after the fetch resolves - * (freshest), and the id mint (`toRow`) runs OUTSIDE the state updater — StrictMode - * replays updaters, so they must be pure. A no-op outside Tauri. */ + * in-session take not yet on disk is kept. Overlapping scans are ordered by request, + * so an older startup scan cannot replace a newer watcher scan. New rows get a stable + * id per filename before the state updater; the updater itself reconciles against + * React's freshest `current` state and remains pure under StrictMode. A no-op outside + * Tauri. */ function reListLibrary< R extends { id: number; state: string; file?: string | null }, E extends { file: string }, >( command: string, - ref: { current: R[] }, + refresh: { current: LibraryRefreshState }, setRows: (next: (current: R[]) => R[]) => void, - toRow: (entry: E) => R, + mintId: () => number, + toRow: (entry: E, id: number) => R, ): void { if (!isTauri()) return + const request = ++refresh.current.issued void (async () => { let entries: E[] try { @@ -245,17 +255,34 @@ function reListLibrary< } catch { return // a failed scan just means no refresh; composing still works } - const byFile = new Map( - ref.current - .map((row) => [fileOf(row), row] as const) - .filter((pair): pair is readonly [string, R] => pair[0] != null), - ) - const restored = entries.map((entry) => byFile.get(entry.file) ?? toRow(entry)) + // A newer successful request already represents a later view of the registry. + // Failed newer requests do not advance `applied`, so an older successful scan + // may still provide the best available view. + if (request < refresh.current.applied) return + refresh.current.applied = request + const restored = entries.map((entry) => { + let id = refresh.current.idsByFile.get(entry.file) + if (id == null) { + id = mintId() + refresh.current.idsByFile.set(entry.file, id) + } + return [entry.file, toRow(entry, id)] as const + }) // Newest-first: in-session takes not yet on disk lead, above the restored // library reversed so the most recently composed file sits at the top (the // registry stores composition order, oldest first), sparing a scroll to the // take you just made. - setRows((current) => [...current.filter((row) => fileOf(row) == null), ...restored.reverse()]) + setRows((current) => { + const byFile = new Map( + current + .map((row) => [fileOf(row), row] as const) + .filter((pair): pair is readonly [string, R] => pair[0] != null), + ) + return [ + ...current.filter((row) => fileOf(row) == null), + ...restored.map(([file, row]) => byFile.get(file) ?? row).reverse(), + ] + }) })() } @@ -414,6 +441,16 @@ export function MediaExplorer({ // A ref, not state: two composes batched into one render (Enter + // click) must not mint the same id. const nextIdRef = useRef(1) + const trackRefreshRef = useRef({ + issued: 0, + applied: 0, + idsByFile: new Map(), + }) + const sampleRefreshRef = useRef({ + issued: 0, + applied: 0, + idsByFile: new Map(), + }) const trackTasksRef = useRef(new Map()) const sampleTasksRef = useRef(new Map()) useEffect( @@ -424,18 +461,6 @@ export function MediaExplorer({ }, [], ) - // The latest lists mirrored in refs (synced after commit). A live re-list (tab - // open, or the folder watcher firing) reads these from its effect/callback to reuse - // a row's id + in-memory wav by filename, so a refresh never churns ids or re-reads - // bytes — and the id mint stays OUTSIDE the state updater (StrictMode replays - // updaters, so they must be pure). At most one render stale, which is fine here. - const tracksRef = useRef([]) - const samplesRef = useRef([]) - useEffect(() => { - tracksRef.current = tracks - samplesRef.current = samples - }, [tracks, samples]) - const filteredTracks = tracks.filter((track) => matchesSearch( search, @@ -667,17 +692,18 @@ export function MediaExplorer({ } // The two libraries' re-list, each a thin {@link reListLibrary} call differing only - // in the command, the ref, the setter, and the registry-entry → row mapping (a + // in the command, the refresh state, the setter, and the registry-entry → row mapping (a // sample carries `oneShot`; a song's model runs through `asTrackEngine`). Used at // startup and by the folder watcher. const refreshSongs = useCallback( () => reListLibrary( 'list_generated_songs', - tracksRef, + trackRefreshRef, setTracks, - (entry) => ({ - id: nextIdRef.current++, + () => nextIdRef.current++, + (entry, id) => ({ + id, state: 'ready', title: entry.title, prompt: entry.prompt, @@ -692,10 +718,11 @@ export function MediaExplorer({ () => reListLibrary( 'list_generated_samples', - samplesRef, + sampleRefreshRef, setSamples, - (entry) => ({ - id: nextIdRef.current++, + () => nextIdRef.current++, + (entry, id) => ({ + id, state: 'ready', title: entry.title, prompt: entry.prompt,