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",