postgres.py 18 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551
  1. """PostgreSQL implementation of the formal acquisition repository."""
  2. from __future__ import annotations
  3. from typing import Any
  4. from uuid import UUID
  5. import psycopg2.extras
  6. from acquisition.domain import (
  7. AcquisitionJob,
  8. AcquisitionRun,
  9. CandidateItem,
  10. ItemClassification,
  11. MediaAsset,
  12. Query,
  13. QueryBatch,
  14. )
  15. Json = psycopg2.extras.Json
  16. psycopg2.extras.register_uuid()
  17. class PostgresAcquisitionRepository:
  18. """Repository backed by the formal cloud PostgreSQL schema.
  19. The repository does not commit by itself; callers own transaction scope via
  20. core.db_session.transaction or an equivalent connection boundary.
  21. """
  22. def __init__(self, conn: Any):
  23. self.conn = conn
  24. def _one(self, sql: str, params: tuple[Any, ...]) -> dict[str, Any]:
  25. with self.conn.cursor(cursor_factory=psycopg2.extras.RealDictCursor) as cur:
  26. cur.execute(sql, params)
  27. row = cur.fetchone()
  28. if row is None:
  29. raise RuntimeError("expected one row, got none")
  30. return dict(row)
  31. def _all(self, sql: str, params: tuple[Any, ...]) -> list[dict[str, Any]]:
  32. with self.conn.cursor(cursor_factory=psycopg2.extras.RealDictCursor) as cur:
  33. cur.execute(sql, params)
  34. return [dict(row) for row in cur.fetchall()]
  35. def create_query_batch(
  36. self,
  37. *,
  38. name: str,
  39. source_type: str = "manual",
  40. generation_method: str | None = None,
  41. target_platforms: list[str] | None = None,
  42. status: str = "draft",
  43. metadata: dict[str, Any] | None = None,
  44. ) -> QueryBatch:
  45. row = self._one(
  46. """
  47. INSERT INTO query_batches(
  48. name, source_type, generation_method, target_platforms,
  49. status, metadata
  50. )
  51. VALUES (%s, %s, %s, %s, %s, %s)
  52. RETURNING *
  53. """,
  54. (
  55. name,
  56. source_type,
  57. generation_method,
  58. target_platforms or [],
  59. status,
  60. Json(metadata or {}),
  61. ),
  62. )
  63. return QueryBatch.model_validate(row)
  64. def add_query(
  65. self,
  66. *,
  67. batch_id: UUID | None,
  68. query_text: str,
  69. axes: dict[str, Any] | None = None,
  70. keep: bool | None = None,
  71. filter_reason: str | None = None,
  72. status: str = "draft",
  73. sort_order: int = 0,
  74. metadata: dict[str, Any] | None = None,
  75. ) -> Query:
  76. row = self._one(
  77. """
  78. INSERT INTO queries(
  79. batch_id, query_text, axes, keep, filter_reason,
  80. status, sort_order, metadata
  81. )
  82. VALUES (%s, %s, %s, %s, %s, %s, %s, %s)
  83. RETURNING *
  84. """,
  85. (
  86. batch_id,
  87. query_text,
  88. Json(axes or {}),
  89. keep,
  90. filter_reason,
  91. status,
  92. sort_order,
  93. Json(metadata or {}),
  94. ),
  95. )
  96. return Query.model_validate(row)
  97. def list_queries_for_batch(
  98. self,
  99. batch_id: UUID,
  100. *,
  101. keep: bool | None = None,
  102. ) -> list[Query]:
  103. if keep is None:
  104. rows = self._all(
  105. """
  106. SELECT * FROM queries
  107. WHERE batch_id = %s
  108. ORDER BY sort_order, created_at
  109. """,
  110. (batch_id,),
  111. )
  112. else:
  113. rows = self._all(
  114. """
  115. SELECT * FROM queries
  116. WHERE batch_id = %s AND keep IS NOT DISTINCT FROM %s
  117. ORDER BY sort_order, created_at
  118. """,
  119. (batch_id, keep),
  120. )
  121. return [Query.model_validate(row) for row in rows]
  122. def get_query_batch(self, batch_id: UUID) -> QueryBatch:
  123. row = self._one("SELECT * FROM query_batches WHERE id = %s", (batch_id,))
  124. return QueryBatch.model_validate(row)
  125. def create_acquisition_run(
  126. self,
  127. *,
  128. batch_id: UUID | None = None,
  129. run_key: str | None = None,
  130. status: str = "pending",
  131. note: str | None = None,
  132. metadata: dict[str, Any] | None = None,
  133. ) -> AcquisitionRun:
  134. row = self._one(
  135. """
  136. INSERT INTO acquisition_runs(batch_id, run_key, status, note, metadata)
  137. VALUES (%s, %s, %s, %s, %s)
  138. ON CONFLICT (run_key) DO UPDATE SET
  139. status = EXCLUDED.status,
  140. note = EXCLUDED.note,
  141. metadata = EXCLUDED.metadata
  142. RETURNING *
  143. """,
  144. (batch_id, run_key, status, note, Json(metadata or {})),
  145. )
  146. return AcquisitionRun.model_validate(row)
  147. def ensure_acquisition_job(
  148. self,
  149. *,
  150. run_id: UUID,
  151. query_id: UUID | None,
  152. platform: str,
  153. search_limit: int | None = None,
  154. display_limit: int | None = None,
  155. status: str = "pending",
  156. metadata: dict[str, Any] | None = None,
  157. ) -> AcquisitionJob:
  158. row = self._one(
  159. """
  160. INSERT INTO acquisition_jobs(
  161. run_id, query_id, platform, search_limit,
  162. display_limit, status, metadata
  163. )
  164. VALUES (%s, %s, %s, %s, %s, %s, %s)
  165. ON CONFLICT (run_id, query_id, platform) DO UPDATE SET
  166. search_limit = EXCLUDED.search_limit,
  167. display_limit = EXCLUDED.display_limit,
  168. metadata = EXCLUDED.metadata
  169. RETURNING *
  170. """,
  171. (
  172. run_id,
  173. query_id,
  174. platform,
  175. search_limit,
  176. display_limit,
  177. status,
  178. Json(metadata or {}),
  179. ),
  180. )
  181. return AcquisitionJob.model_validate(row)
  182. def update_acquisition_job(
  183. self,
  184. job_id: UUID,
  185. *,
  186. status: str,
  187. attempt_count: int | None = None,
  188. error_message: str | None = None,
  189. metadata: dict[str, Any] | None = None,
  190. ) -> AcquisitionJob:
  191. row = self._one(
  192. """
  193. UPDATE acquisition_jobs SET
  194. status = %s,
  195. attempt_count = COALESCE(%s, attempt_count),
  196. error_message = %s,
  197. metadata = CASE WHEN %s THEN %s ELSE metadata END
  198. WHERE id = %s
  199. RETURNING *
  200. """,
  201. (
  202. status,
  203. attempt_count,
  204. error_message,
  205. metadata is not None,
  206. Json(metadata or {}),
  207. job_id,
  208. ),
  209. )
  210. return AcquisitionJob.model_validate(row)
  211. def upsert_candidate_item(
  212. self,
  213. *,
  214. platform: str,
  215. job_id: UUID | None = None,
  216. query_id: UUID | None = None,
  217. platform_item_id: str | None = None,
  218. canonical_url: str | None = None,
  219. content_type: str | None = None,
  220. title: str | None = None,
  221. author_name: str | None = None,
  222. raw_summary: str | None = None,
  223. status: str = "candidate",
  224. source_payload: dict[str, Any] | None = None,
  225. metadata: dict[str, Any] | None = None,
  226. error_message: str | None = None,
  227. ) -> CandidateItem:
  228. existing_id = None
  229. if platform_item_id:
  230. row = self._one_or_none(
  231. """
  232. SELECT id FROM candidate_items
  233. WHERE platform = %s AND platform_item_id = %s
  234. ORDER BY created_at DESC
  235. LIMIT 1
  236. """,
  237. (platform, platform_item_id),
  238. )
  239. existing_id = row["id"] if row else None
  240. if existing_id:
  241. row = self._one(
  242. """
  243. UPDATE candidate_items SET
  244. job_id = %s,
  245. query_id = %s,
  246. canonical_url = %s,
  247. content_type = %s,
  248. title = %s,
  249. author_name = %s,
  250. raw_summary = %s,
  251. status = %s,
  252. source_payload = %s,
  253. metadata = %s,
  254. error_message = %s
  255. WHERE id = %s
  256. RETURNING *
  257. """,
  258. (
  259. job_id,
  260. query_id,
  261. canonical_url,
  262. content_type,
  263. title,
  264. author_name,
  265. raw_summary,
  266. status,
  267. Json(source_payload or {}),
  268. Json(metadata or {}),
  269. error_message,
  270. existing_id,
  271. ),
  272. )
  273. else:
  274. row = self._one(
  275. """
  276. INSERT INTO candidate_items(
  277. job_id, query_id, platform, platform_item_id, canonical_url,
  278. content_type, title, author_name, raw_summary, status,
  279. source_payload, metadata, error_message
  280. )
  281. VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
  282. RETURNING *
  283. """,
  284. (
  285. job_id,
  286. query_id,
  287. platform,
  288. platform_item_id,
  289. canonical_url,
  290. content_type,
  291. title,
  292. author_name,
  293. raw_summary,
  294. status,
  295. Json(source_payload or {}),
  296. Json(metadata or {}),
  297. error_message,
  298. ),
  299. )
  300. return CandidateItem.model_validate(row)
  301. def _one_or_none(self, sql: str, params: tuple[Any, ...]) -> dict[str, Any] | None:
  302. with self.conn.cursor(cursor_factory=psycopg2.extras.RealDictCursor) as cur:
  303. cur.execute(sql, params)
  304. row = cur.fetchone()
  305. return dict(row) if row else None
  306. def add_media_asset(
  307. self,
  308. *,
  309. item_id: UUID,
  310. media_type: str,
  311. source_url: str | None = None,
  312. oss_url: str | None = None,
  313. cdn_url: str | None = None,
  314. position: int = 0,
  315. status: str = "pending",
  316. source_payload: dict[str, Any] | None = None,
  317. metadata: dict[str, Any] | None = None,
  318. ) -> MediaAsset:
  319. row = self._one(
  320. """
  321. INSERT INTO media_assets(
  322. item_id, media_type, source_url, oss_url, cdn_url,
  323. position, status, source_payload, metadata
  324. )
  325. VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s)
  326. RETURNING *
  327. """,
  328. (
  329. item_id,
  330. media_type,
  331. source_url,
  332. oss_url,
  333. cdn_url,
  334. position,
  335. status,
  336. Json(source_payload or {}),
  337. Json(metadata or {}),
  338. ),
  339. )
  340. return MediaAsset.model_validate(row)
  341. def add_item_classification(
  342. self,
  343. *,
  344. item_id: UUID,
  345. is_creation_knowledge: bool | None = None,
  346. label: str | None = None,
  347. confidence: float | None = None,
  348. reason: str | None = None,
  349. model_name: str | None = None,
  350. prompt_version: str | None = None,
  351. result_payload: dict[str, Any] | None = None,
  352. status: str = "pending",
  353. error_message: str | None = None,
  354. ) -> ItemClassification:
  355. row = self._one(
  356. """
  357. INSERT INTO item_classifications(
  358. item_id, is_creation_knowledge, label, confidence, reason,
  359. model_name, prompt_version, result_payload, status, error_message
  360. )
  361. VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
  362. RETURNING *
  363. """,
  364. (
  365. item_id,
  366. is_creation_knowledge,
  367. label,
  368. confidence,
  369. reason,
  370. model_name,
  371. prompt_version,
  372. Json(result_payload or {}),
  373. status,
  374. error_message,
  375. ),
  376. )
  377. return ItemClassification.model_validate(row)
  378. def get_run_summary(self, run_id: UUID) -> dict[str, Any]:
  379. summary = self._one(
  380. """
  381. SELECT
  382. ar.id,
  383. ar.run_key,
  384. ar.batch_id,
  385. ar.status,
  386. ar.started_at,
  387. ar.finished_at,
  388. COUNT(DISTINCT q.id)::int AS query_count,
  389. COUNT(DISTINCT aj.id)::int AS job_count,
  390. COUNT(DISTINCT ci.id)::int AS candidate_count,
  391. COUNT(DISTINCT ic.id) FILTER (
  392. WHERE ic.is_creation_knowledge IS TRUE
  393. )::int AS creation_hit_count
  394. FROM acquisition_runs ar
  395. LEFT JOIN acquisition_jobs aj ON aj.run_id = ar.id
  396. LEFT JOIN queries q ON q.id = aj.query_id
  397. LEFT JOIN candidate_items ci ON ci.job_id = aj.id
  398. LEFT JOIN item_classifications ic ON ic.item_id = ci.id
  399. WHERE ar.id = %s
  400. GROUP BY ar.id
  401. """,
  402. (run_id,),
  403. )
  404. summary["queries"] = self._all(
  405. """
  406. SELECT
  407. q.id AS query_id,
  408. q.query_text,
  409. COUNT(DISTINCT aj.id)::int AS job_count,
  410. COUNT(DISTINCT ci.id)::int AS candidate_count,
  411. COUNT(DISTINCT ic.id) FILTER (
  412. WHERE ic.is_creation_knowledge IS TRUE
  413. )::int AS creation_hit_count,
  414. jsonb_object_agg(
  415. aj.platform,
  416. jsonb_build_object(
  417. 'status', aj.status,
  418. 'attempt_count', aj.attempt_count,
  419. 'display_limit', aj.display_limit,
  420. 'search_limit', aj.search_limit,
  421. 'error_message', aj.error_message
  422. )
  423. ) FILTER (WHERE aj.id IS NOT NULL) AS platforms
  424. FROM queries q
  425. JOIN acquisition_jobs aj ON aj.query_id = q.id
  426. LEFT JOIN candidate_items ci ON ci.job_id = aj.id
  427. LEFT JOIN item_classifications ic ON ic.item_id = ci.id
  428. WHERE aj.run_id = %s
  429. GROUP BY q.id, q.query_text, q.sort_order
  430. ORDER BY q.sort_order, q.query_text
  431. """,
  432. (run_id,),
  433. )
  434. return summary
  435. def get_query_detail(self, *, run_id: UUID, query_id: UUID) -> dict[str, Any]:
  436. query = self._one("SELECT * FROM queries WHERE id = %s", (query_id,))
  437. jobs = self._all(
  438. """
  439. SELECT * FROM acquisition_jobs
  440. WHERE run_id = %s AND query_id = %s
  441. ORDER BY platform
  442. """,
  443. (run_id, query_id),
  444. )
  445. items = self._all(
  446. """
  447. SELECT ci.* FROM candidate_items ci
  448. JOIN acquisition_jobs aj ON aj.id = ci.job_id
  449. WHERE aj.run_id = %s AND ci.query_id = %s
  450. ORDER BY ci.platform, ci.created_at
  451. """,
  452. (run_id, query_id),
  453. )
  454. item_ids = [row["id"] for row in items]
  455. media: list[dict[str, Any]] = []
  456. classifications: list[dict[str, Any]] = []
  457. if item_ids:
  458. media = self._all(
  459. """
  460. SELECT * FROM media_assets
  461. WHERE item_id = ANY(%s)
  462. ORDER BY item_id, position
  463. """,
  464. (item_ids,),
  465. )
  466. classifications = self._all(
  467. """
  468. SELECT * FROM item_classifications
  469. WHERE item_id = ANY(%s)
  470. ORDER BY created_at DESC
  471. """,
  472. (item_ids,),
  473. )
  474. return {
  475. "query": query,
  476. "jobs": jobs,
  477. "items": items,
  478. "media_assets": media,
  479. "classifications": classifications,
  480. }
  481. def list_creation_candidate_items(
  482. self,
  483. *,
  484. run_id: UUID | None = None,
  485. limit: int = 100,
  486. ) -> list[CandidateItem]:
  487. if run_id is None:
  488. rows = self._all(
  489. """
  490. SELECT ci.* FROM candidate_items ci
  491. JOIN item_classifications ic ON ic.item_id = ci.id
  492. WHERE ic.is_creation_knowledge IS TRUE
  493. ORDER BY ci.created_at
  494. LIMIT %s
  495. """,
  496. (limit,),
  497. )
  498. else:
  499. rows = self._all(
  500. """
  501. SELECT ci.* FROM candidate_items ci
  502. JOIN acquisition_jobs aj ON aj.id = ci.job_id
  503. JOIN item_classifications ic ON ic.item_id = ci.id
  504. WHERE aj.run_id = %s AND ic.is_creation_knowledge IS TRUE
  505. ORDER BY ci.created_at
  506. LIMIT %s
  507. """,
  508. (run_id, limit),
  509. )
  510. return [CandidateItem.model_validate(row) for row in rows]
  511. def get_candidate_item(self, item_id: UUID) -> CandidateItem:
  512. row = self._one("SELECT * FROM candidate_items WHERE id = %s", (item_id,))
  513. return CandidateItem.model_validate(row)
  514. def list_media_assets_for_item(self, item_id: UUID) -> list[MediaAsset]:
  515. rows = self._all(
  516. """
  517. SELECT * FROM media_assets
  518. WHERE item_id = %s
  519. ORDER BY position
  520. """,
  521. (item_id,),
  522. )
  523. return [MediaAsset.model_validate(row) for row in rows]