fix(types): resolve pre-existing basedpyright errors in py/ and standalone.py

Fix ~950 basedpyright errors across the backend:
- Convert ineffective # type: ignore comments to # pyright: ignore[rule]
- Add missing generic type arguments (Dict[str, Any], list[Any], ...)
- Annotate dynamic dict literals and runtime-initialized attributes
- Widen CivitAI provider tuple signatures in recipe parsers
- Remove dead LoraRoutes handlers calling nonexistent LoraService methods
- Suppress unavoidable ServiceRegistry import cycles (basedpyright counts
  function-local imports as cycle edges)
This commit is contained in:
Will Miao
2026-08-08 20:12:52 +08:00
parent 6fcdeb799d
commit 8e724538bd
103 changed files with 1184 additions and 1015 deletions

View File

@@ -2,7 +2,7 @@ from __future__ import annotations
import logging
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Callable, Dict, Mapping
from typing import TYPE_CHECKING, Awaitable, Callable, Dict, Mapping
import jinja2
from aiohttp import web
@@ -84,7 +84,7 @@ class BaseModelRoutes(ABC):
self.metadata_progress_callback = WebSocketBroadcastCallback()
self._handler_set: ModelHandlerSet | None = None
self._handler_mapping: Dict[str, Callable[[web.Request], web.StreamResponse]] | None = None
self._handler_mapping: Dict[str, Callable[[web.Request], Awaitable[web.Response]]] | None = None
self._preview_service = PreviewAssetService(
metadata_manager=MetadataManager,
@@ -131,7 +131,7 @@ class BaseModelRoutes(ABC):
self._handler_set = None
self._handler_mapping = None
def _ensure_handler_mapping(self) -> Mapping[str, Callable[[web.Request], web.StreamResponse]]:
def _ensure_handler_mapping(self) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]:
if self._handler_mapping is None:
handler_set = self._create_handler_set()
self._handler_set = handler_set
@@ -220,7 +220,7 @@ class BaseModelRoutes(ABC):
)
@property
def route_handlers(self) -> Mapping[str, Callable[[web.Request], web.StreamResponse]]:
def route_handlers(self) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]:
return self._ensure_handler_mapping()
def setup_routes(self, app: web.Application, prefix: str) -> None:
@@ -237,7 +237,7 @@ class BaseModelRoutes(ABC):
"""Setup model-specific routes."""
raise NotImplementedError
def _parse_specific_params(self, request: web.Request) -> Dict:
def _parse_specific_params(self, request: web.Request) -> Dict[str, Any]:
"""Parse model-specific parameters - to be overridden by subclasses."""
return {}
@@ -253,7 +253,7 @@ class BaseModelRoutes(ABC):
"""Find the appropriate model file from the files list - can be overridden by subclasses."""
return next((file for file in files if file.get("type") in ("Model", "Diffusion Model") and file.get("primary") is True), None)
def get_handler(self, name: str) -> Callable[[web.Request], web.StreamResponse]:
def get_handler(self, name: str) -> Callable[[web.Request], Awaitable[web.StreamResponse]]:
"""Expose handlers for subclasses or tests."""
return self._ensure_handler_mapping()[name]
@@ -285,7 +285,7 @@ class BaseModelRoutes(ABC):
)
return self.model_lifecycle_service
def _make_handler_proxy(self, name: str) -> Callable[[web.Request], web.StreamResponse]:
def _make_handler_proxy(self, name: str) -> Callable[[web.Request], Awaitable[web.StreamResponse]]:
async def proxy(request: web.Request) -> web.StreamResponse:
try:
handler = self.get_handler(name)

View File

@@ -4,7 +4,7 @@ from __future__ import annotations
import logging
import os
from typing import Callable, Mapping
from typing import Awaitable, Callable, Mapping
import jinja2
from aiohttp import web
@@ -61,7 +61,9 @@ class BaseRecipeRoutes:
self._i18n_registered = False
self._startup_hooks_registered = False
self._handler_set: RecipeHandlerSet | None = None
self._handler_mapping: dict[str, Callable] | None = None
self._handler_mapping: Mapping[
str, Callable[[web.Request], Awaitable[web.StreamResponse]]
] | None = None
async def attach_dependencies(self, app: web.Application | None = None) -> None:
"""Resolve shared services from the registry."""
@@ -84,7 +86,9 @@ class BaseRecipeRoutes:
app.on_startup.append(self.attach_dependencies)
self._startup_hooks_registered = True
def to_route_mapping(self) -> Mapping[str, Callable]:
def to_route_mapping(
self,
) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]:
"""Return a mapping of handler name to coroutine for registrar binding."""
if self._handler_mapping is None:
@@ -124,17 +128,17 @@ class BaseRecipeRoutes:
or os.environ.get("HF_HUB_DISABLE_TELEMETRY", "0") == "0"
)
if not standalone_mode:
from ..metadata_collector import get_metadata # type: ignore[import-not-found]
from ..metadata_collector.metadata_processor import ( # type: ignore[import-not-found]
from ..metadata_collector import get_metadata # pyright: ignore[reportMissingImports]
from ..metadata_collector.metadata_processor import ( # pyright: ignore[reportMissingImports]
MetadataProcessor,
)
from ..metadata_collector.metadata_registry import ( # type: ignore[import-not-found]
from ..metadata_collector.metadata_registry import ( # pyright: ignore[reportMissingImports]
MetadataRegistry,
)
else: # pragma: no cover - optional dependency path
get_metadata = None # type: ignore[assignment]
MetadataProcessor = None # type: ignore[assignment]
MetadataRegistry = None # type: ignore[assignment]
get_metadata = None # pyright: ignore[reportAssignmentType]
MetadataProcessor = None # pyright: ignore[reportAssignmentType]
MetadataRegistry = None # pyright: ignore[reportAssignmentType]
analysis_service = RecipeAnalysisService(
exif_utils=ExifUtils,

View File

@@ -1,5 +1,5 @@
import logging
from typing import Dict, List, Set
from typing import Any, Dict, List, Set
from aiohttp import web
from .base_model_routes import BaseModelRoutes
@@ -28,13 +28,13 @@ class CheckpointRoutes(BaseModelRoutes):
# Attach service dependencies
self.attach_service(self.service)
def setup_routes(self, app: web.Application):
def setup_routes(self, app: web.Application, prefix: str = "checkpoints"):
"""Setup Checkpoint routes"""
# Schedule service initialization on app startup
app.on_startup.append(lambda _: self.initialize_services())
# Setup common routes with 'checkpoints' prefix (includes page route)
super().setup_routes(app, 'checkpoints')
super().setup_routes(app, prefix)
def setup_specific_routes(self, registrar: ModelRouteRegistrar, prefix: str):
"""Setup Checkpoint-specific routes"""
@@ -53,9 +53,9 @@ class CheckpointRoutes(BaseModelRoutes):
"""Get expected model types string for error messages"""
return "Checkpoint"
def _parse_specific_params(self, request: web.Request) -> Dict:
def _parse_specific_params(self, request: web.Request) -> Dict[str, Any]:
"""Parse Checkpoint-specific parameters"""
params: Dict = {}
params: Dict[str, Any] = {}
if 'checkpoint_hash' in request.query:
params['hash_filters'] = {'single_hash': request.query['checkpoint_hash'].lower()}
@@ -70,7 +70,7 @@ class CheckpointRoutes(BaseModelRoutes):
"""Get detailed information for a specific checkpoint by name"""
try:
name = request.match_info.get('name', '')
checkpoint_info = await self.service.get_model_info_by_name(name)
checkpoint_info = await self.service.get_model_info_by_name(name) # pyright: ignore[reportAttributeAccessIssue]
if checkpoint_info:
return web.json_response(checkpoint_info)
@@ -89,7 +89,7 @@ class CheckpointRoutes(BaseModelRoutes):
roots.extend(config.checkpoints_roots or [])
roots.extend(config.extra_checkpoints_roots or [])
# Remove duplicates while preserving order
seen: set = set()
seen: set[str] = set()
unique_roots: List[str] = []
for root in roots:
if root and root not in seen:
@@ -114,7 +114,7 @@ class CheckpointRoutes(BaseModelRoutes):
roots.extend(config.unet_roots or [])
roots.extend(config.extra_unet_roots or [])
# Remove duplicates while preserving order
seen: set = set()
seen: set[str] = set()
unique_roots: List[str] = []
for root in roots:
if root and root not in seen:

View File

@@ -26,13 +26,13 @@ class EmbeddingRoutes(BaseModelRoutes):
# Attach service dependencies
self.attach_service(self.service)
def setup_routes(self, app: web.Application):
def setup_routes(self, app: web.Application, prefix: str = "embeddings"):
"""Setup Embedding routes"""
# Schedule service initialization on app startup
app.on_startup.append(lambda _: self.initialize_services())
# Setup common routes with 'embeddings' prefix (includes page route)
super().setup_routes(app, 'embeddings')
super().setup_routes(app, prefix)
def setup_specific_routes(self, registrar: ModelRouteRegistrar, prefix: str):
"""Setup Embedding-specific routes"""
@@ -51,7 +51,7 @@ class EmbeddingRoutes(BaseModelRoutes):
"""Get detailed information for a specific embedding by name"""
try:
name = request.match_info.get('name', '')
embedding_info = await self.service.get_model_info_by_name(name)
embedding_info = await self.service.get_model_info_by_name(name) # pyright: ignore[reportAttributeAccessIssue]
if embedding_info:
return web.json_response(embedding_info)

View File

@@ -1,7 +1,7 @@
from __future__ import annotations
import logging
from typing import Callable, Mapping
from typing import Any, Awaitable, Callable, Mapping
from aiohttp import web
@@ -35,7 +35,7 @@ class ExampleImagesRoutes:
*,
ws_manager,
download_manager: DownloadManager | None = None,
processor=ExampleImagesProcessor,
processor: Any = ExampleImagesProcessor,
file_manager=ExampleImagesFileManager,
cleanup_service: ExampleImagesCleanupService | None = None,
) -> None:
@@ -46,7 +46,9 @@ class ExampleImagesRoutes:
self._file_manager = file_manager
self._cleanup_service = cleanup_service or ExampleImagesCleanupService()
self._handler_set: ExampleImagesHandlerSet | None = None
self._handler_mapping: Mapping[str, Callable[[web.Request], web.StreamResponse]] | None = None
self._handler_mapping: Mapping[
str, Callable[[web.Request], Awaitable[web.StreamResponse]]
] | None = None
@classmethod
def setup_routes(cls, app: web.Application, *, ws_manager) -> None:
@@ -61,7 +63,9 @@ class ExampleImagesRoutes:
registrar = ExampleImagesRouteRegistrar(app)
registrar.register_routes(self.to_route_mapping())
def to_route_mapping(self) -> Mapping[str, Callable[[web.Request], web.StreamResponse]]:
def to_route_mapping(
self,
) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]:
"""Return the registrar-compatible mapping of handler names to callables."""
if self._handler_mapping is None:

View File

@@ -3,7 +3,7 @@ from __future__ import annotations
import logging
from dataclasses import dataclass
from typing import Callable, Mapping
from typing import Awaitable, Callable, Mapping
from aiohttp import web
@@ -170,7 +170,7 @@ class ExampleImagesHandlerSet:
management: ExampleImagesManagementHandler
files: ExampleImagesFileHandler
def to_route_mapping(self) -> Mapping[str, Callable[[web.Request], web.StreamResponse]]:
def to_route_mapping(self) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]:
"""Flatten handler methods into the registrar mapping."""
return {

View File

@@ -276,7 +276,7 @@ def _collect_comfyui_session_logs(
) -> dict[str, Any]:
if log_entries is None:
try:
import app.logger as comfy_logger
import app.logger as comfy_logger # pyright: ignore[reportMissingImports]
log_entries = list(comfy_logger.get_logs() or [])
except Exception as exc: # pragma: no cover - environment dependent
@@ -422,10 +422,10 @@ class PromptServerProtocol(Protocol):
"""Subset of PromptServer used by the handlers."""
instance: "PromptServerProtocol"
sockets: dict # maps clientId (sid) → WebSocketResponse
sockets: dict[str, Any] # maps clientId (sid) → WebSocketResponse
def send_sync(
self, event: str, payload: dict | None = None, sid: str | None = None
self, event: str, payload: dict[str, Any] | None = None, sid: str | None = None
) -> None: # pragma: no cover - protocol
...
@@ -443,7 +443,12 @@ class UsageStatsFactory(Protocol):
class MetadataProviderProtocol(Protocol):
async def get_model_versions(
self, model_id: int
) -> dict | None: # pragma: no cover - protocol
) -> dict[str, Any] | None: # pragma: no cover - protocol
...
async def get_user_models(
self, username: str, cursor: str | None = None
) -> Any: # pragma: no cover - protocol
...
@@ -466,16 +471,16 @@ class MetadataArchiveManagerProtocol(Protocol):
class BackupServiceProtocol(Protocol):
async def create_snapshot(
self, *, snapshot_type: str = "manual", persist: bool = False
) -> dict: # pragma: no cover - protocol
) -> dict[str, Any]: # pragma: no cover - protocol
...
async def restore_snapshot(self, archive_path: str) -> dict: # pragma: no cover - protocol
async def restore_snapshot(self, archive_path: str) -> dict[str, Any]: # pragma: no cover - protocol
...
def get_status(self) -> dict: # pragma: no cover - protocol
def get_status(self) -> dict[str, Any]: # pragma: no cover - protocol
...
def get_available_snapshots(self) -> list[dict]: # pragma: no cover - protocol
def get_available_snapshots(self) -> list[dict[str, Any]]: # pragma: no cover - protocol
...
@@ -491,7 +496,7 @@ class NodeRegistry:
def __init__(self) -> None:
self._lock = asyncio.Lock()
# sid → {unique_id → node_info}
self._tab_nodes: Dict[str, Dict[str, dict]] = {}
self._tab_nodes: Dict[str, Dict[str, dict[str, Any]]] = {}
self._ready = asyncio.Event()
self._waiting_clients: set[str] = set()
@@ -504,7 +509,7 @@ class NodeRegistry:
# Helpers to build one node dict (extracted so it's reused for each tab)
# ------------------------------------------------------------------
@staticmethod
def _build_node_dict(node: dict) -> dict:
def _build_node_dict(node: dict[str, Any]) -> dict[str, Any]:
node_id = node["node_id"]
graph_id = str(node["graph_id"])
unique_id = f"{graph_id}:{node_id}"
@@ -513,11 +518,11 @@ class NodeRegistry:
bgcolor = node.get("bgcolor") or DEFAULT_NODE_COLOR
raw_capabilities = node.get("capabilities")
capabilities: dict = {}
capabilities: dict[str, Any] = {}
if isinstance(raw_capabilities, dict):
capabilities = dict(raw_capabilities)
raw_widget_names: list | None = node.get("widget_names")
raw_widget_names: list[Any] | None = node.get("widget_names")
if not isinstance(raw_widget_names, list):
capability_widget_names = capabilities.get("widget_names")
raw_widget_names = (
@@ -565,9 +570,9 @@ class NodeRegistry:
# ------------------------------------------------------------------
# Public API
# ------------------------------------------------------------------
async def register_nodes(self, sid: str, nodes: list[dict]) -> None:
async def register_nodes(self, sid: str, nodes: list[dict[str, Any]]) -> None:
"""Register/replace the node list for a single ComfyUI tab (identified by *sid*)."""
tab_nodes: dict[str, dict] = {}
tab_nodes: dict[str, dict[str, Any]] = {}
for node in nodes:
nd = self._build_node_dict(node)
tab_nodes[nd["unique_id"]] = nd
@@ -602,7 +607,7 @@ class NodeRegistry:
except asyncio.TimeoutError:
return False
async def get_merged_registry(self, active_sids: set[str] | None = None) -> dict:
async def get_merged_registry(self, active_sids: set[str] | None = None) -> dict[str, Any]:
"""Return the union of all known tab nodes, pruning any tab that is no
longer connected."""
async with self._lock:
@@ -619,8 +624,8 @@ class NodeRegistry:
len(stale_sids), stale_sids,
)
merged: dict[str, dict] = {}
tab_info: dict[str, dict] = {}
merged: dict[str, dict[str, Any]] = {}
tab_info: dict[str, dict[str, Any]] = {}
for sid, nodes in self._tab_nodes.items():
tab_info[sid] = {
"node_count": len(nodes),
@@ -653,7 +658,7 @@ class SupportersHandler:
def __init__(self, logger: logging.Logger | None = None) -> None:
self._logger = logger or logging.getLogger(__name__)
def _load_supporters(self) -> dict:
def _load_supporters(self) -> dict[str, Any]:
"""Load supporters data from JSON file."""
try:
current_file = os.path.abspath(__file__)
@@ -1229,10 +1234,8 @@ class DoctorHandler:
settings_snapshot = _sanitize_sensitive_data(
getattr(self._settings, "settings", {}) or {}
)
startup_messages_getter = getattr(self._settings, "get_startup_messages", None)
startup_messages = (
list(startup_messages_getter()) if callable(startup_messages_getter) else []
)
startup_messages_getter: Any = getattr(self._settings, "get_startup_messages", None)
startup_messages = list(startup_messages_getter()) if startup_messages_getter else []
environment = {
"app_version": app_version,
@@ -1439,7 +1442,7 @@ class SettingsHandler:
*,
settings_service=None,
metadata_provider_updater: Callable[
[], Awaitable[None]
[], Awaitable[Any]
] = update_metadata_providers,
downloader_factory: Callable[
[], Awaitable[DownloaderProtocol]
@@ -1484,8 +1487,8 @@ class SettingsHandler:
settings_file = getattr(self._settings, "settings_file", None)
if settings_file:
response_data["settings_file"] = settings_file
messages_getter = getattr(self._settings, "get_startup_messages", None)
messages = list(messages_getter()) if callable(messages_getter) else []
messages_getter: Any = getattr(self._settings, "get_startup_messages", None)
messages = list(messages_getter()) if messages_getter else []
return web.json_response(
{
"success": True,
@@ -2005,11 +2008,11 @@ async def _noop_backup_service() -> None:
@dataclass
class ServiceRegistryAdapter:
get_lora_scanner: Callable[[], Awaitable]
get_checkpoint_scanner: Callable[[], Awaitable]
get_embedding_scanner: Callable[[], Awaitable]
get_downloaded_version_history_service: Callable[[], Awaitable]
get_backup_service: Callable[[], Awaitable] = _noop_backup_service
get_lora_scanner: Callable[[], Awaitable[Any]]
get_checkpoint_scanner: Callable[[], Awaitable[Any]]
get_embedding_scanner: Callable[[], Awaitable[Any]]
get_downloaded_version_history_service: Callable[[], Awaitable[Any]]
get_backup_service: Callable[[], Awaitable[Any]] = _noop_backup_service
class ModelLibraryHandler:
@@ -2050,8 +2053,8 @@ class ModelLibraryHandler:
return await self._service_registry.get_downloaded_version_history_service()
@staticmethod
def _with_downloaded_flag(versions: list[dict]) -> list[dict]:
enriched: list[dict] = []
def _with_downloaded_flag(versions: list[dict[str, Any]]) -> list[dict[str, Any]]:
enriched: list[dict[str, Any]] = []
for version in versions:
entry = dict(version)
entry.setdefault("hasBeenDownloaded", True)
@@ -2244,7 +2247,7 @@ class ModelLibraryHandler:
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
embedding_scanner = await self._service_registry.get_embedding_scanner()
results: list[dict] = []
results: list[dict[str, Any]] = []
for model_id in model_ids:
lora_versions = await lora_scanner.get_model_versions_by_id(model_id)
if lora_versions:
@@ -2353,7 +2356,7 @@ class ModelLibraryHandler:
)
try:
model_version_id = int(data.get("modelVersionId"))
model_version_id = int(data.get("modelVersionId")) # pyright: ignore[reportArgumentType]
except (TypeError, ValueError):
return web.json_response(
{"success": False, "error": "Parameter modelVersionId must be an integer"},
@@ -2465,10 +2468,10 @@ class ModelLibraryHandler:
"checkpoint": checkpoint_scanner,
"embedding": embedding_scanner,
}
scanner = scanner_map.get(found_type)
scanner = scanner_map.get(found_type or "")
if scanner:
persist = getattr(scanner, "_persist_current_cache", None)
if callable(persist):
persist: Any = getattr(scanner, "_persist_current_cache", None)
if persist:
await persist()
history_service = await self._get_download_history_service()
@@ -2649,13 +2652,13 @@ class ModelLibraryHandler:
}
lora_type_aliases = {model_type.lower() for model_type in VALID_LORA_TYPES}
type_scanner_map: Dict[str, object | None] = {
type_scanner_map: Dict[str, Any] = {
**{alias: lora_scanner for alias in lora_type_aliases},
"checkpoint": checkpoint_scanner,
"textualinversion": embedding_scanner,
}
versions: list[dict] = []
versions: list[dict[str, Any]] = []
history_service = await self._get_download_history_service()
model_ids: list[int] = []
model_count = 0
@@ -2707,6 +2710,8 @@ class ModelLibraryHandler:
tags_value = model.get("tags")
tags = tags_value if isinstance(tags_value, list) else []
model_id = model.get("id")
if model_id is None:
continue
try:
model_id_int = int(model_id)
except (TypeError, ValueError):
@@ -2722,6 +2727,8 @@ class ModelLibraryHandler:
continue
version_id = version.get("id")
if version_id is None:
continue
try:
version_id_int = int(version_id)
except (TypeError, ValueError):
@@ -2783,7 +2790,7 @@ class MetadataArchiveHandler:
] = get_metadata_archive_manager,
settings_service=None,
metadata_provider_updater: Callable[
[], Awaitable[None]
[], Awaitable[Any]
] = update_metadata_providers,
) -> None:
self._metadata_archive_manager_factory = metadata_archive_manager_factory
@@ -2930,7 +2937,7 @@ class BackupHandler:
if request.content_type.startswith("multipart/"):
reader = await request.multipart()
field = await reader.next()
field: Any = await reader.next()
uploaded = False
while field is not None:
if getattr(field, "filename", None):
@@ -3549,7 +3556,7 @@ class NodeRegistryHandler:
except (TypeError, ValueError):
parsed_node_id = node_identifier
payload: dict = {
payload: dict[str, Any] = {
"id": parsed_node_id,
"value": value,
"mode": mode,
@@ -3673,7 +3680,7 @@ class NodeRegistryHandler:
except (TypeError, ValueError):
parsed_node_id = node_identifier
payload: dict = {
payload: dict[str, Any] = {
"id": parsed_node_id,
"value": value,
"mode": mode,
@@ -3740,8 +3747,8 @@ class MiscHandlerSet:
doctor: DoctorHandler,
example_workflows: ExampleWorkflowsHandler,
base_model: BaseModelHandlerSet,
hf_handler: HfHandler | None = None,
agent_handler: AgentHandler | None = None,
hf_handler: Any = None,
agent_handler: Any = None,
) -> None:
self.health = health
self.settings = settings

View File

@@ -71,7 +71,7 @@ class ModelPageView:
self._server_i18n = server_i18n
self._logger = logger
def _load_supporters(self) -> dict:
def _load_supporters(self) -> dict[str, Any]:
"""Load supporters data from JSON file."""
try:
current_file = os.path.abspath(__file__)
@@ -152,7 +152,7 @@ class ModelPageView:
self._template_env.filters["t"] = (
self._server_i18n.create_template_filter()
)
self._template_env._i18n_filter_added = True # type: ignore[attr-defined]
self._template_env._i18n_filter_added = True # pyright: ignore[reportAttributeAccessIssue]
from ...services.llm_service import PROVIDER_PRESETS
@@ -199,7 +199,7 @@ class ModelListingHandler:
self,
*,
service,
parse_specific_params: Callable[[web.Request], Dict],
parse_specific_params: Callable[[web.Request], Dict[str, Any]],
logger: logging.Logger,
) -> None:
self._service = service
@@ -287,7 +287,7 @@ class ModelListingHandler:
)
return web.json_response({"error": str(exc)}, status=500)
def _parse_common_params(self, request: web.Request) -> Dict:
def _parse_common_params(self, request: web.Request) -> Dict[str, Any]:
page = int(request.query.get("page", "1"))
page_size = min(int(request.query.get("page_size", "20")), 100)
sort_by = request.query.get("sort_by", "name")
@@ -658,7 +658,7 @@ class ModelManagementHandler:
try:
reader = await request.multipart()
field = await reader.next()
field: Any = await reader.next()
if field is None or field.name != "preview_file":
raise ValueError("Expected 'preview_file' field")
content_type = field.headers.get("Content-Type", "image/png")
@@ -700,7 +700,7 @@ class ModelManagementHandler:
{
"success": True,
"preview_url": config.get_preview_static_url(
result["preview_path"]
str(result["preview_path"])
),
"preview_nsfw_level": result["preview_nsfw_level"],
}
@@ -781,7 +781,7 @@ class ModelManagementHandler:
result = await self._preview_service.replace_preview(
model_path=model_path,
preview_data=preview_data,
preview_data=preview_bytes,
content_type=content_type,
original_filename=original_filename,
nsfw_level=nsfw_level,
@@ -793,7 +793,7 @@ class ModelManagementHandler:
{
"success": True,
"preview_url": config.get_preview_static_url(
result["preview_path"]
str(result["preview_path"])
),
"preview_nsfw_level": result["preview_nsfw_level"],
}
@@ -2060,7 +2060,7 @@ class ModelCivitaiHandler:
settings_service: SettingsManager,
ws_manager: WebSocketManager,
logger: logging.Logger,
metadata_provider_factory: Callable[[], Awaitable],
metadata_provider_factory: Callable[[], Awaitable[Any]],
validate_model_type: Callable[[str], bool],
expected_model_types: Callable[[], str],
find_model_file: Callable[
@@ -2125,7 +2125,7 @@ class ModelCivitaiHandler:
downloaded_version_ids = set(
await history_service.get_downloaded_version_ids(
self._service.model_type,
model_id,
int(model_id),
)
)
except Exception as exc: # pragma: no cover - defensive logging
@@ -2402,8 +2402,8 @@ class ModelUpdateHandler:
self._logger.error("Failed to fetch license info: %s", exc, exc_info=True)
return web.json_response({"success": False, "error": str(exc)}, status=500)
updated: List[Dict[str, str]] = []
errors: List[Dict[str, str]] = []
updated: List[Dict[str, Any]] = []
errors: List[Dict[str, Any]] = []
for model_id in model_ids:
license_payload = license_map.get(model_id)
if not license_payload:
@@ -2416,6 +2416,7 @@ class ModelUpdateHandler:
model_section = civitai_section.get("model")
if not isinstance(model_section, Mapping):
model_section = {}
model_section = dict(model_section)
model_section.update(resolved_payload)
civitai_section["model"] = model_section
metadata_payload["civitai"] = civitai_section
@@ -2431,7 +2432,7 @@ class ModelUpdateHandler:
)
errors.append({"filePath": metadata_path, "error": str(exc)})
response_payload = {"success": True, "updated": updated}
response_payload: Dict[str, Any] = {"success": True, "updated": updated}
missing_model_ids = [mid for mid in model_ids if mid not in license_map]
if missing_model_ids:
response_payload["missingModelIds"] = missing_model_ids
@@ -2780,6 +2781,7 @@ class ModelUpdateHandler:
civitai_payload = metadata_payload.get("civitai")
if not isinstance(civitai_payload, Mapping):
civitai_payload = {}
civitai_payload = dict(civitai_payload)
model_payload = civitai_payload.get("model")
if not isinstance(model_payload, Mapping):
@@ -2824,7 +2826,7 @@ class ModelUpdateHandler:
return aggregated
def _extract_target_model_ids(self, payload: Dict) -> Optional[List[int]]:
def _extract_target_model_ids(self, payload: Dict[str, Any]) -> Optional[List[int]]:
if not isinstance(payload, Mapping):
return None
@@ -2852,7 +2854,7 @@ class ModelUpdateHandler:
return {}
to_dict = getattr(metadata, "to_dict", None)
if callable(to_dict):
if to_dict:
try:
return to_dict()
except Exception:
@@ -2863,7 +2865,7 @@ class ModelUpdateHandler:
return {}
async def _read_json(self, request: web.Request) -> Dict:
async def _read_json(self, request: web.Request) -> Dict[str, Any]:
if not request.can_read_body:
return {}
try:
@@ -2895,7 +2897,7 @@ class ModelUpdateHandler:
record,
*,
version_context: Optional[Dict[int, Dict[str, Any]]] = None,
) -> Dict:
) -> Dict[str, Any]:
context = version_context or {}
# Check user setting for hiding early access versions
hide_early_access = False
@@ -2924,7 +2926,7 @@ class ModelUpdateHandler:
@staticmethod
def _serialize_version(
version, context: Optional[Dict[str, Any]]
) -> Dict:
) -> Dict[str, Any]:
context = context or {}
preview_override = context.get("preview_override")
preview_url = (

View File

@@ -1082,10 +1082,10 @@ class RecipeManagementHandler:
*,
image_url: str,
name: str,
lora_entries: list,
checkpoint_entry: dict,
gen_params_request: dict,
tags: list,
lora_entries: list[Any],
checkpoint_entry: Dict[str, Any] | None,
gen_params_request: Dict[str, Any] | None,
tags: list[Any],
base_model: str,
source_path: str,
) -> web.Response:
@@ -1678,7 +1678,7 @@ class RecipeManagementHandler:
if not provider:
return ""
version_info = await provider.get_model_version_info(version_id)
version_info = await provider.get_model_version_info(str(version_id))
if isinstance(version_info, tuple):
version_info = version_info[0]
@@ -2391,7 +2391,7 @@ class RecipeAnalysisHandler:
content_type = request.headers.get("Content-Type", "")
if "multipart/form-data" in content_type:
reader = await request.multipart()
field = await reader.next()
field: Any = await reader.next()
if field is None or field.name != "image":
raise RecipeValidationError("No image field found")
image_chunks = bytearray()

View File

@@ -1,8 +1,8 @@
import asyncio
import logging
from aiohttp import web
from typing import Dict
from server import PromptServer # type: ignore
from typing import Any, Dict
from server import PromptServer # pyright: ignore[reportMissingImports]
from .base_model_routes import BaseModelRoutes
from .model_route_registrar import ModelRouteRegistrar
@@ -31,13 +31,13 @@ class LoraRoutes(BaseModelRoutes):
# Attach service dependencies
self.attach_service(self.service)
def setup_routes(self, app: web.Application):
def setup_routes(self, app: web.Application, prefix: str = "loras"):
"""Setup LoRA routes"""
# Schedule service initialization on app startup
app.on_startup.append(lambda _: self.initialize_services())
# Setup common routes with 'loras' prefix (includes page route)
super().setup_routes(app, "loras")
super().setup_routes(app, prefix)
def setup_specific_routes(self, registrar: ModelRouteRegistrar, prefix: str):
"""Setup LoRA-specific routes"""
@@ -73,7 +73,7 @@ class LoraRoutes(BaseModelRoutes):
"POST", "/api/lm/{prefix}/get_trigger_words", prefix, self.get_trigger_words
)
def _parse_specific_params(self, request: web.Request) -> Dict:
def _parse_specific_params(self, request: web.Request) -> Dict[str, Any]:
"""Parse LoRA-specific parameters"""
params = {}
@@ -119,25 +119,6 @@ class LoraRoutes(BaseModelRoutes):
logger.error(f"Error getting letter counts: {e}")
return web.json_response({"success": False, "error": str(e)}, status=500)
async def get_lora_notes(self, request: web.Request) -> web.Response:
"""Get notes for a specific LoRA file"""
try:
lora_name = request.query.get("name")
if not lora_name:
return web.Response(text="Lora file name is required", status=400)
notes = await self.service.get_lora_notes(lora_name)
if notes is not None:
return web.json_response({"success": True, "notes": notes})
else:
return web.json_response(
{"success": False, "error": "LoRA not found in cache"}, status=404
)
except Exception as e:
logger.error(f"Error getting lora notes: {e}", exc_info=True)
return web.json_response({"success": False, "error": str(e)}, status=500)
async def get_lora_trigger_words(self, request: web.Request) -> web.Response:
"""Get trigger words for a specific LoRA file"""
try:
@@ -168,52 +149,6 @@ class LoraRoutes(BaseModelRoutes):
logger.error(f"Error getting lora usage tips by path: {e}", exc_info=True)
return web.json_response({"success": False, "error": str(e)}, status=500)
async def get_lora_preview_url(self, request: web.Request) -> web.Response:
"""Get the static preview URL for a LoRA file"""
try:
lora_name = request.query.get("name")
if not lora_name:
return web.Response(text="Lora file name is required", status=400)
preview_url = await self.service.get_lora_preview_url(lora_name)
if preview_url:
return web.json_response({"success": True, "preview_url": preview_url})
else:
return web.json_response(
{
"success": False,
"error": "No preview URL found for the specified lora",
},
status=404,
)
except Exception as e:
logger.error(f"Error getting lora preview URL: {e}", exc_info=True)
return web.json_response({"success": False, "error": str(e)}, status=500)
async def get_lora_civitai_url(self, request: web.Request) -> web.Response:
"""Get the Civitai URL for a LoRA file"""
try:
lora_name = request.query.get("name")
if not lora_name:
return web.Response(text="Lora file name is required", status=400)
result = await self.service.get_lora_civitai_url(lora_name)
if result["civitai_url"]:
return web.json_response({"success": True, **result})
else:
return web.json_response(
{
"success": False,
"error": "No Civitai data found for the specified lora",
},
status=404,
)
except Exception as e:
logger.error(f"Error getting lora Civitai URL: {e}", exc_info=True)
return web.json_response({"success": False, "error": str(e)}, status=500)
async def get_random_loras(self, request: web.Request) -> web.Response:
"""Get random LoRAs based on filters and strength ranges"""
try:
@@ -337,7 +272,7 @@ class LoraRoutes(BaseModelRoutes):
graph_identifier = entry.get("graph_id")
try:
parsed_node_id = int(node_identifier)
parsed_node_id = int(node_identifier) # pyright: ignore[reportArgumentType]
except (TypeError, ValueError):
parsed_node_id = node_identifier

View File

@@ -5,7 +5,7 @@ miscellaneous endpoints share a consistent registration flow.
"""
from dataclasses import dataclass
from typing import Callable, Iterable, Mapping
from typing import Any, Callable, Iterable, Mapping
from aiohttp import web
@@ -147,7 +147,7 @@ class MiscRouteRegistrar:
handler_lookup[definition.handler_name],
)
def _bind(self, method: str, path: str, handler: Callable) -> None:
def _bind(self, method: str, path: str, handler: Callable[..., Any]) -> None:
add_method_name = self._METHOD_MAP[method.upper()]
add_method = getattr(self._app.router, add_method_name)
add_method(path, handler)

View File

@@ -7,7 +7,7 @@ import os
from typing import Awaitable, Callable, Mapping
from aiohttp import web
from server import PromptServer # type: ignore
from server import PromptServer # pyright: ignore[reportMissingImports]
from ..services.metadata_service import (
get_metadata_archive_manager,

View File

@@ -3,7 +3,7 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Callable, Iterable, Mapping
from typing import Any, Callable, Iterable, Mapping
from aiohttp import web
@@ -174,15 +174,15 @@ class ModelRouteRegistrar:
handler_lookup[definition.handler_name],
)
def add_route(self, method: str, path: str, handler: Callable) -> None:
def add_route(self, method: str, path: str, handler: Callable[..., Any]) -> None:
self._bind_route(method, path, handler)
def add_prefixed_route(
self, method: str, path_template: str, prefix: str, handler: Callable
self, method: str, path_template: str, prefix: str, handler: Callable[..., Any]
) -> None:
self._bind_route(method, path_template.replace("{prefix}", prefix), handler)
def _bind_route(self, method: str, path: str, handler: Callable) -> None:
def _bind_route(self, method: str, path: str, handler: Callable[..., Any]) -> None:
add_method_name = self._METHOD_MAP[method.upper()]
add_method = getattr(self._app.router, add_method_name)
add_method(path, handler)

View File

@@ -3,7 +3,7 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Callable, Mapping
from typing import Any, Callable, Mapping
from aiohttp import web
@@ -105,7 +105,7 @@ class RecipeRouteRegistrar:
handler = handler_lookup[definition.handler_name]
self._bind_route(definition.method, definition.path, handler)
def _bind_route(self, method: str, path: str, handler: Callable) -> None:
def _bind_route(self, method: str, path: str, handler: Callable[..., Any]) -> None:
add_method_name = self._METHOD_MAP[method.upper()]
add_method = getattr(self._app.router, add_method_name)
add_method(path, handler)

View File

@@ -40,10 +40,11 @@ class StatsRoutes:
"""Route handlers for Statistics page and API endpoints"""
def __init__(self):
self.lora_scanner = None
self.checkpoint_scanner = None
self.embedding_scanner = None
self.usage_stats = None
self.lora_scanner: Any = None
self.checkpoint_scanner: Any = None
self.embedding_scanner: Any = None
self.usage_stats: Any = None
self._i18n_filter_added = False
self.template_env = jinja2.Environment(
loader=jinja2.FileSystemLoader(config.templates_path),
autoescape=True
@@ -95,9 +96,9 @@ class StatsRoutes:
server_i18n.set_locale(user_language)
# 为模板环境添加i18n过滤器
if not hasattr(self.template_env, '_i18n_filter_added'):
if not self._i18n_filter_added:
self.template_env.filters['t'] = server_i18n.create_template_filter()
self.template_env._i18n_filter_added = True
self._i18n_filter_added = True
template = self.template_env.get_template('statistics.html')
rendered = template.render(
@@ -549,7 +550,7 @@ class StatsRoutes:
'error': str(e)
}, status=500)
def _count_unused_models(self, models: List[Dict], usage_data: Dict) -> int:
def _count_unused_models(self, models: List[Dict[str, Any]], usage_data: Dict[str, Any]) -> int:
"""Count models that have never been used"""
used_hashes = set(usage_data.keys())
unused_count = 0
@@ -560,7 +561,7 @@ class StatsRoutes:
return unused_count
def _get_top_used_models(self, usage_data: Dict, model_map: Dict, limit: int) -> List[Dict]:
def _get_top_used_models(self, usage_data: Dict[str, Any], model_map: Dict[str, Any], limit: int) -> List[Dict[str, Any]]:
"""Get top used models with their metadata"""
sorted_usage = sorted(usage_data.items(), key=lambda x: x[1].get('total', 0), reverse=True)
@@ -578,7 +579,7 @@ class StatsRoutes:
return top_models
def _get_usage_timeline(self, usage_data: Dict, days: int) -> List[Dict]:
def _get_usage_timeline(self, usage_data: Dict[str, Any], days: int) -> List[Dict[str, Any]]:
"""Get usage timeline for the past N days"""
timeline = []
today = datetime.now()
@@ -614,7 +615,7 @@ class StatsRoutes:
return list(reversed(timeline)) # Oldest to newest
def _format_size(self, size_bytes: int) -> str:
def _format_size(self, size_bytes: float) -> str:
"""Format file size in human readable format"""
for unit in ['B', 'KB', 'MB', 'GB', 'TB']:
if size_bytes < 1024.0:

View File

@@ -6,7 +6,7 @@ import shutil
import tempfile
import asyncio
from aiohttp import web, ClientError
from typing import Dict, List
from typing import Any, Dict, List, cast
from ..utils.settings_paths import ensure_settings_file
from ..services.downloader import get_downloader
@@ -467,9 +467,10 @@ class UpdateRoutes:
if not success:
logger.error(f"Failed to fetch release info: {data}")
return False, ""
zip_url = data.get("zipball_url")
version = data.get("tag_name", "unknown")
release_payload = cast(dict[str, Any], data)
zip_url = release_payload.get("zipball_url", "")
version = release_payload.get("tag_name", "unknown")
# Download ZIP to temporary file
with tempfile.NamedTemporaryFile(delete=False, suffix=".zip") as tmp_zip:
@@ -580,9 +581,10 @@ class UpdateRoutes:
logger.warning("Failed to fetch GitHub commit: %s", data)
return "main", [], 0, ""
commit_sha = data.get('sha', '')[:7]
commit_message = data.get('commit', {}).get('message', '')
commit_date = data.get('commit', {}).get('committer', {}).get('date', '')[:10]
commit_payload = cast(dict[str, Any], data)
commit_sha = commit_payload.get('sha', '')[:7]
commit_message = commit_payload.get('commit', {}).get('message', '')
commit_date = commit_payload.get('commit', {}).get('committer', {}).get('date', '')[:10]
version = f"main-{commit_sha}"
changelog = [commit_message] if commit_message else []
@@ -598,10 +600,11 @@ class UpdateRoutes:
custom_headers={'Accept': 'application/vnd.github+json'}
)
if c_ok:
if c_data.get('status') in ('ahead', 'diverged'):
behind_by = c_data.get('ahead_by', 0)
compare_payload = cast(dict[str, Any], c_data)
if compare_payload.get('status') in ('ahead', 'diverged'):
behind_by = compare_payload.get('ahead_by', 0)
else:
behind_by = c_data.get('behind_by', 0)
behind_by = compare_payload.get('behind_by', 0)
return version, changelog, behind_by, commit_date
@@ -706,7 +709,7 @@ class UpdateRoutes:
logger.info(f"Successfully updated to {new_version}")
return True, new_version
except git.exc.GitError as e:
except git.exc.GitError as e: # pyright: ignore[reportAttributeAccessIssue]
logger.error(f"Git error during update: {e}")
return False, ""
except Exception as e:
@@ -767,7 +770,7 @@ class UpdateRoutes:
return git_info
@staticmethod
async def _get_remote_version() -> tuple[str, List[str], List[Dict]]:
async def _get_remote_version() -> tuple[str, List[str], List[Dict[str, Any]]]:
"""
Fetch remote version from GitHub
Returns:
@@ -789,7 +792,7 @@ class UpdateRoutes:
# Parse releases
releases = []
for i, release in enumerate(data):
for i, release in enumerate(cast(list[dict[str, Any]], data)):
version = release.get('tag_name', '')
if not version.startswith('v'):
version = f"v{version}"