mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-08-08 23:10:15 -03:00
fix(types): resolve pre-existing basedpyright errors in tests
Fix ~790 basedpyright errors across the test suite: - Type stub subclasses of real production classes with super().__init__() - Add missing generic type arguments and Dict[str, Any] annotations - Add None guards before subscript/member access - Adapt tests to production API changes (removed dead handlers, PersistentModelCache.get_default, _i18n_filter_added location)
This commit is contained in:
@@ -8,8 +8,10 @@ response schemas.
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import pytest
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from syrupy import SnapshotAssertion
|
||||
|
||||
from py.routes.handlers.misc_handlers import (
|
||||
@@ -54,13 +56,35 @@ async def noop_async(*_args, **_kwargs):
|
||||
return None
|
||||
|
||||
|
||||
class FakeDownloader:
|
||||
"""Minimal downloader stub satisfying DownloaderProtocol."""
|
||||
|
||||
async def refresh_session(self) -> None:
|
||||
return None
|
||||
|
||||
|
||||
async def fake_downloader_factory() -> FakeDownloader:
|
||||
return FakeDownloader()
|
||||
|
||||
|
||||
async def fake_metadata_provider_factory():
|
||||
return None
|
||||
|
||||
|
||||
def json_payload(response) -> Any:
|
||||
"""Decode the JSON body of a web.Response, asserting it is not null."""
|
||||
text = response.text
|
||||
assert text is not None
|
||||
return json.loads(text)
|
||||
|
||||
|
||||
class FakePromptServer:
|
||||
"""Fake prompt server for testing."""
|
||||
|
||||
sent = []
|
||||
|
||||
class Instance:
|
||||
sockets: dict = {}
|
||||
sockets: dict[str, Any] = {}
|
||||
|
||||
def send_sync(self, event, payload, sid=None):
|
||||
FakePromptServer.sent.append((event, payload))
|
||||
@@ -103,11 +127,11 @@ class TestSettingsHandlerSnapshots:
|
||||
handler = SettingsHandler(
|
||||
settings_service=settings_service,
|
||||
metadata_provider_updater=noop_async,
|
||||
downloader_factory=lambda: None,
|
||||
downloader_factory=fake_downloader_factory,
|
||||
)
|
||||
|
||||
response = await handler.get_settings(FakeRequest())
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.get_settings(FakeRequest()) # pyright: ignore[reportArgumentType]
|
||||
payload = json_payload(response)
|
||||
|
||||
assert payload == snapshot
|
||||
|
||||
@@ -118,12 +142,12 @@ class TestSettingsHandlerSnapshots:
|
||||
handler = SettingsHandler(
|
||||
settings_service=settings_service,
|
||||
metadata_provider_updater=noop_async,
|
||||
downloader_factory=lambda: None,
|
||||
downloader_factory=fake_downloader_factory,
|
||||
)
|
||||
|
||||
request = FakeRequest(json_data={"language": "zh"})
|
||||
response = await handler.update_settings(request)
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.update_settings(request) # pyright: ignore[reportArgumentType]
|
||||
payload = json_payload(response)
|
||||
|
||||
assert payload == snapshot
|
||||
|
||||
@@ -137,7 +161,7 @@ class TestNodeRegistryHandlerSnapshots:
|
||||
node_registry = NodeRegistry()
|
||||
handler = NodeRegistryHandler(
|
||||
node_registry=node_registry,
|
||||
prompt_server=FakePromptServer,
|
||||
prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
|
||||
standalone_mode=False,
|
||||
)
|
||||
|
||||
@@ -155,8 +179,8 @@ class TestNodeRegistryHandlerSnapshots:
|
||||
}
|
||||
)
|
||||
|
||||
response = await handler.register_nodes(request)
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.register_nodes(request) # pyright: ignore[reportArgumentType]
|
||||
payload = json_payload(response)
|
||||
|
||||
assert payload == snapshot
|
||||
|
||||
@@ -166,13 +190,13 @@ class TestNodeRegistryHandlerSnapshots:
|
||||
node_registry = NodeRegistry()
|
||||
handler = NodeRegistryHandler(
|
||||
node_registry=node_registry,
|
||||
prompt_server=FakePromptServer,
|
||||
prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
|
||||
standalone_mode=False,
|
||||
)
|
||||
|
||||
request = FakeRequest(json_data={"nodes": [], "client_id": "test-client-1"})
|
||||
response = await handler.register_nodes(request)
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.register_nodes(request) # pyright: ignore[reportArgumentType]
|
||||
payload = json_payload(response)
|
||||
|
||||
assert payload == snapshot
|
||||
|
||||
@@ -249,10 +273,12 @@ class TestModelLibraryHandlerSnapshots:
|
||||
get_embedding_scanner=scanner_factory,
|
||||
get_downloaded_version_history_service=fake_download_history_service_factory,
|
||||
),
|
||||
metadata_provider_factory=lambda: None,
|
||||
metadata_provider_factory=fake_metadata_provider_factory,
|
||||
)
|
||||
|
||||
response = await handler.check_model_exists(FakeRequest(query={"modelId": "1"}))
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.check_model_exists(
|
||||
FakeRequest(query={"modelId": "1"}) # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
payload = json_payload(response)
|
||||
|
||||
assert payload == snapshot
|
||||
|
||||
@@ -6,10 +6,10 @@ from pathlib import Path
|
||||
|
||||
import types
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
from typing import Any, Optional
|
||||
|
||||
folder_paths_stub = types.SimpleNamespace(get_folder_paths=lambda *_: [])
|
||||
sys.modules.setdefault("folder_paths", folder_paths_stub)
|
||||
sys.modules.setdefault("folder_paths", folder_paths_stub) # pyright: ignore[reportArgumentType]
|
||||
|
||||
import pytest
|
||||
from aiohttp import FormData, web
|
||||
@@ -38,7 +38,9 @@ class DummyRoutes(BaseModelRoutes):
|
||||
|
||||
def __init__(self, service=None):
|
||||
super().__init__(service)
|
||||
self.set_model_update_service(NullModelUpdateService())
|
||||
self.set_model_update_service(
|
||||
NullModelUpdateService() # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -110,7 +112,7 @@ class NullModelUpdateService:
|
||||
return None
|
||||
|
||||
|
||||
async def create_test_client(service) -> TestClient:
|
||||
async def create_test_client(service) -> TestClient[Any, Any]:
|
||||
routes = DummyRoutes(service)
|
||||
app = web.Application()
|
||||
routes.setup_routes(app, "test-models")
|
||||
@@ -457,18 +459,18 @@ def test_fetch_civitai_hydrates_metadata_before_sync(
|
||||
mock_scanner._cache.raw_data = [minimal_cache_entry]
|
||||
|
||||
class FakeMetadata:
|
||||
def __init__(self, payload: dict) -> None:
|
||||
def __init__(self, payload: dict[str, Any]) -> None:
|
||||
self._payload = payload
|
||||
self._unknown_fields = {"legacy_field": "legacy"}
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return self._payload.copy()
|
||||
|
||||
async def fake_load_metadata(path: str, *_args, **_kwargs):
|
||||
assert path == str(model_path)
|
||||
return FakeMetadata(existing_metadata), False
|
||||
|
||||
async def fake_save_metadata(path: str, metadata: dict) -> bool:
|
||||
async def fake_save_metadata(path: str, metadata: dict[str, Any]) -> bool:
|
||||
save_calls.append((path, json.loads(json.dumps(metadata))))
|
||||
return True
|
||||
|
||||
@@ -477,7 +479,7 @@ def test_fetch_civitai_hydrates_metadata_before_sync(
|
||||
*,
|
||||
sha256: str,
|
||||
file_path: str,
|
||||
model_data: dict,
|
||||
model_data: dict[str, Any],
|
||||
update_cache_func,
|
||||
):
|
||||
captured["model_data"] = json.loads(json.dumps(model_data))
|
||||
@@ -490,8 +492,8 @@ def test_fetch_civitai_hydrates_metadata_before_sync(
|
||||
await update_cache_func(file_path, file_path, model_data)
|
||||
return True, None
|
||||
|
||||
save_calls: list[tuple[str, dict]] = []
|
||||
captured: dict[str, dict] = {}
|
||||
save_calls: list[tuple[str, dict[str, Any]]] = []
|
||||
captured: dict[str, dict[str, Any]] = {}
|
||||
|
||||
monkeypatch.setattr(
|
||||
MetadataManager, "load_metadata", staticmethod(fake_load_metadata)
|
||||
|
||||
@@ -24,7 +24,7 @@ class StubEmbeddingService:
|
||||
@pytest.fixture
|
||||
def routes():
|
||||
handler = EmbeddingRoutes()
|
||||
handler.service = StubEmbeddingService()
|
||||
handler.service = StubEmbeddingService() # pyright: ignore[reportAttributeAccessIssue]
|
||||
return handler
|
||||
|
||||
|
||||
|
||||
@@ -3,9 +3,9 @@ from __future__ import annotations
|
||||
import json
|
||||
from contextlib import asynccontextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict
|
||||
from typing import Any, AsyncGenerator, Dict
|
||||
|
||||
from aiohttp import web
|
||||
from aiohttp import ClientResponse, web
|
||||
from aiohttp.test_utils import TestClient, TestServer
|
||||
|
||||
from py.routes.example_images_route_registrar import ExampleImagesRouteRegistrar
|
||||
@@ -140,14 +140,14 @@ class StubFileManager:
|
||||
|
||||
@dataclass
|
||||
class RegistrarHarness:
|
||||
client: TestClient
|
||||
client: TestClient[Any, Any]
|
||||
download_use_case: StubDownloadUseCase
|
||||
download_manager: StubDownloadManager
|
||||
import_use_case: StubImportUseCase
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def registrar_app() -> RegistrarHarness:
|
||||
async def registrar_app() -> AsyncGenerator[RegistrarHarness, None]:
|
||||
app = web.Application()
|
||||
|
||||
download_use_case = StubDownloadUseCase()
|
||||
@@ -158,8 +158,15 @@ async def registrar_app() -> RegistrarHarness:
|
||||
file_manager = StubFileManager()
|
||||
|
||||
handler_set = ExampleImagesHandlerSet(
|
||||
download=ExampleImagesDownloadHandler(download_use_case, download_manager),
|
||||
management=ExampleImagesManagementHandler(import_use_case, processor, cleanup_service),
|
||||
download=ExampleImagesDownloadHandler(
|
||||
download_use_case, # pyright: ignore[reportArgumentType]
|
||||
download_manager,
|
||||
),
|
||||
management=ExampleImagesManagementHandler(
|
||||
import_use_case, # pyright: ignore[reportArgumentType]
|
||||
processor,
|
||||
cleanup_service,
|
||||
),
|
||||
files=ExampleImagesFileHandler(file_manager),
|
||||
)
|
||||
|
||||
@@ -181,7 +188,7 @@ async def registrar_app() -> RegistrarHarness:
|
||||
await client.close()
|
||||
|
||||
|
||||
async def _json(response: web.StreamResponse) -> Dict[str, Any]:
|
||||
async def _json(response: ClientResponse) -> Dict[str, Any]:
|
||||
text = await response.text()
|
||||
return json.loads(text) if text else {}
|
||||
|
||||
@@ -358,7 +365,7 @@ async def test_check_example_images_needed_returns_error_on_exception():
|
||||
# Actually, we need to make the method raise an exception
|
||||
original_method = harness.download_manager.check_pending_models
|
||||
|
||||
async def failing_check(_model_types):
|
||||
async def failing_check(model_types):
|
||||
raise RuntimeError("Database connection failed")
|
||||
|
||||
harness.download_manager.check_pending_models = failing_check
|
||||
|
||||
@@ -3,7 +3,7 @@ from __future__ import annotations
|
||||
import json
|
||||
from contextlib import asynccontextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Tuple
|
||||
from typing import Any, AsyncGenerator, Dict, List, Tuple
|
||||
|
||||
from aiohttp import web
|
||||
from aiohttp.test_utils import TestClient, TestServer
|
||||
@@ -19,11 +19,16 @@ from py.routes.handlers.example_images_handlers import (
|
||||
)
|
||||
|
||||
|
||||
def _json_response(response) -> Any:
|
||||
"""Decode the JSON body of a handler response."""
|
||||
return json.loads(response.text)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ExampleImagesHarness:
|
||||
"""Container exposing the aiohttp client and stubbed collaborators."""
|
||||
|
||||
client: TestClient
|
||||
client: TestClient[Any, Any]
|
||||
download_manager: "StubDownloadManager"
|
||||
processor: "StubExampleImagesProcessor"
|
||||
file_manager: "StubExampleImagesFileManager"
|
||||
@@ -35,27 +40,27 @@ class StubDownloadManager:
|
||||
def __init__(self) -> None:
|
||||
self.calls: List[Tuple[str, Any]] = []
|
||||
|
||||
async def start_download(self, payload: Any) -> dict:
|
||||
async def start_download(self, payload: Any) -> Dict[str, Any]:
|
||||
self.calls.append(("start_download", payload))
|
||||
return {"operation": "start_download", "payload": payload}
|
||||
|
||||
async def get_status(self, request: web.Request) -> dict:
|
||||
async def get_status(self, request: web.Request) -> Dict[str, Any]:
|
||||
self.calls.append(("get_status", dict(request.query)))
|
||||
return {"operation": "get_status"}
|
||||
|
||||
async def pause_download(self, request: web.Request) -> dict:
|
||||
async def pause_download(self, request: web.Request) -> Dict[str, Any]:
|
||||
self.calls.append(("pause_download", None))
|
||||
return {"operation": "pause_download"}
|
||||
|
||||
async def resume_download(self, request: web.Request) -> dict:
|
||||
async def resume_download(self, request: web.Request) -> Dict[str, Any]:
|
||||
self.calls.append(("resume_download", None))
|
||||
return {"operation": "resume_download"}
|
||||
|
||||
async def stop_download(self, request: web.Request) -> dict:
|
||||
async def stop_download(self, request: web.Request) -> Dict[str, Any]:
|
||||
self.calls.append(("stop_download", None))
|
||||
return {"operation": "stop_download"}
|
||||
|
||||
async def start_force_download(self, payload: Any) -> dict:
|
||||
async def start_force_download(self, payload: Any) -> Dict[str, Any]:
|
||||
self.calls.append(("start_force_download", payload))
|
||||
return {"operation": "start_force_download", "payload": payload}
|
||||
|
||||
@@ -64,7 +69,7 @@ class StubExampleImagesProcessor:
|
||||
def __init__(self) -> None:
|
||||
self.calls: List[Tuple[str, Any]] = []
|
||||
|
||||
async def import_images(self, model_hash: str, files: List[str]) -> dict:
|
||||
async def import_images(self, model_hash: str, files: List[str]) -> Dict[str, Any]:
|
||||
payload = {"model_hash": model_hash, "file_paths": files}
|
||||
self.calls.append(("import_images", payload))
|
||||
return {"operation": "import_images", "payload": payload}
|
||||
@@ -122,7 +127,7 @@ class StubWebSocketManager:
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def example_images_app() -> ExampleImagesHarness:
|
||||
async def example_images_app() -> AsyncGenerator[ExampleImagesHarness, None]:
|
||||
"""Yield an ExampleImagesRoutes app wired with stubbed collaborators."""
|
||||
|
||||
download_manager = StubDownloadManager()
|
||||
@@ -133,10 +138,10 @@ async def example_images_app() -> ExampleImagesHarness:
|
||||
|
||||
controller = ExampleImagesRoutes(
|
||||
ws_manager=ws_manager,
|
||||
download_manager=download_manager,
|
||||
download_manager=download_manager, # pyright: ignore[reportArgumentType]
|
||||
processor=processor,
|
||||
file_manager=file_manager,
|
||||
cleanup_service=cleanup_service,
|
||||
file_manager=file_manager, # pyright: ignore[reportArgumentType]
|
||||
cleanup_service=cleanup_service, # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
|
||||
app = web.Application()
|
||||
@@ -323,23 +328,23 @@ async def test_download_handler_methods_delegate() -> None:
|
||||
def __init__(self) -> None:
|
||||
self.calls: List[Tuple[str, Any]] = []
|
||||
|
||||
async def get_status(self, request) -> dict:
|
||||
async def get_status(self, request) -> Dict[str, Any]:
|
||||
self.calls.append(("get_status", request))
|
||||
return {"status": "ok"}
|
||||
|
||||
async def pause_download(self, request) -> dict:
|
||||
async def pause_download(self, request) -> Dict[str, Any]:
|
||||
self.calls.append(("pause_download", request))
|
||||
return {"status": "paused"}
|
||||
|
||||
async def resume_download(self, request) -> dict:
|
||||
async def resume_download(self, request) -> Dict[str, Any]:
|
||||
self.calls.append(("resume_download", request))
|
||||
return {"status": "running"}
|
||||
|
||||
async def stop_download(self, request) -> dict:
|
||||
async def stop_download(self, request) -> Dict[str, Any]:
|
||||
self.calls.append(("stop_download", request))
|
||||
return {"status": "stopping"}
|
||||
|
||||
async def start_force_download(self, payload) -> dict:
|
||||
async def start_force_download(self, payload) -> Dict[str, Any]:
|
||||
self.calls.append(("start_force_download", payload))
|
||||
return {"status": "force", "payload": payload}
|
||||
|
||||
@@ -347,35 +352,50 @@ async def test_download_handler_methods_delegate() -> None:
|
||||
def __init__(self) -> None:
|
||||
self.payloads: List[Any] = []
|
||||
|
||||
async def execute(self, payload: dict) -> dict:
|
||||
async def execute(self, payload: Dict[str, Any]) -> Dict[str, Any]:
|
||||
self.payloads.append(payload)
|
||||
return {"status": "started", "payload": payload}
|
||||
|
||||
class DummyRequest:
|
||||
def __init__(self, payload: dict) -> None:
|
||||
def __init__(self, payload: Dict[str, Any]) -> None:
|
||||
self._payload = payload
|
||||
self.query = {}
|
||||
|
||||
async def json(self) -> dict:
|
||||
async def json(self) -> Dict[str, Any]:
|
||||
return self._payload
|
||||
|
||||
recorder = Recorder()
|
||||
use_case = StubDownloadUseCase()
|
||||
handler = ExampleImagesDownloadHandler(use_case, recorder)
|
||||
handler = ExampleImagesDownloadHandler(
|
||||
use_case, # pyright: ignore[reportArgumentType]
|
||||
recorder,
|
||||
)
|
||||
request = DummyRequest({"foo": "bar"})
|
||||
|
||||
download_response = await handler.download_example_images(request)
|
||||
assert json.loads(download_response.text) == {"status": "started", "payload": {"foo": "bar"}}
|
||||
status_response = await handler.get_example_images_status(request)
|
||||
assert json.loads(status_response.text) == {"status": "ok"}
|
||||
pause_response = await handler.pause_example_images(request)
|
||||
assert json.loads(pause_response.text) == {"status": "paused"}
|
||||
resume_response = await handler.resume_example_images(request)
|
||||
assert json.loads(resume_response.text) == {"status": "running"}
|
||||
stop_response = await handler.stop_example_images(request)
|
||||
assert json.loads(stop_response.text) == {"status": "stopping"}
|
||||
force_response = await handler.force_download_example_images(request)
|
||||
assert json.loads(force_response.text) == {"status": "force", "payload": {"foo": "bar"}}
|
||||
download_response = await handler.download_example_images(
|
||||
request # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
assert _json_response(download_response) == {"status": "started", "payload": {"foo": "bar"}}
|
||||
status_response = await handler.get_example_images_status(
|
||||
request # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
assert _json_response(status_response) == {"status": "ok"}
|
||||
pause_response = await handler.pause_example_images(
|
||||
request # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
assert _json_response(pause_response) == {"status": "paused"}
|
||||
resume_response = await handler.resume_example_images(
|
||||
request # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
assert _json_response(resume_response) == {"status": "running"}
|
||||
stop_response = await handler.stop_example_images(
|
||||
request # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
assert _json_response(stop_response) == {"status": "stopping"}
|
||||
force_response = await handler.force_download_example_images(
|
||||
request # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
assert _json_response(force_response) == {"status": "force", "payload": {"foo": "bar"}}
|
||||
|
||||
assert use_case.payloads == [{"foo": "bar"}]
|
||||
assert recorder.calls == [
|
||||
@@ -393,7 +413,7 @@ async def test_management_handler_methods_delegate() -> None:
|
||||
def __init__(self) -> None:
|
||||
self.requests: List[Any] = []
|
||||
|
||||
async def execute(self, request: Any) -> dict:
|
||||
async def execute(self, request: Any) -> Dict[str, Any]:
|
||||
self.requests.append(request)
|
||||
return {"status": "imported"}
|
||||
|
||||
@@ -412,15 +432,23 @@ async def test_management_handler_methods_delegate() -> None:
|
||||
recorder = Recorder()
|
||||
cleanup_service = StubExampleImagesCleanupService()
|
||||
use_case = StubImportUseCase()
|
||||
handler = ExampleImagesManagementHandler(use_case, recorder, cleanup_service)
|
||||
handler = ExampleImagesManagementHandler(
|
||||
use_case, # pyright: ignore[reportArgumentType]
|
||||
recorder,
|
||||
cleanup_service,
|
||||
)
|
||||
request = object()
|
||||
|
||||
import_response = await handler.import_example_images(request)
|
||||
assert json.loads(import_response.text) == {"status": "imported"}
|
||||
assert await handler.delete_example_image(request) == "delete"
|
||||
import_response = await handler.import_example_images(
|
||||
request # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
assert _json_response(import_response) == {"status": "imported"}
|
||||
assert await handler.delete_example_image(request) == "delete" # pyright: ignore[reportArgumentType]
|
||||
cleanup_service.result = {"success": True}
|
||||
cleanup_response = await handler.cleanup_example_image_folders(request)
|
||||
assert json.loads(cleanup_response.text) == {"success": True}
|
||||
cleanup_response = await handler.cleanup_example_image_folders(
|
||||
request # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
assert _json_response(cleanup_response) == {"success": True}
|
||||
assert use_case.requests == [request]
|
||||
assert recorder.calls == [("delete_custom_image", request)]
|
||||
assert len(cleanup_service.calls) == 1
|
||||
@@ -448,9 +476,9 @@ async def test_file_handler_methods_delegate() -> None:
|
||||
handler = ExampleImagesFileHandler(recorder)
|
||||
request = object()
|
||||
|
||||
assert await handler.open_example_images_folder(request) == "open"
|
||||
assert await handler.get_example_image_files(request) == "files"
|
||||
assert await handler.has_example_images(request) == "has"
|
||||
assert await handler.open_example_images_folder(request) == "open" # pyright: ignore[reportArgumentType]
|
||||
assert await handler.get_example_image_files(request) == "files" # pyright: ignore[reportArgumentType]
|
||||
assert await handler.has_example_images(request) == "has" # pyright: ignore[reportArgumentType]
|
||||
assert recorder.calls == [
|
||||
("open_folder", request),
|
||||
("get_files", request),
|
||||
@@ -483,9 +511,16 @@ def test_handler_set_route_mapping_includes_all_handlers() -> None:
|
||||
async def set_example_image_nsfw_level(self, request):
|
||||
return {}
|
||||
|
||||
download = ExampleImagesDownloadHandler(DummyUseCase(), DummyManager())
|
||||
download = ExampleImagesDownloadHandler(
|
||||
DummyUseCase(), # pyright: ignore[reportArgumentType]
|
||||
DummyManager(),
|
||||
)
|
||||
cleanup_service = StubExampleImagesCleanupService()
|
||||
management = ExampleImagesManagementHandler(DummyUseCase(), DummyProcessor(), cleanup_service)
|
||||
management = ExampleImagesManagementHandler(
|
||||
DummyUseCase(), # pyright: ignore[reportArgumentType]
|
||||
DummyProcessor(),
|
||||
cleanup_service,
|
||||
)
|
||||
files = ExampleImagesFileHandler(object())
|
||||
handler_set = ExampleImagesHandlerSet(
|
||||
download=download,
|
||||
|
||||
@@ -4,6 +4,7 @@ import asyncio
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from aiohttp import web
|
||||
@@ -160,7 +161,7 @@ async def test_lora_manager_lifecycle(monkeypatch: pytest.MonkeyPatch, tmp_path:
|
||||
)
|
||||
|
||||
original_create_task = asyncio.create_task
|
||||
scheduled_tasks: list[asyncio.Task] = []
|
||||
scheduled_tasks: list[asyncio.Task[Any]] = []
|
||||
|
||||
def track_create_task(coro, *, name=None):
|
||||
task = original_create_task(coro, name=name)
|
||||
|
||||
@@ -5,7 +5,7 @@ from unittest.mock import MagicMock
|
||||
import pytest
|
||||
|
||||
from py.routes.lora_routes import LoraRoutes
|
||||
from server import PromptServer
|
||||
from server import PromptServer # pyright: ignore[reportMissingImports]
|
||||
|
||||
|
||||
class DummyRequest:
|
||||
@@ -20,14 +20,8 @@ class DummyRequest:
|
||||
|
||||
class StubLoraService:
|
||||
def __init__(self):
|
||||
self.notes = {}
|
||||
self.trigger_words = {}
|
||||
self.usage_tips = {}
|
||||
self.previews = {}
|
||||
self.civitai = {}
|
||||
|
||||
async def get_lora_notes(self, name):
|
||||
return self.notes.get(name)
|
||||
|
||||
async def get_lora_trigger_words(self, name):
|
||||
return self.trigger_words.get(name, [])
|
||||
@@ -35,57 +29,14 @@ class StubLoraService:
|
||||
async def get_lora_usage_tips_by_relative_path(self, path):
|
||||
return self.usage_tips.get(path)
|
||||
|
||||
async def get_lora_preview_url(self, name):
|
||||
return self.previews.get(name)
|
||||
|
||||
async def get_lora_civitai_url(self, name):
|
||||
return self.civitai.get(name, {"civitai_url": ""})
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def routes():
|
||||
handler = LoraRoutes()
|
||||
handler.service = StubLoraService()
|
||||
handler.service = StubLoraService() # pyright: ignore[reportAttributeAccessIssue]
|
||||
return handler
|
||||
|
||||
|
||||
async def test_get_lora_notes_success(routes):
|
||||
routes.service.notes["demo"] = "Great notes"
|
||||
request = DummyRequest(query={"name": "demo"})
|
||||
|
||||
response = await routes.get_lora_notes(request)
|
||||
payload = json.loads(response.text)
|
||||
|
||||
assert payload == {"success": True, "notes": "Great notes"}
|
||||
|
||||
|
||||
async def test_get_lora_notes_missing_name(routes):
|
||||
response = await routes.get_lora_notes(DummyRequest())
|
||||
assert response.status == 400
|
||||
assert response.text == "Lora file name is required"
|
||||
|
||||
|
||||
async def test_get_lora_notes_not_found(routes):
|
||||
response = await routes.get_lora_notes(DummyRequest(query={"name": "missing"}))
|
||||
payload = json.loads(response.text)
|
||||
assert response.status == 404
|
||||
assert payload == {"success": False, "error": "LoRA not found in cache"}
|
||||
|
||||
|
||||
async def test_get_lora_notes_error(routes, monkeypatch):
|
||||
async def failing(*_args, **_kwargs):
|
||||
raise RuntimeError("boom")
|
||||
|
||||
routes.service.get_lora_notes = failing
|
||||
|
||||
response = await routes.get_lora_notes(DummyRequest(query={"name": "demo"}))
|
||||
payload = json.loads(response.text)
|
||||
|
||||
assert response.status == 500
|
||||
assert payload["success"] is False
|
||||
assert payload["error"] == "boom"
|
||||
|
||||
|
||||
async def test_get_lora_trigger_words_success(routes):
|
||||
routes.service.trigger_words["demo"] = ["trigger"]
|
||||
response = await routes.get_lora_trigger_words(DummyRequest(query={"name": "demo"}))
|
||||
@@ -133,55 +84,6 @@ async def test_get_usage_tips_error(routes):
|
||||
assert payload["success"] is False
|
||||
|
||||
|
||||
async def test_get_preview_url_success(routes):
|
||||
routes.service.previews["demo"] = "http://preview"
|
||||
response = await routes.get_lora_preview_url(DummyRequest(query={"name": "demo"}))
|
||||
payload = json.loads(response.text)
|
||||
assert payload == {"success": True, "preview_url": "http://preview"}
|
||||
|
||||
|
||||
async def test_get_preview_url_missing(routes):
|
||||
response = await routes.get_lora_preview_url(DummyRequest())
|
||||
assert response.status == 400
|
||||
|
||||
|
||||
async def test_get_preview_url_not_found(routes):
|
||||
response = await routes.get_lora_preview_url(DummyRequest(query={"name": "missing"}))
|
||||
payload = json.loads(response.text)
|
||||
assert response.status == 404
|
||||
assert payload["success"] is False
|
||||
|
||||
|
||||
async def test_get_civitai_url_success(routes):
|
||||
routes.service.civitai["demo"] = {"civitai_url": "https://civitai.com"}
|
||||
response = await routes.get_lora_civitai_url(DummyRequest(query={"name": "demo"}))
|
||||
payload = json.loads(response.text)
|
||||
assert payload == {"success": True, "civitai_url": "https://civitai.com"}
|
||||
|
||||
|
||||
async def test_get_civitai_url_missing(routes):
|
||||
response = await routes.get_lora_civitai_url(DummyRequest())
|
||||
assert response.status == 400
|
||||
|
||||
|
||||
async def test_get_civitai_url_not_found(routes):
|
||||
response = await routes.get_lora_civitai_url(DummyRequest(query={"name": "missing"}))
|
||||
payload = json.loads(response.text)
|
||||
assert response.status == 404
|
||||
assert payload["success"] is False
|
||||
|
||||
|
||||
async def test_get_civitai_url_error(routes):
|
||||
async def failing(*_args, **_kwargs):
|
||||
raise RuntimeError("oops")
|
||||
|
||||
routes.service.get_lora_civitai_url = failing
|
||||
response = await routes.get_lora_civitai_url(DummyRequest(query={"name": "demo"}))
|
||||
payload = json.loads(response.text)
|
||||
assert response.status == 500
|
||||
assert payload["success"] is False
|
||||
|
||||
|
||||
async def test_get_trigger_words_broadcasts(monkeypatch, routes):
|
||||
send_mock = MagicMock()
|
||||
PromptServer.instance = SimpleNamespace(send_sync=send_mock)
|
||||
|
||||
@@ -5,6 +5,7 @@ import os
|
||||
import subprocess
|
||||
import zipfile
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
from unittest.mock import patch, MagicMock
|
||||
|
||||
import pytest
|
||||
@@ -34,6 +35,13 @@ from py.routes.misc_route_registrar import MISC_ROUTE_DEFINITIONS, MiscRouteRegi
|
||||
from py.routes.misc_routes import MiscRoutes
|
||||
|
||||
|
||||
def _json_payload(response) -> dict[str, Any]:
|
||||
"""Decode the JSON body of a web.Response, asserting it is not null."""
|
||||
text = response.text
|
||||
assert text is not None
|
||||
return json.loads(text)
|
||||
|
||||
|
||||
class FakeRequest:
|
||||
def __init__(self, *, json_data=None, query=None, method="POST"):
|
||||
self._json_data = json_data or {}
|
||||
@@ -129,8 +137,8 @@ async def test_get_settings_excludes_no_sync_keys():
|
||||
downloader_factory=dummy_downloader_factory,
|
||||
)
|
||||
|
||||
response = await handler.get_settings(FakeRequest())
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.get_settings(FakeRequest()) # pyright: ignore[reportArgumentType]
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert payload["success"] is True
|
||||
# Regular settings should be synced
|
||||
@@ -155,8 +163,8 @@ async def test_update_settings_rejects_missing_example_path(tmp_path):
|
||||
missing_path = tmp_path / "does-not-exist"
|
||||
request = FakeRequest(json_data={"example_images_path": str(missing_path)})
|
||||
|
||||
response = await handler.update_settings(request)
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.update_settings(request) # pyright: ignore[reportArgumentType]
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert payload["success"] is False
|
||||
assert "Path does not exist" in payload["error"]
|
||||
@@ -181,9 +189,9 @@ async def test_doctor_handler_reports_key_cache_and_ui_issues():
|
||||
)
|
||||
|
||||
response = await handler.get_doctor_diagnostics(
|
||||
FakeRequest(query={"clientVersion": "1.2.2-client"}, method="GET")
|
||||
FakeRequest(query={"clientVersion": "1.2.2-client"}, method="GET") # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
payload = json.loads(response.text)
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert payload["success"] is True
|
||||
assert payload["summary"]["status"] == "error"
|
||||
@@ -209,8 +217,8 @@ async def test_doctor_handler_can_repair_cache():
|
||||
scanner_factories=(("lora", "LoRAs", scanner_factory),),
|
||||
)
|
||||
|
||||
response = await handler.repair_doctor_cache(FakeRequest())
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.repair_doctor_cache(FakeRequest()) # pyright: ignore[reportArgumentType]
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert response.status == 200
|
||||
assert payload["success"] is True
|
||||
@@ -230,7 +238,7 @@ async def test_doctor_handler_exports_support_bundle():
|
||||
)
|
||||
|
||||
response = await handler.export_doctor_bundle(
|
||||
FakeRequest(
|
||||
FakeRequest( # pyright: ignore[reportArgumentType]
|
||||
json_data={
|
||||
"summary": {"status": "warning"},
|
||||
"diagnostics": [{"id": "cache_health", "status": "warning"}],
|
||||
@@ -241,6 +249,7 @@ async def test_doctor_handler_exports_support_bundle():
|
||||
)
|
||||
|
||||
assert response.status == 200
|
||||
assert isinstance(response.body, bytes)
|
||||
with zipfile.ZipFile(io.BytesIO(response.body), "r") as archive:
|
||||
names = set(archive.namelist())
|
||||
assert "doctor-report.json" in names
|
||||
@@ -263,7 +272,7 @@ async def test_doctor_handler_redacts_string_secrets_in_bundle():
|
||||
)
|
||||
|
||||
response = await handler.export_doctor_bundle(
|
||||
FakeRequest(
|
||||
FakeRequest( # pyright: ignore[reportArgumentType]
|
||||
json_data={
|
||||
"frontend_logs": [
|
||||
{
|
||||
@@ -276,6 +285,7 @@ async def test_doctor_handler_redacts_string_secrets_in_bundle():
|
||||
)
|
||||
|
||||
assert response.status == 200
|
||||
assert isinstance(response.body, bytes)
|
||||
with zipfile.ZipFile(io.BytesIO(response.body), "r") as archive:
|
||||
frontend_logs = archive.read("frontend-console.json").decode("utf-8")
|
||||
assert "abcdef123456" not in frontend_logs
|
||||
@@ -308,7 +318,7 @@ async def test_doctor_handler_redacts_json_shaped_string_secrets_in_bundle():
|
||||
}
|
||||
|
||||
response = await handler.export_doctor_bundle(
|
||||
FakeRequest(
|
||||
FakeRequest( # pyright: ignore[reportArgumentType]
|
||||
json_data={
|
||||
"frontend_logs": [
|
||||
{
|
||||
@@ -321,6 +331,7 @@ async def test_doctor_handler_redacts_json_shaped_string_secrets_in_bundle():
|
||||
)
|
||||
|
||||
assert response.status == 200
|
||||
assert isinstance(response.body, bytes)
|
||||
with zipfile.ZipFile(io.BytesIO(response.body), "r") as archive:
|
||||
frontend_logs = archive.read("frontend-console.json").decode("utf-8")
|
||||
backend_logs = archive.read("backend-logs.txt").decode("utf-8")
|
||||
@@ -359,9 +370,10 @@ async def test_doctor_handler_exports_backend_session_logs_from_helper():
|
||||
"notes": [],
|
||||
}
|
||||
|
||||
response = await handler.export_doctor_bundle(FakeRequest(json_data={}))
|
||||
response = await handler.export_doctor_bundle(FakeRequest(json_data={})) # pyright: ignore[reportArgumentType]
|
||||
|
||||
assert response.status == 200
|
||||
assert isinstance(response.body, bytes)
|
||||
with zipfile.ZipFile(io.BytesIO(response.body), "r") as archive:
|
||||
backend_logs = archive.read("backend-logs.txt").decode("utf-8")
|
||||
backend_source = json.loads(
|
||||
@@ -449,14 +461,14 @@ async def test_backup_handler_returns_status_and_exports(monkeypatch):
|
||||
|
||||
handler = BackupHandler(backup_service_factory=factory)
|
||||
|
||||
status_response = await handler.get_backup_status(FakeRequest())
|
||||
status_payload = json.loads(status_response.text)
|
||||
status_response = await handler.get_backup_status(FakeRequest()) # pyright: ignore[reportArgumentType]
|
||||
status_payload = _json_payload(status_response)
|
||||
assert status_payload["success"] is True
|
||||
assert status_payload["status"]["backupDir"] == "/tmp/backups"
|
||||
assert status_payload["status"]["enabled"] is True
|
||||
assert status_payload["snapshots"][0]["name"] == "backup.zip"
|
||||
|
||||
export_response = await handler.export_backup(FakeRequest())
|
||||
export_response = await handler.export_backup(FakeRequest()) # pyright: ignore[reportArgumentType]
|
||||
assert export_response.status == 200
|
||||
assert export_response.body == b"zip-bytes"
|
||||
|
||||
@@ -476,8 +488,8 @@ async def test_backup_handler_rejects_missing_import_archive():
|
||||
async def read(self):
|
||||
return b""
|
||||
|
||||
response = await handler.import_backup(EmptyRequest())
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.import_backup(EmptyRequest()) # pyright: ignore[reportArgumentType]
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert response.status == 400
|
||||
assert payload["success"] is False
|
||||
@@ -504,8 +516,8 @@ async def test_open_backup_location_uses_settings_directory(tmp_path, monkeypatc
|
||||
monkeypatch.setattr("py.routes.handlers.misc_handlers._is_docker", lambda: False)
|
||||
monkeypatch.setattr("py.routes.handlers.misc_handlers._is_wsl", lambda: False)
|
||||
|
||||
response = await handler.open_backup_location(FakeRequest())
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.open_backup_location(FakeRequest()) # pyright: ignore[reportArgumentType]
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert response.status == 200
|
||||
assert payload["success"] is True
|
||||
@@ -535,8 +547,8 @@ async def test_open_wildcards_location_creates_and_opens_directory(tmp_path, mon
|
||||
else str(wildcards_dir),
|
||||
)
|
||||
|
||||
response = await handler.open_wildcards_location(FakeRequest())
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.open_wildcards_location(FakeRequest()) # pyright: ignore[reportArgumentType]
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert response.status == 200
|
||||
assert payload["success"] is True
|
||||
@@ -564,7 +576,7 @@ class RecordingRouter:
|
||||
|
||||
def test_misc_route_registrar_registers_all_routes():
|
||||
app = SimpleNamespace(router=RecordingRouter())
|
||||
registrar = MiscRouteRegistrar(app) # type: ignore[arg-type]
|
||||
registrar = MiscRouteRegistrar(app) # pyright: ignore[reportArgumentType]
|
||||
|
||||
async def dummy_handler(_request):
|
||||
return web.Response()
|
||||
@@ -586,7 +598,7 @@ class FakePromptServer:
|
||||
sent = []
|
||||
|
||||
class Instance:
|
||||
sockets: dict = {}
|
||||
sockets: dict[str, Any] = {}
|
||||
|
||||
def send_sync(self, event, payload, sid=None):
|
||||
FakePromptServer.sent.append((event, payload))
|
||||
@@ -599,7 +611,7 @@ async def test_register_nodes_requires_graph_id():
|
||||
node_registry = NodeRegistry()
|
||||
handler = NodeRegistryHandler(
|
||||
node_registry=node_registry,
|
||||
prompt_server=FakePromptServer,
|
||||
prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
|
||||
standalone_mode=False,
|
||||
)
|
||||
|
||||
@@ -609,8 +621,8 @@ async def test_register_nodes_requires_graph_id():
|
||||
"client_id": "test-client-1",
|
||||
}
|
||||
)
|
||||
response = await handler.register_nodes(request)
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.register_nodes(request) # pyright: ignore[reportArgumentType]
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert response.status == 400
|
||||
assert payload["success"] is False
|
||||
@@ -622,7 +634,7 @@ async def test_register_nodes_stores_graph_identifier():
|
||||
node_registry = NodeRegistry()
|
||||
handler = NodeRegistryHandler(
|
||||
node_registry=node_registry,
|
||||
prompt_server=FakePromptServer,
|
||||
prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
|
||||
standalone_mode=False,
|
||||
)
|
||||
|
||||
@@ -641,8 +653,8 @@ async def test_register_nodes_stores_graph_identifier():
|
||||
}
|
||||
)
|
||||
|
||||
response = await handler.register_nodes(request)
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.register_nodes(request) # pyright: ignore[reportArgumentType]
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert payload["success"] is True
|
||||
|
||||
@@ -659,7 +671,7 @@ async def test_register_nodes_defaults_graph_name_to_none():
|
||||
node_registry = NodeRegistry()
|
||||
handler = NodeRegistryHandler(
|
||||
node_registry=node_registry,
|
||||
prompt_server=FakePromptServer,
|
||||
prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
|
||||
standalone_mode=False,
|
||||
)
|
||||
|
||||
@@ -677,8 +689,8 @@ async def test_register_nodes_defaults_graph_name_to_none():
|
||||
}
|
||||
)
|
||||
|
||||
response = await handler.register_nodes(request)
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.register_nodes(request) # pyright: ignore[reportArgumentType]
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert payload["success"] is True
|
||||
|
||||
@@ -692,7 +704,7 @@ async def test_register_nodes_includes_capabilities():
|
||||
node_registry = NodeRegistry()
|
||||
handler = NodeRegistryHandler(
|
||||
node_registry=node_registry,
|
||||
prompt_server=FakePromptServer,
|
||||
prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
|
||||
standalone_mode=False,
|
||||
)
|
||||
|
||||
@@ -714,8 +726,8 @@ async def test_register_nodes_includes_capabilities():
|
||||
}
|
||||
)
|
||||
|
||||
response = await handler.register_nodes(request)
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.register_nodes(request) # pyright: ignore[reportArgumentType]
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert payload["success"] is True
|
||||
|
||||
@@ -734,7 +746,7 @@ async def test_register_nodes_accepts_compound_node_ids():
|
||||
node_registry = NodeRegistry()
|
||||
handler = NodeRegistryHandler(
|
||||
node_registry=node_registry,
|
||||
prompt_server=FakePromptServer,
|
||||
prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
|
||||
standalone_mode=False,
|
||||
)
|
||||
|
||||
@@ -758,8 +770,8 @@ async def test_register_nodes_accepts_compound_node_ids():
|
||||
}
|
||||
)
|
||||
|
||||
response = await handler.register_nodes(request)
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.register_nodes(request) # pyright: ignore[reportArgumentType]
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert response.status == 200
|
||||
assert payload["success"] is True
|
||||
@@ -778,11 +790,11 @@ async def test_register_nodes_accepts_compound_node_ids():
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_node_widget_sends_payload():
|
||||
send_calls: list[tuple[str, dict]] = []
|
||||
send_calls: list[tuple[str, dict[str, Any]]] = []
|
||||
|
||||
class RecordingPromptServer:
|
||||
class Instance:
|
||||
sockets: dict = {}
|
||||
sockets: dict[str, Any] = {}
|
||||
|
||||
def send_sync(self, event, payload, sid=None):
|
||||
send_calls.append((event, payload))
|
||||
@@ -791,7 +803,7 @@ async def test_update_node_widget_sends_payload():
|
||||
|
||||
handler = NodeRegistryHandler(
|
||||
node_registry=NodeRegistry(),
|
||||
prompt_server=RecordingPromptServer,
|
||||
prompt_server=RecordingPromptServer, # pyright: ignore[reportArgumentType]
|
||||
standalone_mode=False,
|
||||
)
|
||||
|
||||
@@ -803,8 +815,8 @@ async def test_update_node_widget_sends_payload():
|
||||
}
|
||||
)
|
||||
|
||||
response = await handler.update_node_widget(request)
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.update_node_widget(request) # pyright: ignore[reportArgumentType]
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert response.status == 200
|
||||
assert payload["success"] is True
|
||||
@@ -824,18 +836,18 @@ async def test_update_node_widget_sends_payload():
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_lora_code_includes_graph_identifier():
|
||||
send_calls: list[tuple[str, dict]] = []
|
||||
send_calls: list[tuple[str, dict[str, Any]]] = []
|
||||
|
||||
class RecordingPromptServer:
|
||||
class Instance:
|
||||
sockets: dict = {}
|
||||
sockets: dict[str, Any] = {}
|
||||
|
||||
def send_sync(self, event, payload, sid=None):
|
||||
send_calls.append((event, payload))
|
||||
|
||||
instance = Instance()
|
||||
|
||||
handler = LoraCodeHandler(RecordingPromptServer)
|
||||
handler = LoraCodeHandler(RecordingPromptServer) # pyright: ignore[reportArgumentType]
|
||||
|
||||
request = FakeRequest(
|
||||
json_data={
|
||||
@@ -845,8 +857,8 @@ async def test_update_lora_code_includes_graph_identifier():
|
||||
}
|
||||
)
|
||||
|
||||
response = await handler.update_lora_code(request)
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.update_lora_code(request) # pyright: ignore[reportArgumentType]
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert payload["success"] is True
|
||||
assert payload["results"] == [
|
||||
@@ -913,7 +925,7 @@ class FakeUserModelsProvider(FakeMetadataProvider):
|
||||
self.next_cursor = next_cursor
|
||||
self.estimated_total = estimated_total
|
||||
self.received_usernames: list[str] = []
|
||||
self.received_cursors: list = []
|
||||
self.received_cursors: list[Any] = []
|
||||
|
||||
async def get_user_models(self, username, cursor=None):
|
||||
self.received_usernames.append(username)
|
||||
@@ -924,7 +936,7 @@ class FakeUserModelsProvider(FakeMetadataProvider):
|
||||
return self.estimated_total
|
||||
|
||||
|
||||
async def fake_metadata_provider_factory():
|
||||
async def fake_metadata_provider_factory() -> Any:
|
||||
return FakeMetadataProvider()
|
||||
|
||||
|
||||
@@ -949,8 +961,8 @@ async def fake_metadata_archive_manager_factory():
|
||||
class FakeDownloadHistoryService:
|
||||
def __init__(self, downloaded_by_type=None):
|
||||
self.downloaded_by_type = downloaded_by_type or {}
|
||||
self.marked_downloaded: list[tuple] = []
|
||||
self.marked_not_downloaded: list[tuple] = []
|
||||
self.marked_downloaded: list[tuple[Any, ...]] = []
|
||||
self.marked_not_downloaded: list[tuple[Any, ...]] = []
|
||||
|
||||
async def has_been_downloaded(self, model_type, version_id):
|
||||
return version_id in self.downloaded_by_type.get(model_type, set())
|
||||
@@ -1012,20 +1024,20 @@ async def test_misc_routes_bind_produces_expected_handlers():
|
||||
|
||||
controller = MiscRoutes(
|
||||
settings_service=DummySettings(),
|
||||
usage_stats_factory=lambda: SimpleNamespace(
|
||||
usage_stats_factory=lambda: SimpleNamespace( # pyright: ignore[reportArgumentType]
|
||||
process_execution=noop_async, get_stats=noop_async
|
||||
),
|
||||
prompt_server=FakePromptServer,
|
||||
prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
|
||||
service_registry_adapter=service_registry_adapter,
|
||||
metadata_provider_factory=fake_metadata_provider_factory,
|
||||
metadata_archive_manager_factory=fake_metadata_archive_manager_factory,
|
||||
metadata_provider_updater=noop_async,
|
||||
downloader_factory=dummy_downloader_factory,
|
||||
registrar_factory=registrar_factory,
|
||||
registrar_factory=registrar_factory, # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
|
||||
app = SimpleNamespace(router=RecordingRouter())
|
||||
controller.bind(app) # type: ignore[arg-type]
|
||||
controller.bind(app) # pyright: ignore[reportArgumentType]
|
||||
|
||||
assert recorded_registrars, "Expected registrar to be created"
|
||||
mapping = recorded_registrars[0].registered_mapping
|
||||
@@ -1106,7 +1118,7 @@ async def test_get_civitai_user_models_marks_library_versions():
|
||||
|
||||
provider = FakeUserModelsProvider(models)
|
||||
|
||||
async def provider_factory():
|
||||
async def provider_factory() -> Any:
|
||||
return provider
|
||||
|
||||
lora_scanner = FakeExistenceScanner({101})
|
||||
@@ -1133,9 +1145,9 @@ async def test_get_civitai_user_models_marks_library_versions():
|
||||
)
|
||||
|
||||
response = await handler.get_civitai_user_models(
|
||||
FakeRequest(query={"username": "pixel"})
|
||||
FakeRequest(query={"username": "pixel"}) # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
payload = json.loads(response.text)
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert payload["success"] is True
|
||||
assert payload["username"] == "pixel"
|
||||
@@ -1239,7 +1251,7 @@ async def test_get_civitai_user_models_rewrites_civitai_previews():
|
||||
|
||||
provider = FakeUserModelsProvider(models)
|
||||
|
||||
async def provider_factory():
|
||||
async def provider_factory() -> Any:
|
||||
return provider
|
||||
|
||||
handler = ModelLibraryHandler(
|
||||
@@ -1253,9 +1265,9 @@ async def test_get_civitai_user_models_rewrites_civitai_previews():
|
||||
)
|
||||
|
||||
response = await handler.get_civitai_user_models(
|
||||
FakeRequest(query={"username": "pixel"})
|
||||
FakeRequest(query={"username": "pixel"}) # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
payload = json.loads(response.text)
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert payload["success"] is True
|
||||
previews_by_version = {
|
||||
@@ -1275,7 +1287,7 @@ async def test_get_civitai_user_models_rewrites_civitai_previews():
|
||||
async def test_get_civitai_user_models_requires_username():
|
||||
provider = FakeUserModelsProvider([])
|
||||
|
||||
async def provider_factory():
|
||||
async def provider_factory() -> Any:
|
||||
return provider
|
||||
|
||||
handler = ModelLibraryHandler(
|
||||
@@ -1288,8 +1300,8 @@ async def test_get_civitai_user_models_requires_username():
|
||||
metadata_provider_factory=provider_factory,
|
||||
)
|
||||
|
||||
response = await handler.get_civitai_user_models(FakeRequest())
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.get_civitai_user_models(FakeRequest()) # pyright: ignore[reportArgumentType]
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert response.status == 400
|
||||
assert payload["success"] is False
|
||||
@@ -1318,7 +1330,7 @@ async def test_get_civitai_user_models_returns_pagination_fields():
|
||||
|
||||
provider = FakeUserModelsProvider(models, next_cursor="cursor-token", estimated_total=2140)
|
||||
|
||||
async def provider_factory():
|
||||
async def provider_factory() -> Any:
|
||||
return provider
|
||||
|
||||
handler = ModelLibraryHandler(
|
||||
@@ -1332,9 +1344,9 @@ async def test_get_civitai_user_models_returns_pagination_fields():
|
||||
)
|
||||
|
||||
response = await handler.get_civitai_user_models(
|
||||
FakeRequest(query={"username": "pixel"})
|
||||
FakeRequest(query={"username": "pixel"}) # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
payload = json.loads(response.text)
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert response.status == 200
|
||||
assert payload["success"] is True
|
||||
@@ -1351,7 +1363,7 @@ async def test_get_civitai_user_models_returns_pagination_fields():
|
||||
async def test_get_civitai_user_models_passes_cursor_and_omits_estimate():
|
||||
provider = FakeUserModelsProvider([], next_cursor=None, estimated_total=999)
|
||||
|
||||
async def provider_factory():
|
||||
async def provider_factory() -> Any:
|
||||
return provider
|
||||
|
||||
handler = ModelLibraryHandler(
|
||||
@@ -1365,9 +1377,9 @@ async def test_get_civitai_user_models_passes_cursor_and_omits_estimate():
|
||||
)
|
||||
|
||||
response = await handler.get_civitai_user_models(
|
||||
FakeRequest(query={"username": "pixel", "cursor": "opaque-token"})
|
||||
FakeRequest(query={"username": "pixel", "cursor": "opaque-token"}) # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
payload = json.loads(response.text)
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert response.status == 200
|
||||
assert payload["success"] is True
|
||||
@@ -1391,10 +1403,10 @@ def test_ensure_handler_mapping_caches_result():
|
||||
|
||||
controller = MiscRoutes(
|
||||
settings_service=DummySettings(),
|
||||
usage_stats_factory=lambda: SimpleNamespace(
|
||||
usage_stats_factory=lambda: SimpleNamespace( # pyright: ignore[reportArgumentType]
|
||||
process_execution=noop_async, get_stats=noop_async
|
||||
),
|
||||
prompt_server=FakePromptServer,
|
||||
prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
|
||||
service_registry_adapter=ServiceRegistryAdapter(
|
||||
get_lora_scanner=fake_scanner_factory,
|
||||
get_checkpoint_scanner=fake_scanner_factory,
|
||||
@@ -1405,7 +1417,7 @@ def test_ensure_handler_mapping_caches_result():
|
||||
metadata_archive_manager_factory=fake_metadata_archive_manager_factory,
|
||||
metadata_provider_updater=noop_async,
|
||||
downloader_factory=dummy_downloader_factory,
|
||||
handler_set_factory=RecordingHandlerSet,
|
||||
handler_set_factory=RecordingHandlerSet, # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
|
||||
first_mapping = controller._ensure_handler_mapping()
|
||||
@@ -1447,8 +1459,8 @@ async def test_check_model_exists_returns_local_versions():
|
||||
metadata_provider_factory=fake_metadata_provider_factory,
|
||||
)
|
||||
|
||||
response = await handler.check_model_exists(FakeRequest(query={"modelId": "5"}))
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.check_model_exists(FakeRequest(query={"modelId": "5"})) # pyright: ignore[reportArgumentType]
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert payload["success"] is True
|
||||
assert payload["modelType"] == "lora"
|
||||
@@ -1461,7 +1473,7 @@ async def test_check_model_exists_returns_local_versions():
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_model_exists_model_id_only_does_not_call_metadata_provider():
|
||||
async def metadata_provider_factory():
|
||||
async def metadata_provider_factory() -> Any:
|
||||
raise AssertionError("metadata provider should not be called for modelId-only checks")
|
||||
|
||||
handler = ModelLibraryHandler(
|
||||
@@ -1474,8 +1486,8 @@ async def test_check_model_exists_model_id_only_does_not_call_metadata_provider(
|
||||
metadata_provider_factory=metadata_provider_factory,
|
||||
)
|
||||
|
||||
response = await handler.check_model_exists(FakeRequest(query={"modelId": "5"}))
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.check_model_exists(FakeRequest(query={"modelId": "5"})) # pyright: ignore[reportArgumentType]
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert payload == {
|
||||
"success": True,
|
||||
@@ -1503,9 +1515,9 @@ async def test_check_model_exists_returns_download_history_when_file_missing():
|
||||
)
|
||||
|
||||
response = await handler.check_model_exists(
|
||||
FakeRequest(query={"modelId": "5", "modelVersionId": "999"})
|
||||
FakeRequest(query={"modelId": "5", "modelVersionId": "999"}) # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
payload = json.loads(response.text)
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert payload == {
|
||||
"success": True,
|
||||
@@ -1533,9 +1545,9 @@ async def test_model_version_download_status_endpoints():
|
||||
)
|
||||
|
||||
get_response = await handler.get_model_version_download_status(
|
||||
FakeRequest(query={"modelType": "lora", "modelVersionId": "123"})
|
||||
FakeRequest(query={"modelType": "lora", "modelVersionId": "123"}) # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
get_payload = json.loads(get_response.text)
|
||||
get_payload = _json_payload(get_response)
|
||||
assert get_payload == {
|
||||
"success": True,
|
||||
"modelType": "lora",
|
||||
@@ -1544,7 +1556,7 @@ async def test_model_version_download_status_endpoints():
|
||||
}
|
||||
|
||||
set_response = await handler.set_model_version_download_status(
|
||||
FakeRequest(
|
||||
FakeRequest( # pyright: ignore[reportArgumentType]
|
||||
json_data={
|
||||
"modelType": "checkpoint",
|
||||
"modelVersionId": 456,
|
||||
@@ -1554,7 +1566,7 @@ async def test_model_version_download_status_endpoints():
|
||||
}
|
||||
)
|
||||
)
|
||||
set_payload = json.loads(set_response.text)
|
||||
set_payload = _json_payload(set_response)
|
||||
assert set_payload == {
|
||||
"success": True,
|
||||
"modelType": "checkpoint",
|
||||
@@ -1566,7 +1578,7 @@ async def test_model_version_download_status_endpoints():
|
||||
]
|
||||
|
||||
set_get_response = await handler.set_model_version_download_status(
|
||||
FakeRequest(
|
||||
FakeRequest( # pyright: ignore[reportArgumentType]
|
||||
method="GET",
|
||||
query={
|
||||
"modelType": "embedding",
|
||||
@@ -1576,7 +1588,7 @@ async def test_model_version_download_status_endpoints():
|
||||
},
|
||||
)
|
||||
)
|
||||
set_get_payload = json.loads(set_get_response.text)
|
||||
set_get_payload = _json_payload(set_get_response)
|
||||
assert set_get_payload == {
|
||||
"success": True,
|
||||
"modelType": "embedding",
|
||||
@@ -1586,7 +1598,7 @@ async def test_model_version_download_status_endpoints():
|
||||
|
||||
|
||||
def test_create_handler_set_uses_provided_dependencies():
|
||||
recorded_handlers: list[dict] = []
|
||||
recorded_handlers: list[dict[str, Any]] = []
|
||||
|
||||
class RecordingHandlerSet:
|
||||
def __init__(self, **handlers):
|
||||
@@ -1609,8 +1621,8 @@ def test_create_handler_set_uses_provided_dependencies():
|
||||
|
||||
controller = MiscRoutes(
|
||||
settings_service=DummySettings(),
|
||||
usage_stats_factory=lambda: FakeUsageStats(),
|
||||
prompt_server=CustomPromptServer,
|
||||
usage_stats_factory=lambda: FakeUsageStats(), # pyright: ignore[reportArgumentType]
|
||||
prompt_server=CustomPromptServer, # pyright: ignore[reportArgumentType]
|
||||
service_registry_adapter=ServiceRegistryAdapter(
|
||||
get_lora_scanner=fake_scanner_factory,
|
||||
get_checkpoint_scanner=fake_scanner_factory,
|
||||
@@ -1621,8 +1633,8 @@ def test_create_handler_set_uses_provided_dependencies():
|
||||
metadata_archive_manager_factory=fake_metadata_archive_manager_factory,
|
||||
metadata_provider_updater=noop_async,
|
||||
downloader_factory=dummy_downloader_factory,
|
||||
handler_set_factory=RecordingHandlerSet,
|
||||
node_registry=fake_node_registry,
|
||||
handler_set_factory=RecordingHandlerSet, # pyright: ignore[reportArgumentType]
|
||||
node_registry=fake_node_registry, # pyright: ignore[reportArgumentType]
|
||||
standalone_mode_flag=True,
|
||||
)
|
||||
|
||||
@@ -1665,7 +1677,7 @@ def test_is_wsl_returns_false_on_read_error():
|
||||
assert result is False
|
||||
|
||||
|
||||
def test_is_wsl_returns_false_on_read_error():
|
||||
def test_is_wsl_returns_false_on_read_error_builtins():
|
||||
with patch("builtins.open", side_effect=OSError()):
|
||||
result = _is_wsl()
|
||||
assert result is False
|
||||
@@ -1688,7 +1700,7 @@ def test_wsl_to_windows_path_returns_none_on_error():
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_wsl_to_windows_path_returns_none_on_subprocess_error():
|
||||
def test_wsl_to_windows_path_returns_none_on_subprocess_error_plain():
|
||||
with patch(
|
||||
"subprocess.run", side_effect=subprocess.CalledProcessError(1, "wslpath")
|
||||
):
|
||||
@@ -1769,8 +1781,8 @@ async def test_check_filename_conflicts_returns_ok_when_no_duplicates():
|
||||
scanner_factories=(("lora", "LoRAs", scanner_factory),),
|
||||
)
|
||||
|
||||
response = await handler.get_doctor_diagnostics(FakeRequest(method="GET"))
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.get_doctor_diagnostics(FakeRequest(method="GET")) # pyright: ignore[reportArgumentType]
|
||||
payload = _json_payload(response)
|
||||
|
||||
diagnostic_map = {item["id"]: item for item in payload["diagnostics"]}
|
||||
assert diagnostic_map["filename_conflicts"]["status"] == "ok"
|
||||
@@ -1798,8 +1810,8 @@ async def test_check_filename_conflicts_detects_duplicates():
|
||||
scanner_factories=(("lora", "LoRAs", scanner_factory),),
|
||||
)
|
||||
|
||||
response = await handler.get_doctor_diagnostics(FakeRequest(method="GET"))
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.get_doctor_diagnostics(FakeRequest(method="GET")) # pyright: ignore[reportArgumentType]
|
||||
payload = _json_payload(response)
|
||||
|
||||
diagnostic_map = {item["id"]: item for item in payload["diagnostics"]}
|
||||
conflict_diag = diagnostic_map["filename_conflicts"]
|
||||
@@ -1826,8 +1838,8 @@ async def test_resolve_filename_conflicts_returns_renamed_list():
|
||||
scanner_factories=(("lora", "LoRAs", scanner_factory),),
|
||||
)
|
||||
|
||||
response = await handler.resolve_filename_conflicts(FakeRequest(method="POST"))
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.resolve_filename_conflicts(FakeRequest(method="POST")) # pyright: ignore[reportArgumentType]
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert payload["success"] is True
|
||||
# Files don't exist on disk, so nothing gets renamed
|
||||
@@ -1850,8 +1862,8 @@ async def test_resolve_filename_conflicts_handles_scanner_error_gracefully():
|
||||
scanner_factories=(("lora", "LoRAs", scanner_factory),),
|
||||
)
|
||||
|
||||
response = await handler.resolve_filename_conflicts(FakeRequest(method="POST"))
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.resolve_filename_conflicts(FakeRequest(method="POST")) # pyright: ignore[reportArgumentType]
|
||||
payload = _json_payload(response)
|
||||
|
||||
assert payload["success"] is True
|
||||
assert payload["count"] == 0
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from types import SimpleNamespace
|
||||
|
||||
import jinja2
|
||||
@@ -48,19 +49,16 @@ async def test_model_page_view_reads_version_per_request():
|
||||
template_env=template_env,
|
||||
template_name="dummy.html",
|
||||
service=DummyService(),
|
||||
settings_service=DummySettings(),
|
||||
settings_service=DummySettings(), # pyright: ignore[reportArgumentType]
|
||||
server_i18n=DummyI18n(),
|
||||
logger=SimpleNamespace(
|
||||
debug=lambda *_args, **_kwargs: None,
|
||||
error=lambda *_args, **_kwargs: None,
|
||||
),
|
||||
logger=logging.getLogger("test_model_page_view"),
|
||||
)
|
||||
|
||||
view._get_app_version = lambda: "1.0.2-old"
|
||||
first = await view.handle(SimpleNamespace())
|
||||
first = await view.handle(SimpleNamespace()) # pyright: ignore[reportArgumentType]
|
||||
|
||||
view._get_app_version = lambda: "1.0.2-new"
|
||||
second = await view.handle(SimpleNamespace())
|
||||
second = await view.handle(SimpleNamespace()) # pyright: ignore[reportArgumentType]
|
||||
|
||||
assert first.text == "1.0.2-old"
|
||||
assert second.text == "1.0.2-new"
|
||||
|
||||
@@ -21,8 +21,12 @@ async def test_model_query_handler_accepts_limit_zero_for_base_models():
|
||||
service = DummyService()
|
||||
handler = ModelQueryHandler(service=service, logger=logging.getLogger(__name__))
|
||||
|
||||
response = await handler.get_base_models(SimpleNamespace(query={"limit": "0"}))
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.get_base_models(
|
||||
SimpleNamespace(query={"limit": "0"}) # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
text = response.text
|
||||
assert text is not None
|
||||
payload = json.loads(text)
|
||||
|
||||
assert payload["success"] is True
|
||||
assert service.received_limit == 0
|
||||
@@ -33,7 +37,9 @@ async def test_model_query_handler_rejects_negative_limit_for_base_models():
|
||||
service = DummyService()
|
||||
handler = ModelQueryHandler(service=service, logger=logging.getLogger(__name__))
|
||||
|
||||
await handler.get_base_models(SimpleNamespace(query={"limit": "-1"}))
|
||||
await handler.get_base_models(
|
||||
SimpleNamespace(query={"limit": "-1"}) # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
|
||||
assert service.received_limit == 20
|
||||
|
||||
@@ -58,9 +64,11 @@ async def test_model_query_handler_search_tags_passes_query_and_limit():
|
||||
handler = ModelQueryHandler(service=service, logger=logging.getLogger(__name__))
|
||||
|
||||
response = await handler.search_tags(
|
||||
SimpleNamespace(query={"q": "ani", "limit": "50"})
|
||||
SimpleNamespace(query={"q": "ani", "limit": "50"}) # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
payload = json.loads(response.text)
|
||||
text = response.text
|
||||
assert text is not None
|
||||
payload = json.loads(text)
|
||||
|
||||
assert payload["success"] is True
|
||||
assert payload["tags"] == [{"tag": "anime", "count": 3}]
|
||||
@@ -73,7 +81,8 @@ async def test_model_query_handler_search_tags_defaults_limit_to_20():
|
||||
service = DummySearchTagsService()
|
||||
handler = ModelQueryHandler(service=service, logger=logging.getLogger(__name__))
|
||||
|
||||
await handler.search_tags(SimpleNamespace(query={}))
|
||||
await handler.search_tags(SimpleNamespace(query={}) # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
|
||||
assert service.received_limit == 20
|
||||
|
||||
@@ -83,6 +92,8 @@ async def test_model_query_handler_search_tags_clamps_negative_limit():
|
||||
service = DummySearchTagsService()
|
||||
handler = ModelQueryHandler(service=service, logger=logging.getLogger(__name__))
|
||||
|
||||
await handler.search_tags(SimpleNamespace(query={"limit": "-5"}))
|
||||
await handler.search_tags(
|
||||
SimpleNamespace(query={"limit": "-5"}) # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
|
||||
assert service.received_limit == 20
|
||||
|
||||
@@ -2,6 +2,7 @@ import copy
|
||||
import json
|
||||
import logging
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -198,22 +199,24 @@ async def test_get_civitai_versions_degrades_when_download_history_unavailable(m
|
||||
|
||||
handler = ModelCivitaiHandler(
|
||||
service=service,
|
||||
settings_service=SimpleNamespace(get=lambda *_: False),
|
||||
ws_manager=SimpleNamespace(),
|
||||
settings_service=SimpleNamespace(get=lambda *_: False), # pyright: ignore[reportArgumentType]
|
||||
ws_manager=SimpleNamespace(), # pyright: ignore[reportArgumentType]
|
||||
logger=logging.getLogger(__name__),
|
||||
metadata_provider_factory=metadata_provider_factory,
|
||||
validate_model_type=lambda *_: True,
|
||||
expected_model_types=lambda: "LoRA",
|
||||
find_model_file=lambda *_: None,
|
||||
metadata_sync=SimpleNamespace(),
|
||||
metadata_refresh_use_case=SimpleNamespace(),
|
||||
metadata_progress_callback=lambda *_args, **_kwargs: None,
|
||||
metadata_sync=SimpleNamespace(), # pyright: ignore[reportArgumentType]
|
||||
metadata_refresh_use_case=SimpleNamespace(), # pyright: ignore[reportArgumentType]
|
||||
metadata_progress_callback=lambda *_args, **_kwargs: None, # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
|
||||
response = await handler.get_civitai_versions(
|
||||
SimpleNamespace(match_info={"model_id": "42"})
|
||||
SimpleNamespace(match_info={"model_id": "42"}) # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
payload = json.loads(response.text)
|
||||
text = response.text
|
||||
assert text is not None
|
||||
payload = json.loads(text)
|
||||
|
||||
assert response.status == 200
|
||||
assert payload[0]["id"] == 7
|
||||
@@ -284,10 +287,14 @@ async def test_refresh_model_updates_filters_records_without_updates():
|
||||
async def json(self):
|
||||
return {}
|
||||
|
||||
response = await handler.refresh_model_updates(DummyRequest())
|
||||
response = await handler.refresh_model_updates(
|
||||
DummyRequest() # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
assert response.status == 200
|
||||
|
||||
payload = json.loads(response.text)
|
||||
text = response.text
|
||||
assert text is not None
|
||||
payload = json.loads(text)
|
||||
assert payload["success"] is True
|
||||
assert len(payload["records"]) == 1
|
||||
assert payload["records"][0]["modelId"] == 1
|
||||
@@ -347,7 +354,9 @@ async def test_refresh_model_updates_with_target_ids():
|
||||
async def json(self):
|
||||
return {"modelIds": [1, "2", None]}
|
||||
|
||||
response = await handler.refresh_model_updates(DummyRequest())
|
||||
response = await handler.refresh_model_updates(
|
||||
DummyRequest() # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
assert response.status == 200
|
||||
|
||||
call = update_service.calls[0]
|
||||
@@ -399,7 +408,9 @@ async def test_refresh_model_updates_accepts_snake_case_ids():
|
||||
async def json(self):
|
||||
return {"model_ids": [3, "4", "abc", None]}
|
||||
|
||||
response = await handler.refresh_model_updates(DummyRequest())
|
||||
response = await handler.refresh_model_updates(
|
||||
DummyRequest() # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
assert response.status == 200
|
||||
|
||||
call = update_service.calls[0]
|
||||
@@ -429,9 +440,9 @@ async def test_fetch_missing_license_data_updates_metadata(monkeypatch):
|
||||
return None, False
|
||||
return SimpleNamespace(to_dict=lambda: copy.deepcopy(data)), False
|
||||
|
||||
saved: list[tuple[str, dict]] = []
|
||||
saved: list[tuple[str, dict[str, Any]]] = []
|
||||
|
||||
async def fake_save(path: str, metadata: dict):
|
||||
async def fake_save(path: str, metadata: dict[str, Any]):
|
||||
saved.append((path, copy.deepcopy(metadata)))
|
||||
return True
|
||||
|
||||
@@ -479,10 +490,14 @@ async def test_fetch_missing_license_data_updates_metadata(monkeypatch):
|
||||
async def json(self):
|
||||
return {}
|
||||
|
||||
response = await handler.fetch_missing_civitai_license_data(DummyRequest())
|
||||
response = await handler.fetch_missing_civitai_license_data(
|
||||
DummyRequest() # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
assert response.status == 200
|
||||
|
||||
payload = json.loads(response.text)
|
||||
text = response.text
|
||||
assert text is not None
|
||||
payload = json.loads(text)
|
||||
assert payload["success"] is True
|
||||
assert len(payload["updated"]) == 3
|
||||
assert provider_calls == [[10, 20]]
|
||||
@@ -516,9 +531,9 @@ async def test_fetch_missing_license_data_filters_model_ids(monkeypatch):
|
||||
return None, False
|
||||
return SimpleNamespace(to_dict=lambda: copy.deepcopy(data)), False
|
||||
|
||||
saved: list[tuple[str, dict]] = []
|
||||
saved: list[tuple[str, dict[str, Any]]] = []
|
||||
|
||||
async def fake_save(path: str, metadata: dict):
|
||||
async def fake_save(path: str, metadata: dict[str, Any]):
|
||||
saved.append((path, copy.deepcopy(metadata)))
|
||||
return True
|
||||
|
||||
@@ -566,10 +581,14 @@ async def test_fetch_missing_license_data_filters_model_ids(monkeypatch):
|
||||
async def json(self):
|
||||
return {"modelIds": [20]}
|
||||
|
||||
response = await handler.fetch_missing_civitai_license_data(DummyRequest())
|
||||
response = await handler.fetch_missing_civitai_license_data(
|
||||
DummyRequest() # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
assert response.status == 200
|
||||
|
||||
payload = json.loads(response.text)
|
||||
text = response.text
|
||||
assert text is not None
|
||||
payload = json.loads(text)
|
||||
assert payload["success"] is True
|
||||
assert len(payload["updated"]) == 1
|
||||
assert provider_calls == [[20]]
|
||||
|
||||
@@ -31,7 +31,7 @@ class StubLoraService:
|
||||
@pytest.fixture
|
||||
def routes():
|
||||
handler = LoraRoutes()
|
||||
handler.service = StubLoraService()
|
||||
handler.service = StubLoraService() # pyright: ignore[reportAttributeAccessIssue]
|
||||
return handler
|
||||
|
||||
|
||||
|
||||
@@ -34,8 +34,12 @@ async def test_recipe_query_handler_base_models_limit_zero_returns_all():
|
||||
logger=logging.getLogger(__name__),
|
||||
)
|
||||
|
||||
response = await handler.get_base_models(SimpleNamespace(query={"limit": "0"}))
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.get_base_models(
|
||||
SimpleNamespace(query={"limit": "0"}) # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
text = response.text
|
||||
assert text is not None
|
||||
payload = json.loads(text)
|
||||
|
||||
assert payload["success"] is True
|
||||
assert payload["base_models"] == [
|
||||
|
||||
@@ -124,7 +124,7 @@ def test_to_route_mapping_uses_handler_set():
|
||||
super().__init__()
|
||||
self.created = 0
|
||||
|
||||
def _create_handler_set(self): # noqa: D401 - simple override for test
|
||||
def _create_handler_set(self): # noqa: D401 - simple override for test # pyright: ignore[reportIncompatibleMethodOverride]
|
||||
self.created += 1
|
||||
return DummyHandlerSet()
|
||||
|
||||
@@ -162,10 +162,12 @@ def test_recipe_route_registrar_binds_every_route():
|
||||
self.router = FakeRouter()
|
||||
|
||||
app = FakeApp()
|
||||
registrar = recipe_route_registrar.RecipeRouteRegistrar(app)
|
||||
registrar = recipe_route_registrar.RecipeRouteRegistrar(
|
||||
app # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
|
||||
handler_mapping = {
|
||||
definition.handler_name: object()
|
||||
definition.handler_name: lambda _request: None
|
||||
for definition in recipe_route_registrar.ROUTE_DEFINITIONS
|
||||
}
|
||||
|
||||
|
||||
@@ -26,7 +26,7 @@ from py.services.service_registry import ServiceRegistry
|
||||
class RecipeRouteHarness:
|
||||
"""Container exposing the aiohttp client and stubbed collaborators."""
|
||||
|
||||
client: TestClient
|
||||
client: TestClient[Any, Any]
|
||||
scanner: "StubRecipeScanner"
|
||||
analysis: "StubAnalysisService"
|
||||
persistence: "StubPersistenceService"
|
||||
@@ -92,6 +92,9 @@ class StubRecipeScanner:
|
||||
candidate = Path(self.recipes_dir) / f"{recipe_id}.recipe.json"
|
||||
return str(candidate) if candidate.exists() else None
|
||||
|
||||
async def get_recipe_syntax_tokens(self, recipe_id: str) -> List[str]:
|
||||
return self.recipes.get(recipe_id, {}).get("syntax", []) # pragma: no cover - overridden per test
|
||||
|
||||
async def remove_recipe(self, recipe_id: str) -> None:
|
||||
self.removed.append(recipe_id)
|
||||
self.recipes.pop(recipe_id, None)
|
||||
@@ -110,7 +113,7 @@ class StubAnalysisService:
|
||||
self.remote_calls: List[Optional[str]] = []
|
||||
self.local_calls: List[Optional[str]] = []
|
||||
self.result = SimpleNamespace(payload={"loras": []}, status=200)
|
||||
self._recipe_parser_factory = None
|
||||
self._recipe_parser_factory: Any = None
|
||||
StubAnalysisService.instances.append(self)
|
||||
|
||||
async def analyze_uploaded_image(
|
||||
@@ -456,7 +459,7 @@ async def test_list_recipes_offloads_dimensions_to_thread(
|
||||
harness.scanner.cached_raw = list(harness.scanner.listing_items)
|
||||
|
||||
real_to_thread = asyncio.to_thread
|
||||
to_thread_calls: list[tuple] = []
|
||||
to_thread_calls: list[tuple[Any, ...]] = []
|
||||
|
||||
async def counting_to_thread(fn, *args, **kwargs):
|
||||
to_thread_calls.append((fn, args, kwargs))
|
||||
@@ -536,6 +539,7 @@ async def test_list_recipes_passes_checkpoint_hash_filter(
|
||||
|
||||
assert response.status == 200
|
||||
assert payload["items"] == []
|
||||
assert harness.scanner.last_paginated_params is not None
|
||||
assert harness.scanner.last_paginated_params["checkpoint_hash"] == "ckpt123"
|
||||
|
||||
|
||||
@@ -1010,7 +1014,9 @@ async def test_get_recipe_syntax(monkeypatch, tmp_path: Path) -> None:
|
||||
return ["<lora:lora1:0.5>"]
|
||||
raise RecipeNotFoundError(f"Recipe {rid} not found")
|
||||
|
||||
harness.scanner.get_recipe_syntax_tokens = fake_get_recipe_syntax_tokens
|
||||
harness.scanner.get_recipe_syntax_tokens = ( # pyright: ignore[reportAttributeAccessIssue]
|
||||
fake_get_recipe_syntax_tokens
|
||||
)
|
||||
|
||||
response = await harness.client.get(f"/api/lm/recipe/{recipe_id}/syntax")
|
||||
payload = await response.json()
|
||||
|
||||
@@ -5,7 +5,7 @@ from __future__ import annotations
|
||||
import asyncio
|
||||
from contextlib import asynccontextmanager
|
||||
from types import SimpleNamespace
|
||||
from typing import AsyncIterator, Dict, Iterable, List, Sequence
|
||||
from typing import Any, AsyncIterator, Dict, Iterable, List, Sequence
|
||||
|
||||
from aiohttp import web
|
||||
from aiohttp.test_utils import TestClient, TestServer
|
||||
@@ -34,7 +34,7 @@ class IntegrationCache:
|
||||
class IntegrationScanner:
|
||||
"""Scanner double that registers with ServiceRegistry expectations."""
|
||||
|
||||
def __init__(self, items: Iterable[Dict[str, object]]) -> None:
|
||||
def __init__(self, items: Iterable[Dict[str, Any]]) -> None:
|
||||
self.model_type = "lora"
|
||||
self._cache = IntegrationCache(list(items))
|
||||
self._hash_index = SimpleNamespace(
|
||||
@@ -68,7 +68,7 @@ class IntegrationScanner:
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def aiohttp_client(app: web.Application) -> AsyncIterator[TestClient]:
|
||||
async def aiohttp_client(app: web.Application) -> AsyncIterator[TestClient[Any, Any]]:
|
||||
"""Spin up a TestClient with lifecycle management."""
|
||||
|
||||
server = TestServer(app)
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import json
|
||||
from typing import Any, Optional
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -17,7 +18,7 @@ class FakeRequest:
|
||||
class DummySettings:
|
||||
def __init__(self):
|
||||
self.activated = None
|
||||
self.should_raise = None
|
||||
self.should_raise: Optional[Exception] = None
|
||||
|
||||
def activate_library(self, name):
|
||||
if self.should_raise:
|
||||
@@ -25,6 +26,13 @@ class DummySettings:
|
||||
self.activated = name
|
||||
|
||||
|
||||
def json_payload(response) -> Any:
|
||||
"""Decode the JSON body of a web.Response, asserting it is not null."""
|
||||
text = response.text
|
||||
assert text is not None
|
||||
return json.loads(text)
|
||||
|
||||
|
||||
class DummyDownloader:
|
||||
async def refresh_session(self): # pragma: no cover - helper
|
||||
return None
|
||||
@@ -53,7 +61,7 @@ async def test_get_libraries_returns_registry(monkeypatch, handler):
|
||||
monkeypatch.setattr(config, "get_library_registry_snapshot", lambda: registry)
|
||||
|
||||
response = await handler.get_libraries(FakeRequest())
|
||||
payload = json.loads(response.text)
|
||||
payload = json_payload(response)
|
||||
|
||||
assert response.status == 200
|
||||
assert payload == {
|
||||
@@ -71,7 +79,7 @@ async def test_get_libraries_handles_errors(monkeypatch, handler):
|
||||
monkeypatch.setattr(config, "get_library_registry_snapshot", boom)
|
||||
|
||||
response = await handler.get_libraries(FakeRequest())
|
||||
payload = json.loads(response.text)
|
||||
payload = json_payload(response)
|
||||
|
||||
assert response.status == 500
|
||||
assert payload["success"] is False
|
||||
@@ -90,8 +98,10 @@ async def test_activate_library_success(monkeypatch):
|
||||
registry = {"libraries": {"alpha": {"name": "Alpha"}}, "active_library": "alpha"}
|
||||
monkeypatch.setattr(config, "get_library_registry_snapshot", lambda: registry)
|
||||
|
||||
response = await handler.activate_library(FakeRequest(json_data={"library": "alpha"}))
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.activate_library(
|
||||
FakeRequest(json_data={"library": "alpha"}) # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
payload = json_payload(response)
|
||||
|
||||
assert response.status == 200
|
||||
assert payload == {
|
||||
@@ -105,7 +115,7 @@ async def test_activate_library_success(monkeypatch):
|
||||
@pytest.mark.asyncio
|
||||
async def test_activate_library_requires_name(handler):
|
||||
response = await handler.activate_library(FakeRequest(json_data={}))
|
||||
payload = json.loads(response.text)
|
||||
payload = json_payload(response)
|
||||
|
||||
assert response.status == 400
|
||||
assert payload["success"] is False
|
||||
@@ -122,8 +132,10 @@ async def test_activate_library_unknown_returns_404(monkeypatch):
|
||||
downloader_factory=dummy_downloader_factory,
|
||||
)
|
||||
|
||||
response = await handler.activate_library(FakeRequest(json_data={"library": "ghost"}))
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.activate_library(
|
||||
FakeRequest(json_data={"library": "ghost"}) # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
payload = json_payload(response)
|
||||
|
||||
assert response.status == 404
|
||||
assert payload["success"] is False
|
||||
@@ -140,8 +152,10 @@ async def test_activate_library_unexpected_error_returns_500(monkeypatch):
|
||||
downloader_factory=dummy_downloader_factory,
|
||||
)
|
||||
|
||||
response = await handler.activate_library(FakeRequest(json_data={"library": "broken"}))
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.activate_library(
|
||||
FakeRequest(json_data={"library": "broken"}) # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
payload = json_payload(response)
|
||||
|
||||
assert response.status == 500
|
||||
assert payload["success"] is False
|
||||
|
||||
@@ -341,7 +341,7 @@ async def test_handle_stats_page_renders_template(stats_routes):
|
||||
assert response.status == 200
|
||||
assert response.text == "rendered"
|
||||
assert stats_routes.server_i18n.locale_calls[-1] == "ja"
|
||||
assert stats_routes.routes.template_env._i18n_filter_added is True
|
||||
assert stats_routes.routes._i18n_filter_added is True
|
||||
assert "t" in stats_routes.routes.template_env.filters
|
||||
assert stats_routes.routes.template_env.filters["t"]("greeting") == "translated:greeting"
|
||||
assert template_context["is_initializing"] is False
|
||||
|
||||
@@ -7,9 +7,10 @@ from aiohttp.test_utils import TestClient, TestServer
|
||||
|
||||
import sys
|
||||
import types
|
||||
from typing import Any
|
||||
|
||||
folder_paths_stub = types.SimpleNamespace(get_folder_paths=lambda *_: [])
|
||||
sys.modules.setdefault("folder_paths", folder_paths_stub)
|
||||
sys.modules.setdefault("folder_paths", folder_paths_stub) # pyright: ignore[reportArgumentType]
|
||||
|
||||
from py.routes.handlers.model_handlers import ModelListingHandler
|
||||
|
||||
@@ -19,6 +20,7 @@ class MockService:
|
||||
|
||||
def __init__(self):
|
||||
self.model_type = "test-model"
|
||||
self.last_call_kwargs: dict[str, Any] = {}
|
||||
|
||||
async def get_paginated_data(self, **kwargs):
|
||||
# Store the kwargs for verification
|
||||
|
||||
@@ -24,7 +24,7 @@ def _fake_request(body=None, query_params=None):
|
||||
async def _json():
|
||||
return body or {}
|
||||
|
||||
req.json = _json
|
||||
req.json = _json # pyright: ignore[reportAttributeAccessIssue]
|
||||
return req
|
||||
|
||||
|
||||
@@ -131,7 +131,7 @@ async def test_perform_git_update_preserves_user_dirs(monkeypatch, tmp_path):
|
||||
class FakeHeads:
|
||||
def __getitem__(self, name):
|
||||
class Head:
|
||||
def checkout(self_inner):
|
||||
def checkout(self):
|
||||
calls.append(("head-checkout", (name,)))
|
||||
return Head()
|
||||
|
||||
|
||||
@@ -32,9 +32,11 @@ async def test_search_wildcards_returns_results():
|
||||
|
||||
handler = WildcardsHandler(service=StubService())
|
||||
response = await handler.search_wildcards(
|
||||
FakeRequest(query={"search": "cat", "limit": "25", "offset": "2"})
|
||||
FakeRequest(query={"search": "cat", "limit": "25", "offset": "2"}) # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
payload = json.loads(response.text)
|
||||
text = response.text
|
||||
assert text is not None
|
||||
payload = json.loads(text)
|
||||
|
||||
assert response.status == 200
|
||||
assert payload == {
|
||||
@@ -62,8 +64,12 @@ async def test_search_wildcards_handles_errors():
|
||||
raise RuntimeError("boom")
|
||||
|
||||
handler = WildcardsHandler(service=StubService())
|
||||
response = await handler.search_wildcards(FakeRequest(query={"search": "cat"}))
|
||||
payload = json.loads(response.text)
|
||||
response = await handler.search_wildcards(
|
||||
FakeRequest(query={"search": "cat"}) # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
text = response.text
|
||||
assert text is not None
|
||||
payload = json.loads(text)
|
||||
|
||||
assert response.status == 500
|
||||
assert payload["error"] == "boom"
|
||||
|
||||
Reference in New Issue
Block a user