shared.py 6.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207
  1. """demand_grade_agent 工具共享辅助函数(非 @tool,不对外暴露为工具)。"""
  2. from __future__ import annotations
  3. import json
  4. from decimal import Decimal
  5. from typing import Any
  6. from supply_infra.db.models.global_tree_category import GlobalTreeCategory
  7. from supply_infra.db.models.multi_demand_pool_di import MultiDemandPoolDi
  8. _VIDEO_LIST_LIMIT = 10
  9. VALID_GRADES: tuple[str, ...] = ("S", "A", "B", "C", "D")
  10. def normalize_biz_dt(biz_dt: str | None) -> tuple[str | None, str | None]:
  11. """校验并规范化 biz_dt(YYYYMMDD);空值返回 (None, None) 表示未指定。"""
  12. if biz_dt is None:
  13. return None, None
  14. text = str(biz_dt).strip()
  15. if not text:
  16. return None, None
  17. if len(text) != 8 or not text.isdigit():
  18. return None, f"biz_dt 格式无效,应为 YYYYMMDD: {biz_dt!r}"
  19. return text, None
  20. def to_float(value: Any) -> float | None:
  21. """将 Decimal/None/数值统一转换为 float,None 原样返回。"""
  22. if value is None:
  23. return None
  24. return float(value)
  25. def format_score(avg: Decimal | float | int | None) -> str:
  26. """格式化分数为易读文本,None 显示为 —。"""
  27. if avg is None:
  28. return "—"
  29. value = float(avg)
  30. if value >= 100:
  31. return f"{value:.0f}"
  32. if value >= 1:
  33. return f"{value:.2f}"
  34. return f"{value:.4f}"
  35. def percentile(sorted_values: list[float], pct: float) -> float | None:
  36. """对已排序(升序)的数值列表求分位数(线性插值),pct 取 0~100。"""
  37. if not sorted_values:
  38. return None
  39. if len(sorted_values) == 1:
  40. return sorted_values[0]
  41. rank = (pct / 100) * (len(sorted_values) - 1)
  42. lower_idx = int(rank)
  43. upper_idx = min(lower_idx + 1, len(sorted_values) - 1)
  44. frac = rank - lower_idx
  45. return sorted_values[lower_idx] + (sorted_values[upper_idx] - sorted_values[lower_idx]) * frac
  46. def distribution_summary(values: list[float]) -> dict[str, float | int | None]:
  47. """返回一组数值的 min/p25/p50/p75/p90/max/count 分布摘要。"""
  48. if not values:
  49. return {"count": 0, "min": None, "p25": None, "p50": None, "p75": None, "p90": None, "max": None}
  50. ordered = sorted(values)
  51. return {
  52. "count": len(ordered),
  53. "min": ordered[0],
  54. "p25": percentile(ordered, 25),
  55. "p50": percentile(ordered, 50),
  56. "p75": percentile(ordered, 75),
  57. "p90": percentile(ordered, 90),
  58. "max": ordered[-1],
  59. }
  60. def _normalize_parent_id(parent_id: int | None) -> int | None:
  61. if parent_id is None or parent_id == 0:
  62. return None
  63. return parent_id
  64. def build_category_path(
  65. category_id: int,
  66. by_id: dict[int, GlobalTreeCategory],
  67. ) -> str | None:
  68. """从根到指定节点的名称路径,如「美妆护肤 > 护肤 > 防晒」。"""
  69. cat = by_id.get(category_id)
  70. if cat is None:
  71. return None
  72. names: list[str] = []
  73. current: GlobalTreeCategory | None = cat
  74. seen: set[int] = set()
  75. while current is not None:
  76. cid = int(current.id)
  77. if cid in seen:
  78. break
  79. seen.add(cid)
  80. names.append(current.name or str(cid))
  81. parent_key = _normalize_parent_id(current.parent_id)
  82. current = by_id.get(parent_key) if parent_key is not None else None
  83. names.reverse()
  84. return " > ".join(names)
  85. def normalize_str_list(raw: Any, field_name: str = "items") -> tuple[list[str], str | None]:
  86. """将 JSON 文本/列表/单个字符串统一解析为非空 str 列表(去重保序)。"""
  87. if raw is None:
  88. return [], f"{field_name} 不能为空"
  89. items: list[Any]
  90. if isinstance(raw, str):
  91. text = raw.strip()
  92. if not text:
  93. return [], f"{field_name} 不能为空"
  94. try:
  95. parsed = json.loads(text)
  96. items = list(parsed) if isinstance(parsed, list) else [text]
  97. except (ValueError, TypeError):
  98. items = [text]
  99. elif isinstance(raw, (list, tuple)):
  100. items = list(raw)
  101. else:
  102. items = [raw]
  103. out: list[str] = []
  104. seen: set[str] = set()
  105. for item in items:
  106. text = str(item).strip() if item is not None else ""
  107. if not text or text in seen:
  108. continue
  109. seen.add(text)
  110. out.append(text)
  111. if not out:
  112. return [], f"{field_name} 不能为空"
  113. return out, None
  114. def parse_int_list(raw: Any) -> list[int]:
  115. """将 JSON 文本/列表统一解析为 int 列表,解析失败返回空列表。"""
  116. if raw is None:
  117. return []
  118. if isinstance(raw, str):
  119. try:
  120. raw = json.loads(raw)
  121. except (ValueError, TypeError):
  122. return []
  123. if not isinstance(raw, list):
  124. return []
  125. out: list[int] = []
  126. for item in raw:
  127. try:
  128. out.append(int(item))
  129. except (TypeError, ValueError):
  130. continue
  131. return out
  132. def dump_int_list(values: list[int] | None) -> str | None:
  133. """将 int 列表序列化为 JSON 文本,空列表/None 返回 None。"""
  134. if not values:
  135. return None
  136. return json.dumps(values, ensure_ascii=False)
  137. def _parse_video_ids(raw: Any) -> list[str]:
  138. """解析单个 multi_demand_pool_di.video_list(JSON数组/逗号分隔文本)为 vid 字符串列表。"""
  139. if raw is None:
  140. return []
  141. items: list[Any]
  142. if isinstance(raw, str):
  143. text = raw.strip()
  144. if not text:
  145. return []
  146. try:
  147. parsed = json.loads(text)
  148. items = list(parsed) if isinstance(parsed, list) else [text]
  149. except (ValueError, TypeError):
  150. items = [part.strip() for part in text.split(",") if part.strip()]
  151. elif isinstance(raw, (list, tuple)):
  152. items = list(raw)
  153. else:
  154. return []
  155. return [str(v).strip() for v in items if v is not None and str(v).strip()]
  156. def merge_video_ids(pool_rows: list[MultiDemandPoolDi], limit: int = _VIDEO_LIST_LIMIT) -> str | None:
  157. """合并多条原始需求行的 video_list,去重保序,最多取前 limit 个,返回 JSON 文本或 None。"""
  158. merged: list[str] = []
  159. seen: set[str] = set()
  160. for row in pool_rows:
  161. for vid in _parse_video_ids(row.video_list):
  162. if vid not in seen:
  163. seen.add(vid)
  164. merged.append(vid)
  165. if not merged:
  166. return None
  167. return json.dumps(merged[:limit], ensure_ascii=False)
  168. def collect_strategies(pool_rows: list[MultiDemandPoolDi]) -> str | None:
  169. """收集多条原始需求行的策略名,去重排序,返回 JSON 文本或 None。"""
  170. strategies = sorted({row.strategy.strip() for row in pool_rows if row.strategy and row.strategy.strip()})
  171. if not strategies:
  172. return None
  173. return json.dumps(strategies, ensure_ascii=False)