test_retrieval_adapters.py 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385
  1. from __future__ import annotations
  2. import json
  3. from types import SimpleNamespace
  4. import httpx
  5. import pytest
  6. from script_build_host.adapters.retrieval import (
  7. ExternalRetrievalAdapter,
  8. FileDecodeRetrievalAdapter,
  9. FileKnowledgeRetrievalAdapter,
  10. HttpDecodeRetrievalAdapter,
  11. PatternRetrievalAdapter,
  12. SafeHttpClient,
  13. SafeImageAdapter,
  14. )
  15. from script_build_host.adapters.uploaded_topic import SqlUploadedTopicGateway
  16. from script_build_host.domain.errors import ProtocolViolation, UnsafeOutboundTarget
  17. from script_build_host.infrastructure.outbound import OutboundPolicy
  18. from script_build_host.repositories.legacy_input import LegacySqlAlchemyInputReader
  19. async def _public_resolver(_: str, __: int) -> tuple[str, ...]:
  20. return ("93.184.216.34",)
  21. class _RawArtifacts:
  22. def __init__(self) -> None:
  23. self.values: list[bytes] = []
  24. async def freeze_bytes(self, content: bytes, *, media_type: str) -> str:
  25. assert media_type == "image/png"
  26. self.values.append(content)
  27. return "script-build://raw-artifacts/sha256/" + "a" * 64
  28. @pytest.mark.asyncio
  29. async def test_http_retrieval_revalidates_redirect_and_freezes_response_metadata() -> None:
  30. calls: list[str] = []
  31. def handler(request: httpx.Request) -> httpx.Response:
  32. calls.append(str(request.url))
  33. return httpx.Response(
  34. 200,
  35. json={"data": [{"channel_content_id": "post-1", "title": "case"}]},
  36. )
  37. async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client:
  38. safe = SafeHttpClient(
  39. client,
  40. OutboundPolicy(frozenset({"api.example.com"}), resolver=_public_resolver),
  41. )
  42. result = await ExternalRetrievalAdapter("https://api.example.com/search", safe).retrieve(
  43. query={"keyword": "case"}, snapshot=SimpleNamespace()
  44. )
  45. assert calls == ["https://api.example.com/search"]
  46. assert result.source_refs == ("external:post-1",)
  47. assert str(result.metadata["response_sha256"]).startswith("sha256:")
  48. def redirect(_: httpx.Request) -> httpx.Response:
  49. return httpx.Response(302, headers={"location": "https://127.0.0.1/private"})
  50. async with httpx.AsyncClient(transport=httpx.MockTransport(redirect)) as client:
  51. safe = SafeHttpClient(
  52. client,
  53. OutboundPolicy(frozenset({"api.example.com"}), resolver=_public_resolver),
  54. )
  55. with pytest.raises(UnsafeOutboundTarget):
  56. await safe.request("GET", "https://api.example.com/start")
  57. def oversized(_: httpx.Request) -> httpx.Response:
  58. return httpx.Response(200, content=b"123456")
  59. async with httpx.AsyncClient(transport=httpx.MockTransport(oversized)) as client:
  60. safe = SafeHttpClient(
  61. client,
  62. OutboundPolicy(frozenset({"api.example.com"}), resolver=_public_resolver),
  63. max_response_bytes=5,
  64. )
  65. with pytest.raises(ProtocolViolation, match="byte limit"):
  66. await safe.request("GET", "https://api.example.com/large")
  67. @pytest.mark.asyncio
  68. async def test_external_response_secrets_are_redacted_before_evidence_summary() -> None:
  69. def handler(_request: httpx.Request) -> httpx.Response:
  70. return httpx.Response(
  71. 200,
  72. json={
  73. "data": [
  74. {
  75. "id": "p1",
  76. "authorization": "Bearer raw-token",
  77. "url": "https://source.example/item?token=secret",
  78. }
  79. ]
  80. },
  81. )
  82. async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client:
  83. safe = SafeHttpClient(
  84. client,
  85. OutboundPolicy(frozenset({"api.example.com"}), resolver=_public_resolver),
  86. )
  87. result = await ExternalRetrievalAdapter("https://api.example.com/search", safe).retrieve(
  88. query={"keyword": "case"}, snapshot=SimpleNamespace()
  89. )
  90. assert "raw-token" not in result.summary
  91. assert "secret" not in result.summary
  92. @pytest.mark.asyncio
  93. async def test_pattern_polling_and_image_limits_use_mock_transport() -> None:
  94. poll_count = 0
  95. def handler(request: httpx.Request) -> httpx.Response:
  96. nonlocal poll_count
  97. if request.url.path == "/pattern":
  98. return httpx.Response(
  99. 200,
  100. json={
  101. "status": "pending",
  102. "session_id": "task-1",
  103. "poll_url": "https://api.example.com/poll/task-1",
  104. },
  105. )
  106. if request.url.path.startswith("/poll"):
  107. poll_count += 1
  108. return httpx.Response(200, json={"status": "done", "results": [{"id": 2}]})
  109. return httpx.Response(
  110. 200,
  111. content=b"\x89PNG\r\n\x1a\nimage",
  112. headers={"content-type": "image/png"},
  113. )
  114. async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client:
  115. safe = SafeHttpClient(
  116. client,
  117. OutboundPolicy(frozenset({"api.example.com"}), resolver=_public_resolver),
  118. )
  119. result = await PatternRetrievalAdapter(
  120. "https://api.example.com/pattern",
  121. safe,
  122. poll_interval_seconds=0.001,
  123. ).retrieve(query={"message": "shape"}, snapshot=SimpleNamespace())
  124. raw_store = _RawArtifacts()
  125. images = await SafeImageAdapter(
  126. safe,
  127. raw_store,
  128. max_images=1,
  129. max_image_bytes=20,
  130. max_total_image_bytes=20,
  131. ).load(urls=["https://api.example.com/image"], snapshot=SimpleNamespace())
  132. assert poll_count == 1
  133. assert result.source_refs == ("pattern:2",)
  134. assert images[0]["digest"].startswith("sha256:")
  135. assert "url" not in images[0]
  136. assert images[0]["raw_artifact_ref"].startswith("script-build://raw-artifacts/")
  137. assert len(raw_store.values) == 1
  138. @pytest.mark.asyncio
  139. async def test_pattern_timeout_is_frozen_as_a_limitation() -> None:
  140. def handler(_request: httpx.Request) -> httpx.Response:
  141. return httpx.Response(200, json={"status": "pending", "session_id": "task-1"})
  142. async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client:
  143. safe = SafeHttpClient(
  144. client,
  145. OutboundPolicy(frozenset({"api.example.com"}), resolver=_public_resolver),
  146. )
  147. result = await PatternRetrievalAdapter(
  148. "https://api.example.com/pattern",
  149. safe,
  150. poll_interval_seconds=0.001,
  151. poll_timeout_seconds=0,
  152. ).retrieve(query={"message": "shape"}, snapshot=SimpleNamespace())
  153. assert result.source_refs == ()
  154. assert result.limitations == ("timeout",)
  155. assert result.metadata["task_id"] == "task-1"
  156. @pytest.mark.asyncio
  157. async def test_pattern_result_keeps_three_candidates_without_duplicate_items() -> None:
  158. candidates = [
  159. {
  160. "itemset": [{"dimension": f"dimension-{index}"}],
  161. "target_items": [{"dimension": f"dimension-{index}"}],
  162. "support": 10 - index,
  163. }
  164. for index in range(6)
  165. ]
  166. async def handler(_request: httpx.Request) -> httpx.Response:
  167. return httpx.Response(
  168. 200,
  169. json=[{"status": "success", "data": {"candidates": candidates}}],
  170. )
  171. async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client:
  172. safe = SafeHttpClient(
  173. client,
  174. OutboundPolicy(frozenset({"api.example.com"}), resolver=_public_resolver),
  175. )
  176. result = await PatternRetrievalAdapter("https://api.example.com/pattern", safe).retrieve(
  177. query={"message": "shape"}, snapshot=SimpleNamespace()
  178. )
  179. payload = json.loads(result.summary)[0]
  180. assert payload["data"]["candidates_total"] == 6
  181. assert payload["data"]["candidates_returned"] == 3
  182. assert len(payload["data"]["candidates"]) == 3
  183. assert all("target_items" not in item for item in payload["data"]["candidates"])
  184. @pytest.mark.asyncio
  185. async def test_pattern_result_has_a_hard_context_size_limit() -> None:
  186. candidate = {
  187. "itemset": [{"dimension": "x" * 10_000}],
  188. "target_items": [{"dimension": "x" * 10_000}],
  189. "examples": ["post-1"],
  190. }
  191. async def handler(_request: httpx.Request) -> httpx.Response:
  192. return httpx.Response(
  193. 200,
  194. json=[
  195. {
  196. "status": "success",
  197. "output": "large result",
  198. "data": {"candidates_total": 193, "candidates": [candidate] * 20},
  199. }
  200. ],
  201. )
  202. async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client:
  203. safe = SafeHttpClient(
  204. client,
  205. OutboundPolicy(frozenset({"api.example.com"}), resolver=_public_resolver),
  206. )
  207. result = await PatternRetrievalAdapter("https://api.example.com/pattern", safe).retrieve(
  208. query={"message": "shape"}, snapshot=SimpleNamespace()
  209. )
  210. assert len(result.summary) <= 20_100
  211. payload = json.loads(result.summary)[0]
  212. assert payload["data"]["candidates_total"] == 193
  213. assert payload["data"]["truncated_for_context"] is True
  214. @pytest.mark.asyncio
  215. async def test_file_decode_and_knowledge_freeze_file_hash_and_rank(tmp_path) -> None:
  216. decode_path = tmp_path / "decode.json"
  217. decode_path.write_text(
  218. json.dumps({"items": [{"post_id": 4, "account": "acct", "text": "focus"}]}),
  219. encoding="utf-8",
  220. )
  221. knowledge_path = tmp_path / "knowledge.json"
  222. knowledge_path.write_text(json.dumps([{"id": "k1", "title": "focus method"}]), encoding="utf-8")
  223. decode_adapter = FileDecodeRetrievalAdapter(decode_path)
  224. snapshot = SimpleNamespace(
  225. model_manifest={"embedding_model": "frozen-embed"},
  226. datasource_manifest={},
  227. )
  228. decode = await decode_adapter.retrieve(
  229. query={
  230. "keyword": "focus",
  231. "account_name": "acct",
  232. "top_k": 3,
  233. "return_field": "主脉络",
  234. },
  235. snapshot=snapshot,
  236. )
  237. knowledge = await FileKnowledgeRetrievalAdapter(knowledge_path).retrieve(
  238. query={"keyword": "focus", "max_count": 3}, snapshot=snapshot
  239. )
  240. assert decode.metadata["embedding_model"] == "frozen-embed"
  241. assert decode.metadata["ranks"] == [1]
  242. assert decode.metadata["scores"] == [1.0]
  243. assert str(knowledge.metadata["file_sha256"]).startswith("sha256:")
  244. frozen_digest = decode.metadata["index_sha256"]
  245. decode_path.write_text(json.dumps({"items": []}), encoding="utf-8")
  246. replay = await decode_adapter.retrieve(
  247. query={"keyword": "focus", "top_k": 3, "return_field": "主脉络"},
  248. snapshot=SimpleNamespace(
  249. model_manifest={},
  250. datasource_manifest={"decode_index": {"sha256": frozen_digest}},
  251. ),
  252. )
  253. assert replay.source_refs == ("decode:4",)
  254. with pytest.raises(ProtocolViolation, match="changed"):
  255. await FileDecodeRetrievalAdapter(decode_path).retrieve(
  256. query={"keyword": "focus", "top_k": 3, "return_field": "主脉络"},
  257. snapshot=SimpleNamespace(
  258. model_manifest={},
  259. datasource_manifest={"decode_index": {"sha256": frozen_digest}},
  260. ),
  261. )
  262. @pytest.mark.asyncio
  263. async def test_http_decode_preserves_legacy_query_and_ranked_scores() -> None:
  264. requests: list[dict[str, object]] = []
  265. def handler(request: httpx.Request) -> httpx.Response:
  266. requests.append(json.loads(request.content))
  267. return httpx.Response(
  268. 200,
  269. json={"items": [{"post_id": "p1", "score": 0.91, "text": "safe"}]},
  270. )
  271. async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client:
  272. safe = SafeHttpClient(
  273. client,
  274. OutboundPolicy(frozenset({"decode.example"}), resolver=_public_resolver),
  275. )
  276. result = await HttpDecodeRetrievalAdapter(
  277. "https://decode.example/search",
  278. safe,
  279. ).retrieve(
  280. query={
  281. "return_field": "主脉络",
  282. "keyword": "focus",
  283. "account_name": "acct",
  284. "match_fields": {},
  285. "top_k": 3,
  286. },
  287. snapshot=SimpleNamespace(
  288. datasource_manifest={"decode_index": {"sha256": "sha256:" + "1" * 64}},
  289. model_manifest={"embedding_model": "embed-v1"},
  290. ),
  291. )
  292. assert requests[0]["top_k"] == 3
  293. assert result.source_refs == ("decode:p1",)
  294. assert result.metadata["scores"] == [0.91]
  295. assert result.metadata["ranks"] == [1]
  296. @pytest.mark.asyncio
  297. async def test_uploaded_topic_parser_keeps_old_shape_without_writing() -> None:
  298. gateway = SqlUploadedTopicGateway(SimpleNamespace()) # type: ignore[arg-type]
  299. parsed = await gateway.parse(
  300. {
  301. "选题融合": "topic",
  302. "target_post": {"channel_account_name": "acct"},
  303. "灵感点": [
  304. {
  305. "点": "point",
  306. "实质": {"具体元素": [{"名称": "element", "说明": "note"}]},
  307. }
  308. ],
  309. }
  310. )
  311. assert parsed["account_name"] == "acct"
  312. assert parsed["item_count"] == 1
  313. assert parsed["points"][0]["items"][0]["dimension"] == "实质"
  314. @pytest.mark.asyncio
  315. async def test_uploaded_topic_create_writes_a_readable_legacy_graph(database) -> None:
  316. _, sessions = database
  317. gateway = SqlUploadedTopicGateway(sessions)
  318. parsed = await gateway.parse(
  319. {
  320. "选题融合": "uploaded topic",
  321. "target_post": {"channel_account_name": "acct"},
  322. "关键点": [
  323. {
  324. "点": "point",
  325. "形式": [{"名称": "element", "说明": "note"}],
  326. }
  327. ],
  328. }
  329. )
  330. created = await gateway.create(parsed, account_name=None)
  331. graph = await LegacySqlAlchemyInputReader(sessions).read_topic_graph(
  332. execution_id=0,
  333. topic_build_id=created["topic_build_id"],
  334. topic_id=created["topic_id"],
  335. )
  336. assert graph["execution"] == {"id": 0, "status": "upload"}
  337. assert graph["topic"]["result"] == "uploaded topic"
  338. assert graph["points"][0]["item_ids"] == [graph["composition_items"][0]["id"]]