Files
AI-Profile-Router/dev/test_router_support.py
T

316 lines
13 KiB
Python

#!/usr/bin/env python3
"""Unit tests for security, runtime persistence and artifact retention."""
from __future__ import annotations
import os
import tempfile
import threading
import time
import unittest
from pathlib import Path
from unittest.mock import patch
import sys
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "router"))
os.environ.setdefault("ROUTER_PROFILES_FILE", "")
from router_support import ( # noqa: E402
AuthPolicy,
ConfigurationError,
RuntimeStore,
enforce_artifact_retention,
load_profile_registry,
)
from ai_profile_router import ( # noqa: E402
RESTORATION_CHAT_MODEL,
VIRTUAL_MODELS,
STATE,
_cap_chat_generation,
_context_matches,
_inject_global_system_policy,
_inject_restoration_system_policy,
_normalize_chat_image,
_normalize_chat_images,
_normalize_llamacpp_reasoning,
_request_has_image,
switch_profile,
)
class AuthPolicyTests(unittest.TestCase):
def test_required_mode_fails_closed_without_key(self) -> None:
with patch.dict(os.environ, {
"ROUTER_AUTH_MODE": "required",
"ROUTER_API_KEY": "",
"ROUTER_API_KEY_FILE": "/definitely/missing"}, clear=False):
with self.assertRaises(ConfigurationError):
AuthPolicy.from_environment()
def test_bearer_and_x_api_key(self) -> None:
key = "k" * 48
policy = AuthPolicy("required", key)
self.assertTrue(policy.accepts(f"Bearer {key}", None))
self.assertTrue(policy.accepts(None, key))
self.assertFalse(policy.accepts("Bearer wrong", None))
class RuntimeStoreTests(unittest.TestCase):
def test_atomic_roundtrip_and_delete(self) -> None:
with tempfile.TemporaryDirectory() as temp:
store = RuntimeStore(str(Path(temp) / "state.json"))
store.save(worker="image", worker_pid=123, last_profile="fast")
self.assertEqual(store.load()["worker_pid"], 123)
store.clear_worker("image")
state = store.load()
self.assertNotIn("worker", state)
self.assertEqual(state["last_profile"], "fast")
def test_concurrent_updates_are_not_lost(self) -> None:
with tempfile.TemporaryDirectory() as temp:
store = RuntimeStore(str(Path(temp) / "state.json"))
threads = [threading.Thread(target=store.save,
kwargs={f"key_{index}": index})
for index in range(20)]
for thread in threads:
thread.start()
for thread in threads:
thread.join()
state = store.load()
for index in range(20):
self.assertEqual(state[f"key_{index}"], index)
class ProfileRegistryTests(unittest.TestCase):
def test_restoration_model_maps_to_fast_instruction_profile(self) -> None:
self.assertEqual(RESTORATION_CHAT_MODEL, "restauration")
self.assertEqual(VIRTUAL_MODELS[RESTORATION_CHAT_MODEL], "fast")
def test_fallback_contains_uncensored_profile(self) -> None:
registry = load_profile_registry(None)
self.assertEqual(registry["uncensored"]["context"], 80000)
def test_explicit_missing_registry_fails_closed(self) -> None:
with self.assertRaises(ConfigurationError):
load_profile_registry("/definitely/missing/profiles.json")
def test_context_match_accepts_small_runtime_overhead(self) -> None:
self.assertTrue(_context_matches(80000, 80128))
self.assertTrue(_context_matches(160000, 160000))
def test_context_match_rejects_other_profiles_and_lower_context(self) -> None:
self.assertFalse(_context_matches(80000, 76800))
self.assertFalse(_context_matches(80000, 82000))
def test_ready_profile_clears_stale_unavailable_flag(self) -> None:
with STATE.avail_lock:
STATE.qwen_unavailable = True
status = {"reachable": True, "model": "qwen-medium", "ctx": 160000}
with patch("ai_profile_router.current_profile", return_value="medium"), \
patch("ai_profile_router.upstream_status", return_value=status):
switch_profile("medium")
with STATE.avail_lock:
self.assertFalse(STATE.qwen_unavailable)
class ChatImageInputTests(unittest.TestCase):
def test_small_png_data_url_is_accepted(self) -> None:
value = "data:image/png;base64,iVBORw0KGgo="
self.assertEqual(_normalize_chat_image(value), value)
def test_remote_url_is_denied_by_default(self) -> None:
with self.assertRaisesRegex(ValueError, "deaktiviert"):
_normalize_chat_image("https://example.com/private.png")
def test_invalid_base64_is_rejected(self) -> None:
with self.assertRaisesRegex(ValueError, "Base64"):
_normalize_chat_image("data:image/png;base64,not!base64")
def test_multimodal_part_is_retained_for_direct_forwarding(self) -> None:
value = "data:image/png;base64,iVBORw0KGgo="
request = {
"model": "qwen-fast",
"messages": [{
"role": "user",
"content": [
{"type": "text", "text": "Was ist zu sehen?"},
{"type": "image_url", "image_url": {"url": value}},
],
}],
}
self.assertTrue(_request_has_image(request))
normalized = _normalize_chat_images(request)
self.assertEqual(
normalized["messages"][0]["content"][1]["image_url"]["url"],
value,
)
self.assertEqual(request, normalized)
class LlamaCppReasoningTests(unittest.TestCase):
def test_disabled_values_really_disable_thinking(self) -> None:
for effort in (None, "none", "off", "disabled", False):
with self.subTest(effort=effort):
request = {"reasoning_effort": effort, "messages": []}
normalized = _normalize_llamacpp_reasoning(request)
self.assertNotIn("reasoning_effort", normalized)
self.assertEqual(
normalized["chat_template_kwargs"],
{"enable_thinking": False},
)
self.assertEqual(normalized["thinking_budget_tokens"], 0)
def test_reasoning_levels_receive_real_per_request_budgets(self) -> None:
expected = {
"minimal": ("low", 256),
"low": ("low", 768),
"medium": ("medium", 2048),
"high": ("xhigh", 4096),
"xhigh": ("xhigh", 8192),
"max": ("xhigh", 8192),
"ultra": ("xhigh", 8192),
}
for effort, (template_effort, budget) in expected.items():
with self.subTest(effort=effort):
request = {"reasoning_effort": effort, "messages": []}
normalized = _normalize_llamacpp_reasoning(request)
self.assertEqual(normalized["chat_template_kwargs"], {
"enable_thinking": True,
"reasoning_effort": template_effort,
})
self.assertEqual(normalized["thinking_budget_tokens"], budget)
def test_existing_template_kwargs_are_preserved(self) -> None:
request = {
"reasoning_effort": "low",
"chat_template_kwargs": {"preserve_thinking": True},
"messages": [],
}
normalized = _normalize_llamacpp_reasoning(request)
self.assertEqual(normalized["chat_template_kwargs"], {
"preserve_thinking": True,
"enable_thinking": True,
"reasoning_effort": "low",
})
self.assertEqual(normalized["thinking_budget_tokens"], 768)
def test_request_without_effort_uses_safe_off_default(self) -> None:
request = {"messages": []}
self.assertIs(_normalize_llamacpp_reasoning(request), request)
self.assertEqual(
request["chat_template_kwargs"],
{"enable_thinking": False},
)
self.assertEqual(request["thinking_budget_tokens"], 0)
def test_native_thinking_budget_is_preserved_without_openai_effort(self) -> None:
request = {"thinking_budget_tokens": 1234, "messages": []}
self.assertIs(_normalize_llamacpp_reasoning(request), request)
self.assertEqual(request["thinking_budget_tokens"], 1234)
self.assertNotIn("chat_template_kwargs", request)
class ChatGenerationLimitTests(unittest.TestCase):
def test_missing_limit_gets_platform_cap(self) -> None:
request = {"messages": []}
self.assertEqual(_cap_chat_generation(request)["max_tokens"], 8192)
def test_smaller_explicit_limit_is_preserved(self) -> None:
request = {"max_tokens": 512}
self.assertEqual(_cap_chat_generation(request)["max_tokens"], 512)
def test_unlimited_and_oversized_values_are_capped(self) -> None:
for value in (-1, 0, 999999):
with self.subTest(value=value):
request = {"max_completion_tokens": value}
self.assertEqual(
_cap_chat_generation(request)["max_completion_tokens"],
8192,
)
class GlobalSystemPolicyTests(unittest.TestCase):
def _inject(self, request: dict, path: str,
policy: str = "Verify facts.") -> dict:
with patch("ai_profile_router._load_global_system_policy",
return_value=policy):
return _inject_global_system_policy(request, path)
def test_chat_policy_precedes_existing_system_prompt(self) -> None:
request = {"messages": [
{"role": "system", "content": "Client policy."},
{"role": "user", "content": "Hello"},
]}
normalized = self._inject(request, "/v1/chat/completions")
self.assertEqual(
normalized["messages"][0]["content"],
"Verify facts.\n\nClient policy.",
)
def test_chat_policy_is_inserted_without_system_prompt(self) -> None:
request = {"messages": [{"role": "user", "content": "Hello"}]}
normalized = self._inject(request, "/v1/chat/completions")
self.assertEqual(normalized["messages"][0], {
"role": "system", "content": "Verify facts.",
})
def test_policy_is_not_duplicated(self) -> None:
request = {"messages": [{
"role": "system", "content": "Verify facts.\n\nClient policy.",
}]}
normalized = self._inject(request, "/v1/chat/completions")
self.assertEqual(
normalized["messages"][0]["content"].count("Verify facts."), 1)
def test_responses_policy_precedes_instructions(self) -> None:
request = {"instructions": "Client policy.", "input": "Hello"}
normalized = self._inject(request, "/v1/responses")
self.assertEqual(
normalized["instructions"],
"Verify facts.\n\nClient policy.",
)
def test_unrelated_endpoint_is_unchanged(self) -> None:
request = {"prompt": "Draw a cat"}
self.assertEqual(
self._inject(request, "/v1/images/generations"), request)
class RestorationSystemPolicyTests(unittest.TestCase):
def test_chat_policy_requires_image_tool_and_preservation(self) -> None:
request = {"messages": [{"role": "user", "content": "Mach schöner"}]}
normalized = _inject_restoration_system_policy(
request, "/v1/chat/completions")
policy = normalized["messages"][0]["content"]
self.assertIn("image generation/editing tool", policy)
self.assertIn("Preserve identity", policy)
self.assertIn("Do not attempt restoration with Python", policy)
def test_responses_policy_keeps_client_instructions(self) -> None:
request = {"instructions": "Client policy.", "input": "Mach schöner"}
normalized = _inject_restoration_system_policy(
request, "/v1/responses")
self.assertIn("Photo-restoration mode", normalized["instructions"])
self.assertTrue(normalized["instructions"].endswith("Client policy."))
class RetentionTests(unittest.TestCase):
def test_oldest_pairs_are_removed(self) -> None:
with tempfile.TemporaryDirectory() as temp:
root = Path(temp)
for index in range(3):
png = root / f"image-{index}.png"
png.write_bytes(b"x" * 10)
png.with_suffix(".json").write_text("{}", encoding="utf-8")
stamp = time.time() - (30 - index)
os.utime(png, (stamp, stamp))
removed = enforce_artifact_retention(
temp, max_files=2, max_bytes=0, max_age_days=0)
self.assertEqual(removed, ["image-0.png"])
self.assertFalse((root / "image-0.json").exists())
if __name__ == "__main__":
unittest.main()