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:
Will Miao
2026-08-08 20:12:59 +08:00
parent 8e724538bd
commit d2f955266d
95 changed files with 953 additions and 666 deletions

View File

@@ -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

View File

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

View File

@@ -24,7 +24,7 @@ class StubEmbeddingService:
@pytest.fixture
def routes():
handler = EmbeddingRoutes()
handler.service = StubEmbeddingService()
handler.service = StubEmbeddingService() # pyright: ignore[reportAttributeAccessIssue]
return handler

View File

@@ -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

View File

@@ -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,

View File

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

View File

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

View File

@@ -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

View File

@@ -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"

View File

@@ -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

View File

@@ -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]]

View File

@@ -31,7 +31,7 @@ class StubLoraService:
@pytest.fixture
def routes():
handler = LoraRoutes()
handler.service = StubLoraService()
handler.service = StubLoraService() # pyright: ignore[reportAttributeAccessIssue]
return handler

View File

@@ -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"] == [

View File

@@ -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
}

View File

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

View File

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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

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

View File

@@ -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"