Forráskód Böngészése

Normalize demand platform seed points

SamLee 3 hete
szülő
commit
cb3f696a58

+ 4 - 0
examples/demand/db_manager.py

@@ -8,6 +8,8 @@ Pattern evidence reads now go through `pg_pattern_repository` and are read-only.
 from __future__ import annotations
 from __future__ import annotations
 
 
 from examples.demand.pg_pattern_repository import (
 from examples.demand.pg_pattern_repository import (
+    DEFAULT_DEMAND_PLATFORM,
+    normalize_platform,
     query_case_ids_by_post_ids,
     query_case_ids_by_post_ids,
     query_category_level,
     query_category_level,
     query_element_bindings_for_items,
     query_element_bindings_for_items,
@@ -50,7 +52,9 @@ def exist_cluster_tree(merge_level2: str) -> bool:
 
 
 __all__ = [
 __all__ = [
     "DatabaseManager",
     "DatabaseManager",
+    "DEFAULT_DEMAND_PLATFORM",
     "exist_cluster_tree",
     "exist_cluster_tree",
+    "normalize_platform",
     "query_case_ids_by_post_ids",
     "query_case_ids_by_post_ids",
     "query_category_level",
     "query_category_level",
     "query_element_bindings_for_items",
     "query_element_bindings_for_items",

+ 13 - 8
examples/demand/evidence_pack_builder.py

@@ -18,6 +18,7 @@ from examples.demand.db_manager import (
     query_seed_points_for_sources,
     query_seed_points_for_sources,
     query_source_elements,
     query_source_elements,
 )
 )
+from examples.demand.pg_pattern_repository import DEFAULT_DEMAND_PLATFORM, normalize_platform
 
 
 
 
 SOURCE_KIND_PATTERN_ITEMSET = "pattern_itemset"
 SOURCE_KIND_PATTERN_ITEMSET = "pattern_itemset"
@@ -178,11 +179,12 @@ def _build_evidence_pack(
 
 
     case_rows = query_case_ids_by_post_ids(matched_post_ids)
     case_rows = query_case_ids_by_post_ids(matched_post_ids)
     decode_case_ids = _unique_strings(row.get("case_id") for row in case_rows)
     decode_case_ids = _unique_strings(row.get("case_id") for row in case_rows)
+    query_seed_points_top_k = min(max(_env_int("DEMAND_QUERY_SEED_POINTS_TOP_K", 30), 0), 30)
     query_seed_points = query_seed_points_for_sources(
     query_seed_points = query_seed_points_for_sources(
         execution_id=int(execution_id),
         execution_id=int(execution_id),
         matched_post_ids=matched_post_ids,
         matched_post_ids=matched_post_ids,
         category_ids=_unique_ints(binding.get("category_id") for binding in category_bindings),
         category_ids=_unique_ints(binding.get("category_id") for binding in category_bindings),
-        top_k=_env_int("DEMAND_QUERY_SEED_POINTS_TOP_K", 30),
+        top_k=query_seed_points_top_k,
     )
     )
     mining_config_ids = _unique_ints(
     mining_config_ids = _unique_ints(
         config_id
         config_id
@@ -576,13 +578,16 @@ def _resolve_demand_scope(
                 evidence_refs.get("merge_level2"),
                 evidence_refs.get("merge_level2"),
             )
             )
         ),
         ),
-        "platform": _clean_str(
-            _first_present(
-                platform,
-                raw_scope.get("platform"),
-                raw_scope.get("platform_type"),
-                evidence_refs.get("platform"),
-                evidence_refs.get("platform_type"),
+        "platform": normalize_platform(
+            _clean_str(
+                _first_present(
+                    platform,
+                    raw_scope.get("platform"),
+                    raw_scope.get("platform_type"),
+                    evidence_refs.get("platform"),
+                    evidence_refs.get("platform_type"),
+                    DEFAULT_DEMAND_PLATFORM,
+                )
             )
             )
         ),
         ),
         "pattern_execution_id": int(execution_id),
         "pattern_execution_id": int(execution_id),

+ 4 - 0
examples/demand/mysql_demand_content_sink.py

@@ -16,6 +16,8 @@ from urllib.parse import unquote, urlparse
 
 
 import pymysql
 import pymysql
 
 
+from examples.demand.pg_pattern_repository import DEFAULT_DEMAND_PLATFORM, normalize_platform
+
 
 
 @dataclass
 @dataclass
 class MySQLDemandContentWriteResult:
 class MySQLDemandContentWriteResult:
@@ -180,6 +182,8 @@ def _validate_row(row: dict[str, Any], ext_data: dict[str, Any]) -> None:
         raise ValueError("evidence_pack.demand_scope must be object")
         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"):
     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")
         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")
     scoped_count = evidence_pack.get("scoped_post_count")
     filtered_support = evidence_pack.get("filtered_absolute_support")
     filtered_support = evidence_pack.get("filtered_absolute_support")

+ 18 - 5
examples/demand/pg_pattern_repository.py

@@ -166,16 +166,29 @@ def _to_str_filter_list(value: Any) -> list[str]:
     return result
     return result
 
 
 
 
+PLATFORM_ALIASES = {
+    "piaoquan": ["piaoquan", "票圈"],
+    "票圈": ["票圈", "piaoquan"],
+}
+DEFAULT_DEMAND_PLATFORM = "piaoquan"
+PLATFORM_CANONICAL = {
+    "票圈": "piaoquan",
+}
+
+
+def normalize_platform(value: Any) -> str:
+    """Return the canonical machine value for a platform label."""
+    text = str(value).strip() if value is not None else ""
+    return PLATFORM_CANONICAL.get(text, text)
+
+
 def _expand_platform_filters(values: list[str]) -> list[str]:
 def _expand_platform_filters(values: list[str]) -> list[str]:
     """Expand known platform aliases used by Hive and PG post rows."""
     """Expand known platform aliases used by Hive and PG post rows."""
-    aliases = {
-        "piaoquan": ["piaoquan", "票圈"],
-        "票圈": ["票圈", "piaoquan"],
-    }
     result: list[str] = []
     result: list[str] = []
     seen: set[str] = set()
     seen: set[str] = set()
     for value in values:
     for value in values:
-        for candidate in aliases.get(value, [value]):
+        canonical = normalize_platform(value)
+        for candidate in PLATFORM_ALIASES.get(canonical, [canonical]):
             if candidate and candidate not in seen:
             if candidate and candidate not in seen:
                 seen.add(candidate)
                 seen.add(candidate)
                 result.append(candidate)
                 result.append(candidate)

+ 11 - 2
examples/demand/run.py

@@ -22,6 +22,8 @@ os.environ.setdefault("AGENT_DISABLE_SIDE_BRANCHES", "1")
 
 
 from dotenv import load_dotenv
 from dotenv import load_dotenv
 from examples.demand.db_manager import (
 from examples.demand.db_manager import (
+    DEFAULT_DEMAND_PLATFORM,
+    normalize_platform,
     query_category_level,
     query_category_level,
     query_latest_success_execution_id,
     query_latest_success_execution_id,
     query_video_ids_by_names,
     query_video_ids_by_names,
@@ -64,6 +66,12 @@ MYSQL_ENTRYPOINT_ENV = "DEMAND_MYSQL_ENTRYPOINT"
 MYSQL_ENTRYPOINT_NAMES = {"run_existing_execution_mysql", "run_hive_gap_mysql"}
 MYSQL_ENTRYPOINT_NAMES = {"run_existing_execution_mysql", "run_hive_gap_mysql"}
 
 
 
 
+def _resolve_run_platform(platform_type: Optional[str], demand_scope: Optional[dict[str, Any]]) -> str:
+    scope = demand_scope if isinstance(demand_scope, dict) else {}
+    raw_platform = platform_type or scope.get("platform") or scope.get("platform_type") or DEFAULT_DEMAND_PLATFORM
+    return normalize_platform(raw_platform)
+
+
 def _is_local_json_mode() -> bool:
 def _is_local_json_mode() -> bool:
     return os.getenv("DEMAND_OUTPUT_MODE", "").strip().lower() == LOCAL_JSON_MODE
     return os.getenv("DEMAND_OUTPUT_MODE", "").strip().lower() == LOCAL_JSON_MODE
 
 
@@ -774,6 +782,7 @@ async def run_once(
     if _is_local_json_mode() or _is_mysql_demand_content_mode():
     if _is_local_json_mode() or _is_mysql_demand_content_mode():
         _ensure_pg_weight_score_files(int(execution_id))
         _ensure_pg_weight_score_files(int(execution_id))
 
 
+    platform_type = _resolve_run_platform(platform_type, demand_scope)
     TopicBuildAgentContext.set_execution_id(execution_id)
     TopicBuildAgentContext.set_execution_id(execution_id)
     TopicBuildAgentContext.set_metadata("result_base_dir", str(_get_result_base_dir()))
     TopicBuildAgentContext.set_metadata("result_base_dir", str(_get_result_base_dir()))
     TopicBuildAgentContext.set_metadata("merge_leve2", merge_level2)
     TopicBuildAgentContext.set_metadata("merge_leve2", merge_level2)
@@ -781,8 +790,7 @@ async def run_once(
     resolved_demand_scope = dict(demand_scope or {})
     resolved_demand_scope = dict(demand_scope or {})
     resolved_demand_scope.setdefault("scope_source", "manual_cli")
     resolved_demand_scope.setdefault("scope_source", "manual_cli")
     resolved_demand_scope.setdefault("merge_leve2", merge_level2)
     resolved_demand_scope.setdefault("merge_leve2", merge_level2)
-    if platform_type:
-        resolved_demand_scope.setdefault("platform", platform_type)
+    resolved_demand_scope["platform"] = _resolve_run_platform(platform_type, resolved_demand_scope)
     resolved_demand_scope.setdefault("pattern_execution_id", int(execution_id))
     resolved_demand_scope.setdefault("pattern_execution_id", int(execution_id))
     TopicBuildAgentContext.set_metadata("demand_scope", resolved_demand_scope)
     TopicBuildAgentContext.set_metadata("demand_scope", resolved_demand_scope)
 
 
@@ -974,6 +982,7 @@ async def main(
         if str(cluster_name).strip() == "全局树":
         if str(cluster_name).strip() == "全局树":
             raise ValueError("mysql_demand_content 模式只写普通 demand_content,不支持全局树 Hive 输出")
             raise ValueError("mysql_demand_content 模式只写普通 demand_content,不支持全局树 Hive 输出")
 
 
+    platform_type = _resolve_run_platform(platform_type, demand_scope)
     if execution_id is None:
     if execution_id is None:
         execution_id = get_execution_id_by_merge_level2(cluster_name)
         execution_id = get_execution_id_by_merge_level2(cluster_name)
     if not execution_id:
     if not execution_id:

+ 10 - 0
examples/demand/tests/test_mysql_demand_content_sink.py

@@ -200,6 +200,16 @@ class MySQLDemandContentSinkTest(unittest.TestCase):
             with self.assertRaisesRegex(ValueError, "demand_scope"):
             with self.assertRaisesRegex(ValueError, "demand_scope"):
                 sink.write_demand_content_rows([row], run_label="batch01")
                 sink.write_demand_content_rows([row], run_label="batch01")
 
 
+    def test_rejects_non_piaoquan_demand_scope_platform(self):
+        from examples.demand import mysql_demand_content_sink as sink
+
+        row = _valid_row()
+        row["ext_data"]["evidence_pack"]["demand_scope"]["platform"] = "xiaohongshu"
+
+        with patch.object(sink, "_connect", return_value=_FakeConnection()):
+            with self.assertRaisesRegex(ValueError, "platform"):
+                sink.write_demand_content_rows([row], run_label="batch01")
+
     def test_accepts_multiple_itemset_ids(self):
     def test_accepts_multiple_itemset_ids(self):
         from examples.demand import mysql_demand_content_sink as sink
         from examples.demand import mysql_demand_content_sink as sink
 
 

+ 48 - 0
examples/demand/tests/test_pg_evidence_builder.py

@@ -115,6 +115,54 @@ class PgEvidenceBuilderTest(unittest.TestCase):
         self.assertIn("query_seed_points", evidence_pack)
         self.assertIn("query_seed_points", evidence_pack)
         self.assertIn("demand_scope", evidence_pack)
         self.assertIn("demand_scope", evidence_pack)
 
 
+    def test_demand_scope_platform_is_canonicalized(self):
+        self._patch_success()
+        result = build_evidence_pack(
+            581,
+            {
+                "element_names": ["综合性腐败"],
+                "evidence_refs": {
+                    "source_kind": "pattern_itemset",
+                    "source_tool": "get_itemset_detail",
+                    "itemset_ids": [1607313],
+                    "source_post_id": "p1",
+                },
+            },
+            trace_id="trace-test",
+            demand_task_id=1,
+            demand_content_id=1,
+            demand_scope={"platform": "票圈"},
+        )
+
+        self.assertTrue(result["success"])
+        self.assertEqual(result["evidence_pack"]["demand_scope"]["platform"], "piaoquan")
+
+    def test_defaults_to_piaoquan_scope_and_caps_query_seed_points(self):
+        with (
+            patch("examples.demand.evidence_pack_builder.query_execution_for_evidence", return_value=_execution()),
+            patch("examples.demand.evidence_pack_builder.query_itemset_evidence", return_value=[_itemset(matched_post_ids=["pq1", "pq2"])]) as itemset_evidence,
+            patch("examples.demand.evidence_pack_builder.query_itemset_items_with_categories", return_value=[_item()]),
+            patch("examples.demand.evidence_pack_builder.query_element_bindings_for_items", return_value=[_binding(matched_post_ids=["pq1"])]),
+            patch("examples.demand.evidence_pack_builder.query_case_ids_by_post_ids", return_value=[]),
+            patch("examples.demand.evidence_pack_builder.query_seed_points_for_sources", return_value=[]) as seed_points,
+            patch.dict("os.environ", {"DEMAND_QUERY_SEED_POINTS_TOP_K": "100"}),
+        ):
+            result = self._build(
+                {
+                    "source_kind": "pattern_itemset",
+                    "source_tool": "get_itemset_detail",
+                    "itemset_ids": [1607313],
+                    "source_post_id": "pq1",
+                }
+            )
+
+        self.assertTrue(result["success"])
+        self.assertEqual(result["evidence_pack"]["demand_scope"]["platform"], "piaoquan")
+        self.assertEqual(itemset_evidence.call_args.kwargs["platform"], "piaoquan")
+        seed_points.assert_called_once()
+        self.assertEqual(seed_points.call_args.kwargs["matched_post_ids"], ["pq1", "pq2"])
+        self.assertEqual(seed_points.call_args.kwargs["top_k"], 30)
+
     def test_non_topic_scope_is_rejected(self):
     def test_non_topic_scope_is_rejected(self):
         self._patch_success(itemsets=[_itemset(scope="topic_element")])
         self._patch_success(itemsets=[_itemset(scope="topic_element")])
         result = self._build(
         result = self._build(

+ 7 - 0
examples/demand/tests/test_pg_pattern_repository.py

@@ -2,7 +2,9 @@ import unittest
 from unittest.mock import patch
 from unittest.mock import patch
 
 
 from examples.demand.pg_pattern_repository import (
 from examples.demand.pg_pattern_repository import (
+    _expand_platform_filters,
     _is_element_type_dimension,
     _is_element_type_dimension,
+    normalize_platform,
     query_elements,
     query_elements,
     query_seed_points_for_itemsets,
     query_seed_points_for_itemsets,
 )
 )
@@ -17,6 +19,11 @@ class PgPatternRepositoryTest(unittest.TestCase):
         self.assertTrue(_is_element_type_dimension("形式"))
         self.assertTrue(_is_element_type_dimension("形式"))
         self.assertTrue(_is_element_type_dimension("意图"))
         self.assertTrue(_is_element_type_dimension("意图"))
 
 
+    def test_platform_normalization_keeps_piaoquan_filter_aliases(self):
+        self.assertEqual(normalize_platform("票圈"), "piaoquan")
+        self.assertEqual(normalize_platform("piaoquan"), "piaoquan")
+        self.assertEqual(_expand_platform_filters(["票圈"]), ["piaoquan", "票圈"])
+
     def test_query_seed_points_filters_and_ranks(self):
     def test_query_seed_points_filters_and_ranks(self):
         rows = [
         rows = [
             {
             {