| from __future__ import annotations |
|
|
| from typing import Any |
| from unittest.mock import patch |
| import subprocess |
| import unittest |
|
|
| from agent_harness.lm_studio_management import LMStudioResidencyManager, LMStudioServer |
|
|
|
|
| class FakeResidencyManager(LMStudioResidencyManager): |
| def __init__(self) -> None: |
| super().__init__("http://127.0.0.1:1234", "UNUSED") |
| self.instances: list[dict[str, Any]] = [] |
|
|
| def models(self) -> tuple[dict[str, Any], ...]: |
| by_key: dict[str, list[dict[str, Any]]] = {} |
| for item in self.instances: |
| by_key.setdefault(item["model_key"], []).append( |
| {"id": item["instance_id"], "config": item["config"]} |
| ) |
| return tuple( |
| {"key": key, "type": "llm", "loaded_instances": values} |
| for key, values in by_key.items() |
| ) |
|
|
| def _unload(self, instance_id: str) -> dict[str, Any]: |
| self.instances = [item for item in self.instances if item["instance_id"] != instance_id] |
| return {"instance_id": instance_id} |
|
|
| def _request( |
| self, method: str, endpoint: str, payload: dict[str, Any] | None = None |
| ) -> dict[str, Any]: |
| if endpoint != "/api/v1/models/load" or payload is None: |
| raise AssertionError((method, endpoint, payload)) |
| key = str(payload["model"]) |
| context = int(payload["context_length"]) |
| self.instances.append( |
| {"model_key": key, "instance_id": key, "config": {"context_length": context}} |
| ) |
| return { |
| "type": "llm", |
| "instance_id": key, |
| "status": "loaded", |
| "load_time_seconds": 1.0, |
| "load_config": {"context_length": context}, |
| } |
|
|
|
|
| class ResidencyManagerTests(unittest.TestCase): |
| def test_server_status_accepts_success_message_on_stderr(self) -> None: |
| server = LMStudioServer(cli_path=__file__) |
| completed = subprocess.CompletedProcess( |
| args=["lms", "server", "status"], |
| returncode=0, |
| stdout="", |
| stderr="The server is running on port 1234.\n", |
| ) |
| with patch.object(server, "_run", return_value=completed): |
| self.assertTrue(server.status()["running"]) |
|
|
| def test_server_status_does_not_treat_not_running_as_running(self) -> None: |
| server = LMStudioServer(cli_path=__file__) |
| completed = subprocess.CompletedProcess( |
| args=["lms", "server", "status"], |
| returncode=0, |
| stdout="", |
| stderr="The server is not running.\n", |
| ) |
| with patch.object(server, "_run", return_value=completed): |
| self.assertFalse(server.status()["running"]) |
|
|
| def test_ensure_running_waits_for_official_api_readiness(self) -> None: |
| server = LMStudioServer(cli_path=__file__) |
| status = {"running": True, "returncode": 0, "stdout": "", "stderr": ""} |
| readiness = iter((False, True)) |
| with patch.object(server, "status", return_value=status), patch.object( |
| server, "_api_ready", side_effect=lambda: next(readiness) |
| ), patch("agent_harness.lm_studio_management.time.sleep"): |
| result = server.ensure_running() |
| self.assertEqual(result["action"], "already_running") |
| self.assertEqual(result["status"], status) |
|
|
| def test_exclusive_switch_unloads_previous_model(self) -> None: |
| manager = FakeResidencyManager() |
| first = manager.ensure_exclusive("embedding", 8192) |
| second = manager.ensure_exclusive("qwen", 65536) |
| self.assertFalse(first.reused) |
| self.assertEqual(second.unloaded_instances, ("embedding",)) |
| self.assertEqual(second.after_instances, ("qwen",)) |
| self.assertEqual(manager.loaded_instances()[0]["model_key"], "qwen") |
|
|
| def test_matching_model_and_context_are_reused(self) -> None: |
| manager = FakeResidencyManager() |
| manager.ensure_exclusive("qwen", 65536) |
| transition = manager.ensure_exclusive("qwen", 65536) |
| self.assertTrue(transition.reused) |
| self.assertEqual(transition.unloaded_instances, ()) |
|
|
| def test_context_change_forces_reload(self) -> None: |
| manager = FakeResidencyManager() |
| manager.ensure_exclusive("qwen", 65536) |
| transition = manager.ensure_exclusive("qwen", 32768) |
| self.assertFalse(transition.reused) |
| self.assertEqual(transition.unloaded_instances, ("qwen",)) |
| self.assertEqual(manager.loaded_instances()[0]["config"]["context_length"], 32768) |
|
|
|
|
| if __name__ == "__main__": |
| unittest.main() |
|
|