From a2fb8133b4dbcbe9f0a05c0925258b8ba8936c86 Mon Sep 17 00:00:00 2001 From: axisrow Date: Fri, 10 Jul 2026 04:47:29 +0800 Subject: [PATCH] chore(mypy): type database layer (42 -> 29 in src/) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Part of #1133. Typing-only changes, no runtime behavior change: - migrations/connection/pool/repositories.filters: cast() on cursor.fetchall()/execute_fetchall() — aiosqlite stubs declare Iterable[Row], the runtime value is a list - facade.transaction(): bind the asserted connection to a local so the narrowed type stays visible inside the busy-retry lambda - facade.record_rename_event: separate variable for the write cursor (read path returns BufferedCursor, write path aiosqlite.Cursor) - repositories.collection_tasks: add ExportTaskPayload to the _deserialize_payload return union and _serialize_payload param union — both code paths already handle it (parse branch + isinstance tuple) - repositories.messages: explicit None/empty check narrows the setting value before int(); assert sq.id after the not-None pre-filter Co-Authored-By: Claude Fable 5 Claude-Session: https://claude.ai/code/session_01G3ExCZyXkbUpTkRRrQrFA2 --- src/database/connection.py | 4 +++- src/database/facade.py | 14 ++++++++------ src/database/migrations.py | 4 +++- src/database/pool.py | 9 +++++---- src/database/repositories/collection_tasks.py | 3 ++- src/database/repositories/filters.py | 5 +++-- src/database/repositories/messages.py | 3 ++- 7 files changed, 26 insertions(+), 16 deletions(-) diff --git a/src/database/connection.py b/src/database/connection.py index 3db556e9..dd1ce412 100644 --- a/src/database/connection.py +++ b/src/database/connection.py @@ -5,6 +5,7 @@ import time from dataclasses import dataclass from pathlib import Path +from typing import cast import aiosqlite @@ -177,4 +178,5 @@ async def execute(self, sql: str, params: tuple = ()) -> aiosqlite.Cursor: async def execute_fetchall(self, sql: str, params: tuple = ()) -> list: assert self.db is not None - return await self.db.execute_fetchall(sql, params) + # aiosqlite stubs declare Iterable[Row]; at runtime this is a list. + return cast(list, await self.db.execute_fetchall(sql, params)) diff --git a/src/database/facade.py b/src/database/facade.py index 6752e4f4..1d99fec4 100644 --- a/src/database/facade.py +++ b/src/database/facade.py @@ -273,20 +273,22 @@ async def transaction(self) -> AsyncIterator[aiosqlite.Connection]: Not reentrant — nesting would deadlock asyncio.Lock. """ assert self._db is not None + # Local binding keeps the narrowed (non-optional) type visible inside the lambda. + db = self._db async with self._write_lock: await self._with_busy_retry( "transaction begin", - lambda: begin_immediate(self._db), + lambda: begin_immediate(db), ) committed = False try: - yield self._db - await self._db.execute("COMMIT") + yield db + await db.execute("COMMIT") committed = True finally: if not committed: try: - await self._db.execute("ROLLBACK") + await db.execute("ROLLBACK") except aiosqlite.OperationalError: pass @@ -578,7 +580,7 @@ async def create_rename_event( if existing: return existing["id"] try: - cur = await self.execute_write( + write_cur = await self.execute_write( """ INSERT INTO channel_rename_events (channel_id, old_title, new_title, old_username, new_username) @@ -586,7 +588,7 @@ async def create_rename_event( """, (channel_id, old_title, new_title, old_username, new_username), ) - return cur.lastrowid or 0 + return write_cur.lastrowid or 0 except sqlite3.IntegrityError: # Concurrent INSERT won the race; re-select the existing row cur = await self._read( diff --git a/src/database/migrations.py b/src/database/migrations.py index d59c1f42..3f50eec3 100644 --- a/src/database/migrations.py +++ b/src/database/migrations.py @@ -3,6 +3,7 @@ import logging from collections.abc import Mapping, Sequence from pathlib import Path +from typing import cast import aiosqlite @@ -242,7 +243,8 @@ async def _dedupe_primary_accounts(db: aiosqlite.Connection) -> None: if "is_primary" not in await table_columns(db, "accounts"): return cur = await db.execute("SELECT id FROM accounts WHERE is_primary = 1 ORDER BY id ASC") - rows = await cur.fetchall() + # aiosqlite stubs declare fetchall() as Iterable[Row]; at runtime it is a list. + rows = cast("list[aiosqlite.Row]", await cur.fetchall()) if len(rows) <= 1: return keep_id = rows[0][0] diff --git a/src/database/pool.py b/src/database/pool.py index 4a84f2e2..adfde4e2 100644 --- a/src/database/pool.py +++ b/src/database/pool.py @@ -21,7 +21,7 @@ import asyncio from collections.abc import Sequence from contextlib import asynccontextmanager -from typing import Any, AsyncIterator, Protocol +from typing import Any, AsyncIterator, Protocol, cast import aiosqlite @@ -167,13 +167,14 @@ async def execute(self, sql: str, params: Sequence[Any] = ()) -> BufferedCursor: _reject_writes(sql) async with self._pool.acquire_read() as conn: cur = await conn.execute(sql, params) - # fetchall() returns a fresh owned list, so BufferedCursor can take it directly. - return BufferedCursor(await cur.fetchall()) + # fetchall() returns a fresh owned list (stubs say Iterable[Row]), so + # BufferedCursor can take it directly. + return BufferedCursor(cast("list[Any]", await cur.fetchall())) async def execute_fetchall(self, sql: str, params: Sequence[Any] = ()) -> list[Any]: _reject_writes(sql) async with self._pool.acquire_read() as conn: - return await conn.execute_fetchall(sql, params) + return cast("list[Any]", await conn.execute_fetchall(sql, params)) async def create_function(self, name: str, narg: int, func: Any, **kwargs: Any) -> None: """Register a UDF on every read connection (filter queries call this lazily).""" diff --git a/src/database/repositories/collection_tasks.py b/src/database/repositories/collection_tasks.py index 0657b4d6..76e166a0 100644 --- a/src/database/repositories/collection_tasks.py +++ b/src/database/repositories/collection_tasks.py @@ -69,7 +69,7 @@ def _deserialize_payload( ) -> ( dict[str, Any] | StatsAllTaskPayload | SqStatsTaskPayload | FilterAnalyzeTaskPayload | PipelineRunTaskPayload | ContentGenerateTaskPayload | ContentPublishTaskPayload - | TranslateBatchTaskPayload | None + | TranslateBatchTaskPayload | ExportTaskPayload | None ): if not raw: return None @@ -109,6 +109,7 @@ def _serialize_payload( | ContentGenerateTaskPayload | ContentPublishTaskPayload | TranslateBatchTaskPayload + | ExportTaskPayload | None ), ) -> str | None: diff --git a/src/database/repositories/filters.py b/src/database/repositories/filters.py index 44b5686f..612c7a19 100644 --- a/src/database/repositories/filters.py +++ b/src/database/repositories/filters.py @@ -5,7 +5,7 @@ import asyncio import logging import re -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, cast import aiosqlite @@ -370,7 +370,8 @@ async def _run_on_conn( params: tuple = (), ) -> list[aiosqlite.Row]: cur = await conn.execute(sql, params) - return await cur.fetchall() + # aiosqlite stubs declare fetchall() as Iterable[Row]; at runtime it is a list. + return cast("list[aiosqlite.Row]", await cur.fetchall()) async def fetch_maps_parallel( self, diff --git a/src/database/repositories/messages.py b/src/database/repositories/messages.py index ad29354c..aebfe805 100644 --- a/src/database/repositories/messages.py +++ b/src/database/repositories/messages.py @@ -145,7 +145,7 @@ async def _set_setting(self, key: str, value: str) -> None: async def get_embedding_dimensions(self) -> int | None: """Размерность векторов индекса эмбеддингов (из настроек), либо ``None`` если индекс ещё не создан.""" raw_value = await self._get_setting(_EMBEDDING_DIMENSIONS_SETTING) - if raw_value in (None, ""): + if raw_value is None or raw_value == "": return None try: return int(raw_value) @@ -1217,6 +1217,7 @@ async def get_fts_daily_stats_batch( union_parts = [] all_params: list = [] for sq in chunk: + assert sq.id is not None # ``valid`` filtered out None ids above fts_query = self._build_fts_match(sq.query, sq.is_fts) extra_conds, extra_params = self._build_extra_conditions(sq) where_parts = [