pg_pattern_service.py 20 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528
  1. """PG Pattern V2 adapter for DemandAgent pattern tools.
  2. The public function names mirror the old MySQL pattern service, but every query
  3. reads PG `open_aigc.public` Pattern V2 facts and only exposes topic-scope
  4. itemsets as exact evidence candidates.
  5. """
  6. from __future__ import annotations
  7. from collections import defaultdict
  8. from typing import Any, Iterable
  9. from examples.demand import pg_pattern_repository as repo
  10. def _clean_ints(values: Iterable[Any] | None) -> list[int]:
  11. result: list[int] = []
  12. seen: set[int] = set()
  13. for raw in values or []:
  14. try:
  15. value = int(raw)
  16. except (TypeError, ValueError):
  17. continue
  18. if value not in seen:
  19. seen.add(value)
  20. result.append(value)
  21. return result
  22. def _last_path_part(path: Any) -> str:
  23. text = str(path or "").strip()
  24. if not text:
  25. return ""
  26. return text.replace(">", "/").split("/")[-1].strip()
  27. def _group_point_types(rows: list[dict[str, Any]]) -> list[str]:
  28. result: list[str] = []
  29. seen: set[str] = set()
  30. for row in rows:
  31. value = str(row.get("point_type") or "").strip()
  32. if value and value not in seen:
  33. seen.add(value)
  34. result.append(value)
  35. return result
  36. def _format_item(item: dict[str, Any]) -> dict[str, Any]:
  37. return {
  38. "id": item.get("itemset_item_id"),
  39. "itemset_id": item.get("itemset_id"),
  40. "layer": item.get("layer"),
  41. "point_type": item.get("point_type"),
  42. "dimension": item.get("dimension"),
  43. "category_id": item.get("category_id"),
  44. "category_path": item.get("category_path") or item.get("category_full_path"),
  45. "category_name": item.get("category_name") or _last_path_part(item.get("category_path")),
  46. "element_name": item.get("element_name"),
  47. "post_count": item.get("post_count"),
  48. }
  49. def get_category_tree_compact(execution_id: int, source_type: str | None = None) -> str:
  50. categories = repo.query_categories(execution_id, source_type=source_type)
  51. if not categories:
  52. return f"未找到 PG Pattern V2 分类树,execution_id={execution_id}, source_type={source_type}"
  53. lines = [
  54. f"PG Pattern V2 分类树快照 execution_id={execution_id}",
  55. "说明:本 DemandAgent MVP 只使用 scope=topic 的分类 Pattern;topic_element 不进入最终证据。",
  56. ]
  57. for cat in categories:
  58. level = int(cat.get("level") or 0)
  59. indent = " " * max(level - 1, 0)
  60. lines.append(
  61. f"{indent}- [{cat.get('id')}] {cat.get('name')} "
  62. f"(type={cat.get('source_type')}, level={level}, elements={cat.get('element_count') or 0}, "
  63. f"path={cat.get('path')})"
  64. )
  65. return "\n".join(lines)
  66. def search_top_itemsets(
  67. execution_id: int,
  68. top_n: int = 20,
  69. category_ids: list | None = None,
  70. dimension_mode: str | None = None,
  71. min_support: int | None = None,
  72. min_item_count: int | None = None,
  73. max_item_count: int | None = None,
  74. sort_by: str = "absolute_support",
  75. account_name=None,
  76. merge_leve2=None,
  77. platform=None,
  78. ) -> dict[str, Any]:
  79. del account_name
  80. rows = repo.query_topic_itemsets(
  81. execution_id=execution_id,
  82. category_ids=category_ids,
  83. dimension_mode=dimension_mode,
  84. min_support=min_support,
  85. min_item_count=min_item_count,
  86. max_item_count=max_item_count,
  87. sort_by=sort_by,
  88. limit=max(int(top_n or 20) * 10, int(top_n or 20)),
  89. merge_leve2=merge_leve2,
  90. platform=platform,
  91. )
  92. itemset_ids = [int(row["id"]) for row in rows]
  93. items_by_itemset: dict[int, list[dict[str, Any]]] = defaultdict(list)
  94. if itemset_ids:
  95. for item in repo.query_itemset_items_with_categories(execution_id, itemset_ids):
  96. items_by_itemset[int(item["itemset_id"])].append(_format_item(item))
  97. groups: dict[str, dict[str, Any]] = {}
  98. for row in rows:
  99. dimension_key = row.get("dimension_mode") or row.get("mining_config_scope") or "topic"
  100. target_depth = row.get("target_depth") or ""
  101. key = f"{dimension_key}/{target_depth}"
  102. group = groups.setdefault(
  103. key,
  104. {
  105. "dimension_mode": dimension_key,
  106. "target_depth": target_depth,
  107. "total": 0,
  108. "itemsets": [],
  109. },
  110. )
  111. if len(group["itemsets"]) >= int(top_n or 20):
  112. continue
  113. group["total"] += 1
  114. group["itemsets"].append(
  115. {
  116. "id": row.get("id"),
  117. "scope": row.get("scope"),
  118. "mining_config_id": row.get("mining_config_id"),
  119. "dimension_mode": dimension_key,
  120. "target_depth": target_depth,
  121. "combination_type": row.get("combination_type"),
  122. "item_count": row.get("item_count"),
  123. "support": row.get("support"),
  124. "absolute_support": row.get("absolute_support"),
  125. "filtered_absolute_support": row.get("filtered_absolute_support"),
  126. "scoped_post_count": row.get("scoped_post_count"),
  127. "scope_filter": row.get("scope_filter") or {},
  128. "dimensions": row.get("dimensions") or [],
  129. "items": items_by_itemset.get(int(row["id"]), []),
  130. }
  131. )
  132. return {
  133. "pattern_source_system": repo.PG_PATTERN_SOURCE_SYSTEM,
  134. "scope": repo.TOPIC_SCOPE,
  135. "scope_filter": rows[0].get("scope_filter") if rows else {},
  136. "total": len(rows),
  137. "showing": sum(len(group["itemsets"]) for group in groups.values()),
  138. "groups": groups,
  139. }
  140. def get_itemset_posts(
  141. itemset_ids: list[int],
  142. execution_id: int | None = None,
  143. *,
  144. merge_leve2=None,
  145. platform=None,
  146. ) -> list[dict[str, Any]]:
  147. clean_ids = _clean_ints(itemset_ids)
  148. if not clean_ids:
  149. return []
  150. if execution_id is None:
  151. raise ValueError("PG Pattern V2 get_itemset_posts requires execution_id")
  152. itemsets = repo.query_itemset_evidence(
  153. execution_id=execution_id,
  154. itemset_ids=clean_ids,
  155. merge_leve2=merge_leve2,
  156. platform=platform,
  157. )
  158. itemset_items = repo.query_itemset_items_with_categories(execution_id=execution_id, itemset_ids=clean_ids)
  159. items_by_itemset: dict[int, list[dict[str, Any]]] = defaultdict(list)
  160. for item in itemset_items:
  161. items_by_itemset[int(item["itemset_id"])].append(_format_item(item))
  162. details: list[dict[str, Any]] = []
  163. for itemset in itemsets:
  164. details.append(
  165. {
  166. "id": itemset.get("id"),
  167. "execution_id": itemset.get("execution_id"),
  168. "scope": itemset.get("scope"),
  169. "mining_config_id": itemset.get("mining_config_id"),
  170. "dimension_mode": itemset.get("mining_config_dimension_mode"),
  171. "target_depth": itemset.get("mining_config_target_depth"),
  172. "combination_type": itemset.get("combination_type"),
  173. "item_count": itemset.get("item_count"),
  174. "support": itemset.get("support"),
  175. "absolute_support": itemset.get("absolute_support"),
  176. "filtered_absolute_support": itemset.get("filtered_absolute_support"),
  177. "scoped_post_count": itemset.get("scoped_post_count"),
  178. "scope_filter": itemset.get("scope_filter") or {},
  179. "dimensions": itemset.get("dimensions") or [],
  180. "items": items_by_itemset.get(int(itemset["id"]), []),
  181. "post_ids": itemset.get("matched_post_ids") or [],
  182. "matched_post_ids": itemset.get("matched_post_ids") or [],
  183. "global_matched_post_ids": itemset.get("global_matched_post_ids") or [],
  184. }
  185. )
  186. return details
  187. def get_post_elements(execution_id: int, post_ids: list) -> dict[str, Any]:
  188. rows = repo.query_elements(execution_id, post_ids=post_ids, limit=5000)
  189. result: dict[str, Any] = {}
  190. grouped: dict[tuple[str, str, str], dict[str, Any]] = {}
  191. for row in rows:
  192. post_id = str(row.get("post_id") or "")
  193. point_type = str(row.get("point_type") or "未分点")
  194. point_text = str(row.get("point_text") or "")
  195. key = (post_id, point_type, point_text)
  196. point = grouped.setdefault(
  197. key,
  198. {
  199. "point_text": point_text,
  200. "elements": {"实质": [], "形式": [], "意图": []},
  201. },
  202. )
  203. element_type = str(row.get("element_type") or "其他")
  204. point["elements"].setdefault(element_type, []).append(
  205. {
  206. "id": row.get("id"),
  207. "source_element_id": row.get("source_element_id"),
  208. "name": row.get("name"),
  209. "description": row.get("description"),
  210. "category_id": row.get("category_id"),
  211. "category_path": row.get("category_path"),
  212. }
  213. )
  214. for (post_id, point_type, _), point in grouped.items():
  215. result.setdefault(post_id, {}).setdefault(point_type, []).append(point)
  216. return result
  217. def search_elements(
  218. execution_id: int,
  219. keyword: str,
  220. element_type: str | None = None,
  221. limit: int = 50,
  222. account_name=None,
  223. merge_leve2=None,
  224. platform=None,
  225. ) -> list[dict[str, Any]]:
  226. del account_name
  227. rows = repo.query_elements(
  228. execution_id,
  229. keyword=keyword,
  230. element_type=element_type,
  231. merge_leve2=merge_leve2,
  232. platform=platform,
  233. limit=5000,
  234. )
  235. grouped: dict[tuple[str, str, int | None, str | None], list[dict[str, Any]]] = defaultdict(list)
  236. for row in rows:
  237. key = (
  238. str(row.get("name") or ""),
  239. str(row.get("element_type") or ""),
  240. row.get("category_id"),
  241. row.get("category_path"),
  242. )
  243. grouped[key].append(row)
  244. items: list[dict[str, Any]] = []
  245. for (name, etype, category_id, category_path), group_rows in grouped.items():
  246. if not name:
  247. continue
  248. posts = {str(row.get("post_id")) for row in group_rows if row.get("post_id") is not None}
  249. items.append(
  250. {
  251. "name": name,
  252. "element_type": etype,
  253. "category_id": category_id,
  254. "category_path": category_path,
  255. "point_types": _group_point_types(group_rows),
  256. "occurrence_count": len(group_rows),
  257. "post_count": len(posts),
  258. }
  259. )
  260. items.sort(key=lambda item: (item["post_count"], item["occurrence_count"]), reverse=True)
  261. return items[: int(limit or 50)]
  262. def get_element_category_chain(
  263. execution_id: int,
  264. element_name: str,
  265. element_type: str | None = None,
  266. ) -> list[dict[str, Any]]:
  267. rows = repo.query_elements(execution_id, keyword=element_name, element_type=element_type, limit=2000)
  268. exact_rows = [row for row in rows if str(row.get("name") or "") == str(element_name)]
  269. category_ids = _clean_ints(row.get("category_id") for row in exact_rows if row.get("category_id"))
  270. categories = {int(row["id"]): row for row in repo.query_categories(execution_id, category_ids=category_ids)}
  271. results: list[dict[str, Any]] = []
  272. seen: set[int] = set()
  273. for row in exact_rows:
  274. category_id = row.get("category_id")
  275. if category_id is None or int(category_id) in seen:
  276. continue
  277. seen.add(int(category_id))
  278. category = categories.get(int(category_id), {})
  279. results.append(
  280. {
  281. "category_id": category_id,
  282. "category_path": row.get("category_path") or category.get("path"),
  283. "element_type": row.get("element_type"),
  284. "point_types": _group_point_types([r for r in exact_rows if r.get("category_id") == category_id]),
  285. "ancestors": _ancestors_from_path(category.get("path") or row.get("category_path")),
  286. }
  287. )
  288. return results
  289. def _ancestors_from_path(path: Any) -> list[dict[str, Any]]:
  290. parts = [part.strip() for part in str(path or "").replace(">", "/").split("/") if part.strip()]
  291. return [{"name": part, "level": index + 1, "path": "/".join(parts[: index + 1])} for index, part in enumerate(parts)]
  292. def get_category_detail_with_context(execution_id: int, category_id: int) -> dict[str, Any] | None:
  293. categories = repo.query_categories(execution_id, category_ids=[category_id])
  294. if not categories:
  295. return None
  296. category = categories[0]
  297. all_categories = repo.query_categories(execution_id)
  298. children = [cat for cat in all_categories if cat.get("parent_id") == int(category_id)]
  299. siblings = [
  300. cat for cat in all_categories
  301. if cat.get("parent_id") == category.get("parent_id") and cat.get("id") != category.get("id")
  302. ][:50]
  303. elements = get_category_elements(category_id, execution_id=execution_id)
  304. return {
  305. "category": category,
  306. "ancestors": _ancestors_from_path(category.get("path")),
  307. "children": children[:100],
  308. "siblings": siblings,
  309. "elements": elements[:100],
  310. }
  311. def search_categories(execution_id: int, keyword: str, source_type: str | None = None) -> list[dict[str, Any]]:
  312. categories = repo.query_categories(execution_id, source_type=source_type, keyword=keyword, limit=100)
  313. if not categories:
  314. return []
  315. category_ids = _clean_ints(cat.get("id") for cat in categories)
  316. element_rows = repo.query_elements(execution_id, category_ids=category_ids, limit=5000)
  317. elements_by_category: dict[int, list[dict[str, Any]]] = defaultdict(list)
  318. for row in element_rows:
  319. if row.get("category_id") is not None:
  320. elements_by_category[int(row["category_id"])].append(row)
  321. result: list[dict[str, Any]] = []
  322. for cat in categories:
  323. rows = elements_by_category.get(int(cat["id"]), [])
  324. result.append(
  325. {
  326. **cat,
  327. "point_types": _group_point_types(rows),
  328. "post_count": len({str(row.get("post_id")) for row in rows if row.get("post_id") is not None}),
  329. }
  330. )
  331. return result
  332. def get_category_elements(
  333. category_id: int,
  334. execution_id: int,
  335. account_name=None,
  336. merge_leve2=None,
  337. platform=None,
  338. ) -> list[dict[str, Any]]:
  339. del account_name
  340. rows = repo.query_elements(
  341. execution_id,
  342. category_ids=[category_id],
  343. merge_leve2=merge_leve2,
  344. platform=platform,
  345. limit=5000,
  346. )
  347. grouped: dict[tuple[str, str], list[dict[str, Any]]] = defaultdict(list)
  348. for row in rows:
  349. name = str(row.get("name") or "").strip()
  350. element_type = str(row.get("element_type") or "").strip()
  351. if name:
  352. grouped[(name, element_type)].append(row)
  353. result: list[dict[str, Any]] = []
  354. for (name, element_type), group_rows in grouped.items():
  355. result.append(
  356. {
  357. "name": name,
  358. "element_type": element_type,
  359. "point_types": _group_point_types(group_rows),
  360. "occurrence_count": len(group_rows),
  361. "post_count": len({str(row.get("post_id")) for row in group_rows if row.get("post_id") is not None}),
  362. }
  363. )
  364. result.sort(key=lambda item: (item["post_count"], item["occurrence_count"]), reverse=True)
  365. return result
  366. def get_category_co_occurrences(
  367. execution_id: int,
  368. category_ids: list,
  369. top_n: int = 30,
  370. account_name=None,
  371. merge_leve2=None,
  372. platform=None,
  373. ) -> dict[str, Any]:
  374. del account_name
  375. clean_ids = _clean_ints(category_ids)
  376. if not clean_ids:
  377. return {"matched_post_count": 0, "co_categories": []}
  378. rows = repo.query_elements(
  379. execution_id,
  380. category_ids=clean_ids,
  381. merge_leve2=merge_leve2,
  382. platform=platform,
  383. limit=10000,
  384. )
  385. posts_by_category: dict[int, set[str]] = defaultdict(set)
  386. for row in rows:
  387. if row.get("category_id") is not None and row.get("post_id") is not None:
  388. posts_by_category[int(row["category_id"])].add(str(row["post_id"]))
  389. matched_posts: set[str] | None = None
  390. for category_id in clean_ids:
  391. current = posts_by_category.get(category_id, set())
  392. matched_posts = current if matched_posts is None else matched_posts & current
  393. matched_posts = matched_posts or set()
  394. if not matched_posts:
  395. return {"matched_post_count": 0, "matched_post_ids": [], "co_categories": []}
  396. co_rows = repo.query_elements(
  397. execution_id,
  398. post_ids=list(matched_posts)[:5000],
  399. merge_leve2=merge_leve2,
  400. platform=platform,
  401. limit=20000,
  402. )
  403. counts: dict[tuple[int, str, str], set[str]] = defaultdict(set)
  404. for row in co_rows:
  405. cat_id = row.get("category_id")
  406. if cat_id is None or int(cat_id) in clean_ids:
  407. continue
  408. counts[(int(cat_id), str(row.get("category_path") or ""), str(row.get("element_type") or ""))].add(
  409. str(row.get("post_id"))
  410. )
  411. ranked = sorted(counts.items(), key=lambda item: len(item[1]), reverse=True)
  412. return {
  413. "matched_post_count": len(matched_posts),
  414. "matched_post_ids": sorted(matched_posts)[:100],
  415. "co_categories": [
  416. {
  417. "category_id": key[0],
  418. "category_path": key[1],
  419. "element_type": key[2],
  420. "post_count": len(posts),
  421. }
  422. for key, posts in ranked[: int(top_n or 30)]
  423. ],
  424. }
  425. def get_element_co_occurrences(
  426. execution_id: int,
  427. element_names: list,
  428. top_n: int = 30,
  429. account_name=None,
  430. merge_leve2=None,
  431. platform=None,
  432. ) -> dict[str, Any]:
  433. del account_name
  434. names = [str(name).strip() for name in element_names or [] if str(name).strip()]
  435. if not names:
  436. return {"matched_post_count": 0, "co_elements": []}
  437. posts_by_name: dict[str, set[str]] = {}
  438. for name in names:
  439. rows = repo.query_elements(
  440. execution_id,
  441. keyword=name,
  442. merge_leve2=merge_leve2,
  443. platform=platform,
  444. limit=10000,
  445. )
  446. posts_by_name[name] = {
  447. str(row.get("post_id"))
  448. for row in rows
  449. if str(row.get("name") or "") == name and row.get("post_id") is not None
  450. }
  451. matched_posts: set[str] | None = None
  452. for posts in posts_by_name.values():
  453. matched_posts = posts if matched_posts is None else matched_posts & posts
  454. matched_posts = matched_posts or set()
  455. if not matched_posts:
  456. return {"matched_post_count": 0, "matched_post_ids": [], "co_elements": []}
  457. co_rows = repo.query_elements(
  458. execution_id,
  459. post_ids=list(matched_posts)[:5000],
  460. merge_leve2=merge_leve2,
  461. platform=platform,
  462. limit=20000,
  463. )
  464. grouped: dict[tuple[str, str, int | None, str | None], list[dict[str, Any]]] = defaultdict(list)
  465. for row in co_rows:
  466. name = str(row.get("name") or "")
  467. if not name or name in names:
  468. continue
  469. grouped[(name, str(row.get("element_type") or ""), row.get("category_id"), row.get("category_path"))].append(row)
  470. ranked = sorted(grouped.items(), key=lambda item: len({r.get("post_id") for r in item[1]}), reverse=True)
  471. return {
  472. "matched_post_count": len(matched_posts),
  473. "matched_post_ids": sorted(matched_posts)[:100],
  474. "co_elements": [
  475. {
  476. "name": key[0],
  477. "element_type": key[1],
  478. "category_id": key[2],
  479. "category_path": key[3],
  480. "point_types": _group_point_types(rows),
  481. "occurrence_count": len(rows),
  482. "post_count": len({str(row.get("post_id")) for row in rows if row.get("post_id") is not None}),
  483. }
  484. for key, rows in ranked[: int(top_n or 30)]
  485. ],
  486. }