validate_v2_demand_content.py 8.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197
  1. """Validate DemandAgent V2 demand_content JSON contracts."""
  2. from __future__ import annotations
  3. import argparse
  4. import json
  5. import sys
  6. from pathlib import Path
  7. from typing import Any
  8. ALLOWED_POINT_TYPES = {"灵感点", "目的点"}
  9. def _load_rows(path: Path) -> list[dict[str, Any]]:
  10. loaded = json.loads(path.read_text(encoding="utf-8"))
  11. if isinstance(loaded, list):
  12. return [row for row in loaded if isinstance(row, dict)]
  13. if isinstance(loaded, dict) and isinstance(loaded.get("rows"), list):
  14. return [row for row in loaded["rows"] if isinstance(row, dict)]
  15. raise ValueError(f"{path} is not a demand_content row array")
  16. def _as_ext_data(row: dict[str, Any]) -> dict[str, Any]:
  17. value = row.get("ext_data")
  18. if isinstance(value, dict):
  19. return value
  20. if isinstance(value, str) and value.strip():
  21. parsed = json.loads(value)
  22. if isinstance(parsed, dict):
  23. return parsed
  24. return {}
  25. def _derive_seed_candidates(itemset_items: list[dict[str, Any]]) -> set[str]:
  26. result: set[str] = set()
  27. for item in itemset_items:
  28. for field in ("element_name", "category_name", "category_path", "category_full_path"):
  29. value = str(item.get(field) or "").strip()
  30. if not value:
  31. continue
  32. result.add(value.replace(" ", ""))
  33. for part in value.replace(">", "/").split("/"):
  34. part = part.strip()
  35. if part:
  36. result.add(part.replace(" ", ""))
  37. return result
  38. def validate_row(row: dict[str, Any], *, require_scope: bool = True) -> list[str]:
  39. errors: list[str] = []
  40. for field in ("name", "merge_leve2", "dt", "ext_data"):
  41. if not row.get(field):
  42. errors.append(f"missing top-level {field}")
  43. ext_data = _as_ext_data(row)
  44. evidence_pack = ext_data.get("evidence_pack")
  45. if not isinstance(evidence_pack, dict):
  46. return [*errors, "missing ext_data.evidence_pack"]
  47. expected = {
  48. "pattern_source_system": "pg_pattern_v2",
  49. "case_id_type": "post_id",
  50. "source_certainty": "db_validated",
  51. "validation_status": "passed",
  52. }
  53. for field, expected_value in expected.items():
  54. if evidence_pack.get(field) != expected_value:
  55. errors.append(f"invalid evidence_pack.{field}: {evidence_pack.get(field)!r}")
  56. required_non_empty = [
  57. "source_post_id",
  58. "pattern_execution_id",
  59. "source_kind",
  60. "evidence_sources",
  61. "category_bindings",
  62. "element_bindings",
  63. "support",
  64. "absolute_support",
  65. "matched_post_ids",
  66. "video_ids",
  67. "case_ids",
  68. "seed_terms",
  69. "trace_id",
  70. ]
  71. for field in required_non_empty:
  72. value = evidence_pack.get(field)
  73. if value is None or value == "" or value == []:
  74. errors.append(f"missing evidence_pack.{field}")
  75. if evidence_pack.get("source_kind") not in {
  76. "high_weight_element",
  77. "high_weight_category",
  78. "element_co_occurrence",
  79. "category_co_occurrence",
  80. "pattern_itemset",
  81. "multi_source",
  82. }:
  83. errors.append(f"invalid source_kind: {evidence_pack.get('source_kind')!r}")
  84. if not isinstance(evidence_pack.get("evidence_sources"), list) or not evidence_pack["evidence_sources"]:
  85. errors.append("missing evidence_sources")
  86. if not isinstance(evidence_pack.get("itemset_ids", []), list):
  87. errors.append("itemset_ids must be an array")
  88. if not isinstance(evidence_pack.get("itemset_items", []), list):
  89. errors.append("itemset_items must be an array")
  90. matched_post_ids = [str(post_id) for post_id in evidence_pack.get("matched_post_ids") or []]
  91. if str(evidence_pack.get("source_post_id") or "") not in set(matched_post_ids):
  92. errors.append("source_post_id not in matched_post_ids")
  93. if [str(post_id) for post_id in evidence_pack.get("video_ids") or []] != matched_post_ids:
  94. errors.append("video_ids must equal matched_post_ids")
  95. if [str(post_id) for post_id in evidence_pack.get("case_ids") or []] != matched_post_ids:
  96. errors.append("case_ids must equal matched_post_ids")
  97. query_seed_points = evidence_pack.get("query_seed_points")
  98. if not isinstance(query_seed_points, list):
  99. errors.append("query_seed_points must exist as array")
  100. else:
  101. previous_coverage: int | None = None
  102. for index, point in enumerate(query_seed_points, start=1):
  103. if not isinstance(point, dict):
  104. errors.append(f"query_seed_points[{index}] must be object")
  105. continue
  106. if point.get("point_type") not in ALLOWED_POINT_TYPES:
  107. errors.append(f"query_seed_points[{index}].point_type invalid")
  108. coverage = int(point.get("coverage_post_count") or 0)
  109. if coverage <= 0:
  110. errors.append(f"query_seed_points[{index}].coverage_post_count invalid")
  111. if int(point.get("rank") or 0) != index:
  112. errors.append(f"query_seed_points[{index}].rank invalid")
  113. if previous_coverage is not None and coverage > previous_coverage:
  114. errors.append("query_seed_points not sorted by coverage_post_count desc")
  115. previous_coverage = coverage
  116. itemset_items = evidence_pack.get("itemset_items") or []
  117. seed_candidates = _derive_seed_candidates(itemset_items)
  118. for source in evidence_pack.get("evidence_sources") or []:
  119. if isinstance(source, dict):
  120. for value in source.get("source_terms") or []:
  121. seed_candidates.add(str(value).replace(" ", ""))
  122. for term in evidence_pack.get("seed_terms") or []:
  123. if str(term).replace(" ", "") not in seed_candidates:
  124. errors.append(f"seed_term not derived from DB evidence sources: {term}")
  125. demand_scope = evidence_pack.get("demand_scope")
  126. if require_scope and not isinstance(demand_scope, dict):
  127. errors.append("missing demand_scope")
  128. elif isinstance(demand_scope, dict):
  129. if demand_scope.get("merge_leve2") and demand_scope.get("merge_leve2") != row.get("merge_leve2"):
  130. errors.append("demand_scope.merge_leve2 mismatch")
  131. if demand_scope.get("scope_source") == "odps_gap":
  132. for field in ("gap_dt", "requested_count", "lack_count"):
  133. if demand_scope.get(field) in (None, ""):
  134. errors.append(f"missing demand_scope.{field}")
  135. scoped_count = evidence_pack.get("scoped_post_count")
  136. if require_scope and scoped_count is None:
  137. errors.append("missing scoped_post_count")
  138. if scoped_count is not None:
  139. scoped_count = int(scoped_count)
  140. if scoped_count <= 0:
  141. errors.append("scoped_post_count must be positive")
  142. if len(matched_post_ids) != scoped_count:
  143. errors.append("matched_post_ids length must equal scoped_post_count")
  144. if scoped_count > int(evidence_pack.get("absolute_support") or 0):
  145. errors.append("scoped_post_count must not exceed absolute_support")
  146. if evidence_pack.get("filtered_absolute_support") is not None:
  147. if int(evidence_pack["filtered_absolute_support"]) != scoped_count:
  148. errors.append("filtered_absolute_support must equal scoped_post_count")
  149. return errors
  150. def validate_file(path: Path, *, require_scope: bool = True) -> list[str]:
  151. errors: list[str] = []
  152. rows = _load_rows(path)
  153. for index, row in enumerate(rows, start=1):
  154. for error in validate_row(row, require_scope=require_scope):
  155. errors.append(f"{path}:{index}: {error}")
  156. return errors
  157. def main() -> None:
  158. parser = argparse.ArgumentParser(description="Validate DemandAgent V2 demand_content JSON.")
  159. parser.add_argument("paths", nargs="+", type=Path)
  160. parser.add_argument("--allow-no-scope", action="store_true")
  161. args = parser.parse_args()
  162. errors: list[str] = []
  163. for path in args.paths:
  164. errors.extend(validate_file(path, require_scope=not args.allow_no_scope))
  165. if errors:
  166. print("\n".join(errors), file=sys.stderr)
  167. raise SystemExit(1)
  168. print(json.dumps({"success": True, "files": [str(path) for path in args.paths]}, ensure_ascii=False))
  169. if __name__ == "__main__":
  170. main()