Add explicit HYPIR restoration profile

This commit is contained in:
Mikei386
2026-09-07 22:52:59 +02:00
parent 2ae61baec7
commit 118e32005e
11 changed files with 593 additions and 32 deletions
+32 -4
View File
@@ -26,6 +26,11 @@ def image_item(state="exited"):
"Labels": {controller.IMAGE_LABEL_KEY: controller.IMAGE_WORKER}}
def restore_item(state="exited"):
return {"Id": "id-restore", "State": state,
"Labels": {controller.IMAGE_LABEL_KEY: controller.RESTORE_WORKER}}
def tts_item(state="running"):
return {"Id": "id-tts", "State": state,
"Labels": {controller.TTS_LABEL_KEY: controller.TTS_WORKER}}
@@ -48,7 +53,7 @@ class ProfileControllerTests(unittest.TestCase):
return 204, b""
with patch.object(controller, "containers", return_value=profiles), \
patch.object(controller, "image_container", return_value=image_item()), \
patch.object(controller, "image_containers", return_value=[image_item()]), \
patch.object(controller, "tts_container", return_value=tts_item()), \
patch.object(controller, "docker_request", side_effect=request):
result = controller.activate("medium")
@@ -62,7 +67,7 @@ class ProfileControllerTests(unittest.TestCase):
def test_fails_if_profile_container_is_missing(self):
profiles = {name: item(name) for name in controller.ALLOWED[:-1]}
with patch.object(controller, "containers", return_value=profiles), \
patch.object(controller, "image_container", return_value=image_item()), \
patch.object(controller, "image_containers", return_value=[image_item()]), \
patch.object(controller, "tts_container", return_value=tts_item()):
with self.assertRaisesRegex(RuntimeError, "missing"):
controller.activate("fast")
@@ -78,6 +83,8 @@ class ProfileControllerTests(unittest.TestCase):
with patch.object(controller, "containers", return_value=profiles), \
patch.object(controller, "image_container", return_value=image_item()), \
patch.object(controller, "image_containers",
return_value=[image_item(), restore_item()]), \
patch.object(controller, "tts_container", return_value=tts_item()), \
patch.object(controller, "docker_request", side_effect=request):
controller.set_image_worker(True)
@@ -87,6 +94,27 @@ class ProfileControllerTests(unittest.TestCase):
("POST", "/containers/id-flux/start"),
])
def test_restore_start_stops_flux_and_starts_restore(self):
profiles = {name: item(name) for name in controller.ALLOWED}
calls = []
def request(method, path):
calls.append((method, path))
return 204, b""
with patch.object(controller, "containers", return_value=profiles), \
patch.object(controller, "image_container", return_value=restore_item()), \
patch.object(controller, "image_containers",
return_value=[image_item("running"), restore_item()]), \
patch.object(controller, "tts_container", return_value=tts_item()), \
patch.object(controller, "docker_request", side_effect=request):
controller.set_image_worker(True, controller.RESTORE_WORKER)
self.assertEqual(calls, [
("POST", "/containers/id-tts/stop?t=30"),
("POST", "/containers/id-flux/stop?t=20"),
("POST", "/containers/id-restore/start"),
])
def test_profile_activation_stops_image_worker_first(self):
profiles = {name: item(name) for name in controller.ALLOWED}
calls = []
@@ -96,8 +124,8 @@ class ProfileControllerTests(unittest.TestCase):
return 204, b""
with patch.object(controller, "containers", return_value=profiles), \
patch.object(controller, "image_container",
return_value=image_item("running")), \
patch.object(controller, "image_containers",
return_value=[image_item("running"), restore_item()]), \
patch.object(controller, "tts_container", return_value=tts_item()), \
patch.object(controller, "docker_request", side_effect=request):
controller.activate("fast")
+25
View File
@@ -24,10 +24,13 @@ from router_support import ( # noqa: E402
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,
@@ -80,6 +83,10 @@ class RuntimeStoreTests(unittest.TestCase):
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)
@@ -270,6 +277,24 @@ class GlobalSystemPolicyTests(unittest.TestCase):
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: