test_legacy_input_and_adapters.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359
  1. from __future__ import annotations
  2. import json
  3. from pathlib import Path
  4. import pytest
  5. from sqlalchemy import insert
  6. from sqlalchemy.ext.asyncio import (
  7. AsyncEngine,
  8. AsyncSession,
  9. async_sessionmaker,
  10. create_async_engine,
  11. )
  12. from script_build_host.adapters.persona import FilePersonaSource
  13. from script_build_host.adapters.prompts import DatabaseFirstPromptSource
  14. from script_build_host.adapters.strategy import SqlAlchemyStrategySource
  15. from script_build_host.domain.errors import InputRelationMismatch
  16. from script_build_host.domain.ports import PromptRequest
  17. from script_build_host.infrastructure.legacy_tables import (
  18. build_strategy,
  19. build_strategy_version,
  20. prompt,
  21. topic_build_composition_item,
  22. topic_build_item_relation,
  23. topic_build_item_source,
  24. topic_build_point,
  25. topic_build_point_item_relation,
  26. topic_build_record,
  27. topic_build_topic,
  28. topic_pattern_execution,
  29. legacy_metadata,
  30. )
  31. from script_build_host.repositories.legacy_input import LegacySqlAlchemyInputReader
  32. async def _seed_topic(sessions: async_sessionmaker[AsyncSession], raw_upload_data: object) -> None:
  33. async with sessions() as session, session.begin():
  34. await session.execute(insert(topic_pattern_execution).values(id=10, status="success"))
  35. await session.execute(
  36. insert(topic_build_record).values(
  37. id=20,
  38. execution_id=10,
  39. demand="make a useful script",
  40. demand_constraints={"audience": "reader"},
  41. agent_type="AigcAgent",
  42. agent_config={"model_name": "test-model", "api_key": "must-redact-later"},
  43. status="success",
  44. is_deleted=False,
  45. strategies_config={"always_on": [7], "on_demand": [8]},
  46. personal_config={"account_name": "每天心理学"},
  47. origin="upload",
  48. raw_upload_data=raw_upload_data,
  49. )
  50. )
  51. await session.execute(
  52. insert(topic_build_topic).values(
  53. id=30,
  54. build_id=20,
  55. execution_id=10,
  56. sort_order=0,
  57. topic_direction='{"title":"direction"}',
  58. result="topic result",
  59. status="mature",
  60. )
  61. )
  62. await session.execute(
  63. insert(topic_build_point),
  64. [
  65. {
  66. "id": 40,
  67. "topic_id": 30,
  68. "build_id": 20,
  69. "point_type": "key",
  70. "point_result": "active",
  71. "is_active": True,
  72. },
  73. {
  74. "id": 41,
  75. "topic_id": 30,
  76. "build_id": 20,
  77. "point_type": "key",
  78. "point_result": "inactive",
  79. "is_active": False,
  80. },
  81. ],
  82. )
  83. await session.execute(
  84. insert(topic_build_composition_item),
  85. [
  86. {
  87. "id": 50,
  88. "topic_id": 30,
  89. "build_id": 20,
  90. "item_level": "element",
  91. "element_name": "first",
  92. "step": 1,
  93. "sort_order": 2,
  94. "is_active": True,
  95. "derivation_type": "user_demand",
  96. },
  97. {
  98. "id": 51,
  99. "topic_id": 30,
  100. "build_id": 20,
  101. "item_level": "element",
  102. "element_name": "second",
  103. "step": 1,
  104. "sort_order": 1,
  105. "is_active": True,
  106. "derivation_type": "agent_reasoning",
  107. },
  108. ],
  109. )
  110. await session.execute(
  111. insert(topic_build_point_item_relation).values(id=60, point_id=40, item_id=50)
  112. )
  113. await session.execute(
  114. insert(topic_build_item_relation).values(
  115. id=70, topic_id=30, source_item_id=51, target_item_id=50, reason="because"
  116. )
  117. )
  118. await session.execute(
  119. insert(topic_build_item_source).values(
  120. id=80,
  121. topic_id=30,
  122. target_item_id=50,
  123. derivation_type="post_extract",
  124. source_type="post",
  125. source_reference_id="post-1",
  126. source_detail={"title": "source"},
  127. dataset_from="account",
  128. is_active=True,
  129. )
  130. )
  131. @pytest.mark.asyncio
  132. async def test_legacy_reader_freezes_relations_sources_and_arbitrary_raw_upload(
  133. database: tuple[AsyncEngine, async_sessionmaker[AsyncSession]],
  134. ) -> None:
  135. _, sessions = database
  136. await _seed_topic(sessions, ["uploaded", {"nested": True}])
  137. graph = await LegacySqlAlchemyInputReader(sessions, include_raw_upload=True).read_topic_graph(
  138. execution_id=10, topic_build_id=20, topic_id=30
  139. )
  140. assert graph["build_record"]["raw_upload_data"] == ["uploaded", {"nested": True}]
  141. assert [point["id"] for point in graph["points"]] == [40]
  142. assert [item["id"] for item in graph["composition_items"]] == [51, 50]
  143. assert graph["points"][0]["item_ids"] == [50]
  144. assert graph["relations"][0]["source_item_id"] == 51
  145. assert graph["sources"][0]["source_reference_id"] == "post-1"
  146. assert graph["sources"][0]["derivation_type"] == "post_extract"
  147. assert graph["sources"][0]["dataset_from"] == "account"
  148. @pytest.mark.asyncio
  149. async def test_legacy_reader_fails_closed_on_relation_mismatch(
  150. database: tuple[AsyncEngine, async_sessionmaker[AsyncSession]],
  151. ) -> None:
  152. _, sessions = database
  153. await _seed_topic(sessions, {})
  154. reader = LegacySqlAlchemyInputReader(sessions)
  155. with pytest.raises(InputRelationMismatch):
  156. await reader.read_topic_graph(execution_id=11, topic_build_id=20, topic_id=30)
  157. async with sessions() as session, session.begin():
  158. await session.execute(
  159. insert(topic_build_point_item_relation).values(
  160. id=61,
  161. point_id=40,
  162. item_id=999,
  163. )
  164. )
  165. with pytest.raises(InputRelationMismatch):
  166. await reader.read_topic_graph(execution_id=10, topic_build_id=20, topic_id=30)
  167. @pytest.mark.asyncio
  168. async def test_legacy_reader_routes_uploaded_topics_to_runtime_database(
  169. database: tuple[AsyncEngine, async_sessionmaker[AsyncSession]],
  170. tmp_path: Path,
  171. ) -> None:
  172. _, input_sessions = database
  173. runtime_engine = create_async_engine(f"sqlite+aiosqlite:///{tmp_path / 'runtime.db'}")
  174. async with runtime_engine.begin() as connection:
  175. await connection.run_sync(legacy_metadata.create_all)
  176. runtime_sessions = async_sessionmaker(runtime_engine, expire_on_commit=False)
  177. try:
  178. async with runtime_sessions() as session, session.begin():
  179. await session.execute(
  180. insert(topic_build_record).values(
  181. id=21,
  182. execution_id=0,
  183. agent_type="UploadedTopic",
  184. status="success",
  185. is_deleted=False,
  186. personal_config={"account_name": "uploaded-account"},
  187. origin="upload",
  188. )
  189. )
  190. await session.execute(
  191. insert(topic_build_topic).values(
  192. id=31,
  193. build_id=21,
  194. execution_id=0,
  195. sort_order=0,
  196. result="uploaded topic",
  197. status="mature",
  198. )
  199. )
  200. await session.execute(
  201. insert(topic_build_point).values(
  202. id=41,
  203. topic_id=31,
  204. build_id=21,
  205. point_type="关键点",
  206. point_result="uploaded point",
  207. is_active=True,
  208. )
  209. )
  210. await session.execute(
  211. insert(topic_build_composition_item).values(
  212. id=51,
  213. topic_id=31,
  214. build_id=21,
  215. item_level="element",
  216. dimension="实质",
  217. point_type="关键点",
  218. element_name="uploaded element",
  219. step=0,
  220. sort_order=0,
  221. is_active=True,
  222. )
  223. )
  224. await session.execute(
  225. insert(topic_build_point_item_relation).values(
  226. id=61,
  227. point_id=41,
  228. item_id=51,
  229. )
  230. )
  231. graph = await LegacySqlAlchemyInputReader(
  232. input_sessions,
  233. uploaded_sessions=runtime_sessions,
  234. ).read_topic_graph(execution_id=0, topic_build_id=21, topic_id=31)
  235. assert graph["build_record"]["origin"] == "upload"
  236. assert graph["topic"]["result"] == "uploaded topic"
  237. assert graph["points"][0]["item_ids"] == [51]
  238. finally:
  239. await runtime_engine.dispose()
  240. @pytest.mark.asyncio
  241. async def test_persona_adapter_applies_alias_filters_and_versions(tmp_path: Path) -> None:
  242. persona_root = tmp_path / "persona"
  243. section_root = tmp_path / "sections"
  244. alias = "每天一点心理学"
  245. (persona_root / alias).mkdir(parents=True)
  246. (section_root / alias).mkdir(parents=True)
  247. (persona_root / alias / "point_records.json").write_text(
  248. json.dumps(
  249. {
  250. "records": [
  251. {"阶段": "创作", "点类型": "维度值", "名称": "keep", "帖子覆盖率": 0.2},
  252. {"阶段": "创作", "点类型": "维度值", "名称": "drop", "帖子覆盖率": 0.19},
  253. {"阶段": "选题", "点类型": "维度名", "名称": "wrong stage", "帖子覆盖率": 1},
  254. ]
  255. },
  256. ensure_ascii=False,
  257. ),
  258. encoding="utf-8",
  259. )
  260. (section_root / alias / "section_main_dimension_template.json").write_text(
  261. json.dumps({"分段规律摘要": "summary", "ignored": "value"}, ensure_ascii=False),
  262. encoding="utf-8",
  263. )
  264. loaded = await FilePersonaSource(
  265. persona_root=persona_root, section_pattern_root=section_root
  266. ).load("每天心理学")
  267. assert loaded.resolved_account_name == alias
  268. assert [item["点名称"] for item in loaded.persona_points] == ["keep"]
  269. assert loaded.section_patterns == ({"分段规律摘要": "summary"},)
  270. assert set(loaded.source_versions) == {"persona_points", "section_patterns"}
  271. @pytest.mark.asyncio
  272. async def test_strategy_adapter_uses_old_integer_ids_and_current_versions(
  273. database: tuple[AsyncEngine, async_sessionmaker[AsyncSession]],
  274. ) -> None:
  275. _, sessions = database
  276. async with sessions() as session, session.begin():
  277. await session.execute(
  278. insert(build_strategy),
  279. [
  280. {
  281. "id": 7,
  282. "strategy_type": "script",
  283. "name": "always",
  284. "description": "A",
  285. "is_active": True,
  286. "current_version": 2,
  287. },
  288. {
  289. "id": 8,
  290. "strategy_type": "script",
  291. "name": "demand",
  292. "description": "B",
  293. "is_active": True,
  294. "current_version": 1,
  295. },
  296. ],
  297. )
  298. await session.execute(
  299. insert(build_strategy_version),
  300. [
  301. {"id": 70, "strategy_id": 7, "version": 1, "content": "old"},
  302. {"id": 71, "strategy_id": 7, "version": 2, "content": "always content"},
  303. {"id": 80, "strategy_id": 8, "version": 1, "content": "demand content"},
  304. ],
  305. )
  306. loaded = await SqlAlchemyStrategySource(sessions).load(always_on=(7,), on_demand=(8,))
  307. assert [(item["strategy_id"], item["mode"], item["version"]) for item in loaded] == [
  308. (7, "always_on", 2),
  309. (8, "on_demand", 1),
  310. ]
  311. assert loaded[0]["content"] == "always content"
  312. assert loaded[1]["content"] == "demand content"
  313. assert loaded[1]["content_sha256"].startswith("sha256:")
  314. @pytest.mark.asyncio
  315. async def test_prompt_adapter_records_actual_db_and_file_sources(
  316. database: tuple[AsyncEngine, async_sessionmaker[AsyncSession]], tmp_path: Path
  317. ) -> None:
  318. _, sessions = database
  319. async with sessions() as session, session.begin():
  320. await session.execute(
  321. insert(prompt).values(
  322. id=1,
  323. biz_type="script_planner",
  324. name="planner",
  325. prompt_content="from database",
  326. current_version=3,
  327. )
  328. )
  329. (tmp_path / "worker.md").write_text("from file", encoding="utf-8")
  330. source = DatabaseFirstPromptSource(sessions, fallback_root=tmp_path)
  331. manifests = await source.load(
  332. (
  333. PromptRequest("script_planner", "planner.md", "script_planner", "planner"),
  334. PromptRequest("worker", "worker.md", "script_worker", "worker"),
  335. )
  336. )
  337. assert [(item["source"], item["version"]) for item in manifests] == [
  338. ("database", 3),
  339. ("file", None),
  340. ]