Files
AI-Profile-Router/router/router_support.py

242 lines
8.5 KiB
Python

#!/usr/bin/env python3
"""Security and runtime helpers for the AI Profile Router.
This module deliberately contains no model-specific logic. It provides the
small, testable building blocks that the HTTP gateway and the GPU orchestrator
share: API-key authentication, crash-state persistence, Linux child-process
cleanup and bounded artifact retention.
"""
from __future__ import annotations
import hmac
import json
import os
import signal
import threading
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Callable, Iterable
class ConfigurationError(RuntimeError):
"""Raised when the router would otherwise start in an unsafe state."""
def load_profile_registry(path: str | None) -> dict[str, dict]:
"""Load and validate the optional profile registry.
Passing no path keeps source-tree compatibility. An explicitly configured
but missing registry is a fatal configuration error: production must never
silently lose alias validation because a deployment forgot one file.
"""
fallback = {
"fast": {"context": 76800, "model_alias": None},
"medium": {"context": 160000, "model_alias": None},
"large": {"context": 192000, "model_alias": None},
"ultra": {"context": 262144, "model_alias": None},
"uncensored": {"context": 80000, "model_alias": None},
}
if not path:
return fallback
try:
raw = json.loads(Path(path).read_text(encoding="utf-8"))
except (OSError, ValueError) as exc:
raise ConfigurationError(f"Profilregister kann nicht gelesen werden: {exc}")
profiles = raw.get("profiles") if isinstance(raw, dict) else None
if not isinstance(profiles, dict) or not profiles:
raise ConfigurationError("Profilregister enthält keine 'profiles'")
validated: dict[str, dict] = {}
for name, definition in profiles.items():
if not isinstance(name, str) or not re_full_profile_name(name):
raise ConfigurationError(f"ungültiger Profilname: {name!r}")
if not isinstance(definition, dict):
raise ConfigurationError(f"Profil {name!r} ist kein Objekt")
context = definition.get("context")
alias = definition.get("model_alias")
if not isinstance(context, int) or context < 1024:
raise ConfigurationError(f"ungültiger Kontext für Profil {name!r}")
if alias is not None and (not isinstance(alias, str) or not alias.strip()):
raise ConfigurationError(f"ungültiger Modellalias für Profil {name!r}")
validated[name] = {"context": context,
"model_alias": alias.strip() if alias else None}
return validated
def re_full_profile_name(value: str) -> bool:
return bool(value) and all(ch.isalnum() or ch in "-_" for ch in value)
@dataclass(frozen=True)
class AuthPolicy:
"""Bearer/X-API-Key authentication policy.
``mode`` is either ``required`` or ``off``. ``off`` is intended only for
loopback development tests. Production startup fails closed when the key
is missing or too short.
"""
mode: str
api_key: str | None
@classmethod
def from_environment(cls) -> "AuthPolicy":
mode = os.environ.get("ROUTER_AUTH_MODE", "required").strip().lower()
if mode not in {"required", "off"}:
raise ConfigurationError(
"ROUTER_AUTH_MODE muss 'required' oder 'off' sein")
if mode == "off":
return cls(mode=mode, api_key=None)
key = os.environ.get("ROUTER_API_KEY", "").strip()
key_file = os.environ.get(
"ROUTER_API_KEY_FILE", "/etc/mike-ai/router-api-key")
if not key and key_file:
try:
key = Path(key_file).read_text(encoding="utf-8").strip()
except OSError:
pass
if len(key) < 32:
raise ConfigurationError(
"Router-Authentifizierung ist aktiv, aber kein API-Key mit "
"mindestens 32 Zeichen vorhanden")
return cls(mode=mode, api_key=key)
@property
def enabled(self) -> bool:
return self.mode == "required"
def accepts(self, authorization: str | None,
x_api_key: str | None) -> bool:
if not self.enabled:
return True
candidate = (x_api_key or "").strip()
auth = (authorization or "").strip()
if not candidate and auth.lower().startswith("bearer "):
candidate = auth[7:].strip()
return bool(candidate and self.api_key
and hmac.compare_digest(candidate, self.api_key))
class RuntimeStore:
"""Tiny atomic JSON store used for crash reconciliation."""
def __init__(self, path: str) -> None:
self.path = Path(path)
self._lock = threading.RLock()
def load(self) -> dict:
with self._lock:
try:
value = json.loads(self.path.read_text(encoding="utf-8"))
except (OSError, ValueError):
return {}
return value if isinstance(value, dict) else {}
def save(self, **updates) -> None:
with self._lock:
state = self.load()
for key, value in updates.items():
if value is None:
state.pop(key, None)
else:
state[key] = value
state["updated_at"] = int(time.time())
self.path.parent.mkdir(parents=True, exist_ok=True)
temp = self.path.with_name(
f".{self.path.name}.{os.getpid()}.{threading.get_ident()}.tmp")
try:
temp.write_text(
json.dumps(state, indent=2, sort_keys=True) + "\n",
encoding="utf-8")
os.chmod(temp, 0o600)
os.replace(temp, self.path)
finally:
try:
temp.unlink()
except FileNotFoundError:
pass
def clear_worker(self, worker: str) -> None:
with self._lock:
state = self.load()
if state.get("worker") == worker:
self.save(worker=None, worker_pid=None)
def terminate_recorded_worker(state: dict, allowed_markers: Iterable[str],
log: Callable[[str], None]) -> bool:
"""Terminate a previously recorded worker, but only after cmdline checks."""
pid = state.get("worker_pid")
if not isinstance(pid, int) or pid <= 1:
return False
cmdline_path = Path(f"/proc/{pid}/cmdline")
try:
cmdline = cmdline_path.read_bytes().replace(b"\0", b" ").decode(
errors="replace")
except OSError:
return False
if not any(marker in cmdline for marker in allowed_markers):
log(f"Recorded PID {pid} nicht beendet: Prozessprüfung fehlgeschlagen")
return False
try:
os.kill(pid, signal.SIGTERM)
except ProcessLookupError:
return False
log(f"Verwaister Router-Worker PID {pid} wurde beendet")
return True
def enforce_artifact_retention(directory: str, max_files: int, max_bytes: int,
max_age_days: int,
protected: Iterable[str] = ()) -> list[str]:
"""Delete oldest PNG + sidecar pairs until all retention limits hold."""
root = Path(directory)
if not root.is_dir():
return []
protected_set = set(protected)
now = time.time()
cutoff = now - max_age_days * 86400 if max_age_days > 0 else None
items: list[tuple[float, int, Path]] = []
for path in root.glob("*.png"):
if path.name in protected_set:
continue
try:
stat = path.stat()
except OSError:
continue
sidecar = path.with_suffix(".json")
size = stat.st_size
try:
size += sidecar.stat().st_size
except OSError:
pass
items.append((stat.st_mtime, size, path))
items.sort(key=lambda item: item[0])
total = sum(item[1] for item in items)
removed: list[str] = []
while items:
mtime, size, path = items[0]
too_old = cutoff is not None and mtime < cutoff
too_many = max_files > 0 and len(items) > max_files
too_large = max_bytes > 0 and total > max_bytes
if not (too_old or too_many or too_large):
break
items.pop(0)
try:
path.unlink()
removed.append(path.name)
except OSError:
continue
try:
path.with_suffix(".json").unlink()
except OSError:
pass
total -= size
return removed