Bladeren bron

chore(类型): 修复 SQLAlchemy 严格类型门禁

将动态 Table 的 INSERT/UPDATE 执行结果显式收窄为 CursorResult,保留 inserted_primary_key 和 rowcount 的运行语义,同时让 strict mypy 能验证这些 CAS 分支。

覆盖旧 Build 状态、收藏删除、Owner fencing、HTTP 幂等命令和 Final publication 的所有行数断言,不改变 SQL、事务边界或异常行为。

移除 SQLAlchemy 与上传选题写入路径上已经失效的 type: ignore,恢复 unused-ignore 检查。

strict mypy 由 18 个错误恢复为 81 个源码文件全部通过。
SamLee 1 dag geleden
bovenliggende
commit
d03e95fb64

+ 4 - 4
script_build_host/src/script_build_host/adapters/uploaded_topic.py

@@ -104,7 +104,7 @@ class SqlUploadedTopicGateway:
         resolved_account = (account_name or parsed.get("account_name") or "").strip() or None
         now = datetime.now(UTC).replace(tzinfo=None)
         async with self._sessions() as session, session.begin():
-            build_result = cast(  # type: ignore[redundant-cast]
+            build_result = cast(
                 CursorResult[Any],
                 await session.execute(
                     insert(topic_build_record).values(
@@ -128,7 +128,7 @@ class SqlUploadedTopicGateway:
                 ),
             )
             topic_build_id = _inserted_id(build_result)
-            topic_result = cast(  # type: ignore[redundant-cast]
+            topic_result = cast(
                 CursorResult[Any],
                 await session.execute(
                     insert(topic_build_topic).values(
@@ -143,7 +143,7 @@ class SqlUploadedTopicGateway:
             topic_id = _inserted_id(topic_result)
             sort_order = 0
             for point_data in parsed["points"]:
-                point_result = cast(  # type: ignore[redundant-cast]
+                point_result = cast(
                     CursorResult[Any],
                     await session.execute(
                         insert(topic_build_point).values(
@@ -160,7 +160,7 @@ class SqlUploadedTopicGateway:
                 )
                 point_id = _inserted_id(point_result)
                 for item in point_data["items"]:
-                    item_result = cast(  # type: ignore[redundant-cast]
+                    item_result = cast(
                         CursorResult[Any],
                         await session.execute(
                             insert(topic_build_composition_item).values(

+ 22 - 15
script_build_host/src/script_build_host/application/legacy_api.py

@@ -3,9 +3,10 @@
 from __future__ import annotations
 
 from collections import Counter
-from typing import Any
+from typing import Any, cast
 
 from sqlalchemy import select, update
+from sqlalchemy.engine import CursorResult
 from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
 
 from script_build_host.domain.errors import BuildNotFound
@@ -105,13 +106,16 @@ class LegacyScriptBuildApiService:
     ) -> dict[str, Any]:
         await self._authorizer.require_access(principal, script_build_id)
         async with self._write() as session, session.begin():
-            result = await session.execute(
-                update(script_build_record)
-                .where(
-                    script_build_record.c.id == script_build_id,
-                    script_build_record.c.is_deleted.is_(False),
-                )
-                .values(is_favorited=value)
+            result = cast(
+                CursorResult[Any],
+                await session.execute(
+                    update(script_build_record)
+                    .where(
+                        script_build_record.c.id == script_build_id,
+                        script_build_record.c.is_deleted.is_(False),
+                    )
+                    .values(is_favorited=value)
+                ),
             )
             if result.rowcount != 1:
                 raise BuildNotFound()
@@ -120,13 +124,16 @@ class LegacyScriptBuildApiService:
     async def delete(self, script_build_id: int, principal: Principal) -> dict[str, Any]:
         await self._authorizer.require_access(principal, script_build_id)
         async with self._write() as session, session.begin():
-            result = await session.execute(
-                update(script_build_record)
-                .where(
-                    script_build_record.c.id == script_build_id,
-                    script_build_record.c.is_deleted.is_(False),
-                )
-                .values(is_deleted=True)
+            result = cast(
+                CursorResult[Any],
+                await session.execute(
+                    update(script_build_record)
+                    .where(
+                        script_build_record.c.id == script_build_id,
+                        script_build_record.c.is_deleted.is_(False),
+                    )
+                    .values(is_deleted=True)
+                ),
             )
             if result.rowcount != 1:
                 raise BuildNotFound()

+ 52 - 40
script_build_host/src/script_build_host/infrastructure/final_publication.py

@@ -127,19 +127,22 @@ class SqlAlchemyFinalPublicationUnitOfWork:
     ) -> None:
         async with self._sessions() as session, session.begin():
             await self._fencing.verify_in_session(session, owner_token)
-            result = await session.execute(
-                update(publication_table)
-                .where(
-                    publication_table.c.id == publication_id,
-                    publication_table.c.script_build_id == owner_token.script_build_id,
-                    publication_table.c.state != PublicationState.PUBLISHED.value,
-                )
-                .values(
-                    state=PublicationState.FAILED.value,
-                    last_error_code=error_code[:64],
-                    last_error_summary=error_summary[:1000],
-                    updated_at=datetime.now(UTC),
-                )
+            result = cast(
+                CursorResult[Any],
+                await session.execute(
+                    update(publication_table)
+                    .where(
+                        publication_table.c.id == publication_id,
+                        publication_table.c.script_build_id == owner_token.script_build_id,
+                        publication_table.c.state != PublicationState.PUBLISHED.value,
+                    )
+                    .values(
+                        state=PublicationState.FAILED.value,
+                        last_error_code=error_code[:64],
+                        last_error_summary=error_summary[:1000],
+                        updated_at=datetime.now(UTC),
+                    ),
+                ),
             )
             if result.rowcount != 1:
                 raise ProtocolViolation("final publication failure CAS did not match")
@@ -279,16 +282,19 @@ class SqlAlchemyFinalPublicationUnitOfWork:
                 )
                 if locked_artifacts != artifact_ids:
                     raise ProtocolViolation("final publication artifacts are missing")
-                pointer_result = await session.execute(
-                    update(mission_binding_table)
-                    .where(
-                        mission_binding_table.c.script_build_id == script_build_id,
-                        mission_binding_table.c.accepted_root_artifact_version_id.is_(None),
-                    )
-                    .values(
-                        accepted_root_artifact_version_id=manifest_version.artifact_version_id,
-                        updated_at=now,
-                    )
+                pointer_result = cast(
+                    CursorResult[Any],
+                    await session.execute(
+                        update(mission_binding_table)
+                        .where(
+                            mission_binding_table.c.script_build_id == script_build_id,
+                            mission_binding_table.c.accepted_root_artifact_version_id.is_(None),
+                        )
+                        .values(
+                            accepted_root_artifact_version_id=manifest_version.artifact_version_id,
+                            updated_at=now,
+                        ),
+                    ),
                 )
                 if pointer_result.rowcount != 1:
                     pointer = await session.scalar(
@@ -298,15 +304,18 @@ class SqlAlchemyFinalPublicationUnitOfWork:
                     )
                     if pointer != manifest_version.artifact_version_id:
                         raise MissionFencingTokenStale()
-                artifact_result = await session.execute(
-                    update(artifact_version_table)
-                    .where(
-                        artifact_version_table.c.id.in_(artifact_ids),
-                        artifact_version_table.c.state.in_(
-                            [ArtifactState.FROZEN.value, ArtifactState.PUBLISHED.value]
-                        ),
-                    )
-                    .values(state=ArtifactState.PUBLISHED.value, published_at=now)
+                artifact_result = cast(
+                    CursorResult[Any],
+                    await session.execute(
+                        update(artifact_version_table)
+                        .where(
+                            artifact_version_table.c.id.in_(artifact_ids),
+                            artifact_version_table.c.state.in_(
+                                [ArtifactState.FROZEN.value, ArtifactState.PUBLISHED.value]
+                            ),
+                        )
+                        .values(state=ArtifactState.PUBLISHED.value, published_at=now)
+                    ),
                 )
                 if artifact_result.rowcount != len(artifact_ids):
                     raise ProtocolViolation("final publication artifacts are not publishable")
@@ -319,14 +328,17 @@ class SqlAlchemyFinalPublicationUnitOfWork:
                         published_at=now,
                     )
                 )
-                build_result = await session.execute(
-                    update(script_build_record)
-                    .where(script_build_record.c.id == script_build_id)
-                    .values(
-                        status="success",
-                        error_message=None,
-                        end_time=now,
-                    )
+                build_result = cast(
+                    CursorResult[Any],
+                    await session.execute(
+                        update(script_build_record)
+                        .where(script_build_record.c.id == script_build_id)
+                        .values(
+                            status="success",
+                            error_message=None,
+                            end_time=now,
+                        )
+                    ),
                 )
                 if build_result.rowcount != 1:
                     raise BuildNotFound()

+ 21 - 17
script_build_host/src/script_build_host/infrastructure/http_commands.py

@@ -5,9 +5,10 @@ from __future__ import annotations
 import asyncio
 from collections.abc import Awaitable, Callable, Mapping
 from datetime import UTC, datetime
-from typing import Any, TypeVar
+from typing import Any, TypeVar, cast
 
 from sqlalchemy import insert, select, update
+from sqlalchemy.engine import CursorResult
 from sqlalchemy.exc import IntegrityError
 from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
 
@@ -132,22 +133,25 @@ class HttpCommandJournal:
         if len(canonical_json_bytes(response)) > MAX_IDEMPOTENCY_RESPONSE_BYTES:
             raise ProtocolViolation("idempotent HTTP response exceeds 64 KiB")
         async with self._sessions() as session, session.begin():
-            result = await session.execute(
-                update(http_command_table)
-                .where(
-                    http_command_table.c.principal_scope_sha256 == scope,
-                    http_command_table.c.route_family == route_family,
-                    http_command_table.c.idempotency_key == key,
-                    http_command_table.c.request_fingerprint == fingerprint,
-                    http_command_table.c.state == "reserved",
-                )
-                .values(
-                    state="completed",
-                    response_status=status,
-                    response_json=response,
-                    resource_id=resource_id,
-                    updated_at=datetime.now(UTC),
-                )
+            result = cast(
+                CursorResult[Any],
+                await session.execute(
+                    update(http_command_table)
+                    .where(
+                        http_command_table.c.principal_scope_sha256 == scope,
+                        http_command_table.c.route_family == route_family,
+                        http_command_table.c.idempotency_key == key,
+                        http_command_table.c.request_fingerprint == fingerprint,
+                        http_command_table.c.state == "reserved",
+                    )
+                    .values(
+                        state="completed",
+                        response_status=status,
+                        response_json=response,
+                        resource_id=resource_id,
+                        updated_at=datetime.now(UTC),
+                    ),
+                ),
             )
             if result.rowcount != 1:
                 raise ProtocolViolation("idempotent HTTP command completion lost its reservation")

+ 23 - 16
script_build_host/src/script_build_host/infrastructure/ownership.py

@@ -13,10 +13,11 @@ from contextlib import AbstractAsyncContextManager
 from datetime import UTC, datetime
 from pathlib import Path
 from types import TracebackType
-from typing import Any
+from typing import Any, cast
 from uuid import uuid4
 
 from sqlalchemy import select, update
+from sqlalchemy.engine import CursorResult
 from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
 
 from script_build_host.domain.errors import (
@@ -178,13 +179,16 @@ class FencedCommandGate:
 
         async with self._sessions() as session, session.begin():
             await self.verify_in_session(session, token)
-            result = await session.execute(
-                update(mission_binding_table)
-                .where(
-                    mission_binding_table.c.script_build_id == token.script_build_id,
-                    mission_binding_table.c.input_snapshot_id == expected_snapshot_id,
-                )
-                .values(input_snapshot_id=new_snapshot_id, updated_at=datetime.now(UTC))
+            result = cast(
+                CursorResult[Any],
+                await session.execute(
+                    update(mission_binding_table)
+                    .where(
+                        mission_binding_table.c.script_build_id == token.script_build_id,
+                        mission_binding_table.c.input_snapshot_id == expected_snapshot_id,
+                    )
+                    .values(input_snapshot_id=new_snapshot_id, updated_at=datetime.now(UTC))
+                ),
             )
             if result.rowcount != 1:
                 current = await session.scalar(
@@ -237,14 +241,17 @@ class FencedCommandGate:
     ) -> None:
         async with self._sessions() as session, session.begin():
             await self.verify_in_session(session, token)
-            result = await session.execute(
-                update(self._runtime_record)
-                .where(
-                    self._runtime_record.c.id == token.script_build_id,
-                    self._runtime_record.c.is_deleted.is_(False),
-                    self._runtime_record.c.status.in_(expected_statuses),
-                )
-                .values(**values)
+            result = cast(
+                CursorResult[Any],
+                await session.execute(
+                    update(self._runtime_record)
+                    .where(
+                        self._runtime_record.c.id == token.script_build_id,
+                        self._runtime_record.c.is_deleted.is_(False),
+                        self._runtime_record.c.status.in_(expected_statuses),
+                    )
+                    .values(**values)
+                ),
             )
             if result.rowcount != 1:
                 current = (

+ 1 - 1
script_build_host/src/script_build_host/infrastructure/tables.py

@@ -18,7 +18,7 @@ from sqlalchemy.dialects import mysql
 metadata = MetaData()
 _primary_key_type = BigInteger().with_variant(Integer, "sqlite")
 _timestamp_type = DateTime(timezone=True).with_variant(
-    mysql.DATETIME(fsp=6),  # type: ignore[no-untyped-call]
+    mysql.DATETIME(fsp=6),
     "mysql",
 )
 

+ 50 - 37
script_build_host/src/script_build_host/repositories/legacy_state.py

@@ -1,10 +1,11 @@
 from __future__ import annotations
 
 from datetime import UTC, datetime
-from typing import Any
+from typing import Any, cast
 from uuid import uuid4
 
 from sqlalchemy import insert, select, update
+from sqlalchemy.engine import CursorResult
 from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
 
 from script_build_host.domain.errors import BuildNotFound, ProtocolViolation
@@ -39,21 +40,24 @@ class SqlAlchemyLegacyBuildStateRepository:
         root_trace_id: str | None = None,
     ) -> int:
         async with self._sessions() as session, session.begin():
-            result = await session.execute(
-                insert(self._runtime_record).values(
-                    execution_id=execution_id,
-                    topic_build_id=topic_build_id,
-                    topic_id=topic_id,
-                    agent_type=agent_type,
-                    agent_config=agent_config,
-                    data_source_url=data_source_url,
-                    strategies_config=strategies_config,
-                    reson_trace_id=root_trace_id or str(uuid4()),
-                    status=BuildStatus.RUNNING.value,
-                    is_deleted=False,
-                    is_favorited=False,
-                    start_time=datetime.now(UTC),
-                )
+            result = cast(
+                CursorResult[Any],
+                await session.execute(
+                    insert(self._runtime_record).values(
+                        execution_id=execution_id,
+                        topic_build_id=topic_build_id,
+                        topic_id=topic_id,
+                        agent_type=agent_type,
+                        agent_config=agent_config,
+                        data_source_url=data_source_url,
+                        strategies_config=strategies_config,
+                        reson_trace_id=root_trace_id or str(uuid4()),
+                        status=BuildStatus.RUNNING.value,
+                        is_deleted=False,
+                        is_favorited=False,
+                        start_time=datetime.now(UTC),
+                    )
+                ),
             )
             primary_key = result.inserted_primary_key
             if primary_key is None or primary_key[0] is None:
@@ -77,13 +81,16 @@ class SqlAlchemyLegacyBuildStateRepository:
         }:
             values["end_time"] = datetime.now(UTC)
         async with self._sessions() as session, session.begin():
-            result = await session.execute(
-                update(self._runtime_record)
-                .where(
-                    self._runtime_record.c.id == script_build_id,
-                    self._runtime_record.c.is_deleted.is_(False),
-                )
-                .values(**values)
+            result = cast(
+                CursorResult[Any],
+                await session.execute(
+                    update(self._runtime_record)
+                    .where(
+                        self._runtime_record.c.id == script_build_id,
+                        self._runtime_record.c.is_deleted.is_(False),
+                    )
+                    .values(**values)
+                ),
             )
             if result.rowcount != 1:
                 raise BuildNotFound()
@@ -102,13 +109,16 @@ class SqlAlchemyLegacyBuildStateRepository:
 
     async def project_direction(self, script_build_id: int, legacy_markdown: str) -> None:
         async with self._sessions() as session, session.begin():
-            result = await session.execute(
-                update(self._runtime_record)
-                .where(
-                    self._runtime_record.c.id == script_build_id,
-                    self._runtime_record.c.is_deleted.is_(False),
-                )
-                .values(script_direction=legacy_markdown)
+            result = cast(
+                CursorResult[Any],
+                await session.execute(
+                    update(self._runtime_record)
+                    .where(
+                        self._runtime_record.c.id == script_build_id,
+                        self._runtime_record.c.is_deleted.is_(False),
+                    )
+                    .values(script_direction=legacy_markdown)
+                ),
             )
             if result.rowcount != 1:
                 raise BuildNotFound()
@@ -139,13 +149,16 @@ class SqlAlchemyLegacyBuildStateRepository:
             "end_time": datetime.now(UTC),
         }
         async with self._sessions() as session, session.begin():
-            result = await session.execute(
-                update(self._runtime_record)
-                .where(
-                    self._runtime_record.c.id == script_build_id,
-                    self._runtime_record.c.is_deleted.is_(False),
-                )
-                .values(**values)
+            result = cast(
+                CursorResult[Any],
+                await session.execute(
+                    update(self._runtime_record)
+                    .where(
+                        self._runtime_record.c.id == script_build_id,
+                        self._runtime_record.c.is_deleted.is_(False),
+                    )
+                    .values(**values)
+                ),
             )
             if result.rowcount != 1:
                 raise BuildNotFound()