pg_pattern_service.py 19 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483
  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. ) -> list[dict[str, Any]]:
  225. del account_name, merge_leve2
  226. rows = repo.query_elements(execution_id, keyword=keyword, element_type=element_type, limit=5000)
  227. grouped: dict[tuple[str, str, int | None, str | None], list[dict[str, Any]]] = defaultdict(list)
  228. for row in rows:
  229. key = (
  230. str(row.get("name") or ""),
  231. str(row.get("element_type") or ""),
  232. row.get("category_id"),
  233. row.get("category_path"),
  234. )
  235. grouped[key].append(row)
  236. items: list[dict[str, Any]] = []
  237. for (name, etype, category_id, category_path), group_rows in grouped.items():
  238. if not name:
  239. continue
  240. posts = {str(row.get("post_id")) for row in group_rows if row.get("post_id") is not None}
  241. items.append(
  242. {
  243. "name": name,
  244. "element_type": etype,
  245. "category_id": category_id,
  246. "category_path": category_path,
  247. "point_types": _group_point_types(group_rows),
  248. "occurrence_count": len(group_rows),
  249. "post_count": len(posts),
  250. }
  251. )
  252. items.sort(key=lambda item: (item["post_count"], item["occurrence_count"]), reverse=True)
  253. return items[: int(limit or 50)]
  254. def get_element_category_chain(
  255. execution_id: int,
  256. element_name: str,
  257. element_type: str | None = None,
  258. ) -> list[dict[str, Any]]:
  259. rows = repo.query_elements(execution_id, keyword=element_name, element_type=element_type, limit=2000)
  260. exact_rows = [row for row in rows if str(row.get("name") or "") == str(element_name)]
  261. category_ids = _clean_ints(row.get("category_id") for row in exact_rows if row.get("category_id"))
  262. categories = {int(row["id"]): row for row in repo.query_categories(execution_id, category_ids=category_ids)}
  263. results: list[dict[str, Any]] = []
  264. seen: set[int] = set()
  265. for row in exact_rows:
  266. category_id = row.get("category_id")
  267. if category_id is None or int(category_id) in seen:
  268. continue
  269. seen.add(int(category_id))
  270. category = categories.get(int(category_id), {})
  271. results.append(
  272. {
  273. "category_id": category_id,
  274. "category_path": row.get("category_path") or category.get("path"),
  275. "element_type": row.get("element_type"),
  276. "point_types": _group_point_types([r for r in exact_rows if r.get("category_id") == category_id]),
  277. "ancestors": _ancestors_from_path(category.get("path") or row.get("category_path")),
  278. }
  279. )
  280. return results
  281. def _ancestors_from_path(path: Any) -> list[dict[str, Any]]:
  282. parts = [part.strip() for part in str(path or "").replace(">", "/").split("/") if part.strip()]
  283. return [{"name": part, "level": index + 1, "path": "/".join(parts[: index + 1])} for index, part in enumerate(parts)]
  284. def get_category_detail_with_context(execution_id: int, category_id: int) -> dict[str, Any] | None:
  285. categories = repo.query_categories(execution_id, category_ids=[category_id])
  286. if not categories:
  287. return None
  288. category = categories[0]
  289. all_categories = repo.query_categories(execution_id)
  290. children = [cat for cat in all_categories if cat.get("parent_id") == int(category_id)]
  291. siblings = [
  292. cat for cat in all_categories
  293. if cat.get("parent_id") == category.get("parent_id") and cat.get("id") != category.get("id")
  294. ][:50]
  295. elements = get_category_elements(category_id, execution_id=execution_id)
  296. return {
  297. "category": category,
  298. "ancestors": _ancestors_from_path(category.get("path")),
  299. "children": children[:100],
  300. "siblings": siblings,
  301. "elements": elements[:100],
  302. }
  303. def search_categories(execution_id: int, keyword: str, source_type: str | None = None) -> list[dict[str, Any]]:
  304. categories = repo.query_categories(execution_id, source_type=source_type, keyword=keyword, limit=100)
  305. if not categories:
  306. return []
  307. category_ids = _clean_ints(cat.get("id") for cat in categories)
  308. element_rows = repo.query_elements(execution_id, category_ids=category_ids, limit=5000)
  309. elements_by_category: dict[int, list[dict[str, Any]]] = defaultdict(list)
  310. for row in element_rows:
  311. if row.get("category_id") is not None:
  312. elements_by_category[int(row["category_id"])].append(row)
  313. result: list[dict[str, Any]] = []
  314. for cat in categories:
  315. rows = elements_by_category.get(int(cat["id"]), [])
  316. result.append(
  317. {
  318. **cat,
  319. "point_types": _group_point_types(rows),
  320. "post_count": len({str(row.get("post_id")) for row in rows if row.get("post_id") is not None}),
  321. }
  322. )
  323. return result
  324. def get_category_elements(
  325. category_id: int,
  326. execution_id: int,
  327. account_name=None,
  328. merge_leve2=None,
  329. ) -> list[dict[str, Any]]:
  330. del account_name, merge_leve2
  331. rows = repo.query_elements(execution_id, category_ids=[category_id], limit=5000)
  332. grouped: dict[tuple[str, str], list[dict[str, Any]]] = defaultdict(list)
  333. for row in rows:
  334. name = str(row.get("name") or "").strip()
  335. element_type = str(row.get("element_type") or "").strip()
  336. if name:
  337. grouped[(name, element_type)].append(row)
  338. result: list[dict[str, Any]] = []
  339. for (name, element_type), group_rows in grouped.items():
  340. result.append(
  341. {
  342. "name": name,
  343. "element_type": element_type,
  344. "point_types": _group_point_types(group_rows),
  345. "occurrence_count": len(group_rows),
  346. "post_count": len({str(row.get("post_id")) for row in group_rows if row.get("post_id") is not None}),
  347. }
  348. )
  349. result.sort(key=lambda item: (item["post_count"], item["occurrence_count"]), reverse=True)
  350. return result
  351. def get_category_co_occurrences(
  352. execution_id: int,
  353. category_ids: list,
  354. top_n: int = 30,
  355. account_name=None,
  356. merge_leve2=None,
  357. ) -> dict[str, Any]:
  358. del account_name, merge_leve2
  359. clean_ids = _clean_ints(category_ids)
  360. if not clean_ids:
  361. return {"matched_post_count": 0, "co_categories": []}
  362. rows = repo.query_elements(execution_id, category_ids=clean_ids, limit=10000)
  363. posts_by_category: dict[int, set[str]] = defaultdict(set)
  364. for row in rows:
  365. if row.get("category_id") is not None and row.get("post_id") is not None:
  366. posts_by_category[int(row["category_id"])].add(str(row["post_id"]))
  367. matched_posts: set[str] | None = None
  368. for category_id in clean_ids:
  369. current = posts_by_category.get(category_id, set())
  370. matched_posts = current if matched_posts is None else matched_posts & current
  371. matched_posts = matched_posts or set()
  372. co_rows = repo.query_elements(execution_id, post_ids=list(matched_posts)[:5000], limit=20000)
  373. counts: dict[tuple[int, str, str], set[str]] = defaultdict(set)
  374. for row in co_rows:
  375. cat_id = row.get("category_id")
  376. if cat_id is None or int(cat_id) in clean_ids:
  377. continue
  378. counts[(int(cat_id), str(row.get("category_path") or ""), str(row.get("element_type") or ""))].add(
  379. str(row.get("post_id"))
  380. )
  381. ranked = sorted(counts.items(), key=lambda item: len(item[1]), reverse=True)
  382. return {
  383. "matched_post_count": len(matched_posts),
  384. "matched_post_ids": sorted(matched_posts)[:100],
  385. "co_categories": [
  386. {
  387. "category_id": key[0],
  388. "category_path": key[1],
  389. "element_type": key[2],
  390. "post_count": len(posts),
  391. }
  392. for key, posts in ranked[: int(top_n or 30)]
  393. ],
  394. }
  395. def get_element_co_occurrences(
  396. execution_id: int,
  397. element_names: list,
  398. top_n: int = 30,
  399. account_name=None,
  400. merge_leve2=None,
  401. ) -> dict[str, Any]:
  402. del account_name, merge_leve2
  403. names = [str(name).strip() for name in element_names or [] if str(name).strip()]
  404. if not names:
  405. return {"matched_post_count": 0, "co_elements": []}
  406. posts_by_name: dict[str, set[str]] = {}
  407. for name in names:
  408. rows = repo.query_elements(execution_id, keyword=name, limit=10000)
  409. posts_by_name[name] = {
  410. str(row.get("post_id"))
  411. for row in rows
  412. if str(row.get("name") or "") == name and row.get("post_id") is not None
  413. }
  414. matched_posts: set[str] | None = None
  415. for posts in posts_by_name.values():
  416. matched_posts = posts if matched_posts is None else matched_posts & posts
  417. matched_posts = matched_posts or set()
  418. co_rows = repo.query_elements(execution_id, post_ids=list(matched_posts)[:5000], limit=20000)
  419. grouped: dict[tuple[str, str, int | None, str | None], list[dict[str, Any]]] = defaultdict(list)
  420. for row in co_rows:
  421. name = str(row.get("name") or "")
  422. if not name or name in names:
  423. continue
  424. grouped[(name, str(row.get("element_type") or ""), row.get("category_id"), row.get("category_path"))].append(row)
  425. ranked = sorted(grouped.items(), key=lambda item: len({r.get("post_id") for r in item[1]}), reverse=True)
  426. return {
  427. "matched_post_count": len(matched_posts),
  428. "matched_post_ids": sorted(matched_posts)[:100],
  429. "co_elements": [
  430. {
  431. "name": key[0],
  432. "element_type": key[1],
  433. "category_id": key[2],
  434. "category_path": key[3],
  435. "point_types": _group_point_types(rows),
  436. "occurrence_count": len(rows),
  437. "post_count": len({str(row.get("post_id")) for row in rows if row.get("post_id") is not None}),
  438. }
  439. for key, rows in ranked[: int(top_n or 30)]
  440. ],
  441. }