mysql_demand_content_sink.py 11 KB

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