240 lines
8.4 KiB
Python
240 lines
8.4 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": 94208, "model_alias": None},
|
|
"long": {"context": 131072, "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
|