147 lines
5.3 KiB
Python
147 lines
5.3 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
|
|
_context_matches,
|
|
_normalize_chat_image,
|
|
_normalize_chat_images,
|
|
_request_has_image,
|
|
)
|
|
|
|
|
|
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_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))
|
|
|
|
|
|
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 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()
|