Procházet zdrojové kódy

fix(输入): 对齐旧库字段并冻结脚本策略选择

按真实旧库字段重构 Topic Build 图读取:先读取主键,再以受限并发逐行回读大字段,绕开旧 RDS 在覆盖范围查询上的长时间阻塞。

对 Point、Composition Item、关系和来源记录执行闭包校验,任何跨 Topic、跨 Build、悬空引用或缺失记录都继续 fail-closed。

默认不冻结体积大且可能包含敏感信息的 raw_upload_data,仅在明确兼容场景开启;缺失旧字段统一投影为稳定 null,保持快照 schema 不漂移。

修正脚本策略来源:不再误用 Topic Build 的 strategies_config,而是只冻结本次 Script Build 请求中显式选择的常驻和按需策略。

运行清单保留真实 HTTP/HTTPS scheme 且不携带凭据,并补充快照策略及旧库关系异常测试。
SamLee před 1 dnem
rodič
revize
394d5df3f8

+ 63 - 46
script_build_host/src/script_build_host/adapters/strategy.py

@@ -2,7 +2,7 @@ from __future__ import annotations
 
 from typing import Any
 
-from sqlalchemy import or_, select
+from sqlalchemy import select
 from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
 
 from script_build_host.domain.errors import ProtocolViolation
@@ -21,52 +21,8 @@ class SqlAlchemyStrategySource:
         requested = _deduplicate((*always_on, *on_demand))
         if not requested:
             return ()
-        requested_ids = [
-            value for value in (_selector_id(item) for item in requested) if value is not None
-        ]
-        requested_names = [
-            value for value in (_selector_name(item) for item in requested) if value is not None
-        ]
         async with self._sessions() as session:
-            rows = (
-                (
-                    await session.execute(
-                        select(
-                            build_strategy.c.id,
-                            build_strategy.c.name,
-                            build_strategy.c.description,
-                            build_strategy.c.tags,
-                            build_strategy.c.current_version,
-                            build_strategy_version.c.content,
-                            build_strategy_version.c.created_at,
-                        )
-                        .join(
-                            build_strategy_version,
-                            (build_strategy_version.c.strategy_id == build_strategy.c.id)
-                            & (
-                                build_strategy_version.c.version == build_strategy.c.current_version
-                            ),
-                        )
-                        .where(
-                            build_strategy.c.strategy_type == "script",
-                            build_strategy.c.is_active.is_(True),
-                            or_(
-                                build_strategy.c.id.in_(requested_ids),
-                                build_strategy.c.name.in_(requested_names),
-                            ),
-                        )
-                    )
-                )
-                .mappings()
-                .all()
-            )
-        by_id = {int(row["id"]): row for row in rows}
-        by_name = {str(row["name"]): row for row in rows}
-        resolved = []
-        for item in requested:
-            identifier = _selector_id(item)
-            row = by_id.get(identifier) if identifier is not None else None
-            resolved.append(row or by_name.get(_selector_name(item) or ""))
+            resolved = [await _load_strategy(session, selector) for selector in requested]
         missing = [str(item) for item, row in zip(requested, resolved, strict=True) if row is None]
         if missing:
             raise ProtocolViolation(f"script strategy is missing or inactive: {', '.join(missing)}")
@@ -97,6 +53,67 @@ class SqlAlchemyStrategySource:
         return tuple(output)
 
 
+async def _load_strategy(session: AsyncSession, selector: StrategySelector) -> Any:
+    identifier = _selector_id(selector)
+    if identifier is None:
+        name = _selector_name(selector)
+        identifier = (
+            await session.execute(
+                select(build_strategy.c.id).where(
+                    build_strategy.c.name == name,
+                    build_strategy.c.strategy_type == "script",
+                    build_strategy.c.is_active.is_(True),
+                )
+            )
+        ).scalar_one_or_none()
+    if identifier is None:
+        return None
+    row = (
+        (
+            await session.execute(
+                select(
+                    build_strategy.c.id,
+                    build_strategy.c.name,
+                    build_strategy.c.description,
+                    build_strategy.c.tags,
+                    build_strategy.c.current_version,
+                ).where(
+                    build_strategy.c.id == identifier,
+                    build_strategy.c.strategy_type == "script",
+                    build_strategy.c.is_active.is_(True),
+                )
+            )
+        )
+        .mappings()
+        .one_or_none()
+    )
+    if row is None:
+        return None
+    version_id = (
+        await session.execute(
+            select(build_strategy_version.c.id).where(
+                build_strategy_version.c.strategy_id == identifier,
+                build_strategy_version.c.version == row["current_version"],
+            )
+        )
+    ).scalar_one_or_none()
+    if version_id is None:
+        return None
+    version = (
+        (
+            await session.execute(
+                select(
+                    build_strategy_version.c.content,
+                    build_strategy_version.c.created_at,
+                ).where(build_strategy_version.c.id == version_id)
+            )
+        )
+        .mappings()
+        .one()
+    )
+    return {**dict(row), **dict(version)}
+
+
 def _selector_id(value: StrategySelector) -> int | None:
     if isinstance(value, int):
         return value

+ 7 - 4
script_build_host/src/script_build_host/application/input_snapshot_service.py

@@ -73,10 +73,13 @@ class ScriptInputSnapshotService:
             personal_config.get("account_name") if isinstance(personal_config, dict) else None
         )
         persona = await self._persona_source.load(str(account_name or ""))
-        strategies_config = build_record.get("strategies_config") or {}
-        always_on = request.strategies_always_on or tuple(strategies_config.get("always_on", []))
-        on_demand = request.strategies_on_demand or tuple(strategies_config.get("on_demand", []))
-        strategies = await self._strategy_source.load(always_on=always_on, on_demand=on_demand)
+        # topic_build_record.strategies_config contains *topic* strategies in the
+        # legacy schema. Script strategies come only from the Script Build request
+        # (including retry/prefill), otherwise an absent selection means none.
+        strategies = await self._strategy_source.load(
+            always_on=request.strategies_always_on,
+            on_demand=request.strategies_on_demand,
+        )
         loaded_prompts = await self._prompt_source.load(request.prompt_requests)
         prompts = (*request.runtime_prompt_manifest, *loaded_prompts)
         sanitized_topic = redact(topic)

+ 10 - 1
script_build_host/src/script_build_host/infrastructure/manifests.py

@@ -28,6 +28,15 @@ class SettingsRuntimeManifestProvider:
             ("decode", self.settings.decode_endpoint),
             ("embedding", self.settings.embedding_endpoint),
             ("external", self.settings.external_endpoint),
+            ("xhs_search", self.settings.xhs_search_endpoint),
+            ("xhs_detail", self.settings.xhs_detail_endpoint),
+            ("zhihu_search", self.settings.zhihu_search_endpoint),
+            (
+                "openrouter_chat",
+                self.settings.openrouter_chat_endpoint
+                if self.settings.openrouter_api_key is not None
+                else None,
+            ),
             ("image", self.settings.image_endpoint),
         ):
             if endpoint:
@@ -48,7 +57,7 @@ def _credential_free_endpoint(value: str) -> str:
     parsed = urlsplit(value)
     host = parsed.hostname or ""
     port = f":{parsed.port}" if parsed.port and parsed.port != 443 else ""
-    return urlunsplit(("https", f"{host}{port}", parsed.path, "", ""))
+    return urlunsplit((parsed.scheme, f"{host}{port}", parsed.path, "", ""))
 
 
 def _path_manifest(path: Path) -> dict[str, object]:

+ 165 - 49
script_build_host/src/script_build_host/repositories/legacy_input.py

@@ -1,9 +1,10 @@
 from __future__ import annotations
 
+import asyncio
 import json
 from typing import Any
 
-from sqlalchemy import or_, select
+from sqlalchemy import select
 from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
 
 from script_build_host.domain.errors import InputRelationMismatch
@@ -44,14 +45,51 @@ def _json_value(value: Any) -> Any:
 
 
 def _row_dict(row: Any, fields: tuple[str, ...]) -> dict[str, Any]:
-    return {field: row[field] for field in fields}
+    return {field: row[field] if field in row else None for field in fields}
+
+
+async def _rows_by_primary_key(
+    sessions: async_sessionmaker[AsyncSession],
+    table: Any,
+    identifiers: list[int],
+    *,
+    columns: tuple[Any, ...] | None = None,
+) -> list[Any]:
+    """Read LOB-bearing legacy rows by PK; old RDS plans stall on covering-range lookups."""
+
+    limiter = asyncio.Semaphore(8)
+
+    async def load(identifier: int) -> Any:
+        async with limiter, sessions() as session:
+            return (
+                (
+                    await session.execute(
+                        select(*(columns or (table,))).where(table.c.id == identifier)
+                    )
+                )
+                .mappings()
+                .one_or_none()
+            )
+
+    rows = []
+    for row in await asyncio.gather(*(load(identifier) for identifier in identifiers)):
+        if row is None:
+            raise InputRelationMismatch()
+        rows.append(row)
+    return rows
 
 
 class LegacySqlAlchemyInputReader:
     """Anti-corruption adapter over the old tables; it never imports the old runtime package."""
 
-    def __init__(self, sessions: async_sessionmaker[AsyncSession]) -> None:
+    def __init__(
+        self,
+        sessions: async_sessionmaker[AsyncSession],
+        *,
+        include_raw_upload: bool = False,
+    ) -> None:
         self._sessions = sessions
+        self._include_raw_upload = include_raw_upload
 
     async def read_topic_graph(
         self, *, execution_id: int, topic_build_id: int, topic_id: int
@@ -71,10 +109,15 @@ class LegacySqlAlchemyInputReader:
                     .mappings()
                     .one_or_none()
                 )
-            build = (
+            build_base = (
                 (
                     await session.execute(
-                        select(topic_build_record).where(
+                        select(
+                            topic_build_record.c.id,
+                            topic_build_record.c.execution_id,
+                            topic_build_record.c.status,
+                            topic_build_record.c.origin,
+                        ).where(
                             topic_build_record.c.id == topic_build_id,
                             topic_build_record.c.execution_id == execution_id,
                             topic_build_record.c.is_deleted.is_(False),
@@ -84,6 +127,26 @@ class LegacySqlAlchemyInputReader:
                 .mappings()
                 .one_or_none()
             )
+            build = dict(build_base) if build_base is not None else None
+            if build is not None:
+                build_columns = (
+                    "demand",
+                    "demand_constraints",
+                    "agent_type",
+                    "agent_config",
+                    "strategies_config",
+                    "personal_config",
+                    "start_time",
+                    "end_time",
+                    *(("raw_upload_data",) if self._include_raw_upload else ()),
+                )
+                for column_name in build_columns:
+                    column = topic_build_record.c[column_name]
+                    build[column_name] = (
+                        await session.execute(
+                            select(column).where(topic_build_record.c.id == topic_build_id)
+                        )
+                    ).scalar_one()
             topic = (
                 (
                     await session.execute(
@@ -108,10 +171,10 @@ class LegacySqlAlchemyInputReader:
             ):
                 raise InputRelationMismatch()
 
-            point_rows = (
+            point_id_rows = (
                 (
                     await session.execute(
-                        select(topic_build_point)
+                        select(topic_build_point.c.id)
                         .where(
                             topic_build_point.c.topic_id == topic_id,
                             topic_build_point.c.build_id == topic_build_id,
@@ -120,13 +183,13 @@ class LegacySqlAlchemyInputReader:
                         .order_by(topic_build_point.c.id)
                     )
                 )
-                .mappings()
+                .scalars()
                 .all()
             )
-            item_rows = (
+            item_id_rows = (
                 (
                     await session.execute(
-                        select(topic_build_composition_item)
+                        select(topic_build_composition_item.c.id)
                         .where(
                             topic_build_composition_item.c.topic_id == topic_id,
                             topic_build_composition_item.c.build_id == topic_build_id,
@@ -138,58 +201,113 @@ class LegacySqlAlchemyInputReader:
                         )
                     )
                 )
-                .mappings()
+                .scalars()
                 .all()
             )
-            point_ids = {int(row["id"]) for row in point_rows}
-            item_ids = {int(row["id"]) for row in item_rows}
-            point_relation_rows = (
-                (
-                    await session.execute(
-                        select(topic_build_point_item_relation).where(
-                            or_(
+            point_id_order = [int(value) for value in point_id_rows]
+            item_id_order = [int(value) for value in item_id_rows]
+            point_rows = await _rows_by_primary_key(
+                self._sessions,
+                topic_build_point,
+                point_id_order,
+                columns=tuple(
+                    topic_build_point.c[name]
+                    for name in (
+                        "id",
+                        "topic_id",
+                        "build_id",
+                        "point_type",
+                        "point_result",
+                        "is_active",
+                        "note",
+                    )
+                ),
+            )
+            item_rows = await _rows_by_primary_key(
+                self._sessions,
+                topic_build_composition_item,
+                item_id_order,
+                columns=tuple(
+                    topic_build_composition_item.c[name]
+                    for name in (
+                        "id",
+                        "topic_id",
+                        "build_id",
+                        "item_level",
+                        "dimension",
+                        "point_type",
+                        "element_name",
+                        "category_path",
+                        "category_id",
+                        "derivation_type",
+                        "step",
+                        "sort_order",
+                        "is_active",
+                        "note",
+                        "created_at",
+                        "updated_at",
+                    )
+                ),
+            )
+            point_ids = set(point_id_order)
+            item_ids = set(item_id_order)
+            point_relation_ids = (
+                list(
+                    (
+                        await session.execute(
+                            select(topic_build_point_item_relation.c.id).where(
                                 topic_build_point_item_relation.c.point_id.in_(point_ids),
-                                topic_build_point_item_relation.c.item_id.in_(item_ids),
-                            ),
+                            )
                         )
-                    )
+                    ).scalars()
                 )
-                .mappings()
-                .all()
-                if point_ids or item_ids
+                if point_ids
                 else []
             )
-            item_relation_rows = (
-                (
-                    await session.execute(
-                        select(topic_build_item_relation)
-                        .where(
-                            topic_build_item_relation.c.topic_id == topic_id,
+            point_relation_rows = await _rows_by_primary_key(
+                self._sessions,
+                topic_build_point_item_relation,
+                [int(value) for value in point_relation_ids],
+            )
+            item_relation_ids = (
+                list(
+                    (
+                        await session.execute(
+                            select(topic_build_item_relation.c.id)
+                            .where(
+                                topic_build_item_relation.c.topic_id == topic_id,
+                            )
+                            .order_by(topic_build_item_relation.c.id)
                         )
-                        .order_by(topic_build_item_relation.c.id)
-                    )
+                    ).scalars()
                 )
-                .mappings()
-                .all()
                 if item_ids
                 else []
             )
-            source_rows = (
-                (
-                    await session.execute(
-                        select(topic_build_item_source)
-                        .where(
-                            topic_build_item_source.c.topic_id == topic_id,
-                            topic_build_item_source.c.is_active.is_(True),
+            item_relation_rows = await _rows_by_primary_key(
+                self._sessions,
+                topic_build_item_relation,
+                [int(value) for value in item_relation_ids],
+            )
+            source_ids = (
+                list(
+                    (
+                        await session.execute(
+                            select(topic_build_item_source.c.id)
+                            .where(
+                                topic_build_item_source.c.topic_id == topic_id,
+                                topic_build_item_source.c.is_active.is_(True),
+                            )
+                            .order_by(topic_build_item_source.c.id)
                         )
-                        .order_by(topic_build_item_source.c.id)
-                    )
+                    ).scalars()
                 )
-                .mappings()
-                .all()
                 if item_ids
                 else []
             )
+            source_rows = await _rows_by_primary_key(
+                self._sessions, topic_build_item_source, [int(value) for value in source_ids]
+            )
 
         if any(
             int(row["point_id"]) not in point_ids or int(row["item_id"]) not in item_ids
@@ -241,11 +359,8 @@ class LegacySqlAlchemyInputReader:
             "build_id",
             "point_type",
             "point_result",
-            "reason",
             "is_active",
             "note",
-            "created_at",
-            "updated_at",
         )
         item_fields = (
             "id",
@@ -259,7 +374,6 @@ class LegacySqlAlchemyInputReader:
             "category_id",
             "derivation_type",
             "step",
-            "reason",
             "sort_order",
             "is_active",
             "note",
@@ -280,6 +394,8 @@ class LegacySqlAlchemyInputReader:
             "created_at",
         )
         build_dict = _row_dict(build, build_fields)
+        if not self._include_raw_upload:
+            build_dict["raw_upload_data"] = None
         build_dict["demand_constraints"] = _mapping(build_dict["demand_constraints"])
         build_dict["agent_config"] = _mapping(build_dict["agent_config"])
         build_dict["strategies_config"] = _mapping(build_dict["strategies_config"])

+ 2 - 0
script_build_host/tests/test_input_snapshot_service.py

@@ -87,6 +87,8 @@ async def test_validate_source_is_read_only_and_assemble_redacts_secrets() -> No
             topic_build_id=20,
             topic_id=30,
             principal=Principal("user-1"),
+            strategies_always_on=(7,),
+            strategies_on_demand=(8,),
             datasource_manifest={"endpoint": "https://api.example?q=ok&token=hidden"},
             model_manifest={"planner": {"model": "fake", "api_key": "hidden"}},
         )

+ 3 - 3
script_build_host/tests/test_legacy_input_and_adapters.py

@@ -135,7 +135,7 @@ async def test_legacy_reader_freezes_relations_sources_and_arbitrary_raw_upload(
 ) -> None:
     _, sessions = database
     await _seed_topic(sessions, ["uploaded", {"nested": True}])
-    graph = await LegacySqlAlchemyInputReader(sessions).read_topic_graph(
+    graph = await LegacySqlAlchemyInputReader(sessions, include_raw_upload=True).read_topic_graph(
         execution_id=10, topic_build_id=20, topic_id=30
     )
     assert graph["build_record"]["raw_upload_data"] == ["uploaded", {"nested": True}]
@@ -162,8 +162,8 @@ async def test_legacy_reader_fails_closed_on_relation_mismatch(
         await session.execute(
             insert(topic_build_point_item_relation).values(
                 id=61,
-                point_id=999,
-                item_id=50,
+                point_id=40,
+                item_id=999,
             )
         )
     with pytest.raises(InputRelationMismatch):