validate_v2_demand_content.py 7.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182
  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. "source_kind": "pattern_itemset",
  50. "case_id_type": "post_id",
  51. "source_certainty": "db_validated",
  52. "validation_status": "passed",
  53. }
  54. for field, expected_value in expected.items():
  55. if evidence_pack.get(field) != expected_value:
  56. errors.append(f"invalid evidence_pack.{field}: {evidence_pack.get(field)!r}")
  57. required_non_empty = [
  58. "source_post_id",
  59. "pattern_execution_id",
  60. "mining_config_id",
  61. "itemset_ids",
  62. "itemset_items",
  63. "category_bindings",
  64. "element_bindings",
  65. "support",
  66. "absolute_support",
  67. "matched_post_ids",
  68. "video_ids",
  69. "case_ids",
  70. "seed_terms",
  71. "trace_id",
  72. ]
  73. for field in required_non_empty:
  74. value = evidence_pack.get(field)
  75. if value is None or value == "" or value == []:
  76. errors.append(f"missing evidence_pack.{field}")
  77. if not isinstance(evidence_pack.get("itemset_ids"), list) or len(evidence_pack["itemset_ids"]) != 1:
  78. errors.append("itemset_ids must contain exactly one itemset_id")
  79. matched_post_ids = [str(post_id) for post_id in evidence_pack.get("matched_post_ids") or []]
  80. if str(evidence_pack.get("source_post_id") or "") not in set(matched_post_ids):
  81. errors.append("source_post_id not in matched_post_ids")
  82. if [str(post_id) for post_id in evidence_pack.get("video_ids") or []] != matched_post_ids:
  83. errors.append("video_ids must equal matched_post_ids")
  84. if [str(post_id) for post_id in evidence_pack.get("case_ids") or []] != matched_post_ids:
  85. errors.append("case_ids must equal matched_post_ids")
  86. query_seed_points = evidence_pack.get("query_seed_points")
  87. if not isinstance(query_seed_points, list):
  88. errors.append("query_seed_points must exist as array")
  89. else:
  90. previous_coverage: int | None = None
  91. for index, point in enumerate(query_seed_points, start=1):
  92. if not isinstance(point, dict):
  93. errors.append(f"query_seed_points[{index}] must be object")
  94. continue
  95. if point.get("point_type") not in ALLOWED_POINT_TYPES:
  96. errors.append(f"query_seed_points[{index}].point_type invalid")
  97. coverage = int(point.get("coverage_post_count") or 0)
  98. if coverage <= 0:
  99. errors.append(f"query_seed_points[{index}].coverage_post_count invalid")
  100. if int(point.get("rank") or 0) != index:
  101. errors.append(f"query_seed_points[{index}].rank invalid")
  102. if previous_coverage is not None and coverage > previous_coverage:
  103. errors.append("query_seed_points not sorted by coverage_post_count desc")
  104. previous_coverage = coverage
  105. itemset_items = evidence_pack.get("itemset_items") or []
  106. seed_candidates = _derive_seed_candidates(itemset_items)
  107. for term in evidence_pack.get("seed_terms") or []:
  108. if str(term).replace(" ", "") not in seed_candidates:
  109. errors.append(f"seed_term not derived from itemset_items: {term}")
  110. demand_scope = evidence_pack.get("demand_scope")
  111. if require_scope and not isinstance(demand_scope, dict):
  112. errors.append("missing demand_scope")
  113. elif isinstance(demand_scope, dict):
  114. if demand_scope.get("merge_leve2") and demand_scope.get("merge_leve2") != row.get("merge_leve2"):
  115. errors.append("demand_scope.merge_leve2 mismatch")
  116. if demand_scope.get("scope_source") == "odps_gap":
  117. for field in ("gap_dt", "requested_count", "lack_count"):
  118. if demand_scope.get(field) in (None, ""):
  119. errors.append(f"missing demand_scope.{field}")
  120. scoped_count = evidence_pack.get("scoped_post_count")
  121. if require_scope and scoped_count is None:
  122. errors.append("missing scoped_post_count")
  123. if scoped_count is not None:
  124. scoped_count = int(scoped_count)
  125. if scoped_count <= 0:
  126. errors.append("scoped_post_count must be positive")
  127. if len(matched_post_ids) != scoped_count:
  128. errors.append("matched_post_ids length must equal scoped_post_count")
  129. if scoped_count > int(evidence_pack.get("absolute_support") or 0):
  130. errors.append("scoped_post_count must not exceed absolute_support")
  131. if evidence_pack.get("filtered_absolute_support") is not None:
  132. if int(evidence_pack["filtered_absolute_support"]) != scoped_count:
  133. errors.append("filtered_absolute_support must equal scoped_post_count")
  134. return errors
  135. def validate_file(path: Path, *, require_scope: bool = True) -> list[str]:
  136. errors: list[str] = []
  137. rows = _load_rows(path)
  138. for index, row in enumerate(rows, start=1):
  139. for error in validate_row(row, require_scope=require_scope):
  140. errors.append(f"{path}:{index}: {error}")
  141. return errors
  142. def main() -> None:
  143. parser = argparse.ArgumentParser(description="Validate DemandAgent V2 demand_content JSON.")
  144. parser.add_argument("paths", nargs="+", type=Path)
  145. parser.add_argument("--allow-no-scope", action="store_true")
  146. args = parser.parse_args()
  147. errors: list[str] = []
  148. for path in args.paths:
  149. errors.extend(validate_file(path, require_scope=not args.allow_no_scope))
  150. if errors:
  151. print("\n".join(errors), file=sys.stderr)
  152. raise SystemExit(1)
  153. print(json.dumps({"success": True, "files": [str(path) for path in args.paths]}, ensure_ascii=False))
  154. if __name__ == "__main__":
  155. main()