agent-harness / tests /test_lm_studio_management.py
cuber12's picture
Publish agent harness research code and paper artifacts
d61821a verified
Raw
History Blame Contribute Delete
4.57 kB
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()