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()