#!/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 Handler, STATE, _cap_chat_generation, _context_matches, _inject_global_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_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 ImageEditMultipartTests(unittest.TestCase): def test_openai_image_array_upload_keeps_all_references(self) -> None: boundary = "OpenClawImageBoundary" parts = [ ( f"--{boundary}\r\n" 'Content-Disposition: form-data; name="model"\r\n\r\n' "Qwen-Image-2.1-int8\r\n" ).encode(), ( f"--{boundary}\r\n" 'Content-Disposition: form-data; name="prompt"\r\n\r\n' "Nur die Farbe ändern\r\n" ).encode(), ( f"--{boundary}\r\n" 'Content-Disposition: form-data; name="image[]"; filename="a.png"\r\n' "Content-Type: image/png\r\n\r\n" ).encode() + b"PNG-A\r\n", ( f"--{boundary}\r\n" 'Content-Disposition: form-data; name="image[]"; filename="b.png"\r\n' "Content-Type: image/png\r\n\r\n" ).encode() + b"PNG-B\r\n", f"--{boundary}--\r\n".encode(), ] handler = object.__new__(Handler) files, fields = handler._parse_multipart_parts( b"".join(parts), f"multipart/form-data; boundary={boundary}") self.assertEqual(fields["model"], "Qwen-Image-2.1-int8") self.assertEqual(fields["prompt"], "Nur die Farbe ändern") self.assertEqual( files, [("image[]", "a.png", b"PNG-A"), ("image[]", "b.png", b"PNG-B")], ) 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 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()