Add explicit HYPIR restoration profile
This commit is contained in:
@@ -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")
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user