dim_constants.py 5.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186
  1. """generate_demand_agent 分类树热度维度共享常量与工具函数。"""
  2. from __future__ import annotations
  3. from collections.abc import Callable
  4. from decimal import Decimal
  5. from supply_infra.db.models.category_tree_weight import CategoryTreeWeight
  6. from supply_infra.db.models.global_tree_category import GlobalTreeCategory
  7. DIM_KEYS: tuple[str, ...] = (
  8. "ext_pop",
  9. "plat_sust_pop",
  10. "plat_ly_pop",
  11. "recent_pop",
  12. )
  13. DIM_LABEL: dict[str, str] = {
  14. "ext_pop": "外部热度",
  15. "plat_sust_pop": "平台持续热度",
  16. "plat_ly_pop": "平台去年同期热度",
  17. "recent_pop": "近期热度",
  18. }
  19. # 节点下存在有数据的叶子节点时的标记
  20. HAS_DATA_LEAF_MARK = "+"
  21. def normalize_biz_dt(biz_dt: str | None) -> tuple[str | None, str | None]:
  22. """校验并规范化 biz_dt(YYYYMMDD);空值表示使用最新业务日。"""
  23. if biz_dt is None:
  24. return None, None
  25. text = str(biz_dt).strip()
  26. if not text:
  27. return None, None
  28. if len(text) != 8 or not text.isdigit():
  29. return None, f"biz_dt 格式无效,应为 YYYYMMDD: {biz_dt!r}"
  30. return text, None
  31. def resolve_biz_dt(
  32. biz_dt: str | None,
  33. *,
  34. get_latest: Callable[[], str | None],
  35. has_data: Callable[[str], bool] | None = None,
  36. table_label: str = "热度数据",
  37. ) -> tuple[str | None, str | None]:
  38. """解析 biz_dt:未传则用最新业务日;传入则校验格式与数据是否存在。"""
  39. normalized, err = normalize_biz_dt(biz_dt)
  40. if err:
  41. return None, err
  42. resolved = normalized or get_latest()
  43. if not resolved:
  44. return None, f"暂无 {table_label}(请指定 biz_dt 或先写入数据)"
  45. if normalized and has_data is not None and not has_data(normalized):
  46. return None, f"biz_dt={normalized} 无 {table_label}"
  47. return resolved, None
  48. def get_latest_common_biz_dt(session) -> str | None:
  49. """返回 category_tree_weight 与 demand_popularity_stats 均有数据的最新 biz_dt。"""
  50. from sqlalchemy import func, select
  51. from supply_infra.db.models.demand_popularity_stats import DemandPopularityStats
  52. stats_dates = select(DemandPopularityStats.biz_dt).distinct()
  53. stmt = select(func.max(CategoryTreeWeight.biz_dt)).where(
  54. CategoryTreeWeight.biz_dt.in_(stats_dates)
  55. )
  56. return session.scalar(stmt)
  57. def normalize_parent_id(parent_id: int | None) -> int | None:
  58. if parent_id is None or parent_id == 0:
  59. return None
  60. return parent_id
  61. def format_score(avg: Decimal | float | int | None) -> str:
  62. if avg is None:
  63. return "—"
  64. value = float(avg)
  65. if value >= 100:
  66. return f"{value:.0f}"
  67. if value >= 1:
  68. return f"{value:.1f}"
  69. return f"{value:.2f}"
  70. def dim_score(
  71. weight: CategoryTreeWeight | None,
  72. dim: str,
  73. ) -> tuple[float | None, int]:
  74. """返回 (avg, count);count>0 即有维度数据(avg 为 0 也算有)。"""
  75. if weight is None:
  76. return None, 0
  77. count = int(getattr(weight, f"{dim}_count", 0) or 0)
  78. if count <= 0:
  79. return None, 0
  80. avg = getattr(weight, f"{dim}_avg", None)
  81. return (float(avg) if avg is not None else 0.0), count
  82. def build_children_map(
  83. categories: list[GlobalTreeCategory],
  84. ) -> dict[int | None, list[GlobalTreeCategory]]:
  85. children_map: dict[int | None, list[GlobalTreeCategory]] = {}
  86. for category in categories:
  87. parent_key = normalize_parent_id(category.parent_id)
  88. children_map.setdefault(parent_key, []).append(category)
  89. for children in children_map.values():
  90. children.sort(key=lambda c: (c.level or 0, c.id))
  91. return children_map
  92. def build_score_by_id(
  93. weights: list[CategoryTreeWeight],
  94. dim: str,
  95. ) -> dict[int, tuple[float | None, int]]:
  96. return {int(row.category_id): dim_score(row, dim) for row in weights}
  97. def collect_descendant_leaves(
  98. root_id: int,
  99. children_map: dict[int | None, list[GlobalTreeCategory]],
  100. ) -> list[int]:
  101. """收集 root 子树内的叶子节点 id(若 root 本身无子节点则含 root)。"""
  102. leaves: list[int] = []
  103. stack = [root_id]
  104. seen: set[int] = set()
  105. while stack:
  106. cid = stack.pop()
  107. if cid in seen:
  108. continue
  109. seen.add(cid)
  110. kids = children_map.get(cid, [])
  111. if not kids:
  112. leaves.append(cid)
  113. else:
  114. stack.extend(int(c.id) for c in kids)
  115. leaves.sort()
  116. return leaves
  117. def collect_descendant_ids(
  118. root_id: int,
  119. children_map: dict[int | None, list[GlobalTreeCategory]],
  120. ) -> list[int]:
  121. """收集 root 及其全部子孙节点 id。"""
  122. ids: list[int] = []
  123. stack = [root_id]
  124. seen: set[int] = set()
  125. while stack:
  126. cid = stack.pop()
  127. if cid in seen:
  128. continue
  129. seen.add(cid)
  130. ids.append(cid)
  131. stack.extend(int(c.id) for c in children_map.get(cid, []))
  132. ids.sort()
  133. return ids
  134. def build_category_path(
  135. category_id: int,
  136. by_id: dict[int, GlobalTreeCategory],
  137. ) -> str | None:
  138. """从根到指定节点的名称路径,如「美妆护肤 > 护肤 > 防晒」。"""
  139. cat = by_id.get(category_id)
  140. if cat is None:
  141. return None
  142. names: list[str] = []
  143. current: GlobalTreeCategory | None = cat
  144. seen: set[int] = set()
  145. while current is not None:
  146. cid = int(current.id)
  147. if cid in seen:
  148. break
  149. seen.add(cid)
  150. names.append(current.name or str(cid))
  151. parent_key = normalize_parent_id(current.parent_id)
  152. current = by_id.get(parent_key) if parent_key is not None else None
  153. names.reverse()
  154. return " > ".join(names)