Просмотр исходного кода

检索输入:冻结完整脱敏原文并压缩模型可见结果

Pattern、Decode、External 和 Knowledge 检索统一先脱敏、规范化 JSON 并写入 RawArtifactStore。Pattern 默认只返回有界候选但保留完整原文,Knowledge 先做词法预选再调用模型,移除不可达图片流程;补充 193 个候选和约 27 万字符压力验证。
SamLee 13 часов назад
Родитель
Сommit
343e5546b9

+ 39 - 23
script_build_host/src/script_build_host/adapters/legacy_retrieval.py

@@ -16,6 +16,7 @@ from pathlib import Path
 from typing import Any, Protocol
 
 from script_build_host.adapters.retrieval import SafeHttpClient
+from script_build_host.domain.context_broker import ContextCard, rank_context_cards
 from script_build_host.domain.errors import ProtocolViolation
 from script_build_host.infrastructure.canonical_json import canonical_sha256
 from script_build_host.infrastructure.redaction import redact
@@ -270,6 +271,18 @@ class LegacyLlmKnowledgeRetrievalAdapter:
     async def retrieve(self, *, query: Mapping[str, Any], snapshot: Any) -> RetrievalResult:
         del snapshot
         max_count = int(query.get("max_count", 3))
+        keyword = str(query.get("keyword", ""))
+        cards = tuple(
+            ContextCard(
+                handle=f"knowledge_{position}",
+                source_type="knowledge",
+                summary=" ".join(str(item.get(key) or "") for key in ("title", "purpose", "steps")),
+                metadata={"position": position},
+            )
+            for position, item in enumerate(self.items)
+        )
+        ranked = rank_context_cards(cards, keyword, limit=20)
+        candidates = [self.items[int(card.metadata["position"])] for card in ranked]
         index = [
             {
                 "id": item.get("id"),
@@ -277,31 +290,34 @@ class LegacyLlmKnowledgeRetrievalAdapter:
                 "purpose": item.get("purpose", ""),
                 "steps": item.get("steps", []),
             }
-            for item in self.items
+            for item in candidates
         ]
-        selected = await self.openrouter.chat_json(
-            model=self.model,
-            messages=[
-                {
-                    "role": "system",
-                    "content": (
-                        "你是一个创作知识检索助手。根据给定的检索关键词,从知识条目列表中"
-                        '找出最相关的条目 ID。只返回 JSON 对象,格式:{"ids": [<id>, ...]}。'
-                        "不要输出任何其他内容。"
-                    ),
-                },
-                {
-                    "role": "user",
-                    "content": (
-                        f"检索关键词:{query.get('keyword', '')}\n最多返回 {max_count} 个最相关的"
-                        f"条目 ID(按相关度从高到低排列)。\n\n知识条目列表:\n"
-                        + json.dumps(index, ensure_ascii=False, indent=2)
-                    ),
-                },
-            ],
-        )
+        try:
+            selected = await self.openrouter.chat_json(
+                model=self.model,
+                messages=[
+                    {
+                        "role": "system",
+                        "content": (
+                            "你是一个创作知识检索助手。根据给定的检索关键词,从知识条目列表中"
+                            '找出最相关的条目 ID。只返回 JSON 对象,格式:{"ids": [<id>, ...]}。'
+                            "不要输出任何其他内容。"
+                        ),
+                    },
+                    {
+                        "role": "user",
+                        "content": (
+                            f"检索关键词:{keyword}\n最多返回 {max_count} 个最相关的"
+                            f"条目 ID(按相关度从高到低排列)。\n\n知识条目列表:\n"
+                            + json.dumps(index, ensure_ascii=False, indent=2)
+                        ),
+                    },
+                ],
+            )
+        except Exception:
+            selected = {"ids": [item.get("id") for item in candidates[:max_count]]}
         ids = selected.get("ids", []) if isinstance(selected, dict) else []
-        by_id = {item.get("id"): item for item in self.items}
+        by_id = {item.get("id"): item for item in candidates}
         output = []
         for identifier in ids[:max_count] if isinstance(ids, list) else []:
             item = by_id.get(identifier)

+ 93 - 20
script_build_host/src/script_build_host/adapters/retrieval.py

@@ -16,7 +16,7 @@ from urllib.parse import urljoin
 import httpx
 
 from script_build_host.domain.errors import ProtocolViolation
-from script_build_host.infrastructure.canonical_json import canonical_sha256
+from script_build_host.infrastructure.canonical_json import canonical_json_bytes, canonical_sha256
 from script_build_host.infrastructure.outbound import OutboundPolicy
 from script_build_host.infrastructure.redaction import redact, redact_text
 from script_build_host.tools.gateway import RetrievalResult
@@ -105,12 +105,14 @@ class PatternRetrievalAdapter:
         self,
         endpoint: str,
         http: SafeHttpClient,
+        raw_store: RawArtifactStore,
         *,
         poll_interval_seconds: float = 1.0,
         poll_timeout_seconds: float = 600.0,
     ) -> None:
         self.endpoint = endpoint
         self.http = http
+        self.raw_store = raw_store
         self.poll_interval_seconds = poll_interval_seconds
         self.poll_timeout_seconds = poll_timeout_seconds
 
@@ -129,12 +131,18 @@ class PatternRetrievalAdapter:
             "queued",
         }:
             if time.monotonic() >= deadline:
+                safe_payload, raw_ref, response_digest = await _freeze_json_response(
+                    self.raw_store,
+                    payload,
+                )
                 return _limited_result(
                     "pattern",
                     "pattern retrieval timed out",
                     "timeout",
                     task_id=task_id,
-                    payload=payload,
+                    payload=safe_payload,
+                    raw_artifact_ref=raw_ref,
+                    response_digest=response_digest,
                 )
             poll_url = payload.get("poll_url")
             if not isinstance(poll_url, str) or not poll_url:
@@ -144,60 +152,80 @@ class PatternRetrievalAdapter:
             await asyncio.sleep(self.poll_interval_seconds)
             payload = await self.http.json("GET", poll_url)
         if isinstance(payload, dict) and payload.get("status") in {"failed", "error"}:
+            safe_payload, raw_ref, response_digest = await _freeze_json_response(
+                self.raw_store,
+                payload,
+            )
             return _limited_result(
                 "pattern",
                 "pattern retrieval failed",
                 "upstream failure",
                 task_id=task_id,
-                payload=payload,
+                payload=safe_payload,
+                raw_artifact_ref=raw_ref,
+                response_digest=response_digest,
             )
-        compact_payload = _compact_pattern_payload(payload)
+        safe_payload, raw_ref, response_digest = await _freeze_json_response(
+            self.raw_store,
+            payload,
+        )
+        compact_payload = _compact_pattern_payload(safe_payload)
         return _result(
             "pattern",
             compact_payload,
             metadata={
                 "task_id": _safe_source_id(task_id, rank=0) if task_id else None,
-                "response_sha256": canonical_sha256(payload).wire,
+                "response_sha256": response_digest,
             },
+            raw_artifact_ref=raw_ref,
         )
 
 
 class ExternalRetrievalAdapter:
     """Pure HTTP client; it contains no legacy-log repository dependency."""
 
-    def __init__(self, endpoint: str, http: SafeHttpClient) -> None:
+    def __init__(self, endpoint: str, http: SafeHttpClient, raw_store: RawArtifactStore) -> None:
         self.endpoint = endpoint
         self.http = http
+        self.raw_store = raw_store
 
     async def retrieve(self, *, query: Mapping[str, Any], snapshot: Any) -> RetrievalResult:
         del snapshot
         payload = await self.http.json("POST", self.endpoint, json_body=query)
+        safe_payload, raw_ref, response_digest = await _freeze_json_response(
+            self.raw_store,
+            payload,
+        )
         return _result(
             "external",
-            payload,
-            metadata={"response_sha256": canonical_sha256(payload).wire},
+            safe_payload,
+            metadata={"response_sha256": response_digest},
+            raw_artifact_ref=raw_ref,
         )
 
 
 class HttpDecodeRetrievalAdapter:
     """Call the configured Decode service with the legacy query contract."""
 
-    def __init__(self, endpoint: str, http: SafeHttpClient) -> None:
+    def __init__(self, endpoint: str, http: SafeHttpClient, raw_store: RawArtifactStore) -> None:
         self.endpoint = endpoint
         self.http = http
+        self.raw_store = raw_store
 
     async def retrieve(self, *, query: Mapping[str, Any], snapshot: Any) -> RetrievalResult:
         payload = await self.http.json("POST", self.endpoint, json_body=query)
+        safe_payload, raw_ref, response_digest = await _freeze_json_response(
+            self.raw_store,
+            payload,
+        )
         rows = (
-            payload.get("data", payload.get("results", payload.get("items", [])))
-            if isinstance(payload, dict)
-            else payload
+            safe_payload.get("data", safe_payload.get("results", safe_payload.get("items", [])))
+            if isinstance(safe_payload, dict)
+            else safe_payload
         )
         if not isinstance(rows, list):
             raise ProtocolViolation("decode response items must be an array")
-        safe_rows = redact(rows)
-        if not isinstance(safe_rows, list):
-            raise ProtocolViolation("redacted decode rows must remain an array")
+        safe_rows = rows
         refs: list[str] = []
         scores: list[float] = []
         for row in safe_rows:
@@ -217,18 +245,19 @@ class HttpDecodeRetrievalAdapter:
         model_manifest = getattr(snapshot, "model_manifest", {})
         return RetrievalResult(
             source_refs=tuple(refs),
-            summary=json.dumps(safe_rows, ensure_ascii=False, sort_keys=True),
+            summary=_bounded_json_summary(safe_rows),
             supports=tuple(str(query["return_field"]) for _ in refs),
             confidence="ranked",
             limitations=(),
             metadata={
-                "response_sha256": canonical_sha256(payload).wire,
+                "response_sha256": response_digest,
                 "decode_index": datasource_manifest.get("decode_index"),
                 "embedding_model": model_manifest.get("embedding_model"),
                 "hit_ids": refs,
                 "ranks": list(range(1, len(refs) + 1)),
                 "scores": scores,
             },
+            raw_artifact_ref=raw_ref,
         )
 
 
@@ -407,6 +436,7 @@ def _result(
     payload: Any,
     *,
     metadata: Mapping[str, Any],
+    raw_artifact_ref: str | None = None,
 ) -> RetrievalResult:
     rows = payload.get("data", payload.get("results", [])) if isinstance(payload, dict) else payload
     if not isinstance(rows, list):
@@ -422,18 +452,58 @@ def _result(
     )
     return RetrievalResult(
         source_refs=refs,
-        summary=json.dumps(safe_rows, ensure_ascii=False, sort_keys=True),
+        summary=_bounded_json_summary(safe_rows),
         supports=refs,
         confidence="upstream",
         limitations=(),
+        raw_artifact_ref=raw_artifact_ref,
         metadata=metadata,
     )
 
 
+async def _freeze_json_response(
+    raw_store: RawArtifactStore,
+    payload: Any,
+) -> tuple[Any, str, str]:
+    """Redact once, then freeze the exact canonical JSON represented by its digest."""
+
+    safe_payload = redact(payload)
+    content = canonical_json_bytes(safe_payload)
+    raw_ref = await raw_store.freeze_bytes(content, media_type="application/json")
+    return safe_payload, raw_ref, "sha256:" + sha256(content).hexdigest()
+
+
+def _bounded_json_summary(payload: Any, *, max_chars: int = 20_000) -> str:
+    rendered = json.dumps(payload, ensure_ascii=False, sort_keys=True)
+    if len(rendered) <= max_chars:
+        return rendered
+    rows = payload if isinstance(payload, list) else [payload]
+    kept: list[Any] = []
+    for row in rows:
+        candidate = [*kept, row]
+        if len(json.dumps(candidate, ensure_ascii=False, sort_keys=True)) > max_chars:
+            break
+        kept.append(row)
+    while True:
+        bounded = json.dumps(
+            {
+                "items": kept,
+                "items_returned": len(kept),
+                "items_total": len(rows),
+                "truncated_for_context": True,
+            },
+            ensure_ascii=False,
+            sort_keys=True,
+        )
+        if len(bounded) <= max_chars or not kept:
+            return bounded
+        kept.pop()
+
+
 def _compact_pattern_payload(
     payload: Any,
     *,
-    max_candidates: int = 3,
+    max_candidates: int = 5,
     max_chars: int = 20_000,
 ) -> Any:
     """Keep a bounded Pattern evidence view while preserving the raw digest upstream."""
@@ -525,6 +595,8 @@ def _limited_result(
     *,
     task_id: str,
     payload: Any,
+    raw_artifact_ref: str | None = None,
+    response_digest: str | None = None,
 ) -> RetrievalResult:
     return RetrievalResult(
         source_refs=(),
@@ -532,9 +604,10 @@ def _limited_result(
         supports=(),
         confidence="unavailable",
         limitations=(limitation,),
+        raw_artifact_ref=raw_artifact_ref,
         metadata={
             "task_id": _safe_source_id(task_id, rank=0) if task_id else None,
-            "response_sha256": canonical_sha256(payload).wire,
+            "response_sha256": response_digest or canonical_sha256(payload).wire,
         },
     )
 

+ 13 - 2
script_build_host/tests/test_legacy_active_inputs.py

@@ -179,7 +179,7 @@ async def test_legacy_decode_rejects_query_model_that_differs_from_index(tmp_pat
 
 
 @pytest.mark.asyncio
-async def test_legacy_knowledge_uses_full_llm_selection_contract(tmp_path) -> None:
+async def test_legacy_knowledge_lexically_preselects_at_most_twenty_before_llm(tmp_path) -> None:
     source = tmp_path / "knowledge.json"
     source.write_text(
         json.dumps(
@@ -195,7 +195,16 @@ async def test_legacy_knowledge_uses_full_llm_selection_contract(tmp_path) -> No
                             "outputs": [{"value": "钩子"}],
                         }
                     ],
-                }
+                },
+                *[
+                    {
+                        "id": 100 + index,
+                        "title": f"unrelated-{index}",
+                        "purpose": "other material",
+                        "steps": [],
+                    }
+                    for index in range(40)
+                ],
             ],
             ensure_ascii=False,
         ),
@@ -225,6 +234,8 @@ async def test_legacy_knowledge_uses_full_llm_selection_contract(tmp_path) -> No
         ).retrieve(query={"keyword": "结构", "max_count": 3}, snapshot=SimpleNamespace())
     assert bodies[0]["model"] == "knowledge-model"
     assert "先给结果" in bodies[0]["messages"][1]["content"]
+    assert bodies[0]["messages"][1]["content"].count('"id"') <= 20
+    assert "unrelated-39" not in bodies[0]["messages"][1]["content"]
     assert "钩子" in result.summary
 
 

+ 69 - 21
script_build_host/tests/test_retrieval_adapters.py

@@ -1,6 +1,7 @@
 from __future__ import annotations
 
 import json
+from hashlib import sha256
 from types import SimpleNamespace
 
 import httpx
@@ -27,12 +28,11 @@ async def _public_resolver(_: str, __: int) -> tuple[str, ...]:
 
 class _RawArtifacts:
     def __init__(self) -> None:
-        self.values: list[bytes] = []
+        self.values: list[tuple[str, bytes]] = []
 
     async def freeze_bytes(self, content: bytes, *, media_type: str) -> str:
-        assert media_type == "image/png"
-        self.values.append(content)
-        return "script-build://raw-artifacts/sha256/" + "a" * 64
+        self.values.append((media_type, content))
+        return "script-build://raw-artifacts/sha256/" + sha256(content).hexdigest()
 
 
 @pytest.mark.asyncio
@@ -51,12 +51,18 @@ async def test_http_retrieval_revalidates_redirect_and_freezes_response_metadata
             client,
             OutboundPolicy(frozenset({"api.example.com"}), resolver=_public_resolver),
         )
-        result = await ExternalRetrievalAdapter("https://api.example.com/search", safe).retrieve(
-            query={"keyword": "case"}, snapshot=SimpleNamespace()
-        )
+        raw_store = _RawArtifacts()
+        result = await ExternalRetrievalAdapter(
+            "https://api.example.com/search", safe, raw_store
+        ).retrieve(query={"keyword": "case"}, snapshot=SimpleNamespace())
     assert calls == ["https://api.example.com/search"]
     assert result.source_refs == ("external:post-1",)
     assert str(result.metadata["response_sha256"]).startswith("sha256:")
+    assert result.raw_artifact_ref is not None
+    assert result.raw_artifact_ref.endswith(
+        str(result.metadata["response_sha256"]).removeprefix("sha256:")
+    )
+    assert raw_store.values[0][0] == "application/json"
 
     def redirect(_: httpx.Request) -> httpx.Response:
         return httpx.Response(302, headers={"location": "https://127.0.0.1/private"})
@@ -103,11 +109,37 @@ async def test_external_response_secrets_are_redacted_before_evidence_summary()
             client,
             OutboundPolicy(frozenset({"api.example.com"}), resolver=_public_resolver),
         )
-        result = await ExternalRetrievalAdapter("https://api.example.com/search", safe).retrieve(
-            query={"keyword": "case"}, snapshot=SimpleNamespace()
-        )
+        raw_store = _RawArtifacts()
+        result = await ExternalRetrievalAdapter(
+            "https://api.example.com/search", safe, raw_store
+        ).retrieve(query={"keyword": "case"}, snapshot=SimpleNamespace())
     assert "raw-token" not in result.summary
     assert "secret" not in result.summary
+    frozen = raw_store.values[0][1].decode("utf-8")
+    assert "raw-token" not in frozen
+    assert "secret" not in frozen
+
+
+@pytest.mark.asyncio
+async def test_external_large_response_keeps_bounded_summary_and_complete_raw_json() -> None:
+    rows = [{"id": f"p{index}", "text": "x" * 2_000} for index in range(20)]
+
+    def handler(_request: httpx.Request) -> httpx.Response:
+        return httpx.Response(200, json={"data": rows})
+
+    async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client:
+        safe = SafeHttpClient(
+            client,
+            OutboundPolicy(frozenset({"api.example.com"}), resolver=_public_resolver),
+        )
+        raw_store = _RawArtifacts()
+        result = await ExternalRetrievalAdapter(
+            "https://api.example.com/search", safe, raw_store
+        ).retrieve(query={"keyword": "case"}, snapshot=SimpleNamespace())
+
+    assert len(result.summary) <= 20_000
+    assert json.loads(result.summary)["truncated_for_context"] is True
+    assert len(json.loads(raw_store.values[0][1])["data"]) == 20
 
 
 @pytest.mark.asyncio
@@ -139,12 +171,13 @@ async def test_pattern_polling_and_image_limits_use_mock_transport() -> None:
             client,
             OutboundPolicy(frozenset({"api.example.com"}), resolver=_public_resolver),
         )
+        raw_store = _RawArtifacts()
         result = await PatternRetrievalAdapter(
             "https://api.example.com/pattern",
             safe,
+            raw_store,
             poll_interval_seconds=0.001,
         ).retrieve(query={"message": "shape"}, snapshot=SimpleNamespace())
-        raw_store = _RawArtifacts()
         images = await SafeImageAdapter(
             safe,
             raw_store,
@@ -154,10 +187,14 @@ async def test_pattern_polling_and_image_limits_use_mock_transport() -> None:
         ).load(urls=["https://api.example.com/image"], snapshot=SimpleNamespace())
     assert poll_count == 1
     assert result.source_refs == ("pattern:2",)
+    assert result.raw_artifact_ref is not None
     assert images[0]["digest"].startswith("sha256:")
     assert "url" not in images[0]
     assert images[0]["raw_artifact_ref"].startswith("script-build://raw-artifacts/")
-    assert len(raw_store.values) == 1
+    assert [media_type for media_type, _ in raw_store.values] == [
+        "application/json",
+        "image/png",
+    ]
 
 
 @pytest.mark.asyncio
@@ -173,6 +210,7 @@ async def test_pattern_timeout_is_frozen_as_a_limitation() -> None:
         result = await PatternRetrievalAdapter(
             "https://api.example.com/pattern",
             safe,
+            _RawArtifacts(),
             poll_interval_seconds=0.001,
             poll_timeout_seconds=0,
         ).retrieve(query={"message": "shape"}, snapshot=SimpleNamespace())
@@ -182,7 +220,7 @@ async def test_pattern_timeout_is_frozen_as_a_limitation() -> None:
 
 
 @pytest.mark.asyncio
-async def test_pattern_result_keeps_three_candidates_without_duplicate_items() -> None:
+async def test_pattern_result_keeps_five_candidates_and_freezes_all_items() -> None:
     candidates = [
         {
             "itemset": [{"dimension": f"dimension-{index}"}],
@@ -203,15 +241,19 @@ async def test_pattern_result_keeps_three_candidates_without_duplicate_items() -
             client,
             OutboundPolicy(frozenset({"api.example.com"}), resolver=_public_resolver),
         )
-        result = await PatternRetrievalAdapter("https://api.example.com/pattern", safe).retrieve(
-            query={"message": "shape"}, snapshot=SimpleNamespace()
-        )
+        raw_store = _RawArtifacts()
+        result = await PatternRetrievalAdapter(
+            "https://api.example.com/pattern", safe, raw_store
+        ).retrieve(query={"message": "shape"}, snapshot=SimpleNamespace())
 
     payload = json.loads(result.summary)[0]
     assert payload["data"]["candidates_total"] == 6
-    assert payload["data"]["candidates_returned"] == 3
-    assert len(payload["data"]["candidates"]) == 3
+    assert payload["data"]["candidates_returned"] == 5
+    assert len(payload["data"]["candidates"]) == 5
     assert all("target_items" not in item for item in payload["data"]["candidates"])
+    frozen = json.loads(raw_store.values[0][1])
+    assert len(frozen[0]["data"]["candidates"]) == 6
+    assert result.raw_artifact_ref is not None
 
 
 @pytest.mark.asyncio
@@ -239,14 +281,16 @@ async def test_pattern_result_has_a_hard_context_size_limit() -> None:
             client,
             OutboundPolicy(frozenset({"api.example.com"}), resolver=_public_resolver),
         )
-        result = await PatternRetrievalAdapter("https://api.example.com/pattern", safe).retrieve(
-            query={"message": "shape"}, snapshot=SimpleNamespace()
-        )
+        raw_store = _RawArtifacts()
+        result = await PatternRetrievalAdapter(
+            "https://api.example.com/pattern", safe, raw_store
+        ).retrieve(query={"message": "shape"}, snapshot=SimpleNamespace())
 
     assert len(result.summary) <= 20_100
     payload = json.loads(result.summary)[0]
     assert payload["data"]["candidates_total"] == 193
     assert payload["data"]["truncated_for_context"] is True
+    assert len(raw_store.values[0][1]) > len(result.summary.encode("utf-8"))
 
 
 @pytest.mark.asyncio
@@ -316,9 +360,11 @@ async def test_http_decode_preserves_legacy_query_and_ranked_scores() -> None:
             client,
             OutboundPolicy(frozenset({"decode.example"}), resolver=_public_resolver),
         )
+        raw_store = _RawArtifacts()
         result = await HttpDecodeRetrievalAdapter(
             "https://decode.example/search",
             safe,
+            raw_store,
         ).retrieve(
             query={
                 "return_field": "主脉络",
@@ -336,6 +382,8 @@ async def test_http_decode_preserves_legacy_query_and_ranked_scores() -> None:
     assert result.source_refs == ("decode:p1",)
     assert result.metadata["scores"] == [0.91]
     assert result.metadata["ranks"] == [1]
+    assert json.loads(raw_store.values[0][1])["items"][0]["post_id"] == "p1"
+    assert result.raw_artifact_ref is not None
 
 
 @pytest.mark.asyncio