From 0d121acda17bb24088aa1fd833ec2836e3d1f355 Mon Sep 17 00:00:00 2001 From: jan <1760229+slomin@users.noreply.github.com> Date: Wed, 8 Apr 2026 15:33:51 +0100 Subject: [PATCH] feat(inferno): move Potato OS to the standalone inferno package Adds the pinned `potato-inferno` dependency and switches Potato OS to import its inference backend, model, runtime, launch, and orchestration logic from the standalone Inferno repo instead of the in-tree `core/inferno` package. Removes the migrated Inferno modules and duplicated tests from Potato, updates startup scripts and import paths for the external package, and keeps deploy, OTA recovery, and manual rsync workflows working by installing the updated Python dependencies before restart. Keeps `.litertlm` models routed to the LiteRT adapter during startup so service restarts continue to select the correct inference process after the repo split. Closes #298 Refs #264 --- AGENTS.md | 4 +- apps/chat/routes.py | 4 +- bin/install_dev.sh | 4 +- bin/start_litert.sh | 2 +- bin/start_llama.sh | 2 +- core/deps.py | 4 +- core/inferno/__init__.py | 227 ------ core/inferno/backend.py | 358 --------- core/inferno/launch_config.py | 87 --- core/inferno/litert_adapter.py | 443 ----------- core/inferno/model_families.py | 189 ----- core/inferno/model_registry.py | 807 --------------------- core/inferno/orchestrator.py | 522 ------------- core/inferno/runtime_manager.py | 755 ------------------- core/main.py | 42 +- core/model_state.py | 139 ++-- core/runtime_state.py | 14 +- docs/recovery.md | 8 +- requirements.txt | 1 + tests/unit/test_app_lifecycle.py | 2 +- tests/unit/test_chat_repository.py | 248 ------- tests/unit/test_inferno_model_registry.py | 559 -------------- tests/unit/test_inferno_orchestrator.py | 680 ----------------- tests/unit/test_inferno_runtime_manager.py | 646 ----------------- tests/unit/test_launch_config.py | 296 -------- tests/unit/test_litert_adapter.py | 515 ------------- tests/unit/test_model_families.py | 271 ------- tests/unit/test_runtime.py | 2 +- tests/unit/test_script_contracts.py | 4 +- 29 files changed, 90 insertions(+), 6745 deletions(-) delete mode 100644 core/inferno/__init__.py delete mode 100644 core/inferno/backend.py delete mode 100644 core/inferno/launch_config.py delete mode 100644 core/inferno/litert_adapter.py delete mode 100644 core/inferno/model_families.py delete mode 100644 core/inferno/model_registry.py delete mode 100644 core/inferno/orchestrator.py delete mode 100644 core/inferno/runtime_manager.py delete mode 100644 tests/unit/test_chat_repository.py delete mode 100644 tests/unit/test_inferno_model_registry.py delete mode 100644 tests/unit/test_inferno_orchestrator.py delete mode 100644 tests/unit/test_inferno_runtime_manager.py delete mode 100644 tests/unit/test_launch_config.py delete mode 100644 tests/unit/test_litert_adapter.py delete mode 100644 tests/unit/test_model_families.py diff --git a/AGENTS.md b/AGENTS.md index 8eda06a..c300557 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -6,7 +6,7 @@ Core application code lives in `core/`: - `core/main.py`: FastAPI entrypoint, API routes, orchestrator loop, lifespan. - `core/model_state.py`: Model registry, settings persistence, projector helpers. - `core/runtime_state.py`: Runtime config, dual-runtime slot discovery, system metrics, power calibration. -- `core/inferno/`: Inference layer — backend proxy, model family classification, LiteRT adapter, launch config builder. See `core/inferno/__init__.py` for boundary contract. +- `inferno` (external package, [`potato-os/inferno`](https://github.com/potato-os/inferno)): Inference layer — backend proxy, model family classification, LiteRT adapter, launch config builder, model registry, runtime management, orchestration. Installed via `requirements.txt`. - `core/assets/`: Frontend — `index.html`, `shell.css`, `shell.js` (platform shell), vendor libs. - `core/rig_envelope.py`: RIG step envelope validation (MS/TS contract checks). @@ -66,6 +66,8 @@ These commands are written for a **macOS dev environment**. `COPYFILE_DISABLE=1` --rsync-path="SUDO_ASKPASS=/opt/potato/bin/askpass.sh sudo -A rsync" \ -e "ssh -o StrictHostKeyChecking=accept-new" \ core/ pi@potato.local:/opt/potato/core/ + sshpass -e ssh pi@potato.local \ + "echo raspberry | sudo -S /opt/potato/venv/bin/pip install -r /opt/potato/core/requirements.txt" ``` - Fast apps deploy: ``` diff --git a/apps/chat/routes.py b/apps/chat/routes.py index f35e3cd..703daf5 100644 --- a/apps/chat/routes.py +++ b/apps/chat/routes.py @@ -9,16 +9,16 @@ from fastapi import APIRouter, Depends, HTTPException, Request, Response from fastapi.responses import JSONResponse, StreamingResponse +from inferno import BackendProxyError, ChatRepositoryManager + try: from core.deps import get_runtime, get_chat_repository from core.model_state import apply_model_chat_defaults - from core.inferno import BackendProxyError, ChatRepositoryManager from core.runtime_state import RuntimeConfig from core.settings import merge_active_model_chat_defaults, merge_chat_defaults except ModuleNotFoundError: from deps import get_runtime, get_chat_repository # type: ignore[no-redef] from model_state import apply_model_chat_defaults # type: ignore[no-redef] - from inferno import BackendProxyError, ChatRepositoryManager # type: ignore[no-redef] from runtime_state import RuntimeConfig # type: ignore[no-redef] from settings import merge_active_model_chat_defaults, merge_chat_defaults # type: ignore[no-redef] diff --git a/bin/install_dev.sh b/bin/install_dev.sh index 4de6383..e53e16e 100755 --- a/bin/install_dev.sh +++ b/bin/install_dev.sh @@ -301,8 +301,8 @@ rm -f "${sudoers_terminal_tmp}" sudoers_ota_tmp="$(mktemp)" cat > "${sudoers_ota_tmp}" <<'SUDOERS' -potato ALL=(root) NOPASSWD: /bin/chown -R potato\:potato /opt/potato/app -potato ALL=(root) NOPASSWD: /usr/bin/chown -R potato\:potato /opt/potato/app +potato ALL=(root) NOPASSWD: /bin/chown -R potato\:potato /opt/potato/core +potato ALL=(root) NOPASSWD: /usr/bin/chown -R potato\:potato /opt/potato/core potato ALL=(root) NOPASSWD: /bin/chown -R potato\:potato /opt/potato/bin potato ALL=(root) NOPASSWD: /usr/bin/chown -R potato\:potato /opt/potato/bin SUDOERS diff --git a/bin/start_litert.sh b/bin/start_litert.sh index f362052..ed405f2 100755 --- a/bin/start_litert.sh +++ b/bin/start_litert.sh @@ -29,7 +29,7 @@ export POTATO_BASE_DIR cd "${POTATO_BASE_DIR}" exec "${POTATO_VENV_DIR}/bin/uvicorn" \ - core.inferno.litert_adapter:app \ + inferno.litert_adapter:app \ --host 0.0.0.0 \ --port "${POTATO_LLAMA_PORT}" \ --workers 1 diff --git a/bin/start_llama.sh b/bin/start_llama.sh index 06c400b..4c6d9c0 100755 --- a/bin/start_llama.sh +++ b/bin/start_llama.sh @@ -1,6 +1,6 @@ #!/usr/bin/env bash # Thin wrapper around llama-server — all business logic lives in Python -# (core.inferno.launch_config.build_llama_server_args). +# (inferno.launch_config.build_llama_server_args). # # This script only handles: # 1. Dynamic library path for GGML backends diff --git a/core/deps.py b/core/deps.py index 3033887..3015d9a 100644 --- a/core/deps.py +++ b/core/deps.py @@ -4,12 +4,12 @@ from fastapi import Request +from inferno import ChatRepositoryManager + try: from core.runtime_state import RuntimeConfig - from core.inferno import ChatRepositoryManager except ModuleNotFoundError: from runtime_state import RuntimeConfig # type: ignore[no-redef] - from inferno import ChatRepositoryManager # type: ignore[no-redef] def get_runtime(request: Request) -> RuntimeConfig: diff --git a/core/inferno/__init__.py b/core/inferno/__init__.py deleted file mode 100644 index 9df2a03..0000000 --- a/core/inferno/__init__.py +++ /dev/null @@ -1,227 +0,0 @@ -"""Inferno -- Potato OS inference layer. - -Inferno owns the runtime-facing side of inference: backend proxying, -model family classification, model registry, settings normalization, -projector management, and adapter processes. Everything between -"Potato says run this model" and "here is the OpenAI-compatible response" -lives here. - -Boundary contract ------------------ -Potato (caller) provides: - - Base URL of the active inference runtime - - Chat backend mode ("llama" | "fake" | "auto") - - Model filename and optional source URL (for projector resolution) - - POTATO_MODEL_PATH env var (for LiteRT adapter startup) - - Hardware/device/OS inputs (memory, device class, runtime binaries) - - ModelStoreConfig with filesystem paths and product-level defaults - -Inferno (this package) provides: - - BackendProxyError exception for proxy failures - - BackendResponse dataclass for HTTP responses (body or stream) - - ChatCompletionRepository protocol for backend implementations - - LlamaCppRepository real llama.cpp HTTP proxy - - FakeLlamaRepository fake backend for dev/test - - ChatRepositoryManager dispatch to named backends - - is_qwen35_filename detect Qwen 3.5 model files - - is_gemma4_filename detect Gemma 4 model files - - projector_repo_for_model resolve HuggingFace projector repo - - recommended_runtime_for_model preferred runtime family for a model - - build_llama_server_args pure function to build CLI args - - ModelStoreConfig filesystem/policy config for registry ops - - Model registry functions ensure/save/register/delete/update state - - Format handling model_format_for_filename, validate_model_url - - Settings normalization normalize_model_settings, build_model_capabilities - - Projector management build_model_projector_status, download, candidates - - litert_adapter standalone FastAPI app (core.inferno.litert_adapter:app) - -Inferno does NOT import from core.model_state, core.runtime_state, -core.settings, core.deps, or any apps/ code. The dependency arrow -points one way: Potato -> Inferno, never Inferno -> Potato. -""" - -from .backend import ( - BackendProxyError, - BackendResponse, - ChatCompletionRepository, - ChatRepositoryManager, - FakeLlamaRepository, - LlamaCppRepository, -) -from .launch_config import build_llama_server_args -from .model_families import ( - build_model_projector_status, - default_projector_candidates_for_model, - is_gemma4_filename, - is_qwen35_filename, - projector_repo_for_model, - recommended_runtime_for_model, -) -from .runtime_manager import ( - DEVICE_CLOCK_LIMITS, - LLAMA_RUNTIME_BUNDLE_MARKER_FILENAME, - LLAMA_SERVER_RUNTIME_FAMILIES, - MODEL_LOADING_INACTIVE, - MODEL_UPLOAD_PI_16GB_MEMORY_THRESHOLD_BYTES, - PI4_8GB_MEMORY_THRESHOLD_BYTES, - PI4_INCOMPATIBLE_RUNTIMES, - SUPPORTED_RUNTIME_FAMILIES, - RuntimeStoreConfig, - build_large_model_compatibility, - build_llama_large_model_override_status, - build_llama_memory_loading_status, - build_llama_runtime_status, - check_runtime_device_compatibility, - classify_runtime_device, - compute_model_loading_progress, - discover_llama_runtime_bundles, - discover_runtime_slots, - ensure_compatible_runtime, - find_llama_runtime_bundle_by_path, - find_runtime_slot_by_family, - get_device_clock_limits, - get_llama_runtime_bundle_roots, - install_llama_runtime_bundle, - llama_memory_loading_no_mmap_env, - normalize_allow_unsupported_large_models, - normalize_llama_memory_loading_mode, - read_llama_runtime_bundle_marker, - read_llama_runtime_settings, - write_llama_runtime_bundle_marker, - write_llama_runtime_settings, -) -from .orchestrator import ( - READY_HEALTH_POLLS_REQUIRED, - MAX_CONSECUTIVE_FAILURES, - InferenceTickResult, - empty_readiness_state, - empty_runtime_switch_state, - reset_readiness, - resolve_readiness, - check_health, - probe_inference_slot, - refresh_readiness, - restart_inference_process, - resolve_mmproj_for_launch, - ensure_mmproj_for_launch, - resolve_no_mmap, - run_inference_tick, - prepare_activation_runtime, -) -from .model_registry import ( - DEFAULT_MODEL_CHAT_SETTINGS, - DEFAULT_MODEL_VISION_SETTINGS, - MODELS_STATE_VERSION, - VALID_MODEL_EXTENSIONS, - ModelSettingsValidationError, - ModelStoreConfig, - any_model_ready, - apply_model_chat_defaults, - build_model_capabilities, - delete_model, - describe_model_storage, - discover_local_model_filenames, - download_default_projector_for_model, - ensure_models_state, - get_model_by_id, - is_qwen35_a3b_filename, - model_file_path, - model_file_present, - model_format_for_filename, - model_supports_vision_filename, - normalize_model_settings, - register_model_url, - resolve_model_runtime_path, - save_models_state, - update_model_settings, - validate_model_url, -) - -__all__ = [ - "BackendProxyError", - "InferenceTickResult", - "MAX_CONSECUTIVE_FAILURES", - "READY_HEALTH_POLLS_REQUIRED", - "BackendResponse", - "ChatCompletionRepository", - "ChatRepositoryManager", - "DEFAULT_MODEL_CHAT_SETTINGS", - "DEFAULT_MODEL_VISION_SETTINGS", - "DEVICE_CLOCK_LIMITS", - "FakeLlamaRepository", - "LLAMA_RUNTIME_BUNDLE_MARKER_FILENAME", - "LLAMA_SERVER_RUNTIME_FAMILIES", - "LlamaCppRepository", - "MODEL_LOADING_INACTIVE", - "MODEL_UPLOAD_PI_16GB_MEMORY_THRESHOLD_BYTES", - "MODELS_STATE_VERSION", - "ModelSettingsValidationError", - "ModelStoreConfig", - "PI4_8GB_MEMORY_THRESHOLD_BYTES", - "PI4_INCOMPATIBLE_RUNTIMES", - "RuntimeStoreConfig", - "SUPPORTED_RUNTIME_FAMILIES", - "VALID_MODEL_EXTENSIONS", - "any_model_ready", - "apply_model_chat_defaults", - "build_large_model_compatibility", - "build_llama_large_model_override_status", - "build_llama_memory_loading_status", - "build_llama_runtime_status", - "build_llama_server_args", - "build_model_capabilities", - "build_model_projector_status", - "check_runtime_device_compatibility", - "classify_runtime_device", - "compute_model_loading_progress", - "default_projector_candidates_for_model", - "delete_model", - "discover_llama_runtime_bundles", - "discover_runtime_slots", - "describe_model_storage", - "discover_local_model_filenames", - "download_default_projector_for_model", - "ensure_compatible_runtime", - "ensure_models_state", - "find_llama_runtime_bundle_by_path", - "find_runtime_slot_by_family", - "get_device_clock_limits", - "get_llama_runtime_bundle_roots", - "get_model_by_id", - "install_llama_runtime_bundle", - "is_gemma4_filename", - "is_qwen35_a3b_filename", - "is_qwen35_filename", - "llama_memory_loading_no_mmap_env", - "model_file_path", - "model_file_present", - "model_format_for_filename", - "model_supports_vision_filename", - "normalize_allow_unsupported_large_models", - "normalize_llama_memory_loading_mode", - "normalize_model_settings", - "projector_repo_for_model", - "read_llama_runtime_bundle_marker", - "read_llama_runtime_settings", - "recommended_runtime_for_model", - "register_model_url", - "resolve_model_runtime_path", - "save_models_state", - "update_model_settings", - "validate_model_url", - "check_health", - "empty_readiness_state", - "empty_runtime_switch_state", - "ensure_mmproj_for_launch", - "prepare_activation_runtime", - "probe_inference_slot", - "refresh_readiness", - "reset_readiness", - "resolve_mmproj_for_launch", - "resolve_no_mmap", - "resolve_readiness", - "restart_inference_process", - "run_inference_tick", - "write_llama_runtime_bundle_marker", - "write_llama_runtime_settings", -] diff --git a/core/inferno/backend.py b/core/inferno/backend.py deleted file mode 100644 index 58270ba..0000000 --- a/core/inferno/backend.py +++ /dev/null @@ -1,358 +0,0 @@ -from __future__ import annotations - -import asyncio -import json -import os -import random -import time -from dataclasses import dataclass -from typing import Any, AsyncIterator, Protocol - -import httpx - -FAKE_PARODY_REPLIES: tuple[str, ...] = ( - "In 2037, every operating system summit was replaced by the Annual Potato OS Bake-Off, where benchmarks are served with sour cream.", - "Potato OS was declared the official scheduler of the universe after it successfully prioritized snacks over meetings in 12 galaxies.", - "Computer science departments now teach only two courses: 'Potato OS Distributed Systems' and 'How to Peel Legacy Monoliths into Microservices.'", - "The cloud is now called 'the pantry,' and Potato OS autoscaling means adding more ovens whenever traffic spikes.", - "A/B tests became A/BBQ tests; Potato OS picks winners by latency, throughput, and crisp-edge consistency.", - "The Turing Award was briefly renamed the Tubering Award after Potato OS proved all bugs are just under-seasoned features.", - "Potato OS observability dashboards now include four golden signals: latency, errors, saturation, and gravy availability.", - "Kubernetes retired and joined a food truck; Potato OS replaced it with 'Spudernetes,' where pods are literally potato pods.", - "CI/CD now stands for Chop, Inspect, Cook, Deploy, and Potato OS enforces it with strict linting and stricter frying times.", - "Quantum researchers admitted Potato OS solved decoherence by wrapping qubits in foil and giving them emotional support logs.", -) - -DEFAULT_FAKE_PREFILL_DELAY_MS = 0 -DEFAULT_FAKE_STREAM_CHUNK_DELAY_MS = 210 -TEST_FAKE_STREAM_CHUNK_DELAY_MS = 10 - - -class BackendProxyError(RuntimeError): - pass - - -@dataclass -class BackendResponse: - status_code: int - headers: dict[str, str] - body: bytes | None = None - stream: AsyncIterator[bytes] | None = None - background: Any | None = None - - -class ChatCompletionRepository(Protocol): - name: str - - async def create_chat_completion( - self, - payload: dict[str, Any], - forward_headers: dict[str, str], - ) -> BackendResponse: ... - - -class LlamaCppRepository: - name = "llama" - - def __init__(self, base_url: str) -> None: - self.base_url = base_url.rstrip("/") - - async def create_chat_completion( - self, - payload: dict[str, Any], - forward_headers: dict[str, str], - ) -> BackendResponse: - target_url = f"{self.base_url}/v1/chat/completions" - if not payload.get("system_prompt"): - payload.pop("system_prompt", None) - - if bool(payload.get("stream")): - stream_timeout = httpx.Timeout(connect=5.0, read=None, write=60.0, pool=60.0) - client = httpx.AsyncClient(timeout=stream_timeout) - try: - upstream_request = client.build_request( - method="POST", - url=target_url, - json=payload, - headers=forward_headers, - ) - upstream = await client.send(upstream_request, stream=True) - except httpx.HTTPError as exc: - await client.aclose() - raise BackendProxyError(str(exc)) from exc - - passthrough_headers = {} - content_type = upstream.headers.get("content-type") - if content_type: - passthrough_headers["content-type"] = content_type - - if upstream.status_code >= 400 or not ( - content_type and "text/event-stream" in content_type.lower() - ): - body = await upstream.aread() - await upstream.aclose() - await client.aclose() - return BackendResponse( - status_code=upstream.status_code, - headers=passthrough_headers, - body=body, - ) - - async def _forward_stream() -> AsyncIterator[bytes]: - try: - async for chunk in upstream.aiter_raw(): - yield chunk - finally: - await upstream.aclose() - await client.aclose() - - return BackendResponse( - status_code=upstream.status_code, - headers=passthrough_headers, - stream=_forward_stream(), - ) - - timeout = httpx.Timeout(connect=5.0, read=None, write=60.0, pool=60.0) - try: - async with httpx.AsyncClient(timeout=timeout) as client: - upstream = await client.post(target_url, json=payload, headers=forward_headers) - except httpx.HTTPError as exc: - raise BackendProxyError(str(exc)) from exc - - passthrough_headers = {} - content_type = upstream.headers.get("content-type") - if content_type: - passthrough_headers["content-type"] = content_type - - return BackendResponse( - status_code=upstream.status_code, - headers=passthrough_headers, - body=upstream.content, - ) - - -class FakeLlamaRepository: - name = "fake" - - async def create_chat_completion( - self, - payload: dict[str, Any], - forward_headers: dict[str, str], - ) -> BackendResponse: - _ = forward_headers - - model = str(payload.get("model") or "qwen3-vl-4b-instruct-q4_k_m") - completion_id = f"chatcmpl-fake-{int(time.time() * 1000)}" - created = int(time.time()) - content = _fake_content(payload) - prefill_delay_seconds, chunk_delay_seconds = _read_fake_timing_config() - - if bool(payload.get("stream")): - return BackendResponse( - status_code=200, - headers={"content-type": "text/event-stream"}, - stream=_fake_stream( - completion_id=completion_id, - created=created, - model=model, - content=content, - prefill_delay_seconds=prefill_delay_seconds, - chunk_delay_seconds=chunk_delay_seconds, - ), - ) - - if prefill_delay_seconds > 0: - await asyncio.sleep(prefill_delay_seconds) - - usage = _estimate_usage(payload, content) - body = { - "id": completion_id, - "object": "chat.completion", - "created": created, - "model": model, - "choices": [ - { - "index": 0, - "message": {"role": "assistant", "content": content}, - "finish_reason": "stop", - } - ], - "usage": usage, - } - - return BackendResponse( - status_code=200, - headers={"content-type": "application/json"}, - body=json.dumps(body).encode("utf-8"), - ) - - -class ChatRepositoryManager: - def __init__(self, llama: LlamaCppRepository, fake: FakeLlamaRepository) -> None: - self._repos: dict[str, ChatCompletionRepository] = { - llama.name: llama, - fake.name: fake, - } - - async def create_chat_completion( - self, - backend: str, - payload: dict[str, Any], - forward_headers: dict[str, str], - ) -> BackendResponse: - repo = self._repos.get(backend) - if repo is None: - raise BackendProxyError(f"unknown backend: {backend}") - return await repo.create_chat_completion(payload, forward_headers) - - -def _fake_content(payload: dict[str, Any]) -> str: - last_user = _extract_last_user_text(payload) - if not last_user: - last_user = "hello from the starch dimension" - seed = _coerce_seed(payload.get("seed")) - if seed is None: - reply = random.choice(FAKE_PARODY_REPLIES) - else: - reply = random.Random(seed).choice(FAKE_PARODY_REPLIES) - return ( - "[fake-llama.cpp] " - f"{reply} " - f"Last user message (dramatically reenacted): {last_user}" - ) - - -def _coerce_seed(raw_seed: Any) -> int | None: - try: - if raw_seed is None: - return None - return int(raw_seed) - except (TypeError, ValueError): - return None - - -def _extract_last_user_text(payload: dict[str, Any]) -> str: - messages = payload.get("messages") - if not isinstance(messages, list): - return "" - - for msg in reversed(messages): - if not isinstance(msg, dict): - continue - if msg.get("role") != "user": - continue - content = msg.get("content") - if isinstance(content, str): - return content.strip() - if isinstance(content, list): - parts: list[str] = [] - for part in content: - if not isinstance(part, dict): - continue - if part.get("type") == "text" and isinstance(part.get("text"), str): - parts.append(part["text"]) - if parts: - return " ".join(parts).strip() - return "" - - -def _estimate_usage(payload: dict[str, Any], content: str) -> dict[str, int]: - messages = payload.get("messages") - if isinstance(messages, list): - prompt_chars = sum(len(str(m.get("content", ""))) for m in messages if isinstance(m, dict)) - else: - prompt_chars = 0 - - prompt_tokens = max(1, prompt_chars // 4) - completion_tokens = max(1, len(content) // 4) - return { - "prompt_tokens": prompt_tokens, - "completion_tokens": completion_tokens, - "total_tokens": prompt_tokens + completion_tokens, - } - - -def _to_sse_line(payload: dict[str, Any] | str) -> bytes: - if isinstance(payload, str): - return f"data: {payload}\n\n".encode("utf-8") - return f"data: {json.dumps(payload, separators=(',', ':'))}\n\n".encode("utf-8") - - -async def _fake_stream( - completion_id: str, - created: int, - model: str, - content: str, - prefill_delay_seconds: float, - chunk_delay_seconds: float, -) -> AsyncIterator[bytes]: - if prefill_delay_seconds > 0: - await asyncio.sleep(prefill_delay_seconds) - - first = { - "id": completion_id, - "object": "chat.completion.chunk", - "created": created, - "model": model, - "choices": [{"index": 0, "delta": {"role": "assistant"}, "finish_reason": None}], - } - yield _to_sse_line(first) - - for token in _tokenize_for_stream(content): - chunk = { - "id": completion_id, - "object": "chat.completion.chunk", - "created": created, - "model": model, - "choices": [{"index": 0, "delta": {"content": token}, "finish_reason": None}], - } - yield _to_sse_line(chunk) - await asyncio.sleep(chunk_delay_seconds) - - end_chunk = { - "id": completion_id, - "object": "chat.completion.chunk", - "created": created, - "model": model, - "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}], - } - yield _to_sse_line(end_chunk) - yield _to_sse_line("[DONE]") - - -def _tokenize_for_stream(content: str) -> list[str]: - words = content.split(" ") - if not words: - return [content] - - chunks: list[str] = [] - for idx, word in enumerate(words): - if idx < len(words) - 1: - chunks.append(word + " ") - else: - chunks.append(word) - return chunks - - -def _read_fake_timing_config() -> tuple[float, float]: - # Keep UI/dev fake mode Manual QA-paced by default, while tests can stay fast. - test_mode = os.getenv("POTATO_TEST_MODE", "0") == "1" - prefill_delay_ms = _safe_delay_ms( - os.getenv("POTATO_FAKE_PREFILL_DELAY_MS"), - default=DEFAULT_FAKE_PREFILL_DELAY_MS, - ) - stream_chunk_delay_ms = _safe_delay_ms( - os.getenv("POTATO_FAKE_STREAM_CHUNK_DELAY_MS"), - default=TEST_FAKE_STREAM_CHUNK_DELAY_MS if test_mode else DEFAULT_FAKE_STREAM_CHUNK_DELAY_MS, - ) - return prefill_delay_ms / 1000.0, stream_chunk_delay_ms / 1000.0 - - -def _safe_delay_ms(raw: str | None, default: int) -> int: - if raw is None: - return default - try: - value = int(raw) - except ValueError: - return default - return max(0, min(value, 60_000)) diff --git a/core/inferno/launch_config.py b/core/inferno/launch_config.py deleted file mode 100644 index b3aedac..0000000 --- a/core/inferno/launch_config.py +++ /dev/null @@ -1,87 +0,0 @@ -"""Llama-server launch configuration builder. - -Assembles the complete CLI argument list for llama-server from -pre-computed configuration values. Pure function — no file I/O, -no environment variable reads, no imports from Potato product code. -""" - -from __future__ import annotations - -from pathlib import Path - - -def build_llama_server_args( - *, - llama_server_bin: str | Path, - model_path: str | Path, - host: str = "0.0.0.0", - port: int = 8080, - ctx_size: int = 16384, - parallel: int = 1, - cache_ram_mib: int = 1024, - slot_save_path: str | Path, - mmproj_path: str | Path | None = None, - cache_type_k: str = "q8_0", - cache_type_v: str = "q8_0", - kv_flags: str | None = None, - flash_attn: bool = True, - jinja: bool = True, - no_warmup: bool = True, - no_mmap: bool = False, - reasoning_format: str = "none", - chat_template_kwargs: str = '{"enable_thinking": false}', - runtime_family: str | None = None, - extra_flags: str | None = None, -) -> list[str]: - """Build the complete llama-server command-line argument list. - - All business decisions (vision, device tuning, runtime family) must be - resolved by the caller before invoking this function. This keeps the - builder free of I/O and Potato-specific imports. - """ - args: list[str] = [ - str(llama_server_bin), - "--model", str(model_path), - "--host", str(host), - "--port", str(port), - "--ctx-size", str(ctx_size), - "--cache-ram", str(cache_ram_mib), - "--parallel", str(parallel), - "--slot-save-path", str(slot_save_path), - ] - - # Vision projector --------------------------------------------------- - if mmproj_path is not None: - args.extend(["--mmproj", str(mmproj_path)]) - - # KV cache ----------------------------------------------------------- - if kv_flags and kv_flags.strip(): - args.extend(kv_flags.strip().split()) - else: - args.extend(["--cache-type-k", cache_type_k, "--cache-type-v", cache_type_v]) - - # Feature toggles ---------------------------------------------------- - if jinja: - args.append("--jinja") - if flash_attn: - args.extend(["--flash-attn", "on"]) - if no_warmup: - args.append("--no-warmup") - if no_mmap: - args.append("--no-mmap") - - # Reasoning / chat template ------------------------------------------ - args.extend(["--reasoning-format", reasoning_format]) - args.extend(["--chat-template-kwargs", chat_template_kwargs]) - - # Runtime-family-specific flags -------------------------------------- - if runtime_family == "ik_llama": - args.extend(["--webui", "none"]) - elif runtime_family == "llama_cpp": - args.append("--no-webui") - - # Admin extra flags -------------------------------------------------- - if extra_flags and extra_flags.strip(): - args.extend(extra_flags.strip().split()) - - return args diff --git a/core/inferno/litert_adapter.py b/core/inferno/litert_adapter.py deleted file mode 100644 index 751f978..0000000 --- a/core/inferno/litert_adapter.py +++ /dev/null @@ -1,443 +0,0 @@ -"""LiteRT adapter — OpenAI-compatible HTTP wrapper around litert-lm-api. - -Runs as a standalone FastAPI process on the same port as llama-server (8080), -exposing /health and /v1/chat/completions so Potato's existing proxy works -unchanged. - -Conversation persistence: a single Conversation is kept alive across requests -so that the KV cache is reused for multi-turn chat (matching the Gallery -pattern). The conversation is reset only when the incoming message history -diverges from what we've already processed. -""" - -from __future__ import annotations - -import asyncio -import logging -import os -import time -import uuid -from contextlib import asynccontextmanager -from typing import Any - -from fastapi import FastAPI, Request -from fastapi.responses import JSONResponse, StreamingResponse - -logger = logging.getLogger("litert_adapter") - -try: - import litert_lm # type: ignore[import-untyped] -except ImportError: - litert_lm = None # type: ignore[assignment] - -_engine: Any = None -_conversation: Any = None -_vision_enabled: bool = False -# Tracks the messages we've already sent to the persistent conversation -# so we can detect continuations vs new sessions. -_conversation_history: list[dict[str, Any]] = [] -_lock = asyncio.Lock() - - -def _probe_vision_support(model_path: str) -> tuple[Any, bool]: - """Try to create an Engine with vision support. - - Returns (engine, vision_enabled). Falls back to text-only if the - current litert-lm build does not support the vision_backend kwarg - or if the runtime vision calculator is missing. - """ - try: - engine = litert_lm.Engine( - model_path, - backend=litert_lm.Backend.CPU, - vision_backend=litert_lm.Backend.CPU, - ) - logger.info("LiteRT vision probe succeeded — multimodal enabled") - return engine, True - except (TypeError, RuntimeError, Exception) as exc: - logger.info("LiteRT vision probe failed (%s) — text-only mode", exc) - engine = litert_lm.Engine(model_path, backend=litert_lm.Backend.CPU) - return engine, False - - -def _reset_conversation() -> None: - """Close the current conversation and create a fresh one.""" - global _conversation, _conversation_history - if _conversation is not None: - try: - _conversation.__exit__(None, None, None) - except Exception: - pass - _conversation = _engine.create_conversation() - _conversation.__enter__() - _conversation_history = [] - logger.info("Conversation reset") - - -@asynccontextmanager -async def _lifespan(app: FastAPI): - global _engine, _vision_enabled - model_path = os.environ.get("POTATO_MODEL_PATH", "") - if not model_path: - logger.error("POTATO_MODEL_PATH not set") - elif litert_lm is None: - logger.error("litert_lm package not installed — pip install litert-lm-api") - else: - logger.info("Loading LiteRT engine from %s", model_path) - try: - _engine, _vision_enabled = _probe_vision_support(model_path) - _reset_conversation() - logger.info("LiteRT engine loaded successfully (vision=%s)", _vision_enabled) - except Exception: - logger.exception("Failed to load LiteRT engine") - _engine = None - yield - global _conversation - if _conversation is not None: - try: - _conversation.__exit__(None, None, None) - except Exception: - pass - _conversation = None - if _engine is not None: - try: - _engine.__exit__(None, None, None) - except Exception: - logger.exception("Error cleaning up LiteRT engine") - _engine = None - - -app = FastAPI(lifespan=_lifespan) - - -@app.get("/health") -async def health(): - if _engine is None: - return JSONResponse(status_code=503, content={"status": "error", "reason": "engine_not_loaded"}) - return {"status": "ok", "vision": _vision_enabled} - - -def _build_openai_response(text: str, model: str, prompt_tokens: int = 0) -> dict[str, Any]: - completion_tokens = max(1, len(text) // 4) - return { - "id": f"chatcmpl-{uuid.uuid4().hex[:12]}", - "object": "chat.completion", - "created": int(time.time()), - "model": model, - "choices": [ - { - "index": 0, - "message": {"role": "assistant", "content": text}, - "finish_reason": "stop", - } - ], - "usage": { - "prompt_tokens": prompt_tokens, - "completion_tokens": completion_tokens, - "total_tokens": prompt_tokens + completion_tokens, - }, - } - - -def _extract_text(response: Any) -> str: - if isinstance(response, dict): - content_list = response.get("content", []) - if content_list and isinstance(content_list[0], dict): - return content_list[0].get("text", "") - return "" - - -def _convert_openai_to_litert_content(content: str | list[dict[str, Any]]) -> str | list[dict[str, Any]]: - """Convert OpenAI-format multimodal content to litert-lm format. - - OpenAI: [{"type": "text", ...}, {"type": "image_url", "image_url": {"url": "data:...;base64,AAAA"}}] - LiteRT: [{"type": "text", ...}, {"type": "image", "blob": "AAAA"}] - - Only base64 data URLs are supported. Remote URLs (https://...) are - rejected — LiteRT expects raw image bytes, not a URL fetch. - """ - if isinstance(content, str): - return content - parts: list[dict[str, Any]] = [] - for part in content: - if part.get("type") == "image_url": - data_url = part.get("image_url", {}).get("url", "") - if not data_url.startswith("data:"): - raise ValueError(f"Only base64 data URLs are supported for LiteRT vision, got: {data_url[:60]}") - blob = data_url.split(",", 1)[1] if "," in data_url else "" - parts.append({"type": "image", "blob": blob}) - else: - parts.append(part) - return parts - - -def _content_equal(a: str | list | None, b: str | list | None) -> bool: - """Compare message content that may be a string or a list of parts.""" - return a == b - - -def _has_image_content(messages: list[dict[str, Any]]) -> bool: - """Return True if any message contains an image_url content part.""" - for msg in messages: - content = msg.get("content") - if isinstance(content, list): - for part in content: - if part.get("type") == "image_url": - return True - return False - - -def _estimate_prompt_chars(messages: list[dict[str, Any]]) -> int: - """Estimate total character count across all messages, handling multimodal.""" - total = 0 - for msg in messages: - content = msg.get("content", "") - if isinstance(content, str): - total += len(content) - elif isinstance(content, list): - for part in content: - if part.get("type") == "text": - total += len(part.get("text", "")) - else: - total += 256 # approximate image token cost in chars - return total - - -def _messages_match(incoming: list[dict[str, Any]], history: list[dict[str, Any]]) -> bool: - """Check if incoming message history is a continuation of our tracked history. - - Returns True only when history is non-empty AND every tracked message - matches the corresponding incoming message (i.e. incoming starts with - the full tracked history). - """ - if not history: - return False - if len(history) > len(incoming): - return False - for prev, inc in zip(history, incoming): - if prev.get("role") != inc.get("role") or not _content_equal(prev.get("content"), inc.get("content")): - return False - return True - - -def _prepare_conversation_sync(messages: list[dict[str, str]]) -> None: - """Ensure the persistent conversation is ready for the final user message. - - If the incoming history matches what we've already sent, this is a no-op - (KV cache hit). Otherwise, reset and replay user/system messages only — - assistant turns are skipped since send_message would generate fresh - (different) output. The model responds implicitly to each replayed user - turn, rebuilding the KV cache with approximate context. - """ - global _conversation_history - - history_messages = messages[:-1] # everything except the new user message - - if _messages_match(history_messages, _conversation_history): - # Continuation — nothing to do, KV cache covers it. - return - - # History diverged (new session, page reload, session switch). - logger.info("Conversation history diverged, resetting and replaying user turns") - _reset_conversation() - - for msg in history_messages: - role = msg.get("role", "user") - content = msg.get("content", "") - if role == "assistant": - # Skip — the model already generated an implicit response to - # the preceding user turn via send_message above. - continue - if role == "system": - _conversation.send_message(f"[System instruction] {content}" if isinstance(content, str) else content) - elif isinstance(content, list): - # Multimodal: send_message expects a dict with role/content - converted = _convert_openai_to_litert_content(content) - _conversation.send_message({"role": "user", "content": converted}) - else: - _conversation.send_message(content) - - _conversation_history = list(history_messages) - - -def _run_inference_sync(messages: list[dict[str, Any]], stream: bool) -> Any: - """Run inference synchronously — must be called via asyncio.to_thread.""" - _prepare_conversation_sync(messages) - - raw_content = messages[-1].get("content", "") if messages else "" - if isinstance(raw_content, list): - # Multimodal: send_message expects {"role": "user", "content": [...]} - final_content = {"role": "user", "content": _convert_openai_to_litert_content(raw_content)} - else: - final_content = raw_content - - if stream: - return _conversation.send_message_async(final_content) - else: - response = _conversation.send_message(final_content) - text = _extract_text(response) - # Track the full exchange in history (store original content for matching) - _conversation_history.append(messages[-1]) - _conversation_history.append({"role": "assistant", "content": text}) - return text - - -@app.post("/v1/chat/completions") -async def chat_completions(request: Request): - if _engine is None: - return JSONResponse( - status_code=503, - content={"error": {"message": "LiteRT engine not loaded", "type": "server_error"}}, - ) - - try: - payload = await request.json() - except Exception: - return JSONResponse( - status_code=400, - content={"error": {"message": "Invalid JSON", "type": "invalid_request_error"}}, - ) - - messages = payload.get("messages", []) - if not messages: - return JSONResponse( - status_code=400, - content={"error": {"message": "messages required", "type": "invalid_request_error"}}, - ) - - # Reject multimodal input when vision is not available - if not _vision_enabled and _has_image_content(messages): - return JSONResponse( - status_code=400, - content={"error": {"message": "Vision input is not supported by this model/runtime configuration", "type": "invalid_request_error"}}, - ) - - stream = payload.get("stream", False) - model_name = payload.get("model", os.environ.get("POTATO_MODEL_PATH", "litert")) - - async with _lock: - try: - if stream: - iterator = await asyncio.to_thread( - _run_inference_sync, messages, True, - ) - - # Bridge sync iterator → async generator via a queue so - # tokens stream to the client as they arrive. - queue: asyncio.Queue[str | None] = asyncio.Queue() - collected_text: list[str] = [] - generation_start = time.monotonic() - first_token_time: float | None = None - - async def _producer(): - nonlocal first_token_time - def _iterate(): - nonlocal first_token_time - for chunk in iterator: - text = _extract_text(chunk) - if text: - if first_token_time is None: - first_token_time = time.monotonic() - collected_text.append(text) - queue.put_nowait(text) - queue.put_nowait(None) # sentinel - try: - await asyncio.to_thread(_iterate) - except Exception: - queue.put_nowait(None) - raise - - producer_task = asyncio.create_task(_producer()) - - async def _stream_chunks(): - response_id = f"chatcmpl-{uuid.uuid4().hex[:12]}" - try: - while True: - text = await queue.get() - if text is None: - break - chunk_data = { - "id": response_id, - "object": "chat.completion.chunk", - "created": int(time.time()), - "model": model_name, - "choices": [ - { - "index": 0, - "delta": {"content": text}, - "finish_reason": None, - } - ], - } - yield f"data: {_json_dumps(chunk_data)}\n\n" - - # Build timings matching llama-server's format so the - # chat UI stats display works identically. - now = time.monotonic() - full_text = "".join(collected_text) - # Estimate token count from text (~4 chars per token) - predicted_n = max(1, len(full_text) // 4) - prompt_ms = ((first_token_time - generation_start) * 1000) if first_token_time else 0 - # Decode time = total minus prompt/prefill time - decode_start = first_token_time or generation_start - predicted_ms = (now - decode_start) * 1000 - per_token_ms = (predicted_ms / predicted_n) if predicted_n > 0 else 0 - per_second = (predicted_n / (predicted_ms / 1000)) if predicted_ms > 0 else 0 - prompt_n = _estimate_prompt_chars(messages) // 4 - - stop_chunk = { - "id": response_id, - "object": "chat.completion.chunk", - "created": int(time.time()), - "model": model_name, - "choices": [ - { - "index": 0, - "delta": {}, - "finish_reason": "stop", - } - ], - "timings": { - "prompt_n": prompt_n, - "prompt_ms": prompt_ms, - "prompt_per_second": (prompt_n / (prompt_ms / 1000)) if prompt_ms > 0 else 0, - "predicted_n": predicted_n, - "predicted_ms": predicted_ms, - "predicted_per_token_ms": per_token_ms, - "predicted_per_second": per_second, - }, - } - yield f"data: {_json_dumps(stop_chunk)}\n\n" - yield "data: [DONE]\n\n" - finally: - await producer_task - # Track the streamed exchange in conversation history - full_response = "".join(collected_text) - _conversation_history.append(messages[-1]) - _conversation_history.append({"role": "assistant", "content": full_response}) - - return StreamingResponse( - _stream_chunks(), - media_type="text/event-stream", - headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}, - ) - else: - text = await asyncio.to_thread(_run_inference_sync, messages, False) - prompt_tokens = _estimate_prompt_chars(messages) // 4 - return JSONResponse(content=_build_openai_response(str(text), model_name, prompt_tokens)) - except ValueError as ve: - return JSONResponse( - status_code=400, - content={"error": {"message": str(ve), "type": "invalid_request_error"}}, - ) - except Exception: - logger.exception("Inference error") - return JSONResponse( - status_code=500, - content={"error": {"message": "Inference failed", "type": "server_error"}}, - ) - - -def _json_dumps(obj: Any) -> str: - import json - return json.dumps(obj, separators=(",", ":")) diff --git a/core/inferno/model_families.py b/core/inferno/model_families.py deleted file mode 100644 index 3824114..0000000 --- a/core/inferno/model_families.py +++ /dev/null @@ -1,189 +0,0 @@ -from __future__ import annotations - -import re -from pathlib import Path -from typing import Any, Final - -QWEN35_PROJECTOR_REPO_RULES: Final[tuple[tuple[tuple[str, ...], str], ...]] = ( - (("35b", "a3b", "hauhaucs"), "HauhauCS/Qwen3.5-35B-A3B-Uncensored-HauhauCS-Aggressive"), - (("35b", "a3b"), "AesSedai/Qwen3.5-35B-A3B-GGUF"), - (("9b", "byteshape"), "byteshape/Qwen3.5-9B-GGUF"), - (("9b",), "unsloth/Qwen3.5-9B-GGUF"), - (("4b", "hauhaucs"), "HauhauCS/Qwen3.5-4B-Uncensored-HauhauCS-Aggressive"), - (("4b",), "unsloth/Qwen3.5-4B-GGUF"), - (("2b",), "unsloth/Qwen3.5-2B-GGUF"), - (("0.8b",), "unsloth/Qwen3.5-0.8B-GGUF"), -) - - -def _normalized_model_name(filename: str | None) -> str: - return str(filename or "").strip().lower() - - -def is_qwen35_filename(filename: str | None) -> bool: - model_name = _normalized_model_name(filename) - return bool(model_name) and "qwen" in model_name and "3.5" in model_name - - -def _token_at_boundary(token: str, text: str) -> bool: - """Match token only when surrounded by non-alphanumeric boundaries. - - Prevents substring collisions like '9b' matching inside '19b' - or '2b' matching inside '3.92bpw'. - """ - return bool(re.search(r"(? str | None: - model_name = _normalized_model_name(filename) - if not is_qwen35_filename(model_name): - return None - - # Match tokens against filename + source URL so publisher-specific rules - # fire even when the filename itself doesn't contain the publisher name. - match_text = model_name - if source_url: - match_text = _normalized_model_name(source_url) + " " + match_text - - for required_tokens, repo in QWEN35_PROJECTOR_REPO_RULES: - if all(_token_at_boundary(token, match_text) for token in required_tokens): - return repo - return None - - -GEMMA4_PROJECTOR_REPO_RULES: Final[tuple[tuple[tuple[str, ...], str], ...]] = ( - (("e2b",), "unsloth/gemma-4-E2B-it-GGUF"), - (("e4b",), "unsloth/gemma-4-E4B-it-GGUF"), - (("26b", "a4b"), "unsloth/gemma-4-26B-A4B-it-GGUF"), -) - - -def is_gemma4_filename(filename: str | None) -> bool: - model_name = _normalized_model_name(filename) - return bool(model_name) and "gemma" in model_name and _token_at_boundary("4", model_name) - - -def _gemma4_projector_repo(filename: str | None, source_url: str | None = None) -> str | None: - model_name = _normalized_model_name(filename) - if not is_gemma4_filename(model_name): - return None - - match_text = model_name - if source_url: - match_text = _normalized_model_name(source_url) + " " + match_text - - for required_tokens, repo in GEMMA4_PROJECTOR_REPO_RULES: - if all(_token_at_boundary(token, match_text) for token in required_tokens): - return repo - return None - - -def _is_gemma4_26b_a4b(filename: str | None) -> bool: - model_name = _normalized_model_name(filename) - return is_gemma4_filename(model_name) and "26b" in model_name and "a4b" in model_name - - -def recommended_runtime_for_model(filename: str | None) -> str | None: - """Return the preferred runtime family for a model, or None for no preference.""" - if filename and _normalized_model_name(filename).endswith(".litertlm"): - return "litert" - if _is_gemma4_26b_a4b(filename): - return "ik_llama" - return None - - -def projector_repo_for_model(filename: str | None, source_url: str | None = None) -> str | None: - return ( - _qwen35_projector_repo(filename, source_url=source_url) - or _gemma4_projector_repo(filename, source_url=source_url) - ) - - -# --------------------------------------------------------------------------- -# Vision family detection & projector candidates -# --------------------------------------------------------------------------- - - -def _is_vision_family(filename: str) -> bool: - """Return True if the filename belongs to a curated vision model family.""" - model_name = filename.strip().lower() - if "qwen" in model_name and "3.5" in model_name: - return True - return is_gemma4_filename(model_name) - - -def default_projector_candidates_for_model(filename: str | None) -> list[str]: - """Return ordered list of mmproj filename candidates for a given model.""" - model_name = str(filename or "").strip() - if not model_name or not _is_vision_family(model_name): - return [] - - stem = Path(model_name).stem - stem_candidates = [stem] - trimmed_stem = stem - while True: - next_stem = re.sub( - r"-(?:UD-)?(?:\d+(?:\.\d+)?bpw|I?Q\d+(?:_[A-Za-z0-9]+)*|MXFP\d+_MOE)$", - "", - trimmed_stem, - flags=re.IGNORECASE, - ) - if next_stem == trimmed_stem or not next_stem: - break - trimmed_stem = next_stem - if trimmed_stem not in stem_candidates: - stem_candidates.append(trimmed_stem) - - candidates: list[str] = [] - # Model-specific candidates first (f16 preferred, bf16 fallback), - # then generic F16 last. This ensures a downloaded model-specific - # bf16 projector is found before a stale generic F16. - for precision in ("f16", "bf16"): - for candidate_stem in stem_candidates: - candidate_name = f"mmproj-{candidate_stem}-{precision}.gguf" - if candidate_name not in candidates: - candidates.append(candidate_name) - if "mmproj-F16.gguf" not in candidates: - candidates.append("mmproj-F16.gguf") - return candidates - - -def build_model_projector_status(models_dir: Path, model: dict[str, Any]) -> dict[str, Any]: - """Build projector status for a model (presence, path, candidates).""" - from .model_registry import normalize_model_settings - - filename = str(model.get("filename") or "") - settings = normalize_model_settings(model.get("settings"), filename=filename) - vision = settings.get("vision", {}) - projector_mode = str(vision.get("projector_mode") or "default").strip().lower() - configured_filename = str(vision.get("projector_filename") or "").strip() or None - default_candidates = default_projector_candidates_for_model(filename) - search_names: list[str] = [] - if projector_mode == "custom": - if configured_filename: - search_names.append(configured_filename) - else: - for candidate in default_candidates: - if candidate not in search_names: - search_names.append(candidate) - if configured_filename and configured_filename not in search_names: - search_names.append(configured_filename) - - resolved_name = configured_filename - present = False - resolved_path = None - for candidate in search_names: - candidate_path = models_dir / candidate - if candidate_path.exists(): - present = True - resolved_name = candidate - resolved_path = candidate_path - break - - return { - "configured_filename": configured_filename, - "filename": resolved_name, - "present": present, - "path": str(resolved_path) if resolved_path is not None else None, - "default_candidates": default_candidates, - } diff --git a/core/inferno/model_registry.py b/core/inferno/model_registry.py deleted file mode 100644 index 18aead0..0000000 --- a/core/inferno/model_registry.py +++ /dev/null @@ -1,807 +0,0 @@ -"""Model registry, format handling, and settings normalization for Inferno. - -This module owns the model lifecycle logic that is inference-owned: -format detection, URL validation, filename sanitization, settings -normalization, registry CRUD, file operations, and projector download. - -Product-specific defaults (device class, default model selection) are -injected via ModelStoreConfig — this module never imports from -core.model_state, core.runtime_state, or any Potato-specific code. -""" - -from __future__ import annotations - -import dataclasses -import json -import os -import re -from pathlib import Path -from typing import Any - -logger = __import__("logging").getLogger("potato") - -# --------------------------------------------------------------------------- -# Constants -# --------------------------------------------------------------------------- - -VALID_MODEL_EXTENSIONS = (".gguf", ".litertlm") - -MODELS_STATE_VERSION = 1 - -DEFAULT_MODEL_CHAT_SETTINGS = { - "temperature": 0.7, - "top_p": 0.8, - "top_k": 20, - "repetition_penalty": 1.0, - "presence_penalty": 1.5, - "max_tokens": 16384, - "stream": True, - "generation_mode": "random", - "seed": 42, - "system_prompt": "", - "cache_prompt": True, -} - -DEFAULT_MODEL_VISION_SETTINGS = { - "enabled": False, - "projector_mode": "default", - "projector_filename": None, -} - - -# --------------------------------------------------------------------------- -# Store configuration -# --------------------------------------------------------------------------- - - -@dataclasses.dataclass(frozen=True) -class ModelStoreConfig: - """Filesystem and product-policy context for model registry operations. - - This bundles the paths and defaults that the Potato layer injects so - that Inferno never needs to import RuntimeConfig. - """ - - models_dir: Path - state_path: Path - default_filename: str - default_url: str - known_default_filenames: tuple[str, ...] - current_model_filename: str - - -# --------------------------------------------------------------------------- -# Exceptions -# --------------------------------------------------------------------------- - - -class ModelSettingsValidationError(ValueError): - def __init__(self, field: str) -> None: - super().__init__(field) - self.field = field - - -# --------------------------------------------------------------------------- -# Format detection & URL validation -# --------------------------------------------------------------------------- - - -def model_format_for_filename(filename: str) -> str: - if filename.lower().endswith(".litertlm"): - return "litertlm" - return "gguf" - - -def _has_valid_model_extension(filename: str) -> bool: - lower = filename.lower() - return any(lower.endswith(ext) for ext in VALID_MODEL_EXTENSIONS) - - -def validate_model_url(source_url: str) -> tuple[bool, str, str]: - from urllib.parse import unquote, urlparse - - parsed = urlparse(source_url.strip()) - if parsed.scheme != "https": - return False, "https_required", "" - basename = unquote(Path(parsed.path).name) - if not basename: - return False, "filename_missing", "" - if not _has_valid_model_extension(basename): - return False, "unsupported_model_format", "" - safe_name = _sanitize_filename(basename) - if not _has_valid_model_extension(safe_name): - safe_name = f"{Path(safe_name).stem}.gguf" - return True, "", safe_name - - -# --------------------------------------------------------------------------- -# Filename / ID utilities -# --------------------------------------------------------------------------- - - -def _sanitize_filename(filename: str) -> str: - candidate = Path(filename).name.strip() - candidate = re.sub(r"[^A-Za-z0-9._-]+", "-", candidate) - candidate = candidate.lstrip(".") - return candidate or "model.gguf" - - -def _slugify_id(raw: str) -> str: - slug = re.sub(r"[^a-z0-9]+", "-", raw.lower()).strip("-") - return slug or "model" - - -def _unique_model_id(base_id: str, existing_ids: set[str]) -> str: - candidate = base_id - idx = 2 - while candidate in existing_ids: - candidate = f"{base_id}-{idx}" - idx += 1 - return candidate - - -def _unique_filename(base_name: str, existing_names: set[str]) -> str: - stem = Path(base_name).stem - suffix = Path(base_name).suffix or ".gguf" - candidate = f"{stem}{suffix}" - idx = 2 - while candidate in existing_names: - candidate = f"{stem}-{idx}{suffix}" - idx += 1 - return candidate - - -# --------------------------------------------------------------------------- -# Model detection -# --------------------------------------------------------------------------- - - -def is_qwen35_a3b_filename(filename: str | None) -> bool: - value = str(filename or "").strip().lower() - return bool(value) and "qwen" in value and "3.5" in value and "35b" in value and "a3b" in value - - -def model_supports_vision_filename(filename: str | None) -> bool: - value = str(filename or "").strip().lower() - if not value: - return False - if "qwen3" in value and "vl" in value: - return True - if "qwen" in value and "3.5" in value: - return True - from .model_families import is_gemma4_filename - - if is_gemma4_filename(value): - return True - return False - - -# --------------------------------------------------------------------------- -# Settings normalization -# --------------------------------------------------------------------------- - - -def _coerce_float_setting(raw_value: Any, *, field: str, default: float) -> float: - value = default if raw_value is None else raw_value - try: - return float(value) - except (TypeError, ValueError) as exc: - raise ModelSettingsValidationError(field) from exc - - -def _coerce_int_setting(raw_value: Any, *, field: str, default: int) -> int: - value = default if raw_value is None else raw_value - try: - return int(value) - except (TypeError, ValueError) as exc: - raise ModelSettingsValidationError(field) from exc - - -def _normalize_chat_settings(raw_value: Any) -> dict[str, Any]: - raw = raw_value if isinstance(raw_value, dict) else {} - return { - "temperature": _coerce_float_setting( - raw.get("temperature"), - field="chat.temperature", - default=DEFAULT_MODEL_CHAT_SETTINGS["temperature"], - ), - "top_p": _coerce_float_setting( - raw.get("top_p"), - field="chat.top_p", - default=DEFAULT_MODEL_CHAT_SETTINGS["top_p"], - ), - "top_k": _coerce_int_setting( - raw.get("top_k"), - field="chat.top_k", - default=DEFAULT_MODEL_CHAT_SETTINGS["top_k"], - ), - "repetition_penalty": _coerce_float_setting( - raw.get("repetition_penalty"), - field="chat.repetition_penalty", - default=DEFAULT_MODEL_CHAT_SETTINGS["repetition_penalty"], - ), - "presence_penalty": _coerce_float_setting( - raw.get("presence_penalty"), - field="chat.presence_penalty", - default=DEFAULT_MODEL_CHAT_SETTINGS["presence_penalty"], - ), - "max_tokens": _coerce_int_setting( - raw.get("max_tokens"), - field="chat.max_tokens", - default=DEFAULT_MODEL_CHAT_SETTINGS["max_tokens"], - ), - "stream": bool(raw.get("stream", DEFAULT_MODEL_CHAT_SETTINGS["stream"])), - "generation_mode": ( - "deterministic" - if str(raw.get("generation_mode", DEFAULT_MODEL_CHAT_SETTINGS["generation_mode"])).strip().lower() == "deterministic" - else "random" - ), - "seed": _coerce_int_setting( - raw.get("seed"), - field="chat.seed", - default=DEFAULT_MODEL_CHAT_SETTINGS["seed"], - ), - "system_prompt": str(raw.get("system_prompt", DEFAULT_MODEL_CHAT_SETTINGS["system_prompt"]) or ""), - "cache_prompt": bool(raw.get("cache_prompt", DEFAULT_MODEL_CHAT_SETTINGS["cache_prompt"])), - } - - -def _normalize_vision_settings(raw_value: Any, *, filename: str) -> dict[str, Any]: - raw = raw_value if isinstance(raw_value, dict) else {} - default_enabled = model_supports_vision_filename(filename) - projector_filename_raw = raw.get("projector_filename") - projector_filename = None - if projector_filename_raw is not None: - value = str(projector_filename_raw).strip() - projector_filename = value or None - projector_mode = str(raw.get("projector_mode", DEFAULT_MODEL_VISION_SETTINGS["projector_mode"]) or "default").strip().lower() - if projector_mode not in {"default", "custom"}: - projector_mode = "default" - return { - "enabled": bool(raw.get("enabled", default_enabled)), - "projector_mode": projector_mode, - "projector_filename": projector_filename, - } - - -def normalize_model_settings(raw_value: Any, *, filename: str) -> dict[str, Any]: - raw = raw_value if isinstance(raw_value, dict) else {} - return { - "chat": _normalize_chat_settings(raw.get("chat")), - "vision": _normalize_vision_settings(raw.get("vision"), filename=filename), - } - - -# --------------------------------------------------------------------------- -# Capabilities -# --------------------------------------------------------------------------- - - -def build_model_capabilities(filename: str | None) -> dict[str, Any]: - return { - "vision": model_supports_vision_filename(filename), - } - - -def apply_model_chat_defaults(payload: dict[str, Any], *, active_model_filename: str | None) -> dict[str, Any]: - if not is_qwen35_a3b_filename(active_model_filename): - return payload - - chat_template_kwargs = payload.get("chat_template_kwargs") - if isinstance(chat_template_kwargs, dict) and "enable_thinking" in chat_template_kwargs: - return payload - - updated = dict(payload) - if isinstance(chat_template_kwargs, dict): - merged = dict(chat_template_kwargs) - else: - merged = {} - merged["enable_thinking"] = False - updated["chat_template_kwargs"] = merged - return updated - - -# --------------------------------------------------------------------------- -# Registry queries (pure) -# --------------------------------------------------------------------------- - - -def get_model_by_id(state: dict[str, Any], model_id: str) -> dict[str, Any] | None: - for item in state.get("models", []): - if isinstance(item, dict) and item.get("id") == model_id: - return item - return None - - -def _is_discoverable_local_model_filename(filename: str) -> bool: - name = _sanitize_filename(filename) - if not _has_valid_model_extension(name): - return False - stem = Path(name).stem.lower() - if stem.startswith("mmproj") or "mmproj" in stem: - return False - return True - - -# --------------------------------------------------------------------------- -# Atomic write utility -# --------------------------------------------------------------------------- - - -def _atomic_write_json(path: Path, payload: dict[str, Any]) -> None: - try: - path.parent.mkdir(parents=True, exist_ok=True) - import tempfile - - fd, tmp_name = tempfile.mkstemp(dir=path.parent, suffix=".tmp") - with os.fdopen(fd, "w", encoding="utf-8") as f: - f.write(json.dumps(payload)) - os.replace(tmp_name, path) - except OSError: - logger.warning("Could not persist JSON state to %s", path, exc_info=True) - - -# --------------------------------------------------------------------------- -# File / path operations -# --------------------------------------------------------------------------- - - -def model_file_path(models_dir: Path, filename: str) -> Path: - return models_dir / filename - - -def model_file_present(models_dir: Path, filename: str) -> bool: - path = model_file_path(models_dir, filename) - try: - return path.exists() and path.stat().st_size > 0 - except OSError: - return False - - -def describe_model_storage(models_dir: Path, filename: str) -> dict[str, Any]: - path = model_file_path(models_dir, filename) - size_bytes = 0 - exists = False - - try: - exists = path.exists() - except OSError: - exists = False - - if exists: - try: - size_bytes = max(0, int(path.stat().st_size)) - except OSError: - size_bytes = 0 - - return { - "location": "local", - "size_bytes": size_bytes, - "exists": exists, - } - - -def resolve_model_runtime_path(models_dir: Path, filename: str) -> Path: - """Return the real filesystem path for a model, resolving symlinks.""" - path = model_file_path(models_dir, filename) - try: - if path.is_symlink(): - return path.resolve(strict=False) - except OSError: - return path - return path - - -def discover_local_model_filenames(models_dir: Path) -> list[str]: - try: - children = list(models_dir.iterdir()) - except OSError: - return [] - names: list[str] = [] - for child in children: - if not child.is_file(): - continue - filename = _sanitize_filename(child.name) - if not _is_discoverable_local_model_filename(filename): - continue - names.append(filename) - return sorted(set(names)) - - -# --------------------------------------------------------------------------- -# Default model record (built from ModelStoreConfig) -# --------------------------------------------------------------------------- - - -def _default_model_record(store: ModelStoreConfig) -> dict[str, Any]: - return { - "id": "default", - "filename": store.default_filename, - "source_url": store.default_url, - "source_type": "url", - "status": "not_downloaded", - "error": None, - "settings": normalize_model_settings(None, filename=store.default_filename), - } - - -# --------------------------------------------------------------------------- -# State normalization and persistence -# --------------------------------------------------------------------------- - - -def _normalize_models_state(store: ModelStoreConfig, raw: dict[str, Any] | None = None) -> dict[str, Any]: - payload = raw or {} - models_raw = payload.get("models") - models: list[dict[str, Any]] = [] - seen_ids: set[str] = set() - seen_filenames: set[str] = set() - - if isinstance(models_raw, list): - for item in models_raw: - if not isinstance(item, dict): - continue - source_url = str(item.get("source_url") or "") - filename = _sanitize_filename(str(item.get("filename") or "")) - if not _has_valid_model_extension(filename): - filename = f"{Path(filename).stem}.gguf" - item_id_raw = str(item.get("id") or _slugify_id(Path(filename).stem)) - item_id = _unique_model_id(_slugify_id(item_id_raw), seen_ids) - filename = _unique_filename(filename, seen_filenames) - seen_ids.add(item_id) - seen_filenames.add(filename) - source_type_raw = str(item.get("source_type") or "").strip().lower() - if source_url: - source_type = source_type_raw or "url" - elif source_type_raw in {"upload", "local_file"}: - source_type = source_type_raw - else: - source_type = "upload" - models.append( - { - "id": item_id, - "filename": filename, - "source_url": source_url or None, - "source_type": source_type, - "status": str(item.get("status") or "not_downloaded"), - "error": item.get("error"), - "settings": normalize_model_settings(item.get("settings"), filename=filename), - } - ) - - if not models: - default_model = _default_model_record(store) - models.append(default_model) - seen_ids.add(default_model["id"]) - seen_filenames.add(default_model["filename"]) - elif "default" not in seen_ids: - default_model = _default_model_record(store) - default_model["id"] = _unique_model_id("default", seen_ids) - default_model["filename"] = _unique_filename(default_model["filename"], seen_filenames) - models.insert(0, default_model) - seen_ids.add(default_model["id"]) - seen_filenames.add(default_model["filename"]) - - for local_filename in discover_local_model_filenames(store.models_dir): - if local_filename in seen_filenames: - continue - local_id = _unique_model_id(_slugify_id(Path(local_filename).stem), seen_ids) - models.append( - { - "id": local_id, - "filename": local_filename, - "source_url": None, - "source_type": "local_file", - "status": "ready", - "error": None, - "settings": normalize_model_settings(None, filename=local_filename), - } - ) - seen_ids.add(local_id) - seen_filenames.add(local_filename) - - runtime_model_name = _sanitize_filename(store.current_model_filename) if store.current_model_filename else "" - active_model_id = str(payload.get("active_model_id") or "").strip() - if active_model_id not in seen_ids: - runtime_match = next( - ( - str(item.get("id") or "") - for item in models - if isinstance(item, dict) and _sanitize_filename(str(item.get("filename") or "")) == runtime_model_name - ), - "", - ) - active_model_id = runtime_match or models[0]["id"] - - default_model_id = str(payload.get("default_model_id") or "default") - if default_model_id not in seen_ids: - default_model_id = "default" if "default" in seen_ids else models[0]["id"] - - current_download_model_id = payload.get("current_download_model_id") - if current_download_model_id not in seen_ids: - current_download_model_id = None - - return { - "version": MODELS_STATE_VERSION, - "countdown_enabled": bool(payload.get("countdown_enabled", True)), - "default_model_downloaded_once": bool(payload.get("default_model_downloaded_once", False)), - "active_model_id": active_model_id, - "default_model_id": default_model_id, - "current_download_model_id": current_download_model_id, - "models": models, - } - - -def ensure_models_state(store: ModelStoreConfig) -> dict[str, Any]: - raw: dict[str, Any] | None = None - if store.state_path.exists(): - try: - loaded = json.loads(store.state_path.read_text(encoding="utf-8")) - if isinstance(loaded, dict): - raw = loaded - except (OSError, json.JSONDecodeError): - raw = None - - normalized = _normalize_models_state(store, raw) - default_model_id = str(normalized.get("default_model_id") or "default") - default_model = get_model_by_id(normalized, default_model_id) - if isinstance(default_model, dict): - default_filename = str(default_model.get("filename") or "") - if default_filename in store.known_default_filenames and model_file_present(store.models_dir, default_filename): - normalized["default_model_downloaded_once"] = True - _atomic_write_json(store.state_path, normalized) - return normalized - - -def save_models_state(store: ModelStoreConfig, state: dict[str, Any]) -> dict[str, Any]: - normalized = _normalize_models_state(store, state) - _atomic_write_json(store.state_path, normalized) - return normalized - - -# --------------------------------------------------------------------------- -# Registry mutations -# --------------------------------------------------------------------------- - - -def update_model_settings( - store: ModelStoreConfig, - *, - model_id: str, - settings: dict[str, Any], -) -> tuple[bool, str, dict[str, Any] | None]: - state = ensure_models_state(store) - model = get_model_by_id(state, model_id) - if model is None: - return False, "model_not_found", None - filename = str(model.get("filename") or "") - try: - model["settings"] = normalize_model_settings(settings, filename=filename) - except ModelSettingsValidationError: - return False, "invalid_settings", None - saved = save_models_state(store, state) - updated = get_model_by_id(saved, model_id) - return True, "updated", updated - - -def register_model_url(store: ModelStoreConfig, source_url: str, alias: str | None = None) -> tuple[bool, str, dict[str, Any] | None]: - ok, reason, filename = validate_model_url(source_url) - if not ok: - return False, reason, None - - state = ensure_models_state(store) - models = state.get("models", []) - assert isinstance(models, list) - existing_ids = {str(item.get("id")) for item in models if isinstance(item, dict)} - existing_names = {str(item.get("filename")) for item in models if isinstance(item, dict)} - - for item in models: - if isinstance(item, dict) and str(item.get("source_url") or "") == source_url: - saved = save_models_state(store, state) - model = get_model_by_id(saved, str(item.get("id") or "")) - return True, "already_exists", model - - preferred_name = filename - if alias: - alias_safe = _sanitize_filename(alias) - if not alias_safe.lower().endswith(".gguf"): - alias_safe = f"{Path(alias_safe).stem}.gguf" - if alias_safe: - preferred_name = alias_safe - - final_name = _unique_filename(preferred_name, existing_names) - model_id = _unique_model_id(_slugify_id(Path(final_name).stem), existing_ids) - model_record = { - "id": model_id, - "filename": final_name, - "source_url": source_url, - "source_type": "url", - "status": "ready" if model_file_present(store.models_dir, final_name) else "not_downloaded", - "error": None, - } - models.append(model_record) - saved = save_models_state(store, state) - created = get_model_by_id(saved, model_id) - return True, "registered", created - - -def delete_model(store: ModelStoreConfig, *, model_id: str) -> tuple[bool, str, bool, int, bool]: - state = ensure_models_state(store) - active_model_id = str(state.get("active_model_id") or "") - target = get_model_by_id(state, model_id) - if target is None: - return False, "model_not_found", False, 0, False - was_active = model_id == active_model_id - - filename = str(target.get("filename") or "") - models = state.get("models", []) - assert isinstance(models, list) - - same_filename_elsewhere = any( - isinstance(item, dict) - and str(item.get("id") or "") != model_id - and str(item.get("filename") or "") == filename - for item in models - ) - - deleted_file = False - freed_bytes = 0 - if filename and not same_filename_elsewhere: - candidate_paths = ( - model_file_path(store.models_dir, filename), - model_file_path(store.models_dir, filename + ".part"), - ) - for candidate_path in candidate_paths: - candidate_is_symlink = False - try: - candidate_is_symlink = candidate_path.is_symlink() - except OSError: - candidate_is_symlink = False - if not candidate_path.exists() and not candidate_is_symlink: - continue - - target_path: Path | None = None - file_size = 0 - if candidate_is_symlink: - try: - target_path = candidate_path.resolve(strict=False) - if target_path.exists(): - file_size = max(0, target_path.stat().st_size) - except OSError: - target_path = None - file_size = 0 - else: - try: - file_size = max(0, candidate_path.stat().st_size) - except OSError: - file_size = 0 - try: - candidate_path.unlink(missing_ok=True) - if target_path is not None and target_path.exists(): - target_path.unlink(missing_ok=True) - deleted_file = True - freed_bytes += file_size - except OSError: - logger.warning("Could not delete model file: %s", candidate_path, exc_info=True) - return False, "delete_failed", False, 0, was_active - - remaining_models = [ - item - for item in models - if not (isinstance(item, dict) and str(item.get("id") or "") == model_id) - ] - state["models"] = remaining_models - if str(state.get("current_download_model_id") or "") == model_id: - state["current_download_model_id"] = None - if was_active: - next_active_id: str | None = None - for item in remaining_models: - if not isinstance(item, dict): - continue - candidate_id = str(item.get("id") or "") - candidate_name = str(item.get("filename") or "") - if candidate_id and model_file_present(store.models_dir, candidate_name): - next_active_id = candidate_id - break - if next_active_id is None: - for item in remaining_models: - if isinstance(item, dict) and item.get("id"): - next_active_id = str(item["id"]) - break - if next_active_id is None: - next_active_id = str(state.get("default_model_id") or "default") - state["active_model_id"] = next_active_id - save_models_state(store, state) - return True, "deleted", deleted_file, freed_bytes, was_active - - -def any_model_ready(store: ModelStoreConfig) -> bool: - """Return True if any model in the models state has a file on disk.""" - state = ensure_models_state(store) - models = state.get("models") or [] - for model in models: - filename = str(model.get("filename") or "").strip() - if filename and model_file_present(store.models_dir, filename): - return True - return False - - -# --------------------------------------------------------------------------- -# Projector download -# --------------------------------------------------------------------------- - - -def download_default_projector_for_model(store: ModelStoreConfig, model_id: str) -> tuple[bool, str, str | None]: - """Download the default vision projector for a model from HuggingFace.""" - import httpx - - from .model_families import default_projector_candidates_for_model, projector_repo_for_model - - state = ensure_models_state(store) - model = get_model_by_id(state, model_id) - if model is None: - return False, "model_not_found", None - filename = str(model.get("filename") or "") - if not model_supports_vision_filename(filename): - return False, "vision_not_supported", None - repo = projector_repo_for_model(filename, source_url=model.get("source_url")) - candidates = default_projector_candidates_for_model(filename) - if not repo or not candidates: - return False, "projector_repo_unknown", None - - models_dir = store.models_dir - models_dir.mkdir(parents=True, exist_ok=True) - - _generics = {"mmproj-F16.gguf", "mmproj-bf16.gguf"} - preferred_local: str | None = None - for c in candidates: - if c == "mmproj-F16.gguf": - break - preferred_local = c - preferred_local_bf16: str | None = None - if preferred_local: - preferred_local_bf16 = preferred_local.replace("-f16.gguf", "-bf16.gguf") - - for candidate in candidates: - if candidate in _generics and preferred_local: - continue - target_path = models_dir / candidate - if target_path.exists(): - return True, "downloaded", candidate - - download_targets = list(candidates) - if "mmproj-bf16.gguf" not in download_targets: - download_targets.append("mmproj-bf16.gguf") - - client = httpx.Client(follow_redirects=True, timeout=120.0) - try: - for candidate in download_targets: - url = f"https://huggingface.co/{repo}/resolve/main/{candidate}" - if candidate == "mmproj-F16.gguf" and preferred_local: - local_name = preferred_local - elif candidate == "mmproj-bf16.gguf" and preferred_local_bf16: - local_name = preferred_local_bf16 - else: - local_name = candidate - target_path = models_dir / local_name - if target_path.exists(): - return True, "downloaded", local_name - part_path = target_path.with_suffix(target_path.suffix + ".part") - try: - with client.stream("GET", url) as response: - response.raise_for_status() - with part_path.open("wb") as handle: - for chunk in response.iter_bytes(): - if chunk: - handle.write(chunk) - part_path.replace(target_path) - if local_name not in _generics: - for g in _generics: - (models_dir / g).unlink(missing_ok=True) - return True, "downloaded", local_name - except Exception: - part_path.unlink(missing_ok=True) - continue - finally: - client.close() - return False, "download_failed", None diff --git a/core/inferno/orchestrator.py b/core/inferno/orchestrator.py deleted file mode 100644 index 432996a..0000000 --- a/core/inferno/orchestrator.py +++ /dev/null @@ -1,522 +0,0 @@ -"""Inferno orchestrator — inference process lifecycle, health, and readiness. - -This module owns the runtime-facing orchestration logic that was previously -embedded in core.main: health probing, readiness state machine, process -restart coordination, and the per-tick inference loop decision logic. - -Functions accept primitives, dicts, and callbacks — never FastAPI app.state. -The Potato layer (core.main) maps between app.state and these interfaces. -""" - -from __future__ import annotations - -import logging -import time -from pathlib import Path -from typing import Any, Awaitable, Callable, NamedTuple - -import httpx - -from .model_families import ( - build_model_projector_status, - is_gemma4_filename, - recommended_runtime_for_model, -) -from .model_registry import ( - is_qwen35_a3b_filename, - model_supports_vision_filename, - normalize_model_settings, - resolve_model_runtime_path, -) -from .runtime_manager import ( - LLAMA_SERVER_RUNTIME_FAMILIES, - check_runtime_device_compatibility, - discover_runtime_slots, -) - -logger = logging.getLogger("potato") - -# --------------------------------------------------------------------------- -# Constants -# --------------------------------------------------------------------------- - -READY_HEALTH_POLLS_REQUIRED: int = 2 -"""Consecutive healthy probes before marking the inference server ready.""" - -MAX_CONSECUTIVE_FAILURES: int = 5 -"""Failure ceiling — stop restarting the process after this many crashes.""" - -# --------------------------------------------------------------------------- -# Tick result -# --------------------------------------------------------------------------- - - -class InferenceTickResult(NamedTuple): - process: Any - consecutive_failures: int - failure_model_key: str | None - failure_runtime_key: str | None - readiness: dict[str, Any] - - -# --------------------------------------------------------------------------- -# State factories -# --------------------------------------------------------------------------- - - -def empty_readiness_state() -> dict[str, Any]: - return { - "generation": 0, - "model_path": None, - "status": "idle", - "transport_healthy": False, - "ready": False, - "healthy_polls": 0, - "last_error": None, - "last_ready_at_unix": None, - } - - -def empty_runtime_switch_state() -> dict[str, Any]: - return { - "active": False, - "target_bundle_path": None, - "started_at_unix": None, - "completed_at_unix": None, - "error": None, - "last_bundle_path": None, - } - - -# --------------------------------------------------------------------------- -# Readiness state transitions (pure) -# --------------------------------------------------------------------------- - - -def reset_readiness( - previous: dict[str, Any] | None, - *, - model_path: str | None = None, - reason: str | None = None, -) -> dict[str, Any]: - """Create a fresh readiness state, incrementing the generation counter.""" - generation = 1 - if isinstance(previous, dict): - generation = max(0, int(previous.get("generation") or 0)) + 1 - state = empty_readiness_state() - state["generation"] = generation - state["model_path"] = str(model_path) if model_path else None - state["status"] = "loading" if model_path else "idle" - state["last_error"] = reason - return state - - -def resolve_readiness( - current: dict[str, Any] | None, - *, - active_model_path: str | None = None, -) -> dict[str, Any]: - """Return existing readiness state, auto-resetting on model change.""" - if not isinstance(current, dict): - current = empty_readiness_state() - target_path = str(active_model_path) if active_model_path is not None else None - if target_path and current.get("model_path") != target_path: - return reset_readiness(current, model_path=active_model_path, reason="model_changed") - if target_path is None and current.get("model_path") is not None: - return reset_readiness(current, reason="no_model") - return dict(current) - - -# --------------------------------------------------------------------------- -# Health probing (async, httpx) -# --------------------------------------------------------------------------- - - -async def check_health(base_url: str, *, busy_is_healthy: bool = True) -> bool: - """Probe the inference server's health endpoints.""" - timeout = httpx.Timeout(2.0, connect=1.0) - async with httpx.AsyncClient(timeout=timeout) as client: - for path in ("/health", "/v1/models"): - try: - response = await client.get(f"{base_url}{path}") - except httpx.ReadTimeout: - if busy_is_healthy: - return True - continue - except httpx.HTTPError: - continue - if response.status_code < 500: - return True - return False - - -async def probe_inference_slot(base_url: str) -> bool: - """Send a minimal inference request to verify the slot is functional.""" - timeout = httpx.Timeout(4.0, connect=1.0) - payload = { - "model": "qwen-local", - "stream": False, - "max_tokens": 1, - "messages": [{"role": "user", "content": "ping"}], - } - try: - async with httpx.AsyncClient(timeout=timeout) as client: - response = await client.post( - f"{base_url}/v1/chat/completions", - json=payload, - ) - return response.status_code < 500 - except httpx.HTTPError: - return False - - -# --------------------------------------------------------------------------- -# Readiness refresh (async) -# --------------------------------------------------------------------------- - - -async def refresh_readiness( - readiness: dict[str, Any], - *, - base_url: str, - process_alive: bool, -) -> dict[str, Any]: - """Advance the readiness state machine by one health-check cycle. - - Returns a new state dict (the input is not mutated). - """ - state = dict(readiness) - target_path = state.get("model_path") - - if target_path is None: - return state - - if not process_alive: - state.update( - { - "status": "loading", - "transport_healthy": False, - "ready": False, - "healthy_polls": 0, - } - ) - return state - - busy_is_healthy = bool(state.get("ready")) - transport_healthy = await check_health(base_url, busy_is_healthy=busy_is_healthy) - state["transport_healthy"] = transport_healthy - if not transport_healthy: - state.update( - { - "status": "loading", - "ready": False, - "healthy_polls": 0, - } - ) - return state - - state["healthy_polls"] = min( - READY_HEALTH_POLLS_REQUIRED, - max(0, int(state.get("healthy_polls") or 0)) + 1, - ) - if int(state["healthy_polls"]) >= READY_HEALTH_POLLS_REQUIRED: - if not state.get("ready"): - state["last_ready_at_unix"] = time.time() - state["ready"] = True - state["status"] = "ready" - state["last_error"] = None - else: - state["ready"] = False - state["status"] = "warming" - return state - - -# --------------------------------------------------------------------------- -# mmproj resolution -# --------------------------------------------------------------------------- - - -def resolve_mmproj_for_launch( - models_dir: Path, - resolved_model_dir: Path, - active_model: dict[str, Any], - installed_family: str, -) -> str | None: - """Resolve the mmproj path for a vision-enabled model. - - Returns the projector path, or ``None`` if vision is not enabled. - Raises ``RuntimeError`` if vision is enabled but no projector is available. - """ - active_filename = str(active_model.get("filename") or "") - active_settings = normalize_model_settings(active_model.get("settings"), filename=active_filename) - vision_settings = active_settings.get("vision", {}) - - if not (model_supports_vision_filename(active_filename) and bool(vision_settings.get("enabled", False))): - return None - - # Suppress Gemma4 vision on ik_llama (clip_init failure). - if is_gemma4_filename(active_filename) and installed_family == "ik_llama": - return None - - projector_mode = str(vision_settings.get("projector_mode") or "default").strip().lower() - projector_filename = str(vision_settings.get("projector_filename") or "").strip() - - if projector_mode == "custom" and projector_filename: - custom_path = models_dir / projector_filename - if custom_path.exists(): - return str(custom_path) - raise RuntimeError(f"Custom projector not found: {custom_path}") - - projector_status = build_model_projector_status(models_dir, active_model) - if projector_status.get("present") and projector_status.get("path"): - return str(projector_status["path"]) - - if resolved_model_dir != models_dir: - for candidate in projector_status.get("default_candidates") or []: - candidate_path = resolved_model_dir / candidate - if candidate_path.exists(): - return str(candidate_path) - - raise RuntimeError(f"Vision enabled but no projector found for {active_filename}") - - -async def ensure_mmproj_for_launch( - models_dir: Path, - active_model: dict[str, Any], - installed_family: str, - *, - download_fn: Callable[[str], Awaitable[tuple[bool, str, str | None]]] | None = None, -) -> str | None: - """Resolve or download the mmproj for launch. - - *download_fn* is an async callback ``(model_id) -> (ok, reason, filename)`` - injected by the Potato layer. Returns the path or ``None``. - """ - resolved_dir = models_dir - try: - active_filename = str(active_model.get("filename") or "") - resolved_dir = resolve_model_runtime_path(models_dir, active_filename).parent - except Exception: - pass - - try: - return resolve_mmproj_for_launch(models_dir, resolved_dir, active_model, installed_family) - except RuntimeError: - pass - - if download_fn is None: - logger.warning("Vision projector unavailable and no download function — skipping launch") - return None - - active_model_id = str(active_model.get("id") or "") - if not active_model_id: - logger.warning("Vision enabled but model has no id — skipping projector download") - return None - try: - ok, _reason, downloaded_name = await download_fn(active_model_id) - except Exception: - logger.warning("Projector download failed — skipping vision launch", exc_info=True) - return None - if ok and downloaded_name: - return str(models_dir / downloaded_name) - - logger.warning("Vision projector unavailable — skipping launch (will retry)") - return None - - -# --------------------------------------------------------------------------- -# no-mmap resolution -# --------------------------------------------------------------------------- - - -def resolve_no_mmap( - memory_loading_status: dict[str, Any], - active_filename: str, - installed_family: str, - *, - device_class: str, - bundle_marker: dict[str, Any] | None, -) -> bool: - """Resolve the ``--no-mmap`` flag, including the 'auto' heuristic.""" - no_mmap_env = str(memory_loading_status.get("no_mmap_env") or "auto") - if no_mmap_env.lower() in ("true", "1"): - return True - if no_mmap_env.lower() in ("false", "0"): - return False - - # Auto mode — replicate the old shell heuristic. - if not is_qwen35_a3b_filename(active_filename): - return False - if device_class != "pi5-16gb": - return False - runtime_profile = str((bundle_marker or {}).get("profile") or "") - return runtime_profile == "pi5-opt" and installed_family == "ik_llama" - - -# --------------------------------------------------------------------------- -# Process restart -# --------------------------------------------------------------------------- - - -async def restart_inference_process( - readiness: dict[str, Any], - process: Any, - *, - model_path: str | None = None, - terminate_fn: Callable[..., Awaitable[None]], - stray_kill_fn: Callable[[], Awaitable[int]], -) -> tuple[dict[str, Any], bool, str]: - """Restart the inference process. - - Returns ``(new_readiness, terminated_any, reason)``. - """ - new_readiness = reset_readiness(readiness, model_path=model_path, reason="restart_requested") - terminated_running = False - terminated_stale = False - - if process is not None and process.returncode is None: - await terminate_fn(process, timeout=3.0) - terminated_running = True - terminated_stale = bool(await stray_kill_fn()) - - if terminated_running and terminated_stale: - return new_readiness, True, "terminated_running_and_stale_processes" - if terminated_running: - return new_readiness, True, "terminated_running_process" - if terminated_stale: - return new_readiness, True, "terminated_stale_processes" - return new_readiness, False, "no_running_process" - - -# --------------------------------------------------------------------------- -# Inference tick -# --------------------------------------------------------------------------- - - -async def run_inference_tick( - process: Any, - consecutive_failures: int, - failure_model_key: str | None, - failure_runtime_key: str | None, - readiness: dict[str, Any], - *, - model_path: Path, - base_url: str, - installed_family: str, - launch_llama_fn: Callable[[], Awaitable[Any | None]], - launch_litert_fn: Callable[[], Awaitable[Any]] | None, - switch_in_progress: bool = False, -) -> InferenceTickResult: - """Run one iteration of the inference process management loop. - - The tick owns the decision logic (should we spawn? count failures? - check readiness?) while the actual process spawning is delegated to - the caller-provided launch callbacks. - """ - active_model_is_present = False - try: - active_model_is_present = model_path.exists() and model_path.stat().st_size > 0 - except OSError: - active_model_is_present = False - - if not active_model_is_present: - new_readiness = reset_readiness(readiness, reason="model_missing") - return InferenceTickResult( - process=process, - consecutive_failures=0, - failure_model_key=failure_model_key, - failure_runtime_key=failure_runtime_key, - readiness=new_readiness, - ) - - # Reset failure counter when active model or runtime changes. - current_model_key = str(model_path) - current_runtime_key = installed_family - if failure_model_key != current_model_key or failure_runtime_key != current_runtime_key: - consecutive_failures = 0 - failure_model_key = current_model_key - failure_runtime_key = current_runtime_key - - llama_process = process - if llama_process is None or llama_process.returncode is not None: - # Count the previous process's failure BEFORE starting a new one. - if llama_process is not None and llama_process.returncode is not None and llama_process.returncode != 0: - consecutive_failures += 1 - llama_process = None - if consecutive_failures == MAX_CONSECUTIVE_FAILURES: - logger.error( - "llama-server failed %d times in a row — stopping restart attempts (model may be corrupt)", - consecutive_failures, - ) - - if consecutive_failures >= MAX_CONSECUTIVE_FAILURES or switch_in_progress: - pass # Limit reached or switch in progress — don't restart. - elif installed_family == "litert" and launch_litert_fn is not None: - llama_process = await launch_litert_fn() - if llama_process is not None: - logger.info("Started litert adapter process") - else: - llama_process = await launch_llama_fn() - if llama_process is not None: - logger.info("Started llama-server process") - - process_alive = llama_process is not None and llama_process.returncode is None - new_readiness = await refresh_readiness( - resolve_readiness(readiness, active_model_path=str(model_path)), - base_url=base_url, - process_alive=process_alive, - ) - if new_readiness.get("ready"): - consecutive_failures = 0 - - return InferenceTickResult( - process=llama_process, - consecutive_failures=consecutive_failures, - failure_model_key=failure_model_key, - failure_runtime_key=failure_runtime_key, - readiness=new_readiness, - ) - - -# --------------------------------------------------------------------------- -# Activation runtime prep -# --------------------------------------------------------------------------- - - -def prepare_activation_runtime( - model_filename: str, - model_format: str, - current_family: str, - device_class: str, - runtimes_dir: Path, -) -> tuple[bool, str, str | None]: - """Decide if a runtime switch is needed for model activation. - - Returns ``(should_switch, reason, target_family)``. - """ - preferred = recommended_runtime_for_model(model_filename) - # GGUF models can't run on litert — fall back to llama_cpp. - if not preferred and current_family == "litert" and model_format == "gguf": - preferred = "llama_cpp" - - if not preferred: - return False, "no_switch_needed", None - - compat = check_runtime_device_compatibility(device_class, preferred) - if not compat.get("compatible", True): - fmt = model_format - if (fmt == "litertlm" and preferred == "litert") or (fmt == "gguf" and preferred in LLAMA_SERVER_RUNTIME_FAMILIES): - return False, "incompatible_runtime", None - # Preferred runtime is incompatible but format doesn't strictly require it. - return False, "no_switch_needed", None - - if current_family == preferred: - return False, "already_on_preferred", None - - # Check if we have a slot for the preferred family. - slots = discover_runtime_slots(runtimes_dir) - for slot in slots: - if slot.get("family") == preferred: - return True, "switch_required", preferred - - return False, "no_slot_available", preferred diff --git a/core/inferno/runtime_manager.py b/core/inferno/runtime_manager.py deleted file mode 100644 index 165d1ac..0000000 --- a/core/inferno/runtime_manager.py +++ /dev/null @@ -1,755 +0,0 @@ -"""Runtime management for Inferno — slot discovery, compatibility, switching, settings. - -This module owns the inference-facing runtime lifecycle: device -classification, family compatibility, slot discovery, runtime -installation/switching, and settings normalization. - -Product-specific hardware probing (psutil, /proc reads, vcgencmd) and -system metrics stay in core.runtime_state. This module receives device -classification results via RuntimeStoreConfig or explicit parameters. -""" - -from __future__ import annotations - -import asyncio -import dataclasses -import json -import os -import shutil -import time -from pathlib import Path -from typing import Any - -logger = __import__("logging").getLogger("potato") - -# --------------------------------------------------------------------------- -# Constants -# --------------------------------------------------------------------------- - -SUPPORTED_RUNTIME_FAMILIES = ("ik_llama", "llama_cpp", "litert") - -# Families that use an external server binary (bin/llama-server). -# LiteRT uses a Python adapter instead — no binary needed in the slot. -LLAMA_SERVER_RUNTIME_FAMILIES = ("ik_llama", "llama_cpp") - -PI4_8GB_MEMORY_THRESHOLD_BYTES = 6 * 1024 * 1024 * 1024 -PI4_INCOMPATIBLE_RUNTIMES = ("ik_llama", "litert") - -DEVICE_CLOCK_LIMITS: dict[str, dict[str, int]] = { - "pi5": {"cpu_max_hz": 2_400_000_000, "gpu_max_hz": 1_000_000_000}, - "pi4": {"cpu_max_hz": 1_800_000_000, "gpu_max_hz": 500_000_000}, -} - -LLAMA_RUNTIME_BUNDLE_MARKER_FILENAME = ".potato-llama-runtime-bundle.json" - -MODEL_LOADING_INACTIVE: dict[str, Any] = { - "active": False, - "progress_percent": None, - "resident_bytes": None, - "model_size_bytes": None, -} - -# Memory threshold for 16GB vs 8GB Pi classification. -MODEL_UPLOAD_PI_16GB_MEMORY_THRESHOLD_BYTES = 12 * 1024 * 1024 * 1024 - - -# --------------------------------------------------------------------------- -# Store configuration -# --------------------------------------------------------------------------- - - -@dataclasses.dataclass(frozen=True) -class RuntimeStoreConfig: - """Filesystem and device context for runtime management operations. - - This bundles the paths and device info that the Potato layer injects - so that Inferno never needs to import RuntimeConfig. - """ - - runtimes_dir: Path - install_dir: Path - settings_path: Path - device_class: str - total_memory_bytes: int - - -# --------------------------------------------------------------------------- -# Device classification -# --------------------------------------------------------------------------- - - -def classify_runtime_device( - *, - pi_model_name: str, - total_memory_bytes: int, -) -> str: - """Classify device hardware. Both params are required — no hardware probing.""" - model_name = (pi_model_name or "").strip().lower() - if not model_name: - return "unknown" - if "raspberry pi" not in model_name: - return "unknown" - if "raspberry pi 5" in model_name: - if total_memory_bytes >= MODEL_UPLOAD_PI_16GB_MEMORY_THRESHOLD_BYTES: - return "pi5-16gb" - return "pi5-8gb" - if "raspberry pi 4" in model_name: - if total_memory_bytes >= MODEL_UPLOAD_PI_16GB_MEMORY_THRESHOLD_BYTES: - return "pi4-16gb" - if total_memory_bytes >= PI4_8GB_MEMORY_THRESHOLD_BYTES: - return "pi4-8gb" - return "pi4-4gb" - return "other-pi" - - -# --------------------------------------------------------------------------- -# Compatibility checks -# --------------------------------------------------------------------------- - - -def check_runtime_device_compatibility( - device_class: str, - runtime_family: str, -) -> dict[str, Any]: - if device_class.startswith("pi4-") and runtime_family in PI4_INCOMPATIBLE_RUNTIMES: - return { - "compatible": False, - "reason": ( - f"{runtime_family} requires ARMv8.2-A dot product instructions (Cortex-A76+). " - f"Pi 4 (Cortex-A72, ARMv8.0-A) must use llama_cpp." - ), - "recommended_family": "llama_cpp", - } - return {"compatible": True, "reason": None, "recommended_family": None} - - -def get_device_clock_limits(device_class: str) -> dict[str, int]: - for prefix, limits in DEVICE_CLOCK_LIMITS.items(): - if device_class.startswith(prefix): - return dict(limits) - return dict(DEVICE_CLOCK_LIMITS["pi5"]) - - -# --------------------------------------------------------------------------- -# Settings normalization (pure) -# --------------------------------------------------------------------------- - - -def normalize_llama_memory_loading_mode(raw_mode: Any) -> str: - value = str(raw_mode or "").strip().lower() - if value in {"full_ram", "no_mmap", "no-mmap", "1", "true", "on"}: - return "full_ram" - if value in {"mmap", "mapped", "0", "false", "off"}: - return "mmap" - return "auto" - - -def llama_memory_loading_no_mmap_env(mode: str) -> str: - normalized = normalize_llama_memory_loading_mode(mode) - if normalized == "full_ram": - return "1" - if normalized == "mmap": - return "0" - return "auto" - - -def normalize_allow_unsupported_large_models(raw_value: Any) -> bool: - if isinstance(raw_value, bool): - return raw_value - if raw_value is None: - return False - value = str(raw_value).strip().lower() - return value in {"1", "true", "yes", "on"} - - -# --------------------------------------------------------------------------- -# Model loading progress (pure) -# --------------------------------------------------------------------------- - - -def compute_model_loading_progress( - *, - state: str, - has_model: bool, - model_size_bytes: int, - no_mmap_env: str, - llama_rss: dict[str, Any], -) -> dict[str, Any]: - if state != "BOOTING" or not has_model or model_size_bytes <= 0: - return dict(MODEL_LOADING_INACTIVE) - if not llama_rss.get("available"): - return dict(MODEL_LOADING_INACTIVE) - if no_mmap_env == "1": - resident_bytes = llama_rss.get("rss_anon_bytes") - elif no_mmap_env == "auto": - anon = llama_rss.get("rss_anon_bytes") or 0 - file = llama_rss.get("rss_file_bytes") or 0 - resident_bytes = max(anon, file) if (anon or file) else None - else: - resident_bytes = llama_rss.get("rss_file_bytes") - if resident_bytes is None or not isinstance(resident_bytes, (int, float)): - return dict(MODEL_LOADING_INACTIVE) - resident_bytes = int(resident_bytes) - progress_percent = min(100, max(0, int(resident_bytes * 100 / model_size_bytes))) - return { - "active": True, - "progress_percent": progress_percent, - "resident_bytes": resident_bytes, - "model_size_bytes": model_size_bytes, - } - - -# --------------------------------------------------------------------------- -# Atomic write utility -# --------------------------------------------------------------------------- - - -def _atomic_write_json(path: Path, payload: dict[str, Any]) -> None: - try: - path.parent.mkdir(parents=True, exist_ok=True) - import tempfile - - fd, tmp_name = tempfile.mkstemp(dir=path.parent, suffix=".tmp") - with os.fdopen(fd, "w", encoding="utf-8") as f: - f.write(json.dumps(payload)) - os.replace(tmp_name, path) - except OSError: - logger.warning("Could not persist JSON state to %s", path, exc_info=True) - - -# --------------------------------------------------------------------------- -# Slot discovery -# --------------------------------------------------------------------------- - - -def discover_runtime_slots(runtimes_dir: Path) -> list[dict[str, Any]]: - """Discover installed runtime slots across all supported families.""" - slots: list[dict[str, Any]] = [] - for family in SUPPORTED_RUNTIME_FAMILIES: - slot_dir = runtimes_dir / family - if not slot_dir.is_dir(): - continue - if family in LLAMA_SERVER_RUNTIME_FAMILIES: - if not (slot_dir / "bin" / "llama-server").exists(): - continue - else: - if not (slot_dir / "runtime.json").exists(): - continue - metadata: dict[str, Any] = {"family": family, "path": str(slot_dir)} - runtime_json = slot_dir / "runtime.json" - if runtime_json.exists(): - try: - meta = json.loads(runtime_json.read_text(encoding="utf-8")) - if isinstance(meta, dict): - metadata.update(meta) - except (OSError, json.JSONDecodeError): - pass - metadata.setdefault("commit", "unknown") - metadata.setdefault("profile", "unknown") - metadata.setdefault("repo", "") - metadata.setdefault("build_timestamp", "") - metadata.setdefault("version", "") - slots.append(metadata) - return slots - - -def find_runtime_slot_by_family(runtimes_dir: Path, family: str) -> dict[str, Any] | None: - """Find a runtime slot by family name.""" - for slot in discover_runtime_slots(runtimes_dir): - if slot.get("family") == family: - return slot - return None - - -# --------------------------------------------------------------------------- -# Marker management -# --------------------------------------------------------------------------- - - -def read_llama_runtime_bundle_marker(install_dir: Path) -> dict[str, Any] | None: - marker_path = install_dir / LLAMA_RUNTIME_BUNDLE_MARKER_FILENAME - try: - raw = json.loads(marker_path.read_text(encoding="utf-8")) - except (OSError, json.JSONDecodeError): - return None - return raw if isinstance(raw, dict) else None - - -def write_llama_runtime_bundle_marker(install_dir: Path, bundle: dict[str, Any]) -> dict[str, Any]: - payload = { - "family": str(bundle.get("family") or ""), - "source_bundle_path": str(bundle.get("path") or ""), - "source_bundle_name": str(bundle.get("name") or bundle.get("family") or ""), - "profile": str(bundle.get("profile") or "unknown"), - "version_summary": bundle.get("version_summary") or bundle.get("version"), - "llama_cpp_commit": bundle.get("llama_cpp_commit") or bundle.get("commit"), - "switched_at_unix": int(time.time()), - } - _atomic_write_json(install_dir / LLAMA_RUNTIME_BUNDLE_MARKER_FILENAME, payload) - return payload - - -def _detect_installed_runtime_family(install_dir: Path) -> str: - """Detect the active runtime family from marker or installed runtime.json.""" - marker = read_llama_runtime_bundle_marker(install_dir) - if isinstance(marker, dict) and marker.get("family"): - return str(marker["family"]) - runtime_json = install_dir / "runtime.json" - if runtime_json.exists(): - try: - meta = json.loads(runtime_json.read_text(encoding="utf-8")) - if isinstance(meta, dict) and meta.get("family"): - return str(meta["family"]) - except (OSError, json.JSONDecodeError): - pass - return "" - - -def _read_installed_runtime_metadata(install_dir: Path) -> dict[str, Any]: - """Read runtime metadata from marker first, then fallback to runtime.json.""" - marker = read_llama_runtime_bundle_marker(install_dir) - if isinstance(marker, dict) and marker.get("family"): - return marker - runtime_json = install_dir / "runtime.json" - if runtime_json.exists(): - try: - meta = json.loads(runtime_json.read_text(encoding="utf-8")) - if isinstance(meta, dict): - return meta - except (OSError, json.JSONDecodeError): - pass - return {} - - -# --------------------------------------------------------------------------- -# Bundle discovery (legacy bundle search) -# --------------------------------------------------------------------------- - - -def _llama_runtime_bundle_profile_from_name(bundle_name: str) -> str | None: - lowered = bundle_name.lower() - if lowered.endswith("_pi5-opt"): - return "pi5-opt" - if lowered.endswith("_baseline"): - return "baseline" - return None - - -def _llama_runtime_bundle_readme_fields(bundle_dir: Path) -> dict[str, str]: - readme = bundle_dir / "README.txt" - try: - text = readme.read_text(encoding="utf-8", errors="replace") - except OSError: - return {} - - fields: dict[str, str] = {} - version_lines: list[str] = [] - in_version = False - for raw_line in text.splitlines(): - line = raw_line.strip() - if not line: - if in_version and version_lines: - break - continue - if line.lower().startswith("profile:"): - fields["profile"] = line.split(":", 1)[1].strip() - continue - if line.lower().startswith("llama.cpp commit:"): - fields["llama_cpp_commit"] = line.split(":", 1)[1].strip() - continue - if line.lower() == "version:": - in_version = True - continue - if in_version and not line.lower().startswith("contents:"): - version_lines.append(line) - continue - if in_version and line.lower().startswith("contents:"): - break - if version_lines: - fields["version_summary"] = version_lines[0] - return fields - - -def _default_llama_runtime_bundle_roots(base_dir: Path) -> list[Path]: - return [ - base_dir / "llama-bundles", - Path("/tmp/potato-qwen35-ab/references/old_reference_design/llama_cpp_binary"), - Path("/tmp/potato-os/references/old_reference_design/llama_cpp_binary"), - ] - - -def get_llama_runtime_bundle_roots(base_dir: Path) -> list[Path]: - raw = os.getenv("POTATO_LLAMA_RUNTIME_BUNDLE_ROOTS", "").strip() - candidates: list[Path] - if raw: - candidates = [Path(part).expanduser() for part in raw.split(os.pathsep) if part.strip()] - else: - candidates = _default_llama_runtime_bundle_roots(base_dir) - - roots: list[Path] = [] - seen: set[str] = set() - for candidate in candidates: - key = str(candidate) - if key in seen: - continue - seen.add(key) - roots.append(candidate) - return roots - - -def discover_llama_runtime_bundles(bundle_roots: list[Path]) -> list[dict[str, Any]]: - bundles: list[dict[str, Any]] = [] - for root in bundle_roots: - try: - if not root.exists() or not root.is_dir(): - continue - except OSError: - continue - try: - children = list(root.iterdir()) - except OSError: - continue - for bundle_dir in children: - name = bundle_dir.name - if not bundle_dir.is_dir() or not name.startswith("llama_server_bundle_"): - continue - server_path = bundle_dir / "bin" / "llama-server" - if not server_path.exists(): - continue - readme_fields = _llama_runtime_bundle_readme_fields(bundle_dir) - profile = ( - str(readme_fields.get("profile") or "").strip() - or _llama_runtime_bundle_profile_from_name(name) - or "unknown" - ) - try: - mtime_unix = int(bundle_dir.stat().st_mtime) - except OSError: - mtime_unix = 0 - bundles.append( - { - "path": str(bundle_dir), - "name": name, - "root": str(root), - "profile": profile, - "is_pi5_optimized": profile == "pi5-opt", - "has_bench": (bundle_dir / "bin" / "llama-bench").exists(), - "has_lib_dir": (bundle_dir / "lib").is_dir(), - "version_summary": readme_fields.get("version_summary"), - "llama_cpp_commit": readme_fields.get("llama_cpp_commit"), - "mtime_unix": mtime_unix, - } - ) - bundles.sort(key=lambda item: (int(item.get("mtime_unix") or 0), str(item.get("name") or "")), reverse=True) - return bundles - - -def _safe_int(value: Any, default: int = 0) -> int: - try: - if value is None: - return default - return int(value) - except (TypeError, ValueError): - return default - - -def find_llama_runtime_bundle_by_path(bundle_roots: list[Path], bundle_path: str) -> dict[str, Any] | None: - candidate = str(bundle_path or "").strip() - if not candidate: - return None - try: - resolved = str(Path(candidate).resolve()) - except OSError: - return None - for bundle in discover_llama_runtime_bundles(bundle_roots): - try: - bundle_resolved = str(Path(str(bundle.get("path") or "")).resolve()) - except OSError: - continue - if bundle_resolved == resolved: - return bundle - return None - - -# --------------------------------------------------------------------------- -# Settings I/O -# --------------------------------------------------------------------------- - - -def read_llama_runtime_settings(settings_path: Path) -> dict[str, Any]: - """Read runtime settings. Normalizes inferno-owned fields; power_calibration passes through.""" - try: - raw = json.loads(settings_path.read_text(encoding="utf-8")) - except (OSError, json.JSONDecodeError): - raw = {} - if not isinstance(raw, dict): - raw = {} - return { - "memory_loading_mode": normalize_llama_memory_loading_mode(raw.get("memory_loading_mode")), - "allow_unsupported_large_models": normalize_allow_unsupported_large_models( - raw.get("allow_unsupported_large_models") - ), - "power_calibration": raw.get("power_calibration") or {}, - "updated_at_unix": _safe_int(raw.get("updated_at_unix"), 0) or None, - } - - -def write_llama_runtime_settings( - settings_path: Path, - *, - memory_loading_mode: str | None = None, - allow_unsupported_large_models: bool | None = None, - power_calibration: dict[str, Any] | None = None, -) -> dict[str, Any]: - current = read_llama_runtime_settings(settings_path) - payload = { - "memory_loading_mode": normalize_llama_memory_loading_mode( - current.get("memory_loading_mode") if memory_loading_mode is None else memory_loading_mode - ), - "allow_unsupported_large_models": normalize_allow_unsupported_large_models( - current.get("allow_unsupported_large_models") - if allow_unsupported_large_models is None - else allow_unsupported_large_models - ), - "power_calibration": power_calibration if power_calibration is not None else current.get("power_calibration", {}), - "updated_at_unix": int(time.time()), - } - _atomic_write_json(settings_path, payload) - return payload - - -# --------------------------------------------------------------------------- -# Status builders -# --------------------------------------------------------------------------- - - -def build_llama_memory_loading_status(settings_path: Path) -> dict[str, Any]: - settings = read_llama_runtime_settings(settings_path) - mode = normalize_llama_memory_loading_mode(settings.get("memory_loading_mode")) - no_mmap_env = llama_memory_loading_no_mmap_env(mode) - return { - "mode": mode, - "no_mmap_env": no_mmap_env, - "label": ( - "Full RAM load (--no-mmap)" - if mode == "full_ram" - else "Memory-mapped (mmap)" - if mode == "mmap" - else "Automatic (profile-based)" - ), - "updated_at_unix": settings.get("updated_at_unix"), - } - - -def build_llama_large_model_override_status(settings_path: Path) -> dict[str, Any]: - settings = read_llama_runtime_settings(settings_path) - enabled = normalize_allow_unsupported_large_models(settings.get("allow_unsupported_large_models")) - return { - "enabled": enabled, - "label": "Try unsupported large model anyway" if enabled else "Use compatibility warnings (default)", - "updated_at_unix": settings.get("updated_at_unix"), - } - - -def build_large_model_compatibility( - store: RuntimeStoreConfig, - *, - model_filename: str = "", - model_size_bytes: int = 0, - allow_override: bool | None = None, - threshold_bytes: int = 0, - storage_free_bytes: int = 0, - pi_model_name: str = "", -) -> dict[str, Any]: - """Build large model compatibility status. - - Potato provides pre-computed values (threshold, storage, pi_model_name) - so inferno doesn't need to probe hardware or read env vars. - """ - override_enabled = ( - normalize_allow_unsupported_large_models(allow_override) - if allow_override is not None - else normalize_allow_unsupported_large_models( - read_llama_runtime_settings(store.settings_path).get("allow_unsupported_large_models") - ) - ) - size_bytes = max(0, model_size_bytes) - - warnings: list[dict[str, Any]] = [] - if size_bytes > threshold_bytes > 0 and store.device_class != "pi5-16gb" and not override_enabled: - filename = model_filename or "model.gguf" - warnings.append( - { - "code": "large_model_unsupported_pi_warning", - "severity": "warning", - "message": ( - f"{filename} is larger than the unsupported-device warning threshold " - f"({threshold_bytes} bytes). Qwen3.5-35B-A3B is validated on Raspberry Pi 5 16GB only." - ), - "model_filename": filename, - "model_size_bytes": size_bytes, - } - ) - - runtime_family = _detect_installed_runtime_family(store.install_dir) - runtime_compat = check_runtime_device_compatibility(store.device_class, runtime_family) - - return { - "device_class": store.device_class, - "pi_model_name": pi_model_name, - "memory_total_bytes": store.total_memory_bytes, - "large_model_warn_threshold_bytes": threshold_bytes, - "storage_free_bytes": storage_free_bytes, - "supported_target": "raspberry-pi-5-16gb", - "override_enabled": override_enabled, - "runtime_compatibility": runtime_compat, - "warnings": warnings, - } - - -def build_llama_runtime_status( - store: RuntimeStoreConfig, - *, - active_model_filename: str = "", - switch_snapshot: dict[str, Any] | None = None, -) -> dict[str, Any]: - """Build runtime status. switch_snapshot is pre-extracted from app.state by caller.""" - install_dir = store.install_dir - metadata = _read_installed_runtime_metadata(install_dir) - available_runtimes = discover_runtime_slots(store.runtimes_dir) - - current_family = str(metadata.get("family") or metadata.get("source_bundle_name") or "").strip() - active_is_gguf = active_model_filename.lower().endswith(".gguf") if active_model_filename else True - active_is_litertlm = active_model_filename.lower().endswith(".litertlm") if active_model_filename else False - for slot in available_runtimes: - slot["is_active"] = slot.get("family") == current_family - compat = check_runtime_device_compatibility( - store.device_class, slot.get("family", "") - ) - family = slot.get("family", "") - if family == "litert" and active_is_gguf: - slot["compatible"] = False - elif family in LLAMA_SERVER_RUNTIME_FAMILIES and active_is_litertlm: - slot["compatible"] = False - else: - slot["compatible"] = compat["compatible"] - - snap = switch_snapshot or {} - switch_section = { - "active": bool(snap.get("active", False)), - "target_family": snap.get("target_family"), - "started_at_unix": snap.get("started_at_unix"), - "completed_at_unix": snap.get("completed_at_unix"), - "error": snap.get("error"), - } - - detected_family = str(metadata.get("family") or "") - runtime_type = "litert_adapter" if detected_family == "litert" else "llama_server" - - current = { - "install_dir": str(install_dir), - "exists": install_dir.exists(), - "has_server_binary": (install_dir / "bin" / "llama-server").exists(), - "runtime_type": runtime_type, - "family": metadata.get("family"), - "source_bundle_path": metadata.get("source_bundle_path"), - "source_bundle_name": metadata.get("source_bundle_name"), - "profile": metadata.get("profile"), - "version_summary": metadata.get("version_summary") or metadata.get("version"), - "llama_cpp_commit": metadata.get("llama_cpp_commit") or metadata.get("commit"), - "switched_at_unix": metadata.get("switched_at_unix"), - } - - return { - "current": current, - "available_runtimes": available_runtimes, - "switch": switch_section, - "memory_loading": build_llama_memory_loading_status(store.settings_path), - "large_model_override": build_llama_large_model_override_status(store.settings_path), - } - - -# --------------------------------------------------------------------------- -# Runtime installation (async) -# --------------------------------------------------------------------------- - - -async def install_llama_runtime_bundle(install_dir: Path, bundle_dir: Path) -> dict[str, Any]: - """Install a runtime bundle to the install directory via rsync.""" - install_dir.mkdir(parents=True, exist_ok=True) - - # LiteRT has no binary to rsync — just ensure install dir exists. - bundle_runtime_json = bundle_dir / "runtime.json" - if bundle_runtime_json.exists(): - try: - meta = json.loads(bundle_runtime_json.read_text(encoding="utf-8")) - if isinstance(meta, dict) and meta.get("family") == "litert": - return {"ok": True, "reason": "litert_no_rsync_needed", "install_dir": str(install_dir)} - except (OSError, json.JSONDecodeError): - pass - - rsync = shutil.which("rsync") - if not rsync: - return {"ok": False, "reason": "rsync_not_available", "install_dir": str(install_dir)} - - proc = await asyncio.create_subprocess_exec( - rsync, - "-a", - "--delete", - f"{bundle_dir}/", - f"{install_dir}/", - stdout=asyncio.subprocess.PIPE, - stderr=asyncio.subprocess.STDOUT, - ) - stdout, _stderr = await proc.communicate() - if stdout: - logger.info("llama runtime rsync: %s", stdout.decode("utf-8", errors="replace").rstrip()) - if proc.returncode != 0: - return { - "ok": False, - "reason": "rsync_failed", - "returncode": proc.returncode, - "install_dir": str(install_dir), - } - - for rel in ("bin/llama-server", "run-llama-server.sh", "run-llama-bench.sh"): - path = install_dir / rel - try: - if path.exists(): - path.chmod(path.stat().st_mode | 0o111) - except OSError: - logger.warning("Could not chmod runtime bundle file: %s", path, exc_info=True) - - return {"ok": True, "reason": "installed", "install_dir": str(install_dir)} - - -# --------------------------------------------------------------------------- -# Compatibility enforcement (async) -# --------------------------------------------------------------------------- - - -async def ensure_compatible_runtime(store: RuntimeStoreConfig) -> tuple[bool, str]: - """Auto-switch runtime if current is incompatible with device hardware.""" - current_family = _detect_installed_runtime_family(store.install_dir) - compat = check_runtime_device_compatibility(store.device_class, current_family) - if compat["compatible"]: - return False, "compatible" - recommended = compat.get("recommended_family") or "" - if not recommended: - logger.warning("Device %s incompatible with %s but no recommendation available", store.device_class, current_family) - return False, "no_recommendation" - slot = find_runtime_slot_by_family(store.runtimes_dir, recommended) - if slot is None: - logger.warning("Recommended runtime %s not available as a slot", recommended) - return False, "slot_unavailable" - slot_path = Path(slot["path"]) - logger.info( - "Auto-switching runtime: %s -> %s (device %s incompatible with %s)", - current_family, recommended, store.device_class, current_family, - ) - result = await install_llama_runtime_bundle(store.install_dir, slot_path) - if not isinstance(result, dict) or not result.get("ok", False): - reason = result.get("reason", "unknown") if isinstance(result, dict) else "unknown" - logger.error("Runtime install failed during auto-switch: %s", reason) - return False, "install_failed" - return True, "pi4_incompatible_runtime" diff --git a/core/main.py b/core/main.py index f74a910..fb7ac68 100644 --- a/core/main.py +++ b/core/main.py @@ -17,32 +17,20 @@ from fastapi.responses import HTMLResponse, JSONResponse, Response, StreamingResponse from fastapi.staticfiles import StaticFiles -try: - from core.inferno import ( - BackendProxyError, - ChatRepositoryManager, - FakeLlamaRepository, - LlamaCppRepository, - build_llama_server_args, - ) - from core.inferno import orchestrator as _orchestrator -except ModuleNotFoundError: - from inferno import ( # type: ignore[no-redef] - BackendProxyError, - ChatRepositoryManager, - FakeLlamaRepository, - LlamaCppRepository, - build_llama_server_args, - ) - from inferno import orchestrator as _orchestrator # type: ignore[no-redef] +from inferno import ( + BackendProxyError, + ChatRepositoryManager, + FakeLlamaRepository, + LlamaCppRepository, + build_llama_server_args, + is_gemma4_filename, + is_qwen35_filename, + projector_repo_for_model, + recommended_runtime_for_model, +) +from inferno import orchestrator as _orchestrator try: - from core.inferno import ( - is_gemma4_filename, - is_qwen35_filename, - projector_repo_for_model, - recommended_runtime_for_model, - ) from core.model_state import ( DEFAULT_MODEL_CHAT_SETTINGS, MODEL_FILENAME, @@ -167,12 +155,6 @@ _run_vcgencmd, ) except ModuleNotFoundError: - from inferno import ( # type: ignore[no-redef] - is_gemma4_filename, - is_qwen35_filename, - projector_repo_for_model, - recommended_runtime_for_model, - ) from model_state import ( # type: ignore[no-redef] DEFAULT_MODEL_CHAT_SETTINGS, MODEL_FILENAME, diff --git a/core/model_state.py b/core/model_state.py index 35711d6..d5800c2 100644 --- a/core/model_state.py +++ b/core/model_state.py @@ -1,10 +1,10 @@ -"""Model state -- Potato-specific adapter over core.inferno.model_registry. +"""Model state -- Potato-specific adapter over inferno.model_registry. This module provides the RuntimeConfig-aware interface that the rest of Potato uses. All registry, format, settings, and projector logic lives -in core.inferno.model_registry and core.inferno.model_families; this -file supplies product-level defaults (device detection, default model -selection) and translates RuntimeConfig into ModelStoreConfig for inferno. +in inferno.model_registry and inferno.model_families; this file supplies +product-level defaults (device detection, default model selection) and +translates RuntimeConfig into ModelStoreConfig for inferno. Activation flow (resolve_active_model, model_present) remains here because it mutates RuntimeConfig.model_path -- extraction is planned @@ -25,94 +25,49 @@ # Re-exports from inferno (pure functions, no RuntimeConfig dependency) # --------------------------------------------------------------------------- -try: - from core.inferno.model_registry import ( # noqa: F401 — re-exports - DEFAULT_MODEL_CHAT_SETTINGS, - DEFAULT_MODEL_VISION_SETTINGS, - MODELS_STATE_VERSION, - VALID_MODEL_EXTENSIONS, - ModelSettingsValidationError, - ModelStoreConfig, - _has_valid_model_extension, - _is_discoverable_local_model_filename, - _sanitize_filename, - _slugify_id, - _unique_filename, - _unique_model_id, - apply_model_chat_defaults, - build_model_capabilities, - get_model_by_id, - is_qwen35_a3b_filename, - model_format_for_filename, - model_supports_vision_filename, - normalize_model_settings, - validate_model_url, - ) - from core.inferno.model_families import ( # noqa: F401 — re-exports - _is_vision_family, - default_projector_candidates_for_model, - ) - from core.inferno.model_registry import ( - model_file_path as _inferno_model_file_path, - model_file_present as _inferno_model_file_present, - describe_model_storage as _inferno_describe_model_storage, - resolve_model_runtime_path as _inferno_resolve_model_runtime_path, - discover_local_model_filenames as _inferno_discover_local_model_filenames, - ensure_models_state as _inferno_ensure_models_state, - save_models_state as _inferno_save_models_state, - register_model_url as _inferno_register_model_url, - delete_model as _inferno_delete_model, - update_model_settings as _inferno_update_model_settings, - any_model_ready as _inferno_any_model_ready, - download_default_projector_for_model as _inferno_download_default_projector, - ) - from core.inferno.model_families import ( - build_model_projector_status as _inferno_build_model_projector_status, - ) -except ModuleNotFoundError: - from inferno.model_registry import ( # type: ignore[no-redef] - DEFAULT_MODEL_CHAT_SETTINGS, - DEFAULT_MODEL_VISION_SETTINGS, - MODELS_STATE_VERSION, - VALID_MODEL_EXTENSIONS, - ModelSettingsValidationError, - ModelStoreConfig, - _has_valid_model_extension, - _is_discoverable_local_model_filename, - _sanitize_filename, - _slugify_id, - _unique_filename, - _unique_model_id, - apply_model_chat_defaults, - build_model_capabilities, - get_model_by_id, - is_qwen35_a3b_filename, - model_format_for_filename, - model_supports_vision_filename, - normalize_model_settings, - validate_model_url, - ) - from inferno.model_families import ( # type: ignore[no-redef] - _is_vision_family, - default_projector_candidates_for_model, - ) - from inferno.model_registry import ( # type: ignore[no-redef] - model_file_path as _inferno_model_file_path, - model_file_present as _inferno_model_file_present, - describe_model_storage as _inferno_describe_model_storage, - resolve_model_runtime_path as _inferno_resolve_model_runtime_path, - discover_local_model_filenames as _inferno_discover_local_model_filenames, - ensure_models_state as _inferno_ensure_models_state, - save_models_state as _inferno_save_models_state, - register_model_url as _inferno_register_model_url, - delete_model as _inferno_delete_model, - update_model_settings as _inferno_update_model_settings, - any_model_ready as _inferno_any_model_ready, - download_default_projector_for_model as _inferno_download_default_projector, - ) - from inferno.model_families import ( # type: ignore[no-redef] - build_model_projector_status as _inferno_build_model_projector_status, - ) +from inferno.model_registry import ( # noqa: F401 — re-exports + DEFAULT_MODEL_CHAT_SETTINGS, + DEFAULT_MODEL_VISION_SETTINGS, + MODELS_STATE_VERSION, + VALID_MODEL_EXTENSIONS, + ModelSettingsValidationError, + ModelStoreConfig, + _has_valid_model_extension, + _is_discoverable_local_model_filename, + _sanitize_filename, + _slugify_id, + _unique_filename, + _unique_model_id, + apply_model_chat_defaults, + build_model_capabilities, + get_model_by_id, + is_qwen35_a3b_filename, + model_format_for_filename, + model_supports_vision_filename, + normalize_model_settings, + validate_model_url, +) +from inferno.model_families import ( # noqa: F401 — re-exports + _is_vision_family, + default_projector_candidates_for_model, +) +from inferno.model_registry import ( + model_file_path as _inferno_model_file_path, + model_file_present as _inferno_model_file_present, + describe_model_storage as _inferno_describe_model_storage, + resolve_model_runtime_path as _inferno_resolve_model_runtime_path, + discover_local_model_filenames as _inferno_discover_local_model_filenames, + ensure_models_state as _inferno_ensure_models_state, + save_models_state as _inferno_save_models_state, + register_model_url as _inferno_register_model_url, + delete_model as _inferno_delete_model, + update_model_settings as _inferno_update_model_settings, + any_model_ready as _inferno_any_model_ready, + download_default_projector_for_model as _inferno_download_default_projector, +) +from inferno.model_families import ( + build_model_projector_status as _inferno_build_model_projector_status, +) # --------------------------------------------------------------------------- diff --git a/core/runtime_state.py b/core/runtime_state.py index 04d4848..ac54b62 100644 --- a/core/runtime_state.py +++ b/core/runtime_state.py @@ -23,10 +23,7 @@ except ModuleNotFoundError: # pragma: no cover - optional on non-Pi dev hosts psutil = None # type: ignore[assignment] -try: - from core.inferno import runtime_manager as _inferno -except ModuleNotFoundError: - from inferno import runtime_manager as _inferno # type: ignore[no-redef] +from inferno import runtime_manager as _inferno # --------------------------------------------------------------------------- # Re-exports from inferno (constants + pure functions) @@ -571,7 +568,14 @@ def write_llama_runtime_bundle_marker(runtime: RuntimeConfig, bundle: dict[str, def _detect_installed_runtime_family(runtime: RuntimeConfig) -> str: - """Detect the active runtime family from marker or installed runtime.json.""" + """Detect the effective runtime family for the active model. + + .litertlm models always need the litert adapter regardless of + what binary runtime is installed in the llama/ slot. + """ + model_name = getattr(runtime.model_path, "name", "") + if model_name.endswith(".litertlm"): + return "litert" return _inferno._detect_installed_runtime_family(runtime.base_dir / "llama") diff --git a/docs/recovery.md b/docs/recovery.md index 761596c..de6b7d0 100644 --- a/docs/recovery.md +++ b/docs/recovery.md @@ -50,12 +50,16 @@ To re-deploy the working code from your dev machine: export SSHPASS=raspberry # Fix ownership if needed sshpass -e ssh -o StrictHostKeyChecking=accept-new pi@potato.local \ - "echo raspberry | sudo -S chown -R pi:pi /opt/potato/app" + "echo raspberry | sudo -S chown -R pi:pi /opt/potato/core" # Rsync known-good core/ from your checkout sshpass -e rsync -az --delete \ -e "ssh -o StrictHostKeyChecking=accept-new" \ - core/ pi@potato.local:/opt/potato/app/ + core/ pi@potato.local:/opt/potato/core/ + +# Install dependencies (required — some packages live outside core/) +sshpass -e ssh -o StrictHostKeyChecking=accept-new pi@potato.local \ + "echo raspberry | sudo -S /opt/potato/venv/bin/pip install -r /opt/potato/core/requirements.txt" # Restart the service sshpass -e ssh -o StrictHostKeyChecking=accept-new pi@potato.local \ diff --git a/requirements.txt b/requirements.txt index 5e19b37..009bdb2 100644 --- a/requirements.txt +++ b/requirements.txt @@ -3,3 +3,4 @@ uvicorn[standard]>=0.30.0,<1.0.0 httpx>=0.27.0,<1.0.0 psutil>=6.0.0,<7.0.0 PyYAML>=6.0.2,<7.0.0 +potato-inferno @ git+https://github.com/potato-os/inferno.git@5f0b6a0 diff --git a/tests/unit/test_app_lifecycle.py b/tests/unit/test_app_lifecycle.py index 0c73160..13d30f9 100644 --- a/tests/unit/test_app_lifecycle.py +++ b/tests/unit/test_app_lifecycle.py @@ -180,7 +180,7 @@ def test_max_consecutive_failures_constant_exists(): def test_orchestrator_loop_has_crash_loop_guard(): """The orchestrator must check consecutive_failures before restarting.""" import inspect - from core.inferno import orchestrator + from inferno import orchestrator source = inspect.getsource(orchestrator.run_inference_tick) assert "consecutive_failures" in source diff --git a/tests/unit/test_chat_repository.py b/tests/unit/test_chat_repository.py deleted file mode 100644 index a46051a..0000000 --- a/tests/unit/test_chat_repository.py +++ /dev/null @@ -1,248 +0,0 @@ -from __future__ import annotations - -from typing import Any - -import httpx -import pytest - -from core.inferno import backend -from core.inferno.backend import BackendProxyError - - -class _FakeUpstream: - def __init__(self) -> None: - self.status_code = 200 - self.headers = {"content-type": "text/event-stream"} - self.closed = False - - async def aiter_raw(self): - yield b"data: hello\n\n" - - async def aclose(self) -> None: - self.closed = True - - -class _FakeAsyncClient: - def __init__(self, upstream: _FakeUpstream) -> None: - self._upstream = upstream - self.closed = False - - def build_request(self, method: str, url: str, json: dict[str, Any], headers: dict[str, str]) -> dict[str, Any]: - return {"method": method, "url": url, "json": json, "headers": headers} - - async def send(self, request: dict[str, Any], stream: bool = False) -> _FakeUpstream: - _ = request - assert stream is True - return self._upstream - - async def aclose(self) -> None: - self.closed = True - - -@pytest.mark.anyio -async def test_llama_stream_closes_upstream_and_client_when_consumer_closes(monkeypatch: pytest.MonkeyPatch): - upstream = _FakeUpstream() - client = _FakeAsyncClient(upstream) - - def _client_factory(*args: Any, **kwargs: Any) -> _FakeAsyncClient: - _ = args, kwargs - return client - - monkeypatch.setattr(backend.httpx, "AsyncClient", _client_factory) - - repo = backend.LlamaCppRepository("http://llama.local") - response = await repo.create_chat_completion( - payload={"stream": True, "messages": [{"role": "user", "content": "ping"}]}, - forward_headers={}, - ) - - assert response.stream is not None - stream = response.stream - - first = await anext(stream) - assert first == b"data: hello\n\n" - - await stream.aclose() - - assert upstream.closed is True - assert client.closed is True - - -@pytest.mark.anyio -async def test_fake_stream_uses_test_mode_prefill_and_chunk_delay(monkeypatch: pytest.MonkeyPatch): - sleep_calls: list[float] = [] - - async def _fake_sleep(delay: float) -> None: - sleep_calls.append(delay) - - monkeypatch.setenv("POTATO_TEST_MODE", "1") - monkeypatch.setenv("POTATO_FAKE_PREFILL_DELAY_MS", "250") - monkeypatch.setenv("POTATO_FAKE_STREAM_CHUNK_DELAY_MS", "40") - monkeypatch.setattr(backend.asyncio, "sleep", _fake_sleep) - - repo = backend.FakeLlamaRepository() - response = await repo.create_chat_completion( - payload={"stream": True, "messages": [{"role": "user", "content": "hello"}]}, - forward_headers={}, - ) - - assert response.stream is not None - chunks = [] - async for chunk in response.stream: - chunks.append(chunk.decode("utf-8")) - if "[DONE]" in chunks[-1]: - break - - assert any('"delta":{"role":"assistant"}' in chunk for chunk in chunks) - assert 0.25 in sleep_calls - assert 0.04 in sleep_calls - - -@pytest.mark.anyio -async def test_fake_stream_honors_prefill_delay_override_without_test_mode(monkeypatch: pytest.MonkeyPatch): - sleep_calls: list[float] = [] - - async def _fake_sleep(delay: float) -> None: - sleep_calls.append(delay) - - monkeypatch.delenv("POTATO_TEST_MODE", raising=False) - monkeypatch.setenv("POTATO_FAKE_PREFILL_DELAY_MS", "250") - monkeypatch.setattr(backend.asyncio, "sleep", _fake_sleep) - - repo = backend.FakeLlamaRepository() - response = await repo.create_chat_completion( - payload={"stream": True, "messages": [{"role": "user", "content": "hello"}]}, - forward_headers={}, - ) - - assert response.stream is not None - async for chunk in response.stream: - if b"[DONE]" in chunk: - break - - assert 0.25 in sleep_calls - - -def test_fake_content_has_fake_marker_and_last_user_message(): - payload = { - "messages": [ - {"role": "system", "content": "be precise"}, - {"role": "user", "content": "what is next for CS?"}, - ] - } - - content = backend._fake_content(payload) - - assert "[fake-llama.cpp]" in content - assert "Potato OS" in content - assert "what is next for CS?" in content - - -def test_fake_reply_pool_has_ten_entries(): - assert len(backend.FAKE_PARODY_REPLIES) == 10 - - -def test_fake_content_uses_random_choice_for_reply(monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr(backend.random, "choice", lambda _items: "RANDOM_POTATO_REPLY") - payload = {"messages": [{"role": "user", "content": "same prompt every time"}]} - content = backend._fake_content(payload) - assert "RANDOM_POTATO_REPLY" in content - - -def test_fake_content_is_deterministic_when_seed_is_provided(monkeypatch: pytest.MonkeyPatch): - def _random_choice_should_not_run(_items): - raise AssertionError("global random.choice should not be used for seeded fake replies") - - monkeypatch.setattr(backend.random, "choice", _random_choice_should_not_run) - payload = { - "seed": 42, - "messages": [{"role": "user", "content": "same prompt every time"}], - } - - first = backend._fake_content(payload) - second = backend._fake_content(payload) - - assert first == second - - -def test_fake_default_timing_targets_about_five_tokens_per_second(monkeypatch: pytest.MonkeyPatch): - monkeypatch.delenv("POTATO_TEST_MODE", raising=False) - monkeypatch.delenv("POTATO_FAKE_PREFILL_DELAY_MS", raising=False) - monkeypatch.delenv("POTATO_FAKE_STREAM_CHUNK_DELAY_MS", raising=False) - - prefill_s, chunk_s = backend._read_fake_timing_config() - assert prefill_s == 0.0 - assert 0.19 <= chunk_s <= 0.23 - - -class _TimeoutCapturingClient: - """Fake httpx.AsyncClient that captures the timeout and raises ReadTimeout.""" - - captured_timeouts: list[httpx.Timeout] = [] - - def __init__(self, *args: Any, **kwargs: Any) -> None: - timeout = kwargs.get("timeout") - if isinstance(timeout, httpx.Timeout): - _TimeoutCapturingClient.captured_timeouts.append(timeout) - - def build_request(self, **kwargs: Any) -> dict[str, Any]: - return kwargs - - async def send(self, request: Any, **kwargs: Any) -> None: - raise httpx.ReadTimeout("read timed out") - - async def post(self, url: str, **kwargs: Any) -> None: - raise httpx.ReadTimeout("read timed out") - - async def aclose(self) -> None: - pass - - async def __aenter__(self) -> "_TimeoutCapturingClient": - return self - - async def __aexit__(self, *args: Any) -> None: - pass - - -@pytest.mark.anyio -async def test_llama_stream_uses_unbounded_read_timeout(monkeypatch: pytest.MonkeyPatch): - _TimeoutCapturingClient.captured_timeouts.clear() - monkeypatch.setattr(backend.httpx, "AsyncClient", _TimeoutCapturingClient) - repo = backend.LlamaCppRepository("http://llama.local") - with pytest.raises(BackendProxyError): - await repo.create_chat_completion( - payload={"stream": True, "messages": [{"role": "user", "content": "hi"}]}, - forward_headers={}, - ) - assert len(_TimeoutCapturingClient.captured_timeouts) == 1 - assert _TimeoutCapturingClient.captured_timeouts[0].read is None - - -@pytest.mark.anyio -async def test_llama_non_stream_uses_unbounded_read_timeout(monkeypatch: pytest.MonkeyPatch): - _TimeoutCapturingClient.captured_timeouts.clear() - monkeypatch.setattr(backend.httpx, "AsyncClient", _TimeoutCapturingClient) - repo = backend.LlamaCppRepository("http://llama.local") - with pytest.raises(BackendProxyError): - await repo.create_chat_completion( - payload={"stream": False, "messages": [{"role": "user", "content": "hi"}]}, - forward_headers={}, - ) - assert len(_TimeoutCapturingClient.captured_timeouts) == 1 - assert _TimeoutCapturingClient.captured_timeouts[0].read is None - - -@pytest.mark.anyio -async def test_llama_both_paths_have_bounded_connect_timeout(monkeypatch: pytest.MonkeyPatch): - _TimeoutCapturingClient.captured_timeouts.clear() - monkeypatch.setattr(backend.httpx, "AsyncClient", _TimeoutCapturingClient) - repo = backend.LlamaCppRepository("http://llama.local") - for stream in (True, False): - with pytest.raises(BackendProxyError): - await repo.create_chat_completion( - payload={"stream": stream, "messages": [{"role": "user", "content": "hi"}]}, - forward_headers={}, - ) - assert len(_TimeoutCapturingClient.captured_timeouts) == 2 - for t in _TimeoutCapturingClient.captured_timeouts: - assert t.connect == 5.0 diff --git a/tests/unit/test_inferno_model_registry.py b/tests/unit/test_inferno_model_registry.py deleted file mode 100644 index ce76798..0000000 --- a/tests/unit/test_inferno_model_registry.py +++ /dev/null @@ -1,559 +0,0 @@ -"""Tests for core.inferno.model_registry — format, settings, file ops, and state management.""" - -from __future__ import annotations - -import json -from pathlib import Path - -import pytest - -try: - from core.inferno.model_registry import ( - VALID_MODEL_EXTENSIONS, - DEFAULT_MODEL_CHAT_SETTINGS, - DEFAULT_MODEL_VISION_SETTINGS, - MODELS_STATE_VERSION, - ModelSettingsValidationError, - ModelStoreConfig, - model_format_for_filename, - validate_model_url, - _has_valid_model_extension, - _sanitize_filename, - _slugify_id, - _unique_model_id, - _unique_filename, - is_qwen35_a3b_filename, - model_supports_vision_filename, - normalize_model_settings, - build_model_capabilities, - apply_model_chat_defaults, - get_model_by_id, - _is_discoverable_local_model_filename, - model_file_path, - model_file_present, - describe_model_storage, - resolve_model_runtime_path, - discover_local_model_filenames, - ensure_models_state, - save_models_state, - register_model_url, - delete_model, - update_model_settings, - any_model_ready, - ) -except ModuleNotFoundError: - from inferno.model_registry import ( # type: ignore[no-redef] - VALID_MODEL_EXTENSIONS, - DEFAULT_MODEL_CHAT_SETTINGS, - DEFAULT_MODEL_VISION_SETTINGS, - MODELS_STATE_VERSION, - ModelSettingsValidationError, - ModelStoreConfig, - model_format_for_filename, - validate_model_url, - _has_valid_model_extension, - _sanitize_filename, - _slugify_id, - _unique_model_id, - _unique_filename, - is_qwen35_a3b_filename, - model_supports_vision_filename, - normalize_model_settings, - build_model_capabilities, - apply_model_chat_defaults, - get_model_by_id, - _is_discoverable_local_model_filename, - model_file_path, - model_file_present, - describe_model_storage, - resolve_model_runtime_path, - discover_local_model_filenames, - ensure_models_state, - save_models_state, - register_model_url, - delete_model, - update_model_settings, - any_model_ready, - ) - - -@pytest.fixture -def store(tmp_path: Path) -> ModelStoreConfig: - """Create a ModelStoreConfig backed by a temp directory.""" - models_dir = tmp_path / "models" - models_dir.mkdir() - state_path = tmp_path / "state" / "models.json" - state_path.parent.mkdir() - return ModelStoreConfig( - models_dir=models_dir, - state_path=state_path, - default_filename="default-model.gguf", - default_url="https://example.com/default-model.gguf", - known_default_filenames=("default-model.gguf",), - current_model_filename="", - ) - - -# -- Constants -- - - -def test_valid_model_extensions(): - assert ".gguf" in VALID_MODEL_EXTENSIONS - assert ".litertlm" in VALID_MODEL_EXTENSIONS - - -def test_state_version_is_int(): - assert isinstance(MODELS_STATE_VERSION, int) - - -def test_default_chat_settings_has_required_keys(): - for key in ("temperature", "top_p", "top_k", "max_tokens", "stream"): - assert key in DEFAULT_MODEL_CHAT_SETTINGS - - -def test_default_vision_settings_has_required_keys(): - for key in ("enabled", "projector_mode", "projector_filename"): - assert key in DEFAULT_MODEL_VISION_SETTINGS - - -# -- Format detection -- - - -@pytest.mark.parametrize( - "filename, expected", - [ - ("model.gguf", "gguf"), - ("MODEL.GGUF", "gguf"), - ("model.litertlm", "litertlm"), - ("Model.LiteRTLM", "litertlm"), - ("anything-else.bin", "gguf"), - ], -) -def test_model_format_for_filename(filename, expected): - assert model_format_for_filename(filename) == expected - - -@pytest.mark.parametrize( - "filename, expected", - [ - ("model.gguf", True), - ("model.GGUF", True), - ("model.litertlm", True), - ("model.bin", False), - ("model.safetensors", False), - ], -) -def test_has_valid_model_extension(filename, expected): - assert _has_valid_model_extension(filename) == expected - - -# -- URL validation -- - - -def test_validate_model_url_valid(): - ok, reason, name = validate_model_url("https://example.com/model.gguf") - assert ok is True - assert reason == "" - assert name == "model.gguf" - - -def test_validate_model_url_http_rejected(): - ok, reason, _ = validate_model_url("http://example.com/model.gguf") - assert ok is False - assert reason == "https_required" - - -def test_validate_model_url_no_filename(): - ok, reason, _ = validate_model_url("https://example.com/") - assert ok is False - assert reason == "filename_missing" - - -def test_validate_model_url_bad_extension(): - ok, reason, _ = validate_model_url("https://example.com/model.bin") - assert ok is False - assert reason == "unsupported_model_format" - - -def test_validate_model_url_litertlm(): - ok, reason, name = validate_model_url("https://example.com/model.litertlm") - assert ok is True - assert name == "model.litertlm" - - -# -- Filename / ID utilities -- - - -def test_sanitize_filename_basic(): - assert _sanitize_filename("model.gguf") == "model.gguf" - - -def test_sanitize_filename_special_chars(): - result = _sanitize_filename("my model (v2).gguf") - assert result.endswith(".gguf") - assert " " not in result - assert "(" not in result - - -def test_sanitize_filename_empty(): - assert _sanitize_filename("") == "model.gguf" - - -def test_slugify_id(): - assert _slugify_id("My Model-V2") == "my-model-v2" - - -def test_slugify_id_empty(): - assert _slugify_id("") == "model" - - -def test_unique_model_id_no_conflict(): - assert _unique_model_id("base", set()) == "base" - - -def test_unique_model_id_with_conflict(): - assert _unique_model_id("base", {"base"}) == "base-2" - assert _unique_model_id("base", {"base", "base-2"}) == "base-3" - - -def test_unique_filename_no_conflict(): - assert _unique_filename("model.gguf", set()) == "model.gguf" - - -def test_unique_filename_with_conflict(): - assert _unique_filename("model.gguf", {"model.gguf"}) == "model-2.gguf" - - -# -- Model detection -- - - -@pytest.mark.parametrize( - "filename, expected", - [ - ("Qwen3.5-35B-A3B-Q4_K_M.gguf", True), - ("qwen-3.5-35b-a3b-q4.gguf", True), - ("Qwen3.5-2B-Q4_K_M.gguf", False), - (None, False), - ("", False), - ], -) -def test_is_qwen35_a3b_filename(filename, expected): - assert is_qwen35_a3b_filename(filename) == expected - - -@pytest.mark.parametrize( - "filename, expected", - [ - ("Qwen3.5-2B-Q4_K_M.gguf", True), - ("Qwen3-VL-2B-Q4.gguf", True), - ("gemma-4-E2B-it-Q4_K_M.gguf", True), - ("llama-3.2-1B.gguf", False), - (None, False), - ("", False), - ], -) -def test_model_supports_vision_filename(filename, expected): - assert model_supports_vision_filename(filename) == expected - - -# -- Settings normalization -- - - -def test_normalize_model_settings_defaults(): - result = normalize_model_settings(None, filename="model.gguf") - assert "chat" in result - assert "vision" in result - assert result["chat"]["temperature"] == DEFAULT_MODEL_CHAT_SETTINGS["temperature"] - - -def test_normalize_model_settings_preserves_values(): - raw = {"chat": {"temperature": 0.5}} - result = normalize_model_settings(raw, filename="model.gguf") - assert result["chat"]["temperature"] == 0.5 - - -def test_normalize_model_settings_invalid_float_raises(): - raw = {"chat": {"temperature": "not_a_number"}} - with pytest.raises(ModelSettingsValidationError): - normalize_model_settings(raw, filename="model.gguf") - - -def test_normalize_model_settings_vision_enabled_for_vision_model(): - result = normalize_model_settings(None, filename="Qwen3.5-2B-Q4_K_M.gguf") - assert result["vision"]["enabled"] is True - - -def test_normalize_model_settings_vision_disabled_for_non_vision(): - result = normalize_model_settings(None, filename="llama-3.2.gguf") - assert result["vision"]["enabled"] is False - - -def test_normalize_vision_projector_mode_validates(): - raw = {"vision": {"projector_mode": "bogus"}} - result = normalize_model_settings(raw, filename="model.gguf") - assert result["vision"]["projector_mode"] == "default" - - -# -- Capabilities -- - - -def test_build_model_capabilities_vision(): - caps = build_model_capabilities("Qwen3.5-2B-Q4_K_M.gguf") - assert caps["vision"] is True - - -def test_build_model_capabilities_no_vision(): - caps = build_model_capabilities("llama-3.2-1B.gguf") - assert caps["vision"] is False - - -# -- Chat defaults -- - - -def test_apply_model_chat_defaults_a3b_adds_thinking(): - payload = {"messages": []} - result = apply_model_chat_defaults(payload, active_model_filename="Qwen3.5-35B-A3B-Q4.gguf") - assert result["chat_template_kwargs"]["enable_thinking"] is False - - -def test_apply_model_chat_defaults_non_a3b_unchanged(): - payload = {"messages": []} - result = apply_model_chat_defaults(payload, active_model_filename="Qwen3.5-2B-Q4.gguf") - assert result is payload - - -def test_apply_model_chat_defaults_preserves_explicit_thinking(): - payload = {"messages": [], "chat_template_kwargs": {"enable_thinking": True}} - result = apply_model_chat_defaults(payload, active_model_filename="Qwen3.5-35B-A3B-Q4.gguf") - assert result["chat_template_kwargs"]["enable_thinking"] is True - - -# -- get_model_by_id -- - - -def test_get_model_by_id_found(): - state = {"models": [{"id": "a"}, {"id": "b"}]} - assert get_model_by_id(state, "b") == {"id": "b"} - - -def test_get_model_by_id_missing(): - state = {"models": [{"id": "a"}]} - assert get_model_by_id(state, "missing") is None - - -def test_get_model_by_id_empty(): - assert get_model_by_id({"models": []}, "x") is None - - -# -- Discoverable filenames -- - - -def test_discoverable_gguf(): - assert _is_discoverable_local_model_filename("model.gguf") is True - - -def test_discoverable_litertlm(): - assert _is_discoverable_local_model_filename("model.litertlm") is True - - -def test_discoverable_rejects_mmproj(): - assert _is_discoverable_local_model_filename("mmproj-F16.gguf") is False - assert _is_discoverable_local_model_filename("mmproj-model-f16.gguf") is False - - -def test_discoverable_rejects_bad_extension(): - assert _is_discoverable_local_model_filename("model.bin") is False - - -# -- File / path operations -- - - -def test_model_file_path(tmp_path): - result = model_file_path(tmp_path / "models", "test.gguf") - assert result == tmp_path / "models" / "test.gguf" - - -def test_model_file_present_true(tmp_path): - models_dir = tmp_path / "models" - models_dir.mkdir() - (models_dir / "model.gguf").write_bytes(b"data") - assert model_file_present(models_dir, "model.gguf") is True - - -def test_model_file_present_false(tmp_path): - models_dir = tmp_path / "models" - models_dir.mkdir() - assert model_file_present(models_dir, "model.gguf") is False - - -def test_model_file_present_empty_file(tmp_path): - models_dir = tmp_path / "models" - models_dir.mkdir() - (models_dir / "model.gguf").write_bytes(b"") - assert model_file_present(models_dir, "model.gguf") is False - - -def test_describe_model_storage_exists(tmp_path): - models_dir = tmp_path / "models" - models_dir.mkdir() - (models_dir / "model.gguf").write_bytes(b"x" * 100) - result = describe_model_storage(models_dir, "model.gguf") - assert result["exists"] is True - assert result["size_bytes"] == 100 - assert result["location"] == "local" - - -def test_describe_model_storage_missing(tmp_path): - models_dir = tmp_path / "models" - models_dir.mkdir() - result = describe_model_storage(models_dir, "model.gguf") - assert result["exists"] is False - assert result["size_bytes"] == 0 - - -def test_resolve_model_runtime_path_regular(tmp_path): - models_dir = tmp_path / "models" - models_dir.mkdir() - (models_dir / "model.gguf").write_bytes(b"data") - result = resolve_model_runtime_path(models_dir, "model.gguf") - assert result == models_dir / "model.gguf" - - -def test_resolve_model_runtime_path_symlink(tmp_path): - models_dir = tmp_path / "models" - models_dir.mkdir() - real = tmp_path / "real.gguf" - real.write_bytes(b"data") - (models_dir / "model.gguf").symlink_to(real) - result = resolve_model_runtime_path(models_dir, "model.gguf") - assert result == real.resolve() - - -def test_discover_local_model_filenames(tmp_path): - models_dir = tmp_path / "models" - models_dir.mkdir() - (models_dir / "alpha.gguf").write_bytes(b"a") - (models_dir / "beta.litertlm").write_bytes(b"b") - (models_dir / "mmproj-F16.gguf").write_bytes(b"p") - (models_dir / "readme.txt").write_bytes(b"r") - result = discover_local_model_filenames(models_dir) - assert "alpha.gguf" in result - assert "beta.litertlm" in result - assert "mmproj-F16.gguf" not in result - assert "readme.txt" not in result - - -def test_discover_local_model_filenames_empty(tmp_path): - models_dir = tmp_path / "models" - models_dir.mkdir() - assert discover_local_model_filenames(models_dir) == [] - - -# -- State management (ModelStoreConfig) -- - - -def test_ensure_models_state_creates_default(store): - state = ensure_models_state(store) - assert state["version"] == MODELS_STATE_VERSION - assert len(state["models"]) >= 1 - default = state["models"][0] - assert default["id"] == "default" - assert default["filename"] == "default-model.gguf" - assert store.state_path.exists() - - -def test_ensure_models_state_reads_existing(store): - raw = { - "version": 1, - "models": [{"id": "custom", "filename": "custom.gguf", "source_url": None, "source_type": "upload", "status": "ready"}], - "active_model_id": "custom", - } - store.state_path.write_text(json.dumps(raw), encoding="utf-8") - state = ensure_models_state(store) - ids = [m["id"] for m in state["models"]] - assert "custom" in ids - - -def test_ensure_models_state_marks_default_downloaded_once(store): - (store.models_dir / "default-model.gguf").write_bytes(b"model-data") - state = ensure_models_state(store) - assert state["default_model_downloaded_once"] is True - - -def test_save_models_state_normalizes_and_persists(store): - ensure_models_state(store) - state = {"models": [{"id": "x", "filename": "x.gguf"}], "active_model_id": "x"} - saved = save_models_state(store, state) - assert saved["version"] == MODELS_STATE_VERSION - reloaded = json.loads(store.state_path.read_text(encoding="utf-8")) - assert reloaded["version"] == MODELS_STATE_VERSION - - -def test_register_model_url_adds_model(store): - ensure_models_state(store) - ok, reason, model = register_model_url(store, "https://example.com/new-model.gguf") - assert ok is True - assert reason == "registered" - assert model["filename"] == "new-model.gguf" - - -def test_register_model_url_detects_duplicate(store): - ensure_models_state(store) - register_model_url(store, "https://example.com/dup.gguf") - ok, reason, _ = register_model_url(store, "https://example.com/dup.gguf") - assert ok is True - assert reason == "already_exists" - - -def test_register_model_url_rejects_http(store): - ensure_models_state(store) - ok, reason, _ = register_model_url(store, "http://example.com/model.gguf") - assert ok is False - assert reason == "https_required" - - -def test_delete_model_removes_file(store): - ensure_models_state(store) - register_model_url(store, "https://example.com/to-delete.gguf") - (store.models_dir / "to-delete.gguf").write_bytes(b"data") - ok, reason, deleted_file, freed, was_active = delete_model(store, model_id="to-delete") - assert ok is True - assert reason == "deleted" - assert deleted_file is True - assert freed > 0 - assert not (store.models_dir / "to-delete.gguf").exists() - - -def test_delete_model_not_found(store): - ensure_models_state(store) - ok, reason, _, _, _ = delete_model(store, model_id="nonexistent") - assert ok is False - assert reason == "model_not_found" - - -def test_update_model_settings_persists(store): - ensure_models_state(store) - ok, reason, updated = update_model_settings( - store, model_id="default", settings={"chat": {"temperature": 0.3}} - ) - assert ok is True - assert reason == "updated" - assert updated["settings"]["chat"]["temperature"] == 0.3 - - -def test_update_model_settings_not_found(store): - ensure_models_state(store) - ok, reason, _ = update_model_settings(store, model_id="nope", settings={}) - assert ok is False - assert reason == "model_not_found" - - -def test_any_model_ready_false(store): - ensure_models_state(store) - assert any_model_ready(store) is False - - -def test_any_model_ready_true(store): - ensure_models_state(store) - (store.models_dir / "default-model.gguf").write_bytes(b"data") - assert any_model_ready(store) is True diff --git a/tests/unit/test_inferno_orchestrator.py b/tests/unit/test_inferno_orchestrator.py deleted file mode 100644 index 4ea1b1e..0000000 --- a/tests/unit/test_inferno_orchestrator.py +++ /dev/null @@ -1,680 +0,0 @@ -"""Tests for core.inferno.orchestrator — health, readiness, process lifecycle, inference tick.""" - -from __future__ import annotations - -import asyncio -import json -from pathlib import Path -from typing import Any -from unittest.mock import AsyncMock - -import httpx -import pytest - -try: - from core.inferno.orchestrator import ( - READY_HEALTH_POLLS_REQUIRED, - MAX_CONSECUTIVE_FAILURES, - InferenceTickResult, - empty_readiness_state, - empty_runtime_switch_state, - reset_readiness, - resolve_readiness, - check_health, - probe_inference_slot, - refresh_readiness, - restart_inference_process, - resolve_mmproj_for_launch, - resolve_no_mmap, - run_inference_tick, - prepare_activation_runtime, - ) -except ModuleNotFoundError: - from inferno.orchestrator import ( # type: ignore[no-redef] - READY_HEALTH_POLLS_REQUIRED, - MAX_CONSECUTIVE_FAILURES, - InferenceTickResult, - empty_readiness_state, - empty_runtime_switch_state, - reset_readiness, - resolve_readiness, - check_health, - probe_inference_slot, - refresh_readiness, - restart_inference_process, - resolve_mmproj_for_launch, - resolve_no_mmap, - run_inference_tick, - prepare_activation_runtime, - ) - - -# --------------------------------------------------------------------------- -# Fixtures -# --------------------------------------------------------------------------- - - -@pytest.fixture -def tick_env(tmp_path: Path) -> dict[str, Any]: - """Minimal filesystem + paths for tick tests.""" - base = tmp_path / "potato" - models_dir = base / "models" - models_dir.mkdir(parents=True) - - model_path = models_dir / "test-model.gguf" - model_path.write_bytes(b"\x00" * 100) - - return { - "model_path": model_path, - "base_url": "http://llama.test:8080", - } - - -class _FakeProcess: - """Minimal process stub for testing orchestrator logic.""" - - def __init__(self, *, returncode: int | None = None, pid: int = 42): - self.returncode = returncode - self.pid = pid - self.terminated = False - self.killed = False - - def terminate(self) -> None: - self.terminated = True - - def kill(self) -> None: - self.killed = True - - async def wait(self) -> int: - return self.returncode if self.returncode is not None else 0 - - -# --------------------------------------------------------------------------- -# 1. Constants and state factories -# --------------------------------------------------------------------------- - - -def test_ready_health_polls_required_positive(): - assert READY_HEALTH_POLLS_REQUIRED > 0 - - -def test_max_consecutive_failures_positive(): - assert MAX_CONSECUTIVE_FAILURES > 0 - - -def test_empty_readiness_state_defaults(): - state = empty_readiness_state() - assert state["generation"] == 0 - assert state["status"] == "idle" - assert state["ready"] is False - assert state["transport_healthy"] is False - assert state["healthy_polls"] == 0 - assert state["model_path"] is None - assert state["last_error"] is None - assert state["last_ready_at_unix"] is None - - -def test_empty_runtime_switch_state_defaults(): - state = empty_runtime_switch_state() - assert state["active"] is False - assert state["target_bundle_path"] is None - assert state["error"] is None - - -# --------------------------------------------------------------------------- -# 2. Readiness state transitions (pure) -# --------------------------------------------------------------------------- - - -def test_reset_readiness_increments_generation(): - prev = empty_readiness_state() - prev["generation"] = 5 - result = reset_readiness(prev, model_path="/m/test.gguf", reason="test") - assert result["generation"] == 6 - - -def test_reset_readiness_with_model_sets_loading(): - result = reset_readiness(None, model_path="/m/test.gguf", reason=None) - assert result["status"] == "loading" - assert result["model_path"] == "/m/test.gguf" - - -def test_reset_readiness_without_model_sets_idle(): - result = reset_readiness(None, model_path=None, reason="no_model") - assert result["status"] == "idle" - assert result["model_path"] is None - assert result["last_error"] == "no_model" - - -def test_resolve_readiness_model_change_triggers_reset(): - current = empty_readiness_state() - current["model_path"] = "/m/old.gguf" - current["ready"] = True - result = resolve_readiness(current, active_model_path="/m/new.gguf") - assert result["model_path"] == "/m/new.gguf" - assert result["ready"] is False - assert result["status"] == "loading" - - -def test_resolve_readiness_no_model_when_had_one_resets(): - current = empty_readiness_state() - current["model_path"] = "/m/old.gguf" - current["ready"] = True - result = resolve_readiness(current, active_model_path=None) - assert result["model_path"] is None - assert result["ready"] is False - - -def test_resolve_readiness_same_model_no_change(): - current = empty_readiness_state() - current["model_path"] = "/m/same.gguf" - current["status"] = "warming" - current["healthy_polls"] = 1 - result = resolve_readiness(current, active_model_path="/m/same.gguf") - assert result["status"] == "warming" - assert result["healthy_polls"] == 1 - - -# --------------------------------------------------------------------------- -# 3. Health check and slot probe (async, httpx) -# --------------------------------------------------------------------------- - - -@pytest.mark.anyio -async def test_check_health_true_on_healthy_endpoint(monkeypatch): - async def _fake_get(self, url, **kw): - return httpx.Response(200) - - monkeypatch.setattr(httpx.AsyncClient, "get", _fake_get) - assert await check_health("http://llama.test:8080") is True - - -@pytest.mark.anyio -async def test_check_health_false_on_all_500(monkeypatch): - async def _fake_get(self, url, **kw): - return httpx.Response(500) - - monkeypatch.setattr(httpx.AsyncClient, "get", _fake_get) - assert await check_health("http://llama.test:8080") is False - - -@pytest.mark.anyio -async def test_check_health_busy_is_healthy_on_timeout(monkeypatch): - async def _fake_get(self, url, **kw): - raise httpx.ReadTimeout("busy") - - monkeypatch.setattr(httpx.AsyncClient, "get", _fake_get) - assert await check_health("http://llama.test:8080", busy_is_healthy=True) is True - - -@pytest.mark.anyio -async def test_check_health_strict_timeout_not_healthy(monkeypatch): - async def _fake_get(self, url, **kw): - raise httpx.ReadTimeout("busy") - - monkeypatch.setattr(httpx.AsyncClient, "get", _fake_get) - assert await check_health("http://llama.test:8080", busy_is_healthy=False) is False - - -@pytest.mark.anyio -async def test_probe_inference_slot_true_on_success(monkeypatch): - async def _fake_post(self, url, **kw): - return httpx.Response(200) - - monkeypatch.setattr(httpx.AsyncClient, "post", _fake_post) - assert await probe_inference_slot("http://llama.test:8080") is True - - -@pytest.mark.anyio -async def test_probe_inference_slot_false_on_error(monkeypatch): - async def _fake_post(self, url, **kw): - raise httpx.ConnectError("down") - - monkeypatch.setattr(httpx.AsyncClient, "post", _fake_post) - assert await probe_inference_slot("http://llama.test:8080") is False - - -# --------------------------------------------------------------------------- -# 4. Readiness refresh (async) -# --------------------------------------------------------------------------- - - -@pytest.mark.anyio -async def test_refresh_no_model_stays_idle(monkeypatch): - state = empty_readiness_state() - result = await refresh_readiness(state, base_url="http://llama.test:8080", process_alive=True) - assert result["status"] == "idle" - - -@pytest.mark.anyio -async def test_refresh_process_dead_resets_to_loading(monkeypatch): - state = empty_readiness_state() - state["model_path"] = "/m/test.gguf" - state["status"] = "warming" - state["healthy_polls"] = 1 - result = await refresh_readiness(state, base_url="http://llama.test:8080", process_alive=False) - assert result["status"] == "loading" - assert result["healthy_polls"] == 0 - assert result["ready"] is False - - -@pytest.mark.anyio -async def test_refresh_healthy_increments_polls(monkeypatch): - async def _fake_get(self, url, **kw): - return httpx.Response(200) - - monkeypatch.setattr(httpx.AsyncClient, "get", _fake_get) - - state = empty_readiness_state() - state["model_path"] = "/m/test.gguf" - state["status"] = "loading" - result = await refresh_readiness(state, base_url="http://llama.test:8080", process_alive=True) - assert result["healthy_polls"] == 1 - assert result["transport_healthy"] is True - - -@pytest.mark.anyio -async def test_refresh_becomes_ready_at_threshold(monkeypatch): - async def _fake_get(self, url, **kw): - return httpx.Response(200) - - monkeypatch.setattr(httpx.AsyncClient, "get", _fake_get) - - state = empty_readiness_state() - state["model_path"] = "/m/test.gguf" - state["healthy_polls"] = READY_HEALTH_POLLS_REQUIRED - 1 - result = await refresh_readiness(state, base_url="http://llama.test:8080", process_alive=True) - assert result["ready"] is True - assert result["status"] == "ready" - assert result["last_error"] is None - - -@pytest.mark.anyio -async def test_refresh_unhealthy_resets_polls(monkeypatch): - async def _fake_get(self, url, **kw): - return httpx.Response(500) - - monkeypatch.setattr(httpx.AsyncClient, "get", _fake_get) - - state = empty_readiness_state() - state["model_path"] = "/m/test.gguf" - state["healthy_polls"] = 1 - state["transport_healthy"] = True - result = await refresh_readiness(state, base_url="http://llama.test:8080", process_alive=True) - assert result["healthy_polls"] == 0 - assert result["transport_healthy"] is False - assert result["ready"] is False - - -# --------------------------------------------------------------------------- -# 5. mmproj resolution -# --------------------------------------------------------------------------- - - -def test_resolve_mmproj_non_vision_returns_none(tmp_path): - model = {"filename": "plain-model.gguf", "settings": {}} - result = resolve_mmproj_for_launch(tmp_path, tmp_path, model, "ik_llama") - assert result is None - - -def test_resolve_mmproj_gemma4_ik_llama_suppressed(tmp_path): - model = { - "filename": "gemma-4-4b-it-Q4_K_M.gguf", - "settings": {"vision": {"enabled": True}}, - } - result = resolve_mmproj_for_launch(tmp_path, tmp_path, model, "ik_llama") - assert result is None - - -# --------------------------------------------------------------------------- -# 6. no-mmap resolution -# --------------------------------------------------------------------------- - - -def test_resolve_no_mmap_explicit_true(): - status = {"no_mmap_env": "true"} - assert resolve_no_mmap(status, "any-model.gguf", "ik_llama", device_class="pi5-8gb", bundle_marker=None) is True - - -def test_resolve_no_mmap_explicit_false(): - status = {"no_mmap_env": "false"} - assert resolve_no_mmap(status, "any-model.gguf", "ik_llama", device_class="pi5-8gb", bundle_marker=None) is False - - -def test_resolve_no_mmap_auto_non_a3b(): - status = {"no_mmap_env": "auto"} - assert resolve_no_mmap(status, "plain-model.gguf", "ik_llama", device_class="pi5-16gb", bundle_marker=None) is False - - -def test_resolve_no_mmap_auto_a3b_heuristic(): - status = {"no_mmap_env": "auto"} - marker = {"profile": "pi5-opt"} - result = resolve_no_mmap( - status, - "Qwen3.5-35B-A3B-Q4_K_M.gguf", - "ik_llama", - device_class="pi5-16gb", - bundle_marker=marker, - ) - assert result is True - - -# --------------------------------------------------------------------------- -# 7. Restart inference process -# --------------------------------------------------------------------------- - - -@pytest.mark.anyio -async def test_restart_terminates_running_process(): - proc = _FakeProcess(returncode=None) - terminate_called = False - - async def _terminate(p, timeout=3.0): - nonlocal terminate_called - terminate_called = True - - readiness = empty_readiness_state() - readiness["model_path"] = "/m/test.gguf" - - new_readiness, terminated, reason = await restart_inference_process( - readiness=readiness, - process=proc, - model_path="/m/test.gguf", - terminate_fn=_terminate, - stray_kill_fn=AsyncMock(return_value=0), - ) - assert terminate_called - assert terminated is True - assert "terminated_running" in reason - - -@pytest.mark.anyio -async def test_restart_cleans_stale_processes(): - stray_kill = AsyncMock(return_value=2) - - new_readiness, terminated, reason = await restart_inference_process( - readiness=empty_readiness_state(), - process=None, - model_path="/m/test.gguf", - terminate_fn=AsyncMock(), - stray_kill_fn=stray_kill, - ) - stray_kill.assert_awaited_once() - assert terminated is True - assert "stale" in reason - - -@pytest.mark.anyio -async def test_restart_resets_readiness_to_loading(): - readiness = empty_readiness_state() - readiness["ready"] = True - readiness["status"] = "ready" - - new_readiness, _, _ = await restart_inference_process( - readiness=readiness, - process=None, - model_path="/m/test.gguf", - terminate_fn=AsyncMock(), - stray_kill_fn=AsyncMock(return_value=0), - ) - assert new_readiness["status"] == "loading" - assert new_readiness["ready"] is False - - -@pytest.mark.anyio -async def test_restart_propagates_termination_failure(): - async def _exploding_terminate(proc, timeout=3.0): - raise OSError("process stuck") - - proc = _FakeProcess(returncode=None) - with pytest.raises(OSError, match="process stuck"): - await restart_inference_process( - readiness=empty_readiness_state(), - process=proc, - model_path="/m/test.gguf", - terminate_fn=_exploding_terminate, - stray_kill_fn=AsyncMock(return_value=0), - ) - - -# --------------------------------------------------------------------------- -# 8. Inference tick -# --------------------------------------------------------------------------- - - -@pytest.mark.anyio -async def test_tick_model_missing_resets_readiness(tmp_path): - missing = tmp_path / "missing.gguf" - readiness = empty_readiness_state() - readiness["model_path"] = str(missing) - readiness["ready"] = True - - result = await run_inference_tick( - process=None, consecutive_failures=3, - failure_model_key="old", failure_runtime_key="old", - readiness=readiness, - model_path=missing, base_url="http://test:8080", installed_family="ik_llama", - launch_llama_fn=AsyncMock(), launch_litert_fn=None, - ) - assert result.readiness["status"] != "ready" - assert result.consecutive_failures == 0 - - -@pytest.mark.anyio -async def test_tick_model_present_dead_process_spawns(tick_env, monkeypatch): - new_proc = _FakeProcess(returncode=None) - launch = AsyncMock(return_value=new_proc) - - async def _fake_get(self, url, **kw): - return httpx.Response(500) - monkeypatch.setattr(httpx.AsyncClient, "get", _fake_get) - - result = await run_inference_tick( - process=None, consecutive_failures=0, - failure_model_key=None, failure_runtime_key=None, - readiness=empty_readiness_state(), - model_path=tick_env["model_path"], base_url=tick_env["base_url"], - installed_family="ik_llama", - launch_llama_fn=launch, launch_litert_fn=None, - ) - launch.assert_awaited_once() - assert result.process is new_proc - - -@pytest.mark.anyio -async def test_tick_model_changed_resets_failure_counter(tick_env, monkeypatch): - async def _fake_get(self, url, **kw): - return httpx.Response(500) - monkeypatch.setattr(httpx.AsyncClient, "get", _fake_get) - - result = await run_inference_tick( - process=_FakeProcess(returncode=None), consecutive_failures=3, - failure_model_key="/m/old.gguf", failure_runtime_key="ik_llama", - readiness=empty_readiness_state(), - model_path=tick_env["model_path"], base_url=tick_env["base_url"], - installed_family="ik_llama", - launch_llama_fn=AsyncMock(), launch_litert_fn=None, - ) - assert result.consecutive_failures == 0 - assert result.failure_model_key == str(tick_env["model_path"]) - - -@pytest.mark.anyio -async def test_tick_runtime_changed_resets_failure_counter(tick_env, monkeypatch): - async def _fake_get(self, url, **kw): - return httpx.Response(500) - monkeypatch.setattr(httpx.AsyncClient, "get", _fake_get) - - result = await run_inference_tick( - process=_FakeProcess(returncode=None), consecutive_failures=3, - failure_model_key=str(tick_env["model_path"]), failure_runtime_key="llama_cpp", - readiness=empty_readiness_state(), - model_path=tick_env["model_path"], base_url=tick_env["base_url"], - installed_family="ik_llama", - launch_llama_fn=AsyncMock(), launch_litert_fn=None, - ) - assert result.consecutive_failures == 0 - assert result.failure_runtime_key == "ik_llama" - - -@pytest.mark.anyio -async def test_tick_failure_increments_on_nonzero_exit(tick_env, monkeypatch): - async def _fake_get(self, url, **kw): - return httpx.Response(500) - monkeypatch.setattr(httpx.AsyncClient, "get", _fake_get) - - result = await run_inference_tick( - process=_FakeProcess(returncode=1), consecutive_failures=0, - failure_model_key=str(tick_env["model_path"]), failure_runtime_key="ik_llama", - readiness=empty_readiness_state(), - model_path=tick_env["model_path"], base_url=tick_env["base_url"], - installed_family="ik_llama", - launch_llama_fn=AsyncMock(return_value=_FakeProcess(returncode=None)), - launch_litert_fn=None, - ) - assert result.consecutive_failures == 1 - - -@pytest.mark.anyio -async def test_tick_max_failures_stops_restart(tick_env, monkeypatch): - launch = AsyncMock() - - async def _fake_get(self, url, **kw): - return httpx.Response(500) - monkeypatch.setattr(httpx.AsyncClient, "get", _fake_get) - - result = await run_inference_tick( - process=_FakeProcess(returncode=1), - consecutive_failures=MAX_CONSECUTIVE_FAILURES - 1, - failure_model_key=str(tick_env["model_path"]), failure_runtime_key="ik_llama", - readiness=empty_readiness_state(), - model_path=tick_env["model_path"], base_url=tick_env["base_url"], - installed_family="ik_llama", - launch_llama_fn=launch, launch_litert_fn=None, - ) - launch.assert_not_awaited() - assert result.consecutive_failures == MAX_CONSECUTIVE_FAILURES - - -@pytest.mark.anyio -async def test_tick_ready_resets_failures(tick_env, monkeypatch): - async def _fake_get(self, url, **kw): - return httpx.Response(200) - monkeypatch.setattr(httpx.AsyncClient, "get", _fake_get) - - readiness = empty_readiness_state() - readiness["model_path"] = str(tick_env["model_path"]) - readiness["healthy_polls"] = READY_HEALTH_POLLS_REQUIRED - 1 - - result = await run_inference_tick( - process=_FakeProcess(returncode=None), consecutive_failures=2, - failure_model_key=str(tick_env["model_path"]), failure_runtime_key="ik_llama", - readiness=readiness, - model_path=tick_env["model_path"], base_url=tick_env["base_url"], - installed_family="ik_llama", - launch_llama_fn=AsyncMock(), launch_litert_fn=None, - ) - assert result.readiness["ready"] is True - assert result.consecutive_failures == 0 - - -@pytest.mark.anyio -async def test_tick_switch_in_progress_skips_spawn(tick_env, monkeypatch): - launch = AsyncMock() - - async def _fake_get(self, url, **kw): - return httpx.Response(500) - monkeypatch.setattr(httpx.AsyncClient, "get", _fake_get) - - result = await run_inference_tick( - process=None, consecutive_failures=0, - failure_model_key=str(tick_env["model_path"]), failure_runtime_key="ik_llama", - readiness=empty_readiness_state(), - model_path=tick_env["model_path"], base_url=tick_env["base_url"], - installed_family="ik_llama", - launch_llama_fn=launch, launch_litert_fn=None, - switch_in_progress=True, - ) - launch.assert_not_awaited() - - -@pytest.mark.anyio -async def test_tick_launch_returns_none_downgrades_readiness(tick_env, monkeypatch): - """When launch_llama_fn returns None (e.g. missing script), readiness must not stay ready.""" - async def _fake_get(self, url, **kw): - return httpx.Response(500) - monkeypatch.setattr(httpx.AsyncClient, "get", _fake_get) - - readiness = empty_readiness_state() - readiness["model_path"] = str(tick_env["model_path"]) - readiness["ready"] = True - readiness["status"] = "ready" - - result = await run_inference_tick( - process=None, consecutive_failures=0, - failure_model_key=str(tick_env["model_path"]), failure_runtime_key="ik_llama", - readiness=readiness, - model_path=tick_env["model_path"], base_url=tick_env["base_url"], - installed_family="ik_llama", - launch_llama_fn=AsyncMock(return_value=None), launch_litert_fn=None, - ) - assert result.process is None - assert result.readiness["ready"] is False - assert result.readiness["status"] == "loading" - - -# --------------------------------------------------------------------------- -# 11. Activation runtime prep -# --------------------------------------------------------------------------- - - -def test_prepare_activation_incompatible_device(tmp_path): - runtimes_dir = tmp_path / "runtimes" - runtimes_dir.mkdir() - should_switch, reason, family = prepare_activation_runtime( - model_filename="test.gguf", - model_format="gguf", - current_family="llama_cpp", - device_class="pi4-8gb", - runtimes_dir=runtimes_dir, - ) - # ik_llama is incompatible with pi4 — should not switch or should fail. - # The function returns the decision; pi4 can run llama_cpp fine so no switch needed. - assert isinstance(should_switch, bool) - assert isinstance(reason, str) - - -def test_prepare_activation_format_requires_incompatible(tmp_path): - runtimes_dir = tmp_path / "runtimes" - runtimes_dir.mkdir() - should_switch, reason, family = prepare_activation_runtime( - model_filename="model.litertlm", - model_format="litertlm", - current_family="llama_cpp", - device_class="pi4-8gb", - runtimes_dir=runtimes_dir, - ) - # LiteRT is incompatible with Pi 4 — should fail. - assert should_switch is False - assert "incompatible" in reason.lower() or "no_switch" in reason.lower() or family is None - - -def test_prepare_activation_finds_slot(tmp_path): - runtimes_dir = tmp_path / "runtimes" - # GGUF on litert falls back to llama_cpp — create a valid slot. - slot_dir = runtimes_dir / "llama_cpp" - (slot_dir / "bin").mkdir(parents=True) - (slot_dir / "bin" / "llama-server").write_text("#!/bin/sh\n") - runtime_json = slot_dir / "runtime.json" - runtime_json.write_text(json.dumps({"family": "llama_cpp", "profile": "pi5-opt"})) - - should_switch, reason, family = prepare_activation_runtime( - model_filename="test.gguf", - model_format="gguf", - current_family="litert", - device_class="pi5-8gb", - runtimes_dir=runtimes_dir, - ) - # GGUF model on litert should want to switch to llama_cpp. - assert should_switch is True - assert family == "llama_cpp" diff --git a/tests/unit/test_inferno_runtime_manager.py b/tests/unit/test_inferno_runtime_manager.py deleted file mode 100644 index 43688cb..0000000 --- a/tests/unit/test_inferno_runtime_manager.py +++ /dev/null @@ -1,646 +0,0 @@ -"""Tests for core.inferno.runtime_manager — classification, compatibility, slots, markers.""" - -from __future__ import annotations - -import json -from pathlib import Path - -import pytest - -try: - from core.inferno.runtime_manager import ( - SUPPORTED_RUNTIME_FAMILIES, - LLAMA_SERVER_RUNTIME_FAMILIES, - PI4_INCOMPATIBLE_RUNTIMES, - PI4_8GB_MEMORY_THRESHOLD_BYTES, - DEVICE_CLOCK_LIMITS, - MODEL_LOADING_INACTIVE, - LLAMA_RUNTIME_BUNDLE_MARKER_FILENAME, - RuntimeStoreConfig, - classify_runtime_device, - check_runtime_device_compatibility, - get_device_clock_limits, - normalize_llama_memory_loading_mode, - normalize_allow_unsupported_large_models, - llama_memory_loading_no_mmap_env, - compute_model_loading_progress, - discover_runtime_slots, - find_runtime_slot_by_family, - read_llama_runtime_bundle_marker, - write_llama_runtime_bundle_marker, - _detect_installed_runtime_family, - _read_installed_runtime_metadata, - read_llama_runtime_settings, - write_llama_runtime_settings, - build_llama_memory_loading_status, - build_llama_large_model_override_status, - build_large_model_compatibility, - build_llama_runtime_status, - install_llama_runtime_bundle, - ensure_compatible_runtime, - ) -except ModuleNotFoundError: - from inferno.runtime_manager import ( # type: ignore[no-redef] - SUPPORTED_RUNTIME_FAMILIES, - LLAMA_SERVER_RUNTIME_FAMILIES, - PI4_INCOMPATIBLE_RUNTIMES, - PI4_8GB_MEMORY_THRESHOLD_BYTES, - DEVICE_CLOCK_LIMITS, - MODEL_LOADING_INACTIVE, - LLAMA_RUNTIME_BUNDLE_MARKER_FILENAME, - RuntimeStoreConfig, - classify_runtime_device, - check_runtime_device_compatibility, - get_device_clock_limits, - normalize_llama_memory_loading_mode, - normalize_allow_unsupported_large_models, - llama_memory_loading_no_mmap_env, - compute_model_loading_progress, - discover_runtime_slots, - find_runtime_slot_by_family, - read_llama_runtime_bundle_marker, - write_llama_runtime_bundle_marker, - _detect_installed_runtime_family, - _read_installed_runtime_metadata, - read_llama_runtime_settings, - write_llama_runtime_settings, - build_llama_memory_loading_status, - build_llama_large_model_override_status, - build_large_model_compatibility, - build_llama_runtime_status, - install_llama_runtime_bundle, - ensure_compatible_runtime, - ) - - -@pytest.fixture -def store(tmp_path: Path) -> RuntimeStoreConfig: - """Create a RuntimeStoreConfig backed by temp directories.""" - runtimes_dir = tmp_path / "runtimes" - runtimes_dir.mkdir() - install_dir = tmp_path / "llama" - install_dir.mkdir() - settings_path = tmp_path / "state" / "llama_runtime.json" - settings_path.parent.mkdir() - return RuntimeStoreConfig( - runtimes_dir=runtimes_dir, - install_dir=install_dir, - settings_path=settings_path, - device_class="pi5-8gb", - total_memory_bytes=8 * 1024**3, - ) - - -def _make_ik_llama_slot(runtimes_dir: Path) -> Path: - """Create a minimal ik_llama slot for testing.""" - slot = runtimes_dir / "ik_llama" - (slot / "bin").mkdir(parents=True) - (slot / "bin" / "llama-server").write_bytes(b"fake") - (slot / "runtime.json").write_text( - json.dumps({"family": "ik_llama", "commit": "abc123", "profile": "pi5-opt"}), - encoding="utf-8", - ) - return slot - - -def _make_litert_slot(runtimes_dir: Path) -> Path: - """Create a minimal litert slot for testing.""" - slot = runtimes_dir / "litert" - slot.mkdir(parents=True) - (slot / "runtime.json").write_text( - json.dumps({"family": "litert", "version": "1.0"}), - encoding="utf-8", - ) - return slot - - -# -- Constants -- - - -def test_supported_runtime_families(): - assert "ik_llama" in SUPPORTED_RUNTIME_FAMILIES - assert "llama_cpp" in SUPPORTED_RUNTIME_FAMILIES - assert "litert" in SUPPORTED_RUNTIME_FAMILIES - - -def test_llama_server_families(): - assert "ik_llama" in LLAMA_SERVER_RUNTIME_FAMILIES - assert "llama_cpp" in LLAMA_SERVER_RUNTIME_FAMILIES - assert "litert" not in LLAMA_SERVER_RUNTIME_FAMILIES - - -def test_pi4_incompatible_runtimes(): - assert "ik_llama" in PI4_INCOMPATIBLE_RUNTIMES - assert "litert" in PI4_INCOMPATIBLE_RUNTIMES - assert "llama_cpp" not in PI4_INCOMPATIBLE_RUNTIMES - - -# -- Device classification -- - - -def test_classify_pi5_8gb(): - result = classify_runtime_device( - pi_model_name="Raspberry Pi 5 Model B Rev 1.0", - total_memory_bytes=8 * 1024**3, - ) - assert result == "pi5-8gb" - - -def test_classify_pi5_16gb(): - result = classify_runtime_device( - pi_model_name="Raspberry Pi 5 Model B Rev 1.0", - total_memory_bytes=16 * 1024**3, - ) - assert result == "pi5-16gb" - - -def test_classify_pi4_4gb(): - result = classify_runtime_device( - pi_model_name="Raspberry Pi 4 Model B Rev 1.4", - total_memory_bytes=4 * 1024**3, - ) - assert result == "pi4-4gb" - - -def test_classify_pi4_8gb(): - result = classify_runtime_device( - pi_model_name="Raspberry Pi 4 Model B Rev 1.4", - total_memory_bytes=8 * 1024**3, - ) - assert result == "pi4-8gb" - - -def test_classify_unknown_no_model_name(): - assert classify_runtime_device(pi_model_name="", total_memory_bytes=8 * 1024**3) == "unknown" - - -def test_classify_unknown_non_pi(): - assert classify_runtime_device(pi_model_name="Some Board", total_memory_bytes=8 * 1024**3) == "unknown" - - -def test_classify_other_pi(): - assert classify_runtime_device(pi_model_name="Raspberry Pi 3", total_memory_bytes=1 * 1024**3) == "other-pi" - - -# -- Compatibility -- - - -def test_compatible_pi5_ik_llama(): - result = check_runtime_device_compatibility("pi5-8gb", "ik_llama") - assert result["compatible"] is True - - -def test_incompatible_pi4_ik_llama(): - result = check_runtime_device_compatibility("pi4-4gb", "ik_llama") - assert result["compatible"] is False - assert result["recommended_family"] == "llama_cpp" - - -def test_incompatible_pi4_litert(): - result = check_runtime_device_compatibility("pi4-8gb", "litert") - assert result["compatible"] is False - - -def test_compatible_pi4_llama_cpp(): - result = check_runtime_device_compatibility("pi4-4gb", "llama_cpp") - assert result["compatible"] is True - - -# -- Clock limits -- - - -def test_clock_limits_pi5(): - limits = get_device_clock_limits("pi5-8gb") - assert "cpu_max_hz" in limits - assert limits["cpu_max_hz"] == 2_400_000_000 - - -def test_clock_limits_pi4(): - limits = get_device_clock_limits("pi4-4gb") - assert limits["cpu_max_hz"] == 1_800_000_000 - - -def test_clock_limits_unknown_defaults_to_pi5(): - limits = get_device_clock_limits("unknown") - assert limits == get_device_clock_limits("pi5-8gb") - - -# -- Memory loading mode normalization -- - - -@pytest.mark.parametrize( - "raw, expected", - [ - ("full_ram", "full_ram"), - ("no_mmap", "full_ram"), - ("no-mmap", "full_ram"), - ("1", "full_ram"), - ("true", "full_ram"), - ("mmap", "mmap"), - ("mapped", "mmap"), - ("0", "mmap"), - ("false", "mmap"), - ("auto", "auto"), - ("", "auto"), - (None, "auto"), - ("bogus", "auto"), - ], -) -def test_normalize_memory_loading_mode(raw, expected): - assert normalize_llama_memory_loading_mode(raw) == expected - - -@pytest.mark.parametrize( - "mode, expected", - [ - ("full_ram", "1"), - ("mmap", "0"), - ("auto", "auto"), - ], -) -def test_memory_loading_no_mmap_env(mode, expected): - assert llama_memory_loading_no_mmap_env(mode) == expected - - -# -- Large model override normalization -- - - -@pytest.mark.parametrize( - "raw, expected", - [ - (True, True), - (False, False), - (None, False), - ("1", True), - ("true", True), - ("yes", True), - ("0", False), - ("false", False), - ("no", False), - ], -) -def test_normalize_allow_unsupported_large_models(raw, expected): - assert normalize_allow_unsupported_large_models(raw) == expected - - -# -- Model loading progress -- - - -def test_loading_progress_during_boot(): - result = compute_model_loading_progress( - state="BOOTING", - has_model=True, - model_size_bytes=1_000_000, - no_mmap_env="1", - llama_rss={"available": True, "rss_anon_bytes": 500_000}, - ) - assert result["active"] is True - assert result["progress_percent"] == 50 - - -def test_loading_progress_not_booting(): - result = compute_model_loading_progress( - state="IDLE", - has_model=True, - model_size_bytes=1_000_000, - no_mmap_env="1", - llama_rss={"available": True, "rss_anon_bytes": 500_000}, - ) - assert result["active"] is False - - -def test_loading_progress_no_model(): - result = compute_model_loading_progress( - state="BOOTING", - has_model=False, - model_size_bytes=0, - no_mmap_env="1", - llama_rss={"available": True}, - ) - assert result["active"] is False - - -def test_loading_progress_mmap_mode(): - result = compute_model_loading_progress( - state="BOOTING", - has_model=True, - model_size_bytes=1_000_000, - no_mmap_env="0", - llama_rss={"available": True, "rss_file_bytes": 750_000}, - ) - assert result["active"] is True - assert result["progress_percent"] == 75 - - -def test_loading_progress_auto_mode_uses_max(): - result = compute_model_loading_progress( - state="BOOTING", - has_model=True, - model_size_bytes=1_000_000, - no_mmap_env="auto", - llama_rss={"available": True, "rss_anon_bytes": 300_000, "rss_file_bytes": 600_000}, - ) - assert result["active"] is True - assert result["progress_percent"] == 60 - - -def test_loading_progress_caps_at_100(): - result = compute_model_loading_progress( - state="BOOTING", - has_model=True, - model_size_bytes=100, - no_mmap_env="1", - llama_rss={"available": True, "rss_anon_bytes": 200}, - ) - assert result["progress_percent"] == 100 - - -def test_loading_inactive_constant(): - assert MODEL_LOADING_INACTIVE["active"] is False - - -# -- Slot discovery -- - - -def test_discover_runtime_slots_finds_ik_llama(tmp_path): - runtimes_dir = tmp_path / "runtimes" - runtimes_dir.mkdir() - _make_ik_llama_slot(runtimes_dir) - slots = discover_runtime_slots(runtimes_dir) - assert len(slots) == 1 - assert slots[0]["family"] == "ik_llama" - assert slots[0]["commit"] == "abc123" - - -def test_discover_runtime_slots_finds_litert(tmp_path): - runtimes_dir = tmp_path / "runtimes" - runtimes_dir.mkdir() - _make_litert_slot(runtimes_dir) - slots = discover_runtime_slots(runtimes_dir) - assert len(slots) == 1 - assert slots[0]["family"] == "litert" - - -def test_discover_runtime_slots_empty(tmp_path): - runtimes_dir = tmp_path / "runtimes" - runtimes_dir.mkdir() - assert discover_runtime_slots(runtimes_dir) == [] - - -def test_discover_runtime_slots_skips_incomplete_llama(tmp_path): - runtimes_dir = tmp_path / "runtimes" - (runtimes_dir / "ik_llama").mkdir(parents=True) - # No bin/llama-server → should be skipped - assert discover_runtime_slots(runtimes_dir) == [] - - -def test_find_runtime_slot_by_family_found(tmp_path): - runtimes_dir = tmp_path / "runtimes" - runtimes_dir.mkdir() - _make_ik_llama_slot(runtimes_dir) - slot = find_runtime_slot_by_family(runtimes_dir, "ik_llama") - assert slot is not None - assert slot["family"] == "ik_llama" - - -def test_find_runtime_slot_by_family_not_found(tmp_path): - runtimes_dir = tmp_path / "runtimes" - runtimes_dir.mkdir() - assert find_runtime_slot_by_family(runtimes_dir, "llama_cpp") is None - - -# -- Marker management -- - - -def test_write_and_read_marker(tmp_path): - install_dir = tmp_path / "llama" - install_dir.mkdir() - bundle = {"family": "ik_llama", "path": "/tmp/slot", "profile": "pi5-opt", "commit": "abc"} - written = write_llama_runtime_bundle_marker(install_dir, bundle) - assert written["family"] == "ik_llama" - assert written["switched_at_unix"] > 0 - - read_back = read_llama_runtime_bundle_marker(install_dir) - assert read_back is not None - assert read_back["family"] == "ik_llama" - - -def test_read_marker_missing(tmp_path): - install_dir = tmp_path / "llama" - install_dir.mkdir() - assert read_llama_runtime_bundle_marker(install_dir) is None - - -def test_detect_installed_family_from_marker(tmp_path): - install_dir = tmp_path / "llama" - install_dir.mkdir() - write_llama_runtime_bundle_marker(install_dir, {"family": "llama_cpp"}) - assert _detect_installed_runtime_family(install_dir) == "llama_cpp" - - -def test_detect_installed_family_from_runtime_json(tmp_path): - install_dir = tmp_path / "llama" - install_dir.mkdir() - (install_dir / "runtime.json").write_text( - json.dumps({"family": "ik_llama"}), encoding="utf-8" - ) - assert _detect_installed_runtime_family(install_dir) == "ik_llama" - - -def test_detect_installed_family_empty(tmp_path): - install_dir = tmp_path / "llama" - install_dir.mkdir() - assert _detect_installed_runtime_family(install_dir) == "" - - -def test_read_installed_runtime_metadata_prefers_marker(tmp_path): - install_dir = tmp_path / "llama" - install_dir.mkdir() - write_llama_runtime_bundle_marker(install_dir, {"family": "llama_cpp", "commit": "xyz"}) - (install_dir / "runtime.json").write_text( - json.dumps({"family": "ik_llama", "commit": "old"}), encoding="utf-8" - ) - meta = _read_installed_runtime_metadata(install_dir) - assert meta["family"] == "llama_cpp" - - -def test_read_installed_runtime_metadata_falls_back_to_json(tmp_path): - install_dir = tmp_path / "llama" - install_dir.mkdir() - (install_dir / "runtime.json").write_text( - json.dumps({"family": "ik_llama"}), encoding="utf-8" - ) - meta = _read_installed_runtime_metadata(install_dir) - assert meta["family"] == "ik_llama" - - -def test_read_installed_runtime_metadata_empty(tmp_path): - install_dir = tmp_path / "llama" - install_dir.mkdir() - assert _read_installed_runtime_metadata(install_dir) == {} - - -# -- Settings I/O -- - - -def test_read_settings_defaults(store): - settings = read_llama_runtime_settings(store.settings_path) - assert settings["memory_loading_mode"] == "auto" - assert settings["allow_unsupported_large_models"] is False - - -def test_write_and_read_settings(store): - write_llama_runtime_settings(store.settings_path, memory_loading_mode="full_ram") - settings = read_llama_runtime_settings(store.settings_path) - assert settings["memory_loading_mode"] == "full_ram" - assert settings["updated_at_unix"] is not None - - -def test_write_settings_preserves_unset_fields(store): - write_llama_runtime_settings(store.settings_path, memory_loading_mode="mmap") - write_llama_runtime_settings(store.settings_path, allow_unsupported_large_models=True) - settings = read_llama_runtime_settings(store.settings_path) - assert settings["memory_loading_mode"] == "mmap" - assert settings["allow_unsupported_large_models"] is True - - -def test_write_settings_passes_through_power_calibration(store): - cal = {"mode": "custom", "a": 1.5, "b": 0.3} - write_llama_runtime_settings(store.settings_path, power_calibration=cal) - settings = read_llama_runtime_settings(store.settings_path) - assert settings["power_calibration"]["mode"] == "custom" - assert settings["power_calibration"]["a"] == 1.5 - - -# -- Status builders -- - - -def test_memory_loading_status_auto(store): - status = build_llama_memory_loading_status(store.settings_path) - assert status["mode"] == "auto" - assert status["no_mmap_env"] == "auto" - assert "Automatic" in status["label"] - - -def test_memory_loading_status_full_ram(store): - write_llama_runtime_settings(store.settings_path, memory_loading_mode="full_ram") - status = build_llama_memory_loading_status(store.settings_path) - assert status["mode"] == "full_ram" - assert status["no_mmap_env"] == "1" - - -def test_large_model_override_status_default(store): - status = build_llama_large_model_override_status(store.settings_path) - assert status["enabled"] is False - assert "default" in status["label"] - - -def test_large_model_override_status_enabled(store): - write_llama_runtime_settings(store.settings_path, allow_unsupported_large_models=True) - status = build_llama_large_model_override_status(store.settings_path) - assert status["enabled"] is True - - -def test_large_model_compatibility_no_warnings(store): - result = build_large_model_compatibility( - store, - model_filename="small.gguf", - model_size_bytes=100, - threshold_bytes=5_000_000_000, - storage_free_bytes=10_000_000_000, - ) - assert result["device_class"] == "pi5-8gb" - assert result["warnings"] == [] - - -def test_large_model_compatibility_with_warning(store): - result = build_large_model_compatibility( - store, - model_filename="huge.gguf", - model_size_bytes=10_000_000_000, - threshold_bytes=5_000_000_000, - storage_free_bytes=20_000_000_000, - ) - assert len(result["warnings"]) == 1 - assert result["warnings"][0]["code"] == "large_model_unsupported_pi_warning" - - -def test_large_model_compatibility_no_warning_on_16gb_pi5(): - store_16gb = RuntimeStoreConfig( - runtimes_dir=Path("/tmp"), - install_dir=Path("/tmp"), - settings_path=Path("/tmp/nonexistent.json"), - device_class="pi5-16gb", - total_memory_bytes=16 * 1024**3, - ) - result = build_large_model_compatibility( - store_16gb, - model_filename="huge.gguf", - model_size_bytes=10_000_000_000, - threshold_bytes=5_000_000_000, - ) - assert result["warnings"] == [] - - -def test_runtime_status_basic(store): - _make_ik_llama_slot(store.runtimes_dir) - write_llama_runtime_bundle_marker(store.install_dir, {"family": "ik_llama", "commit": "abc"}) - status = build_llama_runtime_status(store, active_model_filename="model.gguf") - assert status["current"]["family"] == "ik_llama" - assert len(status["available_runtimes"]) >= 1 - assert "memory_loading" in status - assert "large_model_override" in status - assert status["switch"]["active"] is False - - -def test_runtime_status_with_switch_snapshot(store): - _make_ik_llama_slot(store.runtimes_dir) - snap = {"active": True, "target_family": "llama_cpp"} - status = build_llama_runtime_status(store, switch_snapshot=snap) - assert status["switch"]["active"] is True - assert status["switch"]["target_family"] == "llama_cpp" - - -# -- Async: install_llama_runtime_bundle -- - - -@pytest.mark.anyio -async def test_install_litert_skips_rsync(tmp_path): - install_dir = tmp_path / "llama" - install_dir.mkdir() - bundle_dir = tmp_path / "bundle" - bundle_dir.mkdir() - (bundle_dir / "runtime.json").write_text(json.dumps({"family": "litert"}), encoding="utf-8") - result = await install_llama_runtime_bundle(install_dir, bundle_dir) - assert result["ok"] is True - assert result["reason"] == "litert_no_rsync_needed" - - -# -- Async: ensure_compatible_runtime -- - - -@pytest.mark.anyio -async def test_ensure_compatible_runtime_already_compatible(store): - write_llama_runtime_bundle_marker(store.install_dir, {"family": "ik_llama"}) - switched, reason = await ensure_compatible_runtime(store) - assert switched is False - assert reason == "compatible" - - -@pytest.mark.anyio -async def test_ensure_compatible_runtime_no_slot_available(tmp_path): - install_dir = tmp_path / "llama" - install_dir.mkdir() - runtimes_dir = tmp_path / "runtimes" - runtimes_dir.mkdir() - write_llama_runtime_bundle_marker(install_dir, {"family": "ik_llama"}) - pi4_store = RuntimeStoreConfig( - runtimes_dir=runtimes_dir, - install_dir=install_dir, - settings_path=tmp_path / "settings.json", - device_class="pi4-4gb", - total_memory_bytes=4 * 1024**3, - ) - switched, reason = await ensure_compatible_runtime(pi4_store) - assert switched is False - assert reason == "slot_unavailable" diff --git a/tests/unit/test_launch_config.py b/tests/unit/test_launch_config.py deleted file mode 100644 index 3c24d2a..0000000 --- a/tests/unit/test_launch_config.py +++ /dev/null @@ -1,296 +0,0 @@ -"""Unit tests for core.inferno.launch_config — llama-server CLI arg builder.""" - -from __future__ import annotations - -import pytest - -from core.inferno.launch_config import build_llama_server_args - - -# --------------------------------------------------------------------------- -# Basic arg construction -# --------------------------------------------------------------------------- - - -def test_basic_args_include_all_mandatory_flags(): - args = build_llama_server_args( - llama_server_bin="/opt/potato/llama/bin/llama-server", - model_path="/opt/potato/models/Qwen3.5-2B-Q4_K_M.gguf", - slot_save_path="/opt/potato/state/llama-slots", - ) - assert args[0] == "/opt/potato/llama/bin/llama-server" - assert "--model" in args - assert "/opt/potato/models/Qwen3.5-2B-Q4_K_M.gguf" in args - assert "--host" in args - assert "--port" in args - assert "--ctx-size" in args - assert "--cache-ram" in args - assert "--parallel" in args - assert "--slot-save-path" in args - - -def test_defaults_match_shell_script_defaults(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - ) - idx = args.index - assert args[idx("--host") + 1] == "0.0.0.0" - assert args[idx("--port") + 1] == "8080" - assert args[idx("--ctx-size") + 1] == "16384" - assert args[idx("--cache-ram") + 1] == "1024" - assert args[idx("--parallel") + 1] == "1" - - -# --------------------------------------------------------------------------- -# Vision projector (mmproj) -# --------------------------------------------------------------------------- - - -def test_mmproj_path_adds_flag(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - mmproj_path="/models/mmproj-F16.gguf", - ) - assert "--mmproj" in args - assert "/models/mmproj-F16.gguf" in args - - -def test_no_mmproj_omits_flag(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - ) - assert "--mmproj" not in args - - -# --------------------------------------------------------------------------- -# KV cache configuration -# --------------------------------------------------------------------------- - - -def test_kv_cache_defaults(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - ) - assert "--cache-type-k" in args - assert "--cache-type-v" in args - idx = args.index - assert args[idx("--cache-type-k") + 1] == "q8_0" - assert args[idx("--cache-type-v") + 1] == "q8_0" - - -def test_custom_kv_flags_override_defaults(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - kv_flags="--cache-type-k f16 --cache-type-v f16", - ) - assert "--cache-type-k" in args - assert "f16" in args - # Should NOT have the default q8_0 values - assert "q8_0" not in args - - -# --------------------------------------------------------------------------- -# Feature toggles -# --------------------------------------------------------------------------- - - -def test_flash_attn_enabled(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - flash_attn=True, - ) - assert "--flash-attn" in args - assert args[args.index("--flash-attn") + 1] == "on" - - -def test_flash_attn_disabled(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - flash_attn=False, - ) - assert "--flash-attn" not in args - - -def test_jinja_enabled(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - jinja=True, - ) - assert "--jinja" in args - - -def test_jinja_disabled(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - jinja=False, - ) - assert "--jinja" not in args - - -def test_no_warmup_enabled(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - no_warmup=True, - ) - assert "--no-warmup" in args - - -def test_no_warmup_disabled(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - no_warmup=False, - ) - assert "--no-warmup" not in args - - -def test_no_mmap_enabled(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - no_mmap=True, - ) - assert "--no-mmap" in args - - -def test_no_mmap_disabled(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - no_mmap=False, - ) - assert "--no-mmap" not in args - - -# --------------------------------------------------------------------------- -# Reasoning and chat template -# --------------------------------------------------------------------------- - - -def test_reasoning_format_default(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - ) - assert "--reasoning-format" in args - assert args[args.index("--reasoning-format") + 1] == "none" - - -def test_reasoning_format_custom(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - reasoning_format="deepseek", - ) - assert args[args.index("--reasoning-format") + 1] == "deepseek" - - -def test_chat_template_kwargs_default(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - ) - assert "--chat-template-kwargs" in args - assert '{"enable_thinking": false}' in args - - -def test_chat_template_kwargs_custom(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - chat_template_kwargs='{"enable_thinking": true}', - ) - assert args[args.index("--chat-template-kwargs") + 1] == '{"enable_thinking": true}' - - -# --------------------------------------------------------------------------- -# Runtime family — WebUI suppression -# --------------------------------------------------------------------------- - - -def test_ik_llama_uses_webui_none(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - runtime_family="ik_llama", - ) - assert "--webui" in args - assert args[args.index("--webui") + 1] == "none" - assert "--no-webui" not in args - - -def test_llama_cpp_uses_no_webui(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - runtime_family="llama_cpp", - ) - assert "--no-webui" in args - assert "--webui" not in args or args[args.index("--webui") + 1] != "none" - - -def test_unknown_family_no_webui_flag(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - runtime_family=None, - ) - assert "--webui" not in args - assert "--no-webui" not in args - - -# --------------------------------------------------------------------------- -# Extra flags passthrough -# --------------------------------------------------------------------------- - - -def test_extra_flags_split_and_appended(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - extra_flags="--verbose --log-timestamps", - ) - assert "--verbose" in args - assert "--log-timestamps" in args - - -def test_empty_extra_flags_ignored(): - args = build_llama_server_args( - llama_server_bin="/bin/llama-server", - model_path="/model.gguf", - slot_save_path="/slots", - extra_flags="", - ) - # Should not crash or add empty strings - assert "" not in args[1:] # first element is the binary path diff --git a/tests/unit/test_litert_adapter.py b/tests/unit/test_litert_adapter.py deleted file mode 100644 index b001564..0000000 --- a/tests/unit/test_litert_adapter.py +++ /dev/null @@ -1,515 +0,0 @@ -"""Tests for the LiteRT adapter with mocked litert_lm engine.""" - -from __future__ import annotations - -import sys -import types - -import pytest -from fastapi.testclient import TestClient - - -class _FakeConversation: - """Mock LiteRT-LM Conversation with send_message / send_message_async.""" - - def __init__(self): - self._history: list = [] - - def __enter__(self): - return self - - def __exit__(self, *_args): - pass - - def send_message(self, msg) -> dict: - self._history.append(msg) - label = str(msg)[:20] if isinstance(msg, str) else "multimodal" - return {"content": [{"type": "text", "text": f"Reply to: {label}"}]} - - def send_message_async(self, msg): - self._history.append(msg) - label = str(msg)[:10] if isinstance(msg, str) else "image" - for word in f"Reply to {label}".split(): - yield {"content": [{"type": "text", "text": word + " "}]} - - -class _FakeEngine: - """Mock LiteRT-LM Engine.""" - - def __init__(self, model_path: str = "", backend=None): - self.model_path = model_path - - def __enter__(self): - return self - - def __exit__(self, *_args): - pass - - def create_conversation(self, **kwargs): - return _FakeConversation() - - -class _FakeBackend: - CPU = "cpu" - - -@pytest.fixture(autouse=True) -def _mock_litert_lm(monkeypatch): - """Inject a fake litert_lm module before importing the adapter.""" - fake_module = types.ModuleType("litert_lm") - fake_module.Engine = _FakeEngine # type: ignore[attr-defined] - fake_module.Backend = _FakeBackend # type: ignore[attr-defined] - monkeypatch.setitem(sys.modules, "litert_lm", fake_module) - - # Force re-import of the adapter so it picks up the mock - for key in list(sys.modules): - if "litert_adapter" in key: - del sys.modules[key] - - import core.inferno.litert_adapter as adapter - adapter.litert_lm = fake_module # type: ignore[attr-defined] - adapter._engine = _FakeEngine("test.litertlm", backend=_FakeBackend.CPU) - adapter._conversation = _FakeConversation() - adapter._conversation.__enter__() - adapter._conversation_history = [] - yield adapter - adapter._engine = None - adapter._conversation = None - adapter._conversation_history = [] - - -@pytest.fixture -def client(_mock_litert_lm): - with TestClient(_mock_litert_lm.app) as c: - yield c - - -def test_health_ok_when_engine_loaded(client): - response = client.get("/health") - assert response.status_code == 200 - assert response.json()["status"] == "ok" - - -def test_health_503_when_no_engine(_mock_litert_lm, client): - _mock_litert_lm._engine = None - response = client.get("/health") - assert response.status_code == 503 - - -def test_chat_completion_non_streaming_openai_format(client): - response = client.post( - "/v1/chat/completions", - json={"messages": [{"role": "user", "content": "Hi"}], "stream": False}, - ) - assert response.status_code == 200 - body = response.json() - assert body["object"] == "chat.completion" - assert len(body["choices"]) == 1 - assert body["choices"][0]["message"]["role"] == "assistant" - assert body["choices"][0]["message"]["content"] # non-empty - assert body["choices"][0]["finish_reason"] == "stop" - assert "usage" in body - - -def test_chat_completion_streaming_sse_chunks(client): - response = client.post( - "/v1/chat/completions", - json={"messages": [{"role": "user", "content": "Hi"}], "stream": True}, - ) - assert response.status_code == 200 - assert "text/event-stream" in response.headers.get("content-type", "") - text = response.text - assert "data:" in text - assert "[DONE]" in text - - -def test_chat_completion_replays_user_turns_on_divergence(_mock_litert_lm, client): - """On history divergence, user turns are replayed (assistant turns skipped).""" - conv = _mock_litert_lm._conversation - response = client.post( - "/v1/chat/completions", - json={ - "messages": [ - {"role": "user", "content": "First"}, - {"role": "assistant", "content": "Response 1"}, - {"role": "user", "content": "Second"}, - ], - "stream": False, - }, - ) - assert response.status_code == 200 - # Conversation should have received "First" (replay) + "Second" (final). - # "Response 1" (assistant) should NOT have been sent through send_message. - history = _mock_litert_lm._conversation._history - assert "First" in history - assert "Second" in history - # No assistant content should appear in conversation history - assert not any("Response 1" in h for h in history) - - -def test_chat_completion_continuation_reuses_conversation(_mock_litert_lm, client): - """Second turn with matching history should NOT reset the conversation.""" - # First turn - client.post( - "/v1/chat/completions", - json={"messages": [{"role": "user", "content": "Hello"}], "stream": False}, - ) - # After first turn, history should have the exchange - assert len(_mock_litert_lm._conversation_history) == 2 # user + assistant - - # Capture the conversation object - conv_before = _mock_litert_lm._conversation - - # Second turn — includes first turn in history (continuation) - resp = client.post( - "/v1/chat/completions", - json={ - "messages": [ - {"role": "user", "content": "Hello"}, - {"role": "assistant", "content": _mock_litert_lm._conversation_history[1]["content"]}, - {"role": "user", "content": "Follow up"}, - ], - "stream": False, - }, - ) - assert resp.status_code == 200 - # Same conversation object — no reset - assert _mock_litert_lm._conversation is conv_before - - -def test_chat_completion_new_session_resets_conversation(_mock_litert_lm, client): - """A completely different message history should reset the conversation.""" - # First turn - client.post( - "/v1/chat/completions", - json={"messages": [{"role": "user", "content": "Hello"}], "stream": False}, - ) - conv_before = _mock_litert_lm._conversation - - # New session — different first message - resp = client.post( - "/v1/chat/completions", - json={"messages": [{"role": "user", "content": "Totally different"}], "stream": False}, - ) - assert resp.status_code == 200 - # Conversation was reset — new object - assert _mock_litert_lm._conversation is not conv_before - - -def test_chat_completion_handles_system_prompt(client): - response = client.post( - "/v1/chat/completions", - json={ - "messages": [ - {"role": "system", "content": "You are helpful"}, - {"role": "user", "content": "Hi"}, - ], - "stream": False, - }, - ) - assert response.status_code == 200 - body = response.json() - assert body["choices"][0]["message"]["content"] - - -def test_chat_completion_returns_500_on_engine_error(_mock_litert_lm, client): - """If inference fails, adapter returns 500.""" - class _BrokenConversation: - def __enter__(self): - return self - def __exit__(self, *_args): - pass - def send_message(self, msg): - raise RuntimeError("Engine crashed") - - class _BrokenEngine: - def create_conversation(self, **kwargs): - return _BrokenConversation() - def __exit__(self, *_args): - pass - - _mock_litert_lm._engine = _BrokenEngine() - _mock_litert_lm._conversation = _BrokenConversation() - _mock_litert_lm._conversation_history = [] - response = client.post( - "/v1/chat/completions", - json={"messages": [{"role": "user", "content": "Hi"}], "stream": False}, - ) - assert response.status_code == 500 - assert "error" in response.json() - - -def test_chat_completion_returns_503_when_no_engine(_mock_litert_lm, client): - _mock_litert_lm._engine = None - response = client.post( - "/v1/chat/completions", - json={"messages": [{"role": "user", "content": "Hi"}]}, - ) - assert response.status_code == 503 - - -def test_chat_completion_rejects_empty_messages(client): - response = client.post( - "/v1/chat/completions", - json={"messages": []}, - ) - assert response.status_code == 400 - - -# -- LiteRT vision support tests ------------------------------------------------ - - -def test_health_reports_vision_false_by_default(client): - """When engine loaded without vision, /health reports vision=false.""" - response = client.get("/health") - assert response.status_code == 200 - body = response.json() - assert body["status"] == "ok" - assert body["vision"] is False - - -def test_health_reports_vision_true_when_probe_succeeds(_mock_litert_lm, client): - """When vision probe succeeded, /health reports vision=true.""" - _mock_litert_lm._vision_enabled = True - response = client.get("/health") - assert response.status_code == 200 - body = response.json() - assert body["vision"] is True - - -def test_vision_probe_falls_back_to_text_only_engine(_mock_litert_lm): - """_probe_vision_support catches failures and returns (engine, False).""" - # Make Engine raise TypeError when vision_backend is passed (simulates - # litert_lm version that doesn't support the kwarg). - original_engine_cls = _mock_litert_lm.litert_lm.Engine - - class _VisionUnsupportedEngine: - def __init__(self, model_path="", backend=None, **kwargs): - if "vision_backend" in kwargs: - raise TypeError("unexpected keyword argument 'vision_backend'") - self.model_path = model_path - - def __enter__(self): - return self - - def __exit__(self, *_args): - pass - - def create_conversation(self, **kwargs): - return _FakeConversation() - - _mock_litert_lm.litert_lm.Engine = _VisionUnsupportedEngine - try: - engine, vision = _mock_litert_lm._probe_vision_support("test.litertlm") - assert engine is not None - assert vision is False - finally: - _mock_litert_lm.litert_lm.Engine = original_engine_cls - - -def test_vision_probe_succeeds_when_engine_accepts_vision_backend(_mock_litert_lm): - """_probe_vision_support returns (engine, True) when Engine accepts vision_backend.""" - - class _VisionEngine: - def __init__(self, model_path="", backend=None, vision_backend=None): - self.model_path = model_path - self.vision_backend = vision_backend - - def __enter__(self): - return self - - def __exit__(self, *_args): - pass - - def create_conversation(self, **kwargs): - return _FakeConversation() - - _mock_litert_lm.litert_lm.Engine = _VisionEngine - try: - engine, vision = _mock_litert_lm._probe_vision_support("test.litertlm") - assert engine is not None - assert vision is True - finally: - _mock_litert_lm.litert_lm.Engine = _FakeEngine - - -def test_multimodal_message_converted_to_litert_format(_mock_litert_lm): - """OpenAI image_url content parts are converted to litert-lm blob format.""" - openai_content = [ - {"type": "text", "text": "Describe this image:"}, - {"type": "image_url", "image_url": {"url": "data:image/jpeg;base64,AAAA"}}, - ] - result = _mock_litert_lm._convert_openai_to_litert_content(openai_content) - assert isinstance(result, list) - assert result[0] == {"type": "text", "text": "Describe this image:"} - assert result[1] == {"type": "image", "blob": "AAAA"} - - -def test_multimodal_inference_sends_dict_message_to_engine(_mock_litert_lm, client): - """Multimodal content is wrapped as {role, content} dict for send_message.""" - _mock_litert_lm._vision_enabled = True - client.post( - "/v1/chat/completions", - json={ - "messages": [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What is this?"}, - {"type": "image_url", "image_url": {"url": "data:image/jpeg;base64,AAAA"}}, - ], - } - ], - "stream": False, - }, - ) - # The last send_message call should receive a dict with role/content - last_sent = _mock_litert_lm._conversation._history[-1] - assert isinstance(last_sent, dict) - assert last_sent["role"] == "user" - assert isinstance(last_sent["content"], list) - assert last_sent["content"][1] == {"type": "image", "blob": "AAAA"} - - -def test_base64_extracted_from_data_url(_mock_litert_lm): - """data:image/png;base64, is correctly split to extract the blob.""" - content = [ - {"type": "image_url", "image_url": {"url": "data:image/png;base64,iVBORw0KGgo="}}, - ] - result = _mock_litert_lm._convert_openai_to_litert_content(content) - assert result[0]["blob"] == "iVBORw0KGgo=" - - -def test_remote_image_url_rejected(_mock_litert_lm): - """Remote https:// image URLs are rejected with ValueError.""" - content = [ - {"type": "image_url", "image_url": {"url": "https://example.com/photo.jpg"}}, - ] - with pytest.raises(ValueError, match="Only base64 data URLs"): - _mock_litert_lm._convert_openai_to_litert_content(content) - - -def test_structured_text_only_content_accepted_without_vision(_mock_litert_lm, client): - """Text-only structured content [{"type":"text","text":"hi"}] must not be rejected.""" - _mock_litert_lm._vision_enabled = False - response = client.post( - "/v1/chat/completions", - json={ - "messages": [ - { - "role": "user", - "content": [{"type": "text", "text": "Hello"}], - } - ], - "stream": False, - }, - ) - assert response.status_code == 200 - - -def test_text_only_content_passes_through_unchanged(_mock_litert_lm): - """Plain string content is returned as-is.""" - result = _mock_litert_lm._convert_openai_to_litert_content("Hello, world!") - assert result == "Hello, world!" - - -def test_multimodal_message_rejected_when_vision_disabled(_mock_litert_lm, client): - """When vision is not available, multimodal messages return 400.""" - _mock_litert_lm._vision_enabled = False - response = client.post( - "/v1/chat/completions", - json={ - "messages": [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What is this?"}, - {"type": "image_url", "image_url": {"url": "data:image/jpeg;base64,AAAA"}}, - ], - } - ], - "stream": False, - }, - ) - assert response.status_code == 400 - assert "vision" in response.json()["error"]["message"].lower() - - -def test_multimodal_message_accepted_when_vision_enabled(_mock_litert_lm, client): - """When vision is enabled, multimodal messages are accepted.""" - _mock_litert_lm._vision_enabled = True - response = client.post( - "/v1/chat/completions", - json={ - "messages": [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What is this?"}, - {"type": "image_url", "image_url": {"url": "data:image/jpeg;base64,AAAA"}}, - ], - } - ], - "stream": False, - }, - ) - assert response.status_code == 200 - - -def test_messages_match_handles_multimodal_content(_mock_litert_lm): - """_messages_match works correctly with content-as-list.""" - multimodal_msg = { - "role": "user", - "content": [ - {"type": "text", "text": "What is this?"}, - {"type": "image_url", "image_url": {"url": "data:image/jpeg;base64,AAAA"}}, - ], - } - history = [multimodal_msg, {"role": "assistant", "content": "It's a photo."}] - incoming = [ - multimodal_msg, - {"role": "assistant", "content": "It's a photo."}, - {"role": "user", "content": "Tell me more."}, - ] - assert _mock_litert_lm._messages_match(incoming, history) is True - - -def test_conversation_history_tracks_multimodal_messages(_mock_litert_lm, client): - """After a multimodal turn, history stores content list correctly.""" - _mock_litert_lm._vision_enabled = True - multimodal_msg = [ - {"type": "text", "text": "What is this?"}, - {"type": "image_url", "image_url": {"url": "data:image/jpeg;base64,AAAA"}}, - ] - client.post( - "/v1/chat/completions", - json={ - "messages": [{"role": "user", "content": multimodal_msg}], - "stream": False, - }, - ) - assert len(_mock_litert_lm._conversation_history) == 2 - assert isinstance(_mock_litert_lm._conversation_history[0]["content"], list) - - -def test_prompt_token_estimate_handles_multimodal_content(_mock_litert_lm, client): - """Token estimation doesn't crash when content is a list of parts.""" - _mock_litert_lm._vision_enabled = True - response = client.post( - "/v1/chat/completions", - json={ - "messages": [ - { - "role": "user", - "content": [ - {"type": "text", "text": "Describe this."}, - {"type": "image_url", "image_url": {"url": "data:image/jpeg;base64,AAAA"}}, - ], - } - ], - "stream": False, - }, - ) - assert response.status_code == 200 - body = response.json() - assert body["usage"]["prompt_tokens"] >= 0 diff --git a/tests/unit/test_model_families.py b/tests/unit/test_model_families.py deleted file mode 100644 index 5108c32..0000000 --- a/tests/unit/test_model_families.py +++ /dev/null @@ -1,271 +0,0 @@ -from __future__ import annotations - -import pytest - -from core.inferno.model_families import ( - build_model_projector_status, - default_projector_candidates_for_model, - is_gemma4_filename, - projector_repo_for_model, - recommended_runtime_for_model, -) - - -def test_projector_repo_for_qwen35_9b(): - """Qwen3.5-9B generic model resolves to unsloth 9B projector repo.""" - assert projector_repo_for_model("Qwen3.5-9B-Q4_K_S-3.92bpw.gguf") == "unsloth/Qwen3.5-9B-GGUF" - - -def test_projector_repo_for_qwen35_9b_byteshape_by_filename(): - """ByteShape in filename resolves to ByteShape repo.""" - assert projector_repo_for_model("byteshape-Qwen3.5-9B-Q4_K_S.gguf") == "byteshape/Qwen3.5-9B-GGUF" - - -def test_projector_repo_for_qwen35_9b_byteshape_by_source_url(): - """ByteShape in source_url resolves to ByteShape repo even when filename - has no publisher prefix — this is the real HuggingFace scenario.""" - assert projector_repo_for_model( - "Qwen3.5-9B-Q4_K_S-3.92bpw.gguf", - source_url="https://huggingface.co/byteshape/Qwen3.5-9B-GGUF/resolve/main/Qwen3.5-9B-Q4_K_S-3.92bpw.gguf", - ) == "byteshape/Qwen3.5-9B-GGUF" - - -def test_projector_repo_for_qwen35_9b_without_source_url(): - """Generic 9B without source_url falls back to unsloth.""" - assert projector_repo_for_model("Qwen3.5-9B-Q4_K_S-3.92bpw.gguf") == "unsloth/Qwen3.5-9B-GGUF" - - -def test_projector_repo_returns_none_for_unknown_qwen35_size(): - """Unrecognized Qwen3.5 sizes must return None, not silently fall back.""" - assert projector_repo_for_model("Qwen3.5-7B-Q4_K_M.gguf") is None - - -def test_projector_repo_no_substring_collision_19b(): - """19B must NOT match the 9B rule — '9b' is a substring of '19b'.""" - assert projector_repo_for_model("Qwen3.5-Creative-19B-A3B-REAP.Q4_K_S.gguf") is None - - -def test_projector_repo_no_substring_collision_bpw(): - """3.92bpw must NOT match the 2B rule — '2b' is a substring of '3.92bpw'.""" - result = projector_repo_for_model("Qwen3.5-9B-Q4_K_S-3.92bpw.gguf") - assert result == "unsloth/Qwen3.5-9B-GGUF", f"got {result!r} (2B collision via '3.92bpw')" - - -def test_projector_repo_dot_delimited_9b(): - """Dot-delimited filenames like Qwen3.5-9B.gguf must still match.""" - assert projector_repo_for_model("Qwen3.5-9B.gguf") == "unsloth/Qwen3.5-9B-GGUF" - assert projector_repo_for_model("Qwen3.5-9B.Q4_K_S.gguf") == "unsloth/Qwen3.5-9B-GGUF" - - -def test_projector_repo_dot_delimited_4b(): - """Dot-delimited filenames like Qwen3.5-4B.Q4_K_M.gguf must still match.""" - assert projector_repo_for_model("Qwen3.5-4B.Q4_K_M.gguf") == "unsloth/Qwen3.5-4B-GGUF" - - -def test_default_candidates_include_model_specific_bf16_but_not_generic(): - """default_projector_candidates_for_model must include model-specific bf16 - names but NOT generic mmproj-bf16.gguf (unsafe cross-model reuse).""" - candidates = default_projector_candidates_for_model("Qwen3.5-9B-Q4_K_S.gguf") - assert "mmproj-F16.gguf" in candidates - assert "mmproj-Qwen3.5-9B-f16.gguf" in candidates - assert "mmproj-Qwen3.5-9B-bf16.gguf" in candidates - # Generic bf16 must NOT be in the list — it's unsafe for cross-model reuse - assert "mmproj-bf16.gguf" not in candidates - - -def test_projector_status_finds_model_specific_bf16_on_disk(runtime): - """build_model_projector_status must detect a model-specific bf16 projector.""" - models_dir = runtime.base_dir / "models" - (models_dir / "mmproj-Qwen3.5-9B-bf16.gguf").write_bytes(b"bf16-projector") - - model = { - "filename": "Qwen3.5-9B-Q4_K_S.gguf", - "settings": { - "vision": {"enabled": True, "projector_mode": "default", "projector_filename": None}, - }, - } - status = build_model_projector_status(models_dir, model) - assert status["present"] is True - assert "bf16" in status["filename"] - - -def test_projector_status_ignores_stale_generic_bf16_for_wrong_model(runtime): - """A generic mmproj-bf16.gguf left over from a 9B download must NOT be - reported as present for a 4B model — wrong dimensions would crash. #136.""" - models_dir = runtime.base_dir / "models" - (models_dir / "mmproj-bf16.gguf").write_bytes(b"stale-9b-bf16") - - model = { - "filename": "Qwen3.5-4B-Q4_K_M.gguf", - "settings": { - "vision": {"enabled": True, "projector_mode": "default", "projector_filename": None}, - }, - } - status = build_model_projector_status(models_dir, model) - assert status["present"] is False, ( - "Generic mmproj-bf16.gguf must not be accepted for a different model size" - ) - - -@pytest.mark.parametrize( - "filename, expected_repo", - [ - ("Qwen3.5-2B-Q4_K_M.gguf", "unsloth/Qwen3.5-2B-GGUF"), - ("Qwen3.5-4B-Q4_K_M.gguf", "unsloth/Qwen3.5-4B-GGUF"), - ("Qwen3.5-0.8B-Q4_K_M.gguf", "unsloth/Qwen3.5-0.8B-GGUF"), - ("Qwen3.5-35B-A3B-Q2_K_L.gguf", "AesSedai/Qwen3.5-35B-A3B-GGUF"), - ], -) -def test_projector_repo_existing_sizes_unchanged(filename: str, expected_repo: str): - """Existing size mappings must not regress.""" - assert projector_repo_for_model(filename) == expected_repo - - -# ── Gemma 4 detection ────────────────────────────────────────────── - - -@pytest.mark.parametrize( - "filename", - [ - "gemma-4-E2B-it-Q4_K_M.gguf", - "gemma-4-E2B-it-UD-Q4_K_XL.gguf", - "gemma-4-E4B-it-Q4_0.gguf", - "gemma-4-E4B-it-Q8_0.gguf", - "gemma-4-26B-A4B-it-UD-IQ2_M.gguf", - "gemma-4-26B-A4B-it-UD-Q4_K_XL.gguf", - ], -) -def test_is_gemma4_filename_positive(filename: str): - """All Gemma 4 variant filenames are detected.""" - assert is_gemma4_filename(filename) is True - - -@pytest.mark.parametrize( - "filename", - [ - "gemma-3-9b-it-Q4_K_M.gguf", - "Qwen3.5-4B-Q4_K_M.gguf", - "llama-4-scout-Q4_K_M.gguf", - None, - "", - ], -) -def test_is_gemma4_filename_negative(filename): - """Non-Gemma-4 filenames must not match.""" - assert is_gemma4_filename(filename) is False - - -# ── Gemma 4 projector repo resolution ────────────────────────────── - - -@pytest.mark.parametrize( - "filename, expected_repo", - [ - ("gemma-4-E2B-it-Q4_K_M.gguf", "unsloth/gemma-4-E2B-it-GGUF"), - ("gemma-4-E2B-it-UD-Q4_K_XL.gguf", "unsloth/gemma-4-E2B-it-GGUF"), - ("gemma-4-E2B-it-Q8_0.gguf", "unsloth/gemma-4-E2B-it-GGUF"), - ("gemma-4-E2B-it-IQ4_NL.gguf", "unsloth/gemma-4-E2B-it-GGUF"), - ("gemma-4-E4B-it-Q4_0.gguf", "unsloth/gemma-4-E4B-it-GGUF"), - ("gemma-4-E4B-it-Q8_0.gguf", "unsloth/gemma-4-E4B-it-GGUF"), - ("gemma-4-26B-A4B-it-UD-IQ2_M.gguf", "unsloth/gemma-4-26B-A4B-it-GGUF"), - ("gemma-4-26B-A4B-it-UD-Q4_K_XL.gguf", "unsloth/gemma-4-26B-A4B-it-GGUF"), - ], -) -def test_projector_repo_for_gemma4(filename: str, expected_repo: str): - """Each Gemma 4 variant resolves to its Unsloth projector repo.""" - assert projector_repo_for_model(filename) == expected_repo - - -def test_projector_repo_returns_none_for_unknown_gemma4_variant(): - """Unrecognized Gemma 4 variants must return None.""" - assert projector_repo_for_model("gemma-4-99B-it-Q4_K_M.gguf") is None - - -def test_projector_repo_gemma4_no_collision_with_quant_4(): - """'4' in Q4_K_M must not trigger Gemma 4 detection on non-Gemma models.""" - assert projector_repo_for_model("some-model-Q4_K_M.gguf") is None - - -def test_projector_repo_gemma4_does_not_break_qwen35(): - """Qwen3.5 repos must still resolve correctly with Gemma 4 code present.""" - assert projector_repo_for_model("Qwen3.5-2B-Q4_K_M.gguf") == "unsloth/Qwen3.5-2B-GGUF" - assert projector_repo_for_model("Qwen3.5-9B-Q4_K_S.gguf") == "unsloth/Qwen3.5-9B-GGUF" - - -# ── Gemma 4 projector candidates ─────────────────────────────────── - - -def test_default_candidates_gemma4_e2b(): - """Gemma 4 E2B produces model-specific f16/bf16 then generic F16.""" - candidates = default_projector_candidates_for_model("gemma-4-E2B-it-Q4_K_M.gguf") - assert len(candidates) > 0 - assert "mmproj-gemma-4-E2B-it-f16.gguf" in candidates - assert "mmproj-gemma-4-E2B-it-bf16.gguf" in candidates - assert "mmproj-F16.gguf" in candidates - # Generic bf16 must NOT be in the list - assert "mmproj-bf16.gguf" not in candidates - # f16 model-specific should come before generic - assert candidates.index("mmproj-gemma-4-E2B-it-f16.gguf") < candidates.index("mmproj-F16.gguf") - - -def test_default_candidates_gemma4_26b_a4b(): - """Gemma 4 26B-A4B stem-trimming strips the UD-IQ2_M quant suffix.""" - candidates = default_projector_candidates_for_model("gemma-4-26B-A4B-it-UD-IQ2_M.gguf") - assert "mmproj-gemma-4-26B-A4B-it-f16.gguf" in candidates - assert "mmproj-F16.gguf" in candidates - - -def test_projector_status_gemma4_finds_model_specific_on_disk(runtime): - """build_model_projector_status detects a Gemma 4 model-specific projector.""" - models_dir = runtime.base_dir / "models" - (models_dir / "mmproj-gemma-4-E2B-it-f16.gguf").write_bytes(b"g4-projector") - - model = { - "filename": "gemma-4-E2B-it-Q4_K_M.gguf", - "settings": { - "vision": {"enabled": True, "projector_mode": "default", "projector_filename": None}, - }, - } - status = build_model_projector_status(models_dir, model) - assert status["present"] is True - assert "gemma-4-E2B-it" in status["filename"] - - -# ── Recommended runtime ──────────────────────────────────────────── - - -def test_recommended_runtime_gemma4_26b_a4b_is_ik_llama(): - """Only Gemma 4 26B-A4B routes to ik_llama (E2B/E4B not yet supported upstream).""" - assert recommended_runtime_for_model("gemma-4-26B-A4B-it-UD-IQ2_M.gguf") == "ik_llama" - assert recommended_runtime_for_model("gemma-4-26B-A4B-it-UD-IQ4_NL.gguf") == "ik_llama" - - -def test_recommended_runtime_gemma4_e2b_e4b_no_preference(): - """Gemma 4 E2B/E4B use default runtime (ik_llama WIP doesn't support them yet).""" - assert recommended_runtime_for_model("gemma-4-E2B-it-Q4_K_M.gguf") is None - assert recommended_runtime_for_model("gemma-4-E4B-it-Q4_0.gguf") is None - - -def test_recommended_runtime_qwen35_has_no_preference(): - """Qwen3.5 models should have no runtime preference (None).""" - assert recommended_runtime_for_model("Qwen3.5-2B-Q4_K_M.gguf") is None - - -def test_recommended_runtime_unknown_model_has_no_preference(): - """Unknown models should have no runtime preference.""" - assert recommended_runtime_for_model("some-random-model.gguf") is None - - -# ── LiteRT model routing ───────────────────────────────────────────── - - -def test_recommended_runtime_for_litertlm_is_litert(): - """All .litertlm files route to the litert runtime.""" - assert recommended_runtime_for_model("gemma-4-E2B-it.litertlm") == "litert" - assert recommended_runtime_for_model("some-model.litertlm") == "litert" - - -def test_recommended_runtime_for_gguf_unchanged(): - """GGUF files should not route to litert.""" - assert recommended_runtime_for_model("gemma-4-E2B-it-Q4_K_M.gguf") is None - assert recommended_runtime_for_model("Qwen3.5-2B-Q4_K_M.gguf") is None diff --git a/tests/unit/test_runtime.py b/tests/unit/test_runtime.py index 3081d30..6e1c3f7 100644 --- a/tests/unit/test_runtime.py +++ b/tests/unit/test_runtime.py @@ -814,7 +814,7 @@ async def test_ensure_compatible_runtime_returns_false_when_install_fails(monkey async def _fake_install_fail(_install_dir, _path): return {"ok": False, "reason": "rsync_not_available"} - monkeypatch.setattr("core.inferno.runtime_manager.install_llama_runtime_bundle", _fake_install_fail) + monkeypatch.setattr("inferno.runtime_manager.install_llama_runtime_bundle", _fake_install_fail) switched, reason = await ensure_compatible_runtime(runtime) assert switched is False diff --git a/tests/unit/test_script_contracts.py b/tests/unit/test_script_contracts.py index 780f5fe..c28fc70 100644 --- a/tests/unit/test_script_contracts.py +++ b/tests/unit/test_script_contracts.py @@ -8,7 +8,7 @@ def test_start_llama_is_thin_wrapper(): """start_llama.sh must be a thin wrapper that execs $@ — all business - logic lives in core.inferno.launch_config (tested in test_launch_config.py).""" + logic lives in inferno.launch_config (tested in test_launch_config.py).""" script = Path("bin/start_llama.sh").read_text(encoding="utf-8") assert 'exec "$@"' in script @@ -19,7 +19,7 @@ def test_start_llama_is_thin_wrapper(): def test_start_llama_launch_config_has_required_defaults(): """The Python launch config builder must use the same defaults that the old shell script used (q8_0 KV cache, 16384 ctx, etc.).""" - from core.inferno.launch_config import build_llama_server_args + from inferno.launch_config import build_llama_server_args args = build_llama_server_args( llama_server_bin="/bin/llama-server",