File size: 4,573 Bytes
d61821a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 | 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()
|