"""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 @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", "source_kind": "pattern_itemset", "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", "mining_config_id", "itemset_ids", "itemset_items", "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 not isinstance(evidence_pack.get("itemset_ids"), list) or len(evidence_pack["itemset_ids"]) != 1: raise ValueError("evidence_pack.itemset_ids must contain exactly one itemset_id") 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") 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