| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272 |
- """MySQL sink for DB-validated DemandAgent demand_content rows.
- This module is intentionally narrow: it only writes the test demand_content
- table. It must not create demand_task rows, Hive rows, Pattern rows, or any
- other side effects.
- """
- from __future__ import annotations
- import copy
- import json
- import os
- from dataclasses import dataclass, field
- from typing import Any
- from urllib.parse import unquote, urlparse
- import pymysql
- from examples.demand.pg_pattern_repository import DEFAULT_DEMAND_PLATFORM, normalize_platform
- @dataclass
- class MySQLDemandContentWriteResult:
- inserted_count: int = 0
- skipped_count: int = 0
- inserted_ids: list[int] = field(default_factory=list)
- skipped: list[dict[str, Any]] = field(default_factory=list)
- def _env_value(*names: str, default: str = "") -> str:
- for name in names:
- value = os.getenv(name)
- if value is not None and str(value).strip():
- return str(value).strip()
- return default
- def _resolve_mysql_config() -> dict[str, Any]:
- dsn = _env_value("DEMAND_CONTENT_MYSQL_DSN")
- if dsn:
- parsed = urlparse(dsn)
- return {
- "host": parsed.hostname or "127.0.0.1",
- "port": int(parsed.port or 3306),
- "user": unquote(parsed.username or ""),
- "password": unquote(parsed.password or ""),
- "database": unquote(parsed.path.lstrip("/") or ""),
- }
- return {
- "host": _env_value("DEMAND_CONTENT_MYSQL_HOST", "CONTENT_SUPPLY_DB_HOST", default="127.0.0.1"),
- "port": int(_env_value("DEMAND_CONTENT_MYSQL_PORT", "CONTENT_SUPPLY_DB_PORT", default="3306")),
- "user": _env_value("DEMAND_CONTENT_MYSQL_USER", "CONTENT_SUPPLY_DB_USER", default="content_rw"),
- "password": _env_value("DEMAND_CONTENT_MYSQL_PASSWORD", "CONTENT_SUPPLY_DB_PASSWORD"),
- "database": _env_value(
- "DEMAND_CONTENT_MYSQL_DB",
- "CONTENT_SUPPLY_DB_NAME",
- default="content-deconstruction-supply",
- ),
- }
- def _connect():
- cfg = _resolve_mysql_config()
- missing = [key for key in ("host", "user", "database") if not cfg.get(key)]
- if missing:
- raise RuntimeError(f"MySQL demand_content config missing: {missing}")
- return pymysql.connect(
- host=cfg["host"],
- port=int(cfg["port"]),
- user=cfg["user"],
- password=cfg["password"],
- database=cfg["database"],
- charset="utf8mb4",
- cursorclass=pymysql.cursors.DictCursor,
- autocommit=False,
- )
- def _as_ext_data(value: Any) -> dict[str, Any]:
- if isinstance(value, dict):
- return copy.deepcopy(value)
- if isinstance(value, str) and value.strip():
- parsed = json.loads(value)
- if isinstance(parsed, dict):
- return parsed
- return {}
- def _dedupe_key(row: dict[str, Any], ext_data: dict[str, Any]) -> tuple[str, str, str]:
- evidence_pack = ext_data.get("evidence_pack") or {}
- source_post_id = str(evidence_pack.get("source_post_id") or "")
- return (
- str(row.get("name") or ""),
- source_post_id,
- "",
- )
- def _row_exists_for_run(cursor, *, run_label: str, row: dict[str, Any], ext_data: dict[str, Any]) -> bool:
- if not run_label:
- return False
- name, source_post_id, _ = _dedupe_key(row, ext_data)
- cursor.execute(
- """
- SELECT id
- FROM demand_content
- WHERE JSON_UNQUOTE(JSON_EXTRACT(ext_data, '$.run_label')) = %s
- AND name = %s
- AND JSON_UNQUOTE(JSON_EXTRACT(ext_data, '$.evidence_pack.source_post_id')) = %s
- LIMIT 1
- """,
- (run_label, name, source_post_id),
- )
- return cursor.fetchone() is not None
- def _validate_row(row: dict[str, Any], ext_data: dict[str, Any]) -> None:
- evidence_pack = ext_data.get("evidence_pack")
- if not isinstance(evidence_pack, dict):
- raise ValueError("missing ext_data.evidence_pack")
- expected = {
- "pattern_source_system": "pg_pattern_v2",
- "case_id_type": "post_id",
- "source_certainty": "db_validated",
- "validation_status": "passed",
- }
- for field, value in expected.items():
- if evidence_pack.get(field) != value:
- raise ValueError(f"invalid evidence_pack.{field}: {evidence_pack.get(field)!r}")
- required = [
- "source_post_id",
- "pattern_execution_id",
- "source_kind",
- "evidence_sources",
- "category_bindings",
- "element_bindings",
- "support",
- "absolute_support",
- "matched_post_ids",
- "video_ids",
- "case_ids",
- "seed_terms",
- "trace_id",
- "query_seed_points",
- "demand_scope",
- ]
- for field in required:
- value = evidence_pack.get(field)
- if field == "query_seed_points":
- if value is None or not isinstance(value, list):
- raise ValueError("missing evidence_pack.query_seed_points")
- continue
- if value is None or value == "" or value == []:
- raise ValueError(f"missing evidence_pack.{field}")
- if evidence_pack.get("source_kind") not in {
- "high_weight_element",
- "high_weight_category",
- "element_co_occurrence",
- "category_co_occurrence",
- "pattern_itemset",
- "multi_source",
- }:
- raise ValueError(f"invalid evidence_pack.source_kind: {evidence_pack.get('source_kind')!r}")
- if not isinstance(evidence_pack.get("evidence_sources"), list) or not evidence_pack["evidence_sources"]:
- raise ValueError("missing evidence_pack.evidence_sources")
- if not isinstance(evidence_pack.get("itemset_ids", []), list):
- raise ValueError("evidence_pack.itemset_ids must be an array")
- if not isinstance(evidence_pack.get("itemset_items", []), list):
- raise ValueError("evidence_pack.itemset_items must be an array")
- matched_post_ids = [str(post_id) for post_id in evidence_pack.get("matched_post_ids") or []]
- if str(evidence_pack["source_post_id"]) not in set(matched_post_ids):
- raise ValueError("evidence_pack.source_post_id must be in matched_post_ids")
- if [str(post_id) for post_id in evidence_pack.get("video_ids") or []] != matched_post_ids:
- raise ValueError("evidence_pack.video_ids must equal matched_post_ids")
- if [str(post_id) for post_id in evidence_pack.get("case_ids") or []] != matched_post_ids:
- raise ValueError("evidence_pack.case_ids must equal matched_post_ids")
- demand_scope = evidence_pack.get("demand_scope") or {}
- if not isinstance(demand_scope, dict):
- raise ValueError("evidence_pack.demand_scope must be object")
- if demand_scope.get("merge_leve2") and row.get("merge_leve2") != demand_scope.get("merge_leve2"):
- raise ValueError("evidence_pack.demand_scope.merge_leve2 must equal row.merge_leve2")
- if normalize_platform(demand_scope.get("platform")) != DEFAULT_DEMAND_PLATFORM:
- raise ValueError("evidence_pack.demand_scope.platform must be piaoquan")
- scoped_count = evidence_pack.get("scoped_post_count")
- filtered_support = evidence_pack.get("filtered_absolute_support")
- if scoped_count is not None:
- scoped_count = int(scoped_count)
- if scoped_count <= 0:
- raise ValueError("evidence_pack.scoped_post_count must be positive")
- if len(matched_post_ids) != scoped_count:
- raise ValueError("evidence_pack.matched_post_ids must cover scoped_post_count")
- if scoped_count > int(evidence_pack["absolute_support"]):
- raise ValueError("evidence_pack.scoped_post_count must not exceed absolute_support")
- elif len(matched_post_ids) < int(evidence_pack["absolute_support"]):
- raise ValueError("evidence_pack.matched_post_ids must cover absolute_support")
- if filtered_support is not None and scoped_count is not None and int(filtered_support) != scoped_count:
- raise ValueError("evidence_pack.filtered_absolute_support must equal scoped_post_count")
- expected_rank = 1
- for point in evidence_pack.get("query_seed_points") or []:
- if not isinstance(point, dict):
- raise ValueError("evidence_pack.query_seed_points items must be objects")
- if point.get("point_type") not in {"灵感点", "目的点"}:
- raise ValueError("evidence_pack.query_seed_points point_type must be 灵感点 or 目的点")
- if int(point.get("coverage_post_count") or 0) <= 0:
- raise ValueError("evidence_pack.query_seed_points coverage_post_count must be positive")
- if int(point.get("rank") or 0) != expected_rank:
- raise ValueError("evidence_pack.query_seed_points rank must be continuous")
- expected_rank += 1
- if not row.get("merge_leve2") or not row.get("name") or not row.get("dt"):
- raise ValueError("missing demand_content merge_leve2/name/dt")
- def write_demand_content_rows(
- rows: list[dict[str, Any]],
- *,
- run_label: str = "",
- ) -> MySQLDemandContentWriteResult:
- result = MySQLDemandContentWriteResult()
- if not rows:
- return result
- conn = _connect()
- try:
- with conn.cursor() as cursor:
- for row in rows:
- ext_data = _as_ext_data(row.get("ext_data"))
- if run_label:
- ext_data["run_label"] = run_label
- _validate_row(row, ext_data)
- if _row_exists_for_run(cursor, run_label=run_label, row=row, ext_data=ext_data):
- result.skipped_count += 1
- result.skipped.append({"name": row.get("name"), "reason": "duplicate_run_label_key"})
- continue
- cursor.execute(
- """
- INSERT INTO demand_content
- (merge_leve2, name, reason, suggestion, ext_data, score, dt)
- VALUES
- (%s, %s, %s, %s, %s, %s, %s)
- """,
- (
- row.get("merge_leve2"),
- row.get("name"),
- row.get("reason"),
- row.get("suggestion"),
- json.dumps(ext_data, ensure_ascii=False),
- row.get("score"),
- row.get("dt"),
- ),
- )
- demand_content_id = int(cursor.lastrowid)
- ext_data.setdefault("evidence_pack", {})["demand_content_id"] = demand_content_id
- cursor.execute(
- "UPDATE demand_content SET ext_data=%s WHERE id=%s",
- (json.dumps(ext_data, ensure_ascii=False), demand_content_id),
- )
- result.inserted_count += 1
- result.inserted_ids.append(demand_content_id)
- conn.commit()
- except Exception:
- conn.rollback()
- raise
- finally:
- conn.close()
- return result
|