mysql_demand_content_sink.py 9.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257
  1. """MySQL sink for DB-validated DemandAgent demand_content rows.
  2. This module is intentionally narrow: it only writes the test demand_content
  3. table. It must not create demand_task rows, Hive rows, Pattern rows, or any
  4. other side effects.
  5. """
  6. from __future__ import annotations
  7. import copy
  8. import json
  9. import os
  10. from dataclasses import dataclass, field
  11. from typing import Any
  12. from urllib.parse import unquote, urlparse
  13. import pymysql
  14. @dataclass
  15. class MySQLDemandContentWriteResult:
  16. inserted_count: int = 0
  17. skipped_count: int = 0
  18. inserted_ids: list[int] = field(default_factory=list)
  19. skipped: list[dict[str, Any]] = field(default_factory=list)
  20. def _env_value(*names: str, default: str = "") -> str:
  21. for name in names:
  22. value = os.getenv(name)
  23. if value is not None and str(value).strip():
  24. return str(value).strip()
  25. return default
  26. def _resolve_mysql_config() -> dict[str, Any]:
  27. dsn = _env_value("DEMAND_CONTENT_MYSQL_DSN")
  28. if dsn:
  29. parsed = urlparse(dsn)
  30. return {
  31. "host": parsed.hostname or "127.0.0.1",
  32. "port": int(parsed.port or 3306),
  33. "user": unquote(parsed.username or ""),
  34. "password": unquote(parsed.password or ""),
  35. "database": unquote(parsed.path.lstrip("/") or ""),
  36. }
  37. return {
  38. "host": _env_value("DEMAND_CONTENT_MYSQL_HOST", "CONTENT_SUPPLY_DB_HOST", default="127.0.0.1"),
  39. "port": int(_env_value("DEMAND_CONTENT_MYSQL_PORT", "CONTENT_SUPPLY_DB_PORT", default="3306")),
  40. "user": _env_value("DEMAND_CONTENT_MYSQL_USER", "CONTENT_SUPPLY_DB_USER", default="content_rw"),
  41. "password": _env_value("DEMAND_CONTENT_MYSQL_PASSWORD", "CONTENT_SUPPLY_DB_PASSWORD"),
  42. "database": _env_value(
  43. "DEMAND_CONTENT_MYSQL_DB",
  44. "CONTENT_SUPPLY_DB_NAME",
  45. default="content-deconstruction-supply",
  46. ),
  47. }
  48. def _connect():
  49. cfg = _resolve_mysql_config()
  50. missing = [key for key in ("host", "user", "database") if not cfg.get(key)]
  51. if missing:
  52. raise RuntimeError(f"MySQL demand_content config missing: {missing}")
  53. return pymysql.connect(
  54. host=cfg["host"],
  55. port=int(cfg["port"]),
  56. user=cfg["user"],
  57. password=cfg["password"],
  58. database=cfg["database"],
  59. charset="utf8mb4",
  60. cursorclass=pymysql.cursors.DictCursor,
  61. autocommit=False,
  62. )
  63. def _as_ext_data(value: Any) -> dict[str, Any]:
  64. if isinstance(value, dict):
  65. return copy.deepcopy(value)
  66. if isinstance(value, str) and value.strip():
  67. parsed = json.loads(value)
  68. if isinstance(parsed, dict):
  69. return parsed
  70. return {}
  71. def _dedupe_key(row: dict[str, Any], ext_data: dict[str, Any]) -> tuple[str, str, str]:
  72. evidence_pack = ext_data.get("evidence_pack") or {}
  73. source_post_id = str(evidence_pack.get("source_post_id") or "")
  74. return (
  75. str(row.get("name") or ""),
  76. source_post_id,
  77. "",
  78. )
  79. def _row_exists_for_run(cursor, *, run_label: str, row: dict[str, Any], ext_data: dict[str, Any]) -> bool:
  80. if not run_label:
  81. return False
  82. name, source_post_id, _ = _dedupe_key(row, ext_data)
  83. cursor.execute(
  84. """
  85. SELECT id
  86. FROM demand_content
  87. WHERE JSON_UNQUOTE(JSON_EXTRACT(ext_data, '$.run_label')) = %s
  88. AND name = %s
  89. AND JSON_UNQUOTE(JSON_EXTRACT(ext_data, '$.evidence_pack.source_post_id')) = %s
  90. LIMIT 1
  91. """,
  92. (run_label, name, source_post_id),
  93. )
  94. return cursor.fetchone() is not None
  95. def _validate_row(row: dict[str, Any], ext_data: dict[str, Any]) -> None:
  96. evidence_pack = ext_data.get("evidence_pack")
  97. if not isinstance(evidence_pack, dict):
  98. raise ValueError("missing ext_data.evidence_pack")
  99. expected = {
  100. "pattern_source_system": "pg_pattern_v2",
  101. "source_kind": "pattern_itemset",
  102. "case_id_type": "post_id",
  103. "source_certainty": "db_validated",
  104. "validation_status": "passed",
  105. }
  106. for field, value in expected.items():
  107. if evidence_pack.get(field) != value:
  108. raise ValueError(f"invalid evidence_pack.{field}: {evidence_pack.get(field)!r}")
  109. required = [
  110. "source_post_id",
  111. "pattern_execution_id",
  112. "mining_config_id",
  113. "itemset_ids",
  114. "itemset_items",
  115. "category_bindings",
  116. "element_bindings",
  117. "support",
  118. "absolute_support",
  119. "matched_post_ids",
  120. "video_ids",
  121. "case_ids",
  122. "seed_terms",
  123. "trace_id",
  124. "query_seed_points",
  125. "demand_scope",
  126. ]
  127. for field in required:
  128. value = evidence_pack.get(field)
  129. if field == "query_seed_points":
  130. if value is None or not isinstance(value, list):
  131. raise ValueError("missing evidence_pack.query_seed_points")
  132. continue
  133. if value is None or value == "" or value == []:
  134. raise ValueError(f"missing evidence_pack.{field}")
  135. if not isinstance(evidence_pack.get("itemset_ids"), list) or len(evidence_pack["itemset_ids"]) != 1:
  136. raise ValueError("evidence_pack.itemset_ids must contain exactly one itemset_id")
  137. matched_post_ids = [str(post_id) for post_id in evidence_pack.get("matched_post_ids") or []]
  138. if str(evidence_pack["source_post_id"]) not in set(matched_post_ids):
  139. raise ValueError("evidence_pack.source_post_id must be in matched_post_ids")
  140. if [str(post_id) for post_id in evidence_pack.get("video_ids") or []] != matched_post_ids:
  141. raise ValueError("evidence_pack.video_ids must equal matched_post_ids")
  142. if [str(post_id) for post_id in evidence_pack.get("case_ids") or []] != matched_post_ids:
  143. raise ValueError("evidence_pack.case_ids must equal matched_post_ids")
  144. demand_scope = evidence_pack.get("demand_scope") or {}
  145. if not isinstance(demand_scope, dict):
  146. raise ValueError("evidence_pack.demand_scope must be object")
  147. if demand_scope.get("merge_leve2") and row.get("merge_leve2") != demand_scope.get("merge_leve2"):
  148. raise ValueError("evidence_pack.demand_scope.merge_leve2 must equal row.merge_leve2")
  149. scoped_count = evidence_pack.get("scoped_post_count")
  150. filtered_support = evidence_pack.get("filtered_absolute_support")
  151. if scoped_count is not None:
  152. scoped_count = int(scoped_count)
  153. if scoped_count <= 0:
  154. raise ValueError("evidence_pack.scoped_post_count must be positive")
  155. if len(matched_post_ids) != scoped_count:
  156. raise ValueError("evidence_pack.matched_post_ids must cover scoped_post_count")
  157. if scoped_count > int(evidence_pack["absolute_support"]):
  158. raise ValueError("evidence_pack.scoped_post_count must not exceed absolute_support")
  159. elif len(matched_post_ids) < int(evidence_pack["absolute_support"]):
  160. raise ValueError("evidence_pack.matched_post_ids must cover absolute_support")
  161. if filtered_support is not None and scoped_count is not None and int(filtered_support) != scoped_count:
  162. raise ValueError("evidence_pack.filtered_absolute_support must equal scoped_post_count")
  163. expected_rank = 1
  164. for point in evidence_pack.get("query_seed_points") or []:
  165. if not isinstance(point, dict):
  166. raise ValueError("evidence_pack.query_seed_points items must be objects")
  167. if point.get("point_type") not in {"灵感点", "目的点"}:
  168. raise ValueError("evidence_pack.query_seed_points point_type must be 灵感点 or 目的点")
  169. if int(point.get("coverage_post_count") or 0) <= 0:
  170. raise ValueError("evidence_pack.query_seed_points coverage_post_count must be positive")
  171. if int(point.get("rank") or 0) != expected_rank:
  172. raise ValueError("evidence_pack.query_seed_points rank must be continuous")
  173. expected_rank += 1
  174. if not row.get("merge_leve2") or not row.get("name") or not row.get("dt"):
  175. raise ValueError("missing demand_content merge_leve2/name/dt")
  176. def write_demand_content_rows(
  177. rows: list[dict[str, Any]],
  178. *,
  179. run_label: str = "",
  180. ) -> MySQLDemandContentWriteResult:
  181. result = MySQLDemandContentWriteResult()
  182. if not rows:
  183. return result
  184. conn = _connect()
  185. try:
  186. with conn.cursor() as cursor:
  187. for row in rows:
  188. ext_data = _as_ext_data(row.get("ext_data"))
  189. if run_label:
  190. ext_data["run_label"] = run_label
  191. _validate_row(row, ext_data)
  192. if _row_exists_for_run(cursor, run_label=run_label, row=row, ext_data=ext_data):
  193. result.skipped_count += 1
  194. result.skipped.append({"name": row.get("name"), "reason": "duplicate_run_label_key"})
  195. continue
  196. cursor.execute(
  197. """
  198. INSERT INTO demand_content
  199. (merge_leve2, name, reason, suggestion, ext_data, score, dt)
  200. VALUES
  201. (%s, %s, %s, %s, %s, %s, %s)
  202. """,
  203. (
  204. row.get("merge_leve2"),
  205. row.get("name"),
  206. row.get("reason"),
  207. row.get("suggestion"),
  208. json.dumps(ext_data, ensure_ascii=False),
  209. row.get("score"),
  210. row.get("dt"),
  211. ),
  212. )
  213. demand_content_id = int(cursor.lastrowid)
  214. ext_data.setdefault("evidence_pack", {})["demand_content_id"] = demand_content_id
  215. cursor.execute(
  216. "UPDATE demand_content SET ext_data=%s WHERE id=%s",
  217. (json.dumps(ext_data, ensure_ascii=False), demand_content_id),
  218. )
  219. result.inserted_count += 1
  220. result.inserted_ids.append(demand_content_id)
  221. conn.commit()
  222. except Exception:
  223. conn.rollback()
  224. raise
  225. finally:
  226. conn.close()
  227. return result