300 lines
9.3 KiB
Python
300 lines
9.3 KiB
Python
# -*- coding: utf-8 -*-
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
from pathlib import Path
|
|
from typing import cast
|
|
|
|
import pytest
|
|
|
|
from pineagents.local_models import manager as local_model_manager_module
|
|
from pineagents.local_models.llamacpp import LlamaCppServerSetupResult
|
|
from pineagents.local_models.manager import (
|
|
DownloadSource,
|
|
LocalModelManager,
|
|
)
|
|
from pineagents.local_models.llamacpp import LlamaCppBackend
|
|
from pineagents.local_models.model_manager import ModelManager
|
|
from pineagents.providers.provider import ModelInfo
|
|
|
|
|
|
class _FakeLlamaCppBackend:
|
|
def __init__(self) -> None:
|
|
self.calls: list[tuple[str, object | None]] = []
|
|
self.server_running = False
|
|
|
|
def check_llamacpp_installability(self) -> tuple[bool, str]:
|
|
self.calls.append(("installability", None))
|
|
return True, ""
|
|
|
|
def check_llamacpp_installation(self) -> tuple[bool, str]:
|
|
self.calls.append(("check", None))
|
|
return True, ""
|
|
|
|
def get_server_status(self) -> dict[str, object]:
|
|
self.calls.append(("server_status", None))
|
|
return {
|
|
"running": self.server_running,
|
|
"port": 8080 if self.server_running else None,
|
|
"model_name": "demo" if self.server_running else None,
|
|
"pid": 123 if self.server_running else None,
|
|
}
|
|
|
|
def download(self, base_url: str, tag: str) -> None:
|
|
self.calls.append(("download", (base_url, tag)))
|
|
|
|
def get_download_progress(self) -> dict[str, object]:
|
|
self.calls.append(("progress", None))
|
|
return {"status": "downloading"}
|
|
|
|
def cancel_download(self) -> None:
|
|
self.calls.append(("cancel", None))
|
|
|
|
async def server_ready(self, timeout: float = 120.0) -> bool:
|
|
self.calls.append(("server_ready", timeout))
|
|
return True
|
|
|
|
async def setup_server(
|
|
self,
|
|
model_path: Path,
|
|
model_name: str,
|
|
max_context_length: int | None = None,
|
|
port: int | None = None,
|
|
) -> LlamaCppServerSetupResult:
|
|
self.calls.append(
|
|
(
|
|
"setup",
|
|
tuple(
|
|
value
|
|
for value in (
|
|
model_path,
|
|
model_name,
|
|
max_context_length,
|
|
port,
|
|
)
|
|
if value is not None
|
|
),
|
|
),
|
|
)
|
|
return LlamaCppServerSetupResult(
|
|
port=8080,
|
|
model_info=ModelInfo(
|
|
id=model_name,
|
|
name=model_name,
|
|
supports_multimodal=False,
|
|
supports_image=False,
|
|
supports_video=False,
|
|
probe_source="documentation",
|
|
),
|
|
)
|
|
|
|
async def shutdown_server(self) -> None:
|
|
self.calls.append(("shutdown", None))
|
|
self.server_running = False
|
|
|
|
|
|
class _FakeModelManager:
|
|
def __init__(self) -> None:
|
|
self.calls: list[tuple[str, object | None]] = []
|
|
|
|
def get_recommended_models(self) -> list[str]:
|
|
self.calls.append(("recommended", None))
|
|
return ["demo-model"]
|
|
|
|
def is_downloaded(self, model_name: str) -> bool:
|
|
self.calls.append(("is_downloaded", model_name))
|
|
return model_name == "downloaded-model"
|
|
|
|
def list_downloaded_models(self) -> list[str]:
|
|
self.calls.append(("list_downloaded", None))
|
|
return ["downloaded-model"]
|
|
|
|
def get_model_dir(self, model_name: str) -> Path:
|
|
return Path(f"/fake/path/{model_name}")
|
|
|
|
def download_model(
|
|
self,
|
|
model_name: str,
|
|
source: DownloadSource | None = None,
|
|
) -> None:
|
|
self.calls.append(("download_model", (model_name, source)))
|
|
|
|
def get_download_progress(self) -> dict[str, object]:
|
|
self.calls.append(("progress", None))
|
|
return {"status": "pending"}
|
|
|
|
def cancel_download(self) -> None:
|
|
self.calls.append(("cancel", None))
|
|
|
|
def remove_downloaded_model(self, model_name: str) -> None:
|
|
self.calls.append(("remove", model_name))
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_local_model_manager_forwards_sync_calls() -> None:
|
|
fake_model_manager = _FakeModelManager()
|
|
fake_llamacpp_backend = _FakeLlamaCppBackend()
|
|
manager = LocalModelManager(
|
|
model_manager=cast(ModelManager, fake_model_manager),
|
|
llamacpp_backend=cast(LlamaCppBackend, fake_llamacpp_backend),
|
|
)
|
|
|
|
assert manager.check_llamacpp_installability() == (True, "")
|
|
assert manager.check_llamacpp_installation() == (True, "")
|
|
server_stopped = await manager.start_llamacpp_download()
|
|
assert manager.get_llamacpp_download_progress() == {
|
|
"status": "downloading",
|
|
}
|
|
manager.cancel_llamacpp_download()
|
|
assert server_stopped is False
|
|
|
|
assert manager.get_recommended_models() == ["demo-model"]
|
|
assert manager.is_model_downloaded("downloaded-model") is True
|
|
assert manager.list_downloaded_models() == ["downloaded-model"]
|
|
manager.start_model_download(
|
|
"demo-model",
|
|
source=DownloadSource.MODELSCOPE,
|
|
)
|
|
assert manager.get_model_download_progress() == {"status": "pending"}
|
|
manager.cancel_model_download()
|
|
manager.remove_downloaded_model("downloaded-model")
|
|
|
|
assert fake_llamacpp_backend.calls == [
|
|
("installability", None),
|
|
("check", None),
|
|
("server_status", None),
|
|
(
|
|
"download",
|
|
(
|
|
LocalModelManager.DEFAULT_LLAMA_CPP_BASE_URL,
|
|
LocalModelManager.DEFAULT_LLAMA_CPP_RELEASE_TAG,
|
|
),
|
|
),
|
|
("progress", None),
|
|
("cancel", None),
|
|
]
|
|
assert fake_model_manager.calls == [
|
|
("recommended", None),
|
|
("is_downloaded", "downloaded-model"),
|
|
("list_downloaded", None),
|
|
("download_model", ("demo-model", DownloadSource.MODELSCOPE)),
|
|
("progress", None),
|
|
("cancel", None),
|
|
("remove", "downloaded-model"),
|
|
]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_local_model_manager_forwards_async_server_calls(
|
|
tmp_path: Path,
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
) -> None:
|
|
monkeypatch.setattr(
|
|
local_model_manager_module,
|
|
"DEFAULT_LOCAL_PROVIDER_DIR",
|
|
tmp_path,
|
|
)
|
|
fake_llamacpp_backend = _FakeLlamaCppBackend()
|
|
manager = LocalModelManager(
|
|
model_manager=cast(ModelManager, _FakeModelManager()),
|
|
llamacpp_backend=cast(LlamaCppBackend, fake_llamacpp_backend),
|
|
)
|
|
await manager.set_max_context_length(131072)
|
|
await manager.set_port(43110)
|
|
|
|
ready = await manager.check_llamacpp_server_ready(timeout=7.5)
|
|
setup_result = await manager.setup_server("demo")
|
|
await manager.shutdown_server()
|
|
|
|
assert ready is True
|
|
assert setup_result.port == 8080
|
|
assert setup_result.model_info == ModelInfo(
|
|
id="demo",
|
|
name="demo",
|
|
supports_multimodal=False,
|
|
supports_image=False,
|
|
supports_video=False,
|
|
probe_source="documentation",
|
|
)
|
|
assert fake_llamacpp_backend.calls == [
|
|
("server_ready", 7.5),
|
|
("setup", (Path("/fake/path/demo"), "demo", 131072, 43110)),
|
|
("shutdown", None),
|
|
]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_set_max_context_length_persists_local_model_config(
|
|
tmp_path: Path,
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
) -> None:
|
|
monkeypatch.setattr(
|
|
local_model_manager_module,
|
|
"DEFAULT_LOCAL_PROVIDER_DIR",
|
|
tmp_path,
|
|
)
|
|
manager = LocalModelManager(
|
|
model_manager=cast(ModelManager, _FakeModelManager()),
|
|
llamacpp_backend=cast(LlamaCppBackend, _FakeLlamaCppBackend()),
|
|
)
|
|
|
|
await manager.set_max_context_length(131072)
|
|
|
|
assert manager.get_config().max_context_length == 131072
|
|
config_path = tmp_path / LocalModelManager.CONFIG_FILE_NAME
|
|
assert json.loads(config_path.read_text(encoding="utf-8")) == {
|
|
"max_context_length": 131072,
|
|
"port": None,
|
|
}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_set_port_persists_local_model_config(
|
|
tmp_path: Path,
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
) -> None:
|
|
monkeypatch.setattr(
|
|
local_model_manager_module,
|
|
"DEFAULT_LOCAL_PROVIDER_DIR",
|
|
tmp_path,
|
|
)
|
|
manager = LocalModelManager(
|
|
model_manager=cast(ModelManager, _FakeModelManager()),
|
|
llamacpp_backend=cast(LlamaCppBackend, _FakeLlamaCppBackend()),
|
|
)
|
|
|
|
await manager.set_port(43110)
|
|
|
|
assert manager.get_config().port == 43110
|
|
config_path = tmp_path / LocalModelManager.CONFIG_FILE_NAME
|
|
assert json.loads(config_path.read_text(encoding="utf-8")) == {
|
|
"max_context_length": 65536,
|
|
"port": 43110,
|
|
}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_start_llamacpp_download_stops_running_server_first() -> None:
|
|
fake_llamacpp_backend = _FakeLlamaCppBackend()
|
|
fake_llamacpp_backend.server_running = True
|
|
manager = LocalModelManager(
|
|
model_manager=cast(ModelManager, _FakeModelManager()),
|
|
llamacpp_backend=cast(LlamaCppBackend, fake_llamacpp_backend),
|
|
)
|
|
|
|
server_stopped = await manager.start_llamacpp_download()
|
|
|
|
assert server_stopped is True
|
|
assert fake_llamacpp_backend.calls == [
|
|
("server_status", None),
|
|
("shutdown", None),
|
|
(
|
|
"download",
|
|
(
|
|
LocalModelManager.DEFAULT_LLAMA_CPP_BASE_URL,
|
|
LocalModelManager.DEFAULT_LLAMA_CPP_RELEASE_TAG,
|
|
),
|
|
),
|
|
]
|