sync.py 19 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597
  1. """
  2. Demand pool ODPS → MySQL sync stages used by the durable pipeline.
  3. Stages owned by this module:
  4. 1. source row sync (`_sync_pool_rows`)
  5. 2. demand word classification (`_classify_words`)
  6. 3. real rov/vov enrichment (`enrich_real_rov_vov_7d`)
  7. 4. popularity stats (`compute_popularity_stats`)
  8. Related pipeline stages live in sibling modules (belong_rel / tree_weight / videos).
  9. """
  10. from __future__ import annotations
  11. import hashlib
  12. import json
  13. import logging
  14. from concurrent.futures import ThreadPoolExecutor
  15. from datetime import datetime, timedelta
  16. from decimal import Decimal
  17. from typing import Any
  18. from supply_infra.category_match import (
  19. CategoryMatchClient,
  20. get_category_match_client,
  21. )
  22. from supply_infra.config import get_infra_settings
  23. from supply_infra.db.repositories.demand_belong_category_repo import (
  24. DemandBelongCategoryRepository,
  25. )
  26. from supply_infra.db.repositories.demand_belong_pool_rel_repo import (
  27. DemandBelongPoolRelRepository,
  28. )
  29. from supply_infra.db.repositories.demand_popularity_stats_repo import (
  30. DemandPopularityStatsRepository,
  31. )
  32. from supply_infra.db.repositories.global_tree_category_repo import (
  33. GlobalTreeCategoryRepository,
  34. )
  35. from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPoolDiRepository
  36. from supply_infra.db.session import get_session
  37. from supply_infra.odps.client import get_odps_client
  38. logger = logging.getLogger(__name__)
  39. _STATS_UPSERT_BATCH = 200
  40. _REAL_METRIC_LIMIT = 1000
  41. _REAL_METRIC_LOOKBACK_DAYS = 7
  42. _GLOBAL_FEATURE_VALUE = "全局SUM"
  43. _VIDEO_LIST_LIMIT = 10
  44. # 策略名 → 统计字段前缀;去年同期阳历/阴历合并为 plat_ly_pop
  45. _STRATEGY_METRIC: dict[str, str] = {
  46. "新热事件": "ext_pop",
  47. "逐月": "plat_sust_pop",
  48. "去年同期阳历": "plat_ly_pop",
  49. "去年同期阴历": "plat_ly_pop",
  50. "当下供需gap": "recent_pop",
  51. }
  52. _TARGET_STRATEGIES = list(_STRATEGY_METRIC.keys())
  53. _METRIC_KEYS = ("ext_pop", "plat_sust_pop", "plat_ly_pop", "recent_pop")
  54. _REAL_METRIC_KEYS = ("real_rov_7d", "real_vov_7d")
  55. _ALL_METRIC_KEYS = _METRIC_KEYS + _REAL_METRIC_KEYS
  56. _METRIC_DECIMAL_PLACES: dict[str, int] = {
  57. "ext_pop": 2,
  58. "plat_sust_pop": 2,
  59. "plat_ly_pop": 2,
  60. "recent_pop": 2,
  61. "real_rov_7d": 4,
  62. "real_vov_7d": 4,
  63. }
  64. RowKey = tuple[str, str]
  65. def _row_key(row: dict[str, Any]) -> RowKey:
  66. return (row["strategy"], row["demand_id"])
  67. _SOURCE_COMPARE_FIELDS = (
  68. "strategy",
  69. "demand_id",
  70. "demand_name",
  71. "weight",
  72. "type",
  73. "video_count",
  74. "video_list",
  75. "extend",
  76. "biz_dt",
  77. )
  78. def _source_signature(row: dict[str, Any]) -> tuple[Any, ...]:
  79. """生成稳定的源字段比较值;排除数据库 id 和下游回填字段。"""
  80. return tuple(row.get(field) for field in _SOURCE_COMPARE_FIELDS)
  81. def _source_snapshot_hash(rows: list[dict[str, Any]]) -> str:
  82. """生成与行顺序无关的内容哈希,供运行审计和水位比较。"""
  83. canonical = [
  84. {field: row.get(field) for field in _SOURCE_COMPARE_FIELDS}
  85. for row in sorted(rows, key=_row_key)
  86. ]
  87. payload = json.dumps(
  88. canonical,
  89. ensure_ascii=False,
  90. sort_keys=True,
  91. separators=(",", ":"),
  92. default=str,
  93. ).encode("utf-8")
  94. return hashlib.sha256(payload).hexdigest()
  95. def _normalize_video_list(raw: Any) -> tuple[str | None, int]:
  96. """取前 N 个 video_id,返回 (JSON 文本, count),二者保持一致。"""
  97. if raw is None:
  98. return None, 0
  99. items: list[Any]
  100. if isinstance(raw, str):
  101. text = raw.strip()
  102. if not text:
  103. return None, 0
  104. try:
  105. parsed = json.loads(text)
  106. items = list(parsed) if isinstance(parsed, list) else [text]
  107. except json.JSONDecodeError:
  108. items = [part.strip() for part in text.split(",") if part.strip()]
  109. elif isinstance(raw, (list, tuple)):
  110. items = list(raw)
  111. else:
  112. try:
  113. items = list(raw)
  114. except TypeError:
  115. return None, 0
  116. truncated = [
  117. str(v).strip()
  118. for v in items[:_VIDEO_LIST_LIMIT]
  119. if v is not None and str(v).strip()
  120. ]
  121. if not truncated:
  122. return None, 0
  123. return json.dumps(truncated, ensure_ascii=False), len(truncated)
  124. def _normalize_weight(raw: Any) -> float | None:
  125. """weight 为 0 或无效时视为无分数,存 null。"""
  126. if raw is None:
  127. return None
  128. try:
  129. value = float(raw)
  130. except (TypeError, ValueError):
  131. return None
  132. return None if value == 0 else value
  133. def _to_mysql_rows(raw_rows: list[dict[str, Any]], biz_dt: str) -> list[dict[str, Any]]:
  134. """转换为 MySQL 行,按 (strategy, demand_id) 去重(保留最后一条)。"""
  135. by_key: dict[RowKey, dict[str, Any]] = {}
  136. for row in raw_rows:
  137. strategy = row.get("strategy")
  138. demand_id = row.get("demand_id")
  139. demand_name = row.get("demand_name")
  140. if strategy is None or demand_id is None or demand_name is None:
  141. logger.warning("Skip row with missing required fields: %s", row)
  142. continue
  143. video_list, video_count = _normalize_video_list(row.get("video_list"))
  144. mapped = {
  145. "strategy": str(strategy),
  146. "demand_id": str(demand_id),
  147. "demand_name": str(demand_name),
  148. "weight": _normalize_weight(row.get("weight")),
  149. "type": str(row["type"]) if row.get("type") is not None else None,
  150. "video_count": video_count,
  151. "video_list": video_list,
  152. "extend": str(row["extend"]) if row.get("extend") is not None else None,
  153. "biz_dt": biz_dt,
  154. }
  155. by_key[_row_key(mapped)] = mapped
  156. return list(by_key.values())
  157. def _build_word_set(demand_names: list[str]) -> set[str]:
  158. """demand_name 按空格分词,全部落入同一个 set。"""
  159. words: set[str] = set()
  160. for name in demand_names:
  161. for token in str(name).split():
  162. word = token.strip()
  163. if word:
  164. words.add(word)
  165. return words
  166. def _classify_words(
  167. biz_dt: str,
  168. *,
  169. client: CategoryMatchClient | None = None,
  170. workers: int | None = None,
  171. ) -> dict:
  172. """逐词调用分类路径接口,将 stable_id 映射到本地分类后写入。"""
  173. with get_session() as session:
  174. demand_names = MultiDemandPoolDiRepository(
  175. session
  176. ).list_demand_names_by_biz_dt(biz_dt)
  177. word_set = _build_word_set(demand_names)
  178. existing = DemandBelongCategoryRepository(session).get_existing_names(word_set)
  179. pending = word_set - existing
  180. source_id_to_mysql_id = GlobalTreeCategoryRepository(
  181. session
  182. ).get_active_source_id_map()
  183. word_list = sorted(pending)
  184. matcher = client or get_category_match_client()
  185. worker_count = min(
  186. max(1, workers or get_infra_settings().category_match_workers),
  187. max(1, len(word_list)),
  188. )
  189. logger.info(
  190. "Category match prepare: demand_names=%d words=%d existing=%d "
  191. "pending=%d workers=%d",
  192. len(demand_names),
  193. len(word_set),
  194. len(existing),
  195. len(pending),
  196. worker_count,
  197. )
  198. rows: list[dict[str, Any]] = []
  199. unmatched_terms: list[str] = []
  200. failed_items: list[dict[str, Any]] = []
  201. def match_term(item: tuple[int, str]) -> tuple[str, Any, Exception | None]:
  202. idx, term = item
  203. logger.info("Matching category %d/%d: %s", idx, len(word_list), term)
  204. try:
  205. return term, matcher.match_one(term, description=""), None
  206. except Exception as exc:
  207. logger.exception(
  208. "Category match failed; continue remaining items: "
  209. "biz_dt=%s item=%s/%s term=%s",
  210. biz_dt,
  211. idx,
  212. len(word_list),
  213. term,
  214. )
  215. return term, None, exc
  216. indexed_terms = list(enumerate(word_list, start=1))
  217. with ThreadPoolExecutor(
  218. max_workers=worker_count,
  219. thread_name_prefix="category-match",
  220. ) as executor:
  221. outcomes = list(executor.map(match_term, indexed_terms))
  222. for term, match, error_value in outcomes:
  223. if error_value is not None:
  224. failed_items.append(
  225. {
  226. "term": term,
  227. "error": str(error_value),
  228. }
  229. )
  230. continue
  231. if match is None:
  232. unmatched_terms.append(term)
  233. continue
  234. category_id = source_id_to_mysql_id.get(match.stable_id)
  235. if category_id is None:
  236. error = (
  237. f"matched stable_id={match.stable_id} is absent from "
  238. "global_tree_category"
  239. )
  240. logger.error("Category match cannot be persisted: term=%s %s", term, error)
  241. failed_items.append({"term": term, "error": error})
  242. continue
  243. rows.append(
  244. {
  245. "name": term,
  246. "category_id": category_id,
  247. "reason": match.reason,
  248. "is_delete": 0,
  249. }
  250. )
  251. with get_session() as session:
  252. persisted = DemandBelongCategoryRepository(
  253. session
  254. ).upsert_category_matches(rows)
  255. result = {
  256. "success": not failed_items,
  257. "demand_names": len(demand_names),
  258. "words": len(word_set),
  259. "existing_filtered": len(existing),
  260. "pending": len(pending),
  261. "api_calls": len(word_list),
  262. "matched": len(rows),
  263. "persisted": persisted,
  264. "unmatched": len(unmatched_terms),
  265. "unmatched_terms": unmatched_terms,
  266. "failed_items": failed_items,
  267. }
  268. if failed_items:
  269. result["error_code"] = "category_match_failed_items"
  270. result["error"] = f"{len(failed_items)} category match item(s) failed"
  271. return result
  272. def _sync_diff(partition_date: str) -> dict[str, Any]:
  273. """按主键和完整源内容同步;行数相同也会识别修订。"""
  274. odps = get_odps_client()
  275. raw_rows = odps.fetch_multi_demand_pool(partition_date)
  276. mysql_rows = _to_mysql_rows(raw_rows, partition_date)
  277. odps_by_key = {_row_key(r): r for r in mysql_rows}
  278. odps_keys = set(odps_by_key)
  279. with get_session() as session:
  280. repo = MultiDemandPoolDiRepository(session)
  281. existing_rows = repo.list_source_rows_by_biz_dt(partition_date)
  282. mysql_by_key = {_row_key(row): row for row in existing_rows}
  283. mysql_keys = set(mysql_by_key)
  284. to_insert_keys = odps_keys - mysql_keys
  285. to_delete_keys = mysql_keys - odps_keys
  286. shared_keys = odps_keys & mysql_keys
  287. to_update_keys = {
  288. key
  289. for key in shared_keys
  290. if _source_signature(odps_by_key[key])
  291. != _source_signature(mysql_by_key[key])
  292. }
  293. insert_rows = [odps_by_key[k] for k in to_insert_keys]
  294. update_rows = [odps_by_key[k] for k in to_update_keys]
  295. deleted_pool_ids = [mysql_by_key[key]["id"] for key in to_delete_keys]
  296. deleted_relations = DemandBelongPoolRelRepository(session).delete_by_pool_ids(
  297. deleted_pool_ids
  298. )
  299. deleted = repo.delete_by_keys(partition_date, list(to_delete_keys))
  300. inserted = repo.bulk_insert(insert_rows)
  301. updated = (
  302. repo.update_source_fields(partition_date, update_rows)
  303. if update_rows
  304. else 0
  305. )
  306. logger.info(
  307. "Diff sync: odps=%d mysql_before=%d insert=%d delete=%d update=%d unchanged=%d",
  308. len(odps_keys),
  309. len(mysql_keys),
  310. inserted,
  311. deleted,
  312. updated,
  313. len(shared_keys) - len(to_update_keys),
  314. )
  315. return {
  316. "fetched": len(raw_rows),
  317. "odps_unique": len(odps_keys),
  318. "mysql_before": len(mysql_keys),
  319. "inserted": inserted,
  320. "deleted": deleted,
  321. "relations_deleted": deleted_relations,
  322. "updated": updated,
  323. "unchanged": len(shared_keys) - len(to_update_keys),
  324. "source_snapshot_hash": _source_snapshot_hash(mysql_rows),
  325. "skipped_invalid": len(raw_rows) - len(mysql_rows),
  326. }
  327. def _to_float(value: Any) -> float | None:
  328. if value is None:
  329. return None
  330. try:
  331. return float(value)
  332. except (TypeError, ValueError):
  333. return None
  334. def _pick_better_real_metric(
  335. current: dict[str, Any] | None,
  336. candidate: dict[str, Any],
  337. ) -> dict[str, Any]:
  338. """相同特征值时保留 rov_diff、vov_diff 更高的一条。"""
  339. if current is None:
  340. return candidate
  341. current_key = (
  342. _to_float(current.get("rov_diff")) or 0.0,
  343. _to_float(current.get("vov_diff")) or 0.0,
  344. )
  345. candidate_key = (
  346. _to_float(candidate.get("rov_diff")) or 0.0,
  347. _to_float(candidate.get("vov_diff")) or 0.0,
  348. )
  349. return candidate if candidate_key > current_key else current
  350. def _build_real_metrics_by_feature(
  351. raw_rows: list[dict[str, Any]],
  352. ) -> dict[str, tuple[float | None, float | None]]:
  353. """按特征值去重,映射为 demand_name → (real_rov_7d, real_vov_7d)。"""
  354. best_by_feature: dict[str, dict[str, Any]] = {}
  355. for row in raw_rows:
  356. feature = row.get("特征值")
  357. if feature is None:
  358. continue
  359. feature_name = str(feature).strip()
  360. if not feature_name or feature_name == _GLOBAL_FEATURE_VALUE:
  361. continue
  362. best_by_feature[feature_name] = _pick_better_real_metric(
  363. best_by_feature.get(feature_name),
  364. row,
  365. )
  366. return {
  367. feature: (_to_float(row.get("rov_diff")), _to_float(row.get("vov_diff")))
  368. for feature, row in best_by_feature.items()
  369. }
  370. def enrich_real_rov_vov_7d(biz_dt: str) -> dict[str, Any]:
  371. """
  372. 从 ODPS 拉取执行日及往前 7 日的 rov_diff/vov_diff,按特征值匹配回填 MySQL。
  373. 相同特征值保留 rov_diff、vov_diff 更高的记录。
  374. """
  375. dt_right = biz_dt
  376. dt_left = (
  377. datetime.strptime(biz_dt, "%Y%m%d") - timedelta(days=_REAL_METRIC_LOOKBACK_DAYS)
  378. ).strftime("%Y%m%d")
  379. logger.info(
  380. "Enrich real rov/vov: biz_dt=%s range=%s~%s limit=%d",
  381. biz_dt,
  382. dt_left,
  383. dt_right,
  384. _REAL_METRIC_LIMIT,
  385. )
  386. odps = get_odps_client()
  387. raw_rows = odps.fetch_real_rov_vov_7d(
  388. dt_left=dt_left,
  389. dt_right=dt_right,
  390. limit=_REAL_METRIC_LIMIT,
  391. )
  392. metrics_by_name = _build_real_metrics_by_feature(raw_rows)
  393. with get_session() as session:
  394. updated = MultiDemandPoolDiRepository(session).update_real_metrics_by_demand_name(
  395. biz_dt,
  396. metrics_by_name,
  397. )
  398. result = {
  399. "dt_left": dt_left,
  400. "dt_right": dt_right,
  401. "odps_rows": len(raw_rows),
  402. "unique_features": len(metrics_by_name),
  403. "updated_rows": updated,
  404. }
  405. logger.info("Enrich real rov/vov completed: %s", result)
  406. return result
  407. def _calc_metric_stats(
  408. weights: list[float],
  409. places: int = 2,
  410. ) -> tuple[Decimal | None, int]:
  411. """
  412. 计算单策略 avg / count。
  413. 权重为 0 不参与平均值;全部为 0(或无有效权重)时 avg=null、count=0。
  414. """
  415. nonzero = [w for w in weights if w != 0]
  416. if not nonzero:
  417. return None, 0
  418. avg_val = sum(nonzero) / len(nonzero)
  419. q = Decimal("0." + "0" * places)
  420. return (
  421. Decimal(str(round(avg_val, places))).quantize(q),
  422. len(nonzero),
  423. )
  424. def _empty_metric_stats() -> dict[str, Decimal | int | None]:
  425. result: dict[str, Decimal | int | None] = {}
  426. for key in _ALL_METRIC_KEYS:
  427. result[f"{key}_avg"] = None
  428. result[f"{key}_count"] = 0
  429. return result
  430. def _build_stats_row(
  431. demand_category_id: int,
  432. demand_word_name: str,
  433. biz_dt: str,
  434. strategy_weights: list[tuple[str, float | None]],
  435. real_metrics: list[tuple[float | None, float | None]],
  436. ) -> dict[str, Any]:
  437. """按策略分组后汇总为 demand_popularity_stats 一行(含 rov_diff/vov_diff)。"""
  438. grouped: dict[str, list[float]] = {key: [] for key in _METRIC_KEYS}
  439. for strategy, weight in strategy_weights:
  440. metric = _STRATEGY_METRIC.get(strategy)
  441. if metric is None or weight is None or weight == 0:
  442. continue
  443. grouped[metric].append(float(weight))
  444. rov_values: list[float] = []
  445. vov_values: list[float] = []
  446. for rov, vov in real_metrics:
  447. if rov is not None:
  448. rov_values.append(float(rov))
  449. if vov is not None:
  450. vov_values.append(float(vov))
  451. row: dict[str, Any] = {
  452. "demand_category_id": demand_category_id,
  453. "demand_word_name": demand_word_name[:128],
  454. "biz_dt": biz_dt,
  455. **_empty_metric_stats(),
  456. }
  457. for key in _METRIC_KEYS:
  458. avg_val, count = _calc_metric_stats(
  459. grouped[key], places=_METRIC_DECIMAL_PLACES[key]
  460. )
  461. row[f"{key}_avg"] = avg_val
  462. row[f"{key}_count"] = count
  463. rov_avg, rov_count = _calc_metric_stats(
  464. rov_values, places=_METRIC_DECIMAL_PLACES["real_rov_7d"]
  465. )
  466. vov_avg, vov_count = _calc_metric_stats(
  467. vov_values, places=_METRIC_DECIMAL_PLACES["real_vov_7d"]
  468. )
  469. row["real_rov_7d_avg"] = rov_avg
  470. row["real_rov_7d_count"] = rov_count
  471. row["real_vov_7d_avg"] = vov_avg
  472. row["real_vov_7d_count"] = vov_count
  473. return row
  474. def compute_popularity_stats(biz_dt: str) -> dict[str, Any]:
  475. """
  476. 遍历 demand_belong_category 全部词,各自 LIKE 查询当天指定策略权重,
  477. 并汇总 rov_diff/vov_diff,写入 demand_popularity_stats(avg/count)。
  478. Args:
  479. biz_dt: 业务日期 (YYYYMMDD)
  480. """
  481. with get_session() as session:
  482. categories = DemandBelongCategoryRepository(session).list_active_id_name()
  483. if not categories:
  484. logger.info("Popularity stats: no demand_belong_category rows, skip")
  485. return {"words": 0, "upserted": 0}
  486. logger.info("Popularity stats: processing %d words for biz_dt=%s", len(categories), biz_dt)
  487. rows: list[dict[str, Any]] = []
  488. upserted = 0
  489. with get_session() as session:
  490. pool_repo = MultiDemandPoolDiRepository(session)
  491. stats_repo = DemandPopularityStatsRepository(session)
  492. for idx, (category_id, name) in enumerate(categories, start=1):
  493. matches = pool_repo.list_weights_by_name_like(
  494. biz_dt, name, _TARGET_STRATEGIES
  495. )
  496. real_matches = pool_repo.list_real_metrics_by_name_like(biz_dt, name)
  497. rows.append(
  498. _build_stats_row(category_id, name, biz_dt, matches, real_matches)
  499. )
  500. if len(rows) >= _STATS_UPSERT_BATCH:
  501. upserted += stats_repo.upsert_rows(rows)
  502. logger.info(
  503. "Popularity stats progress: %d/%d words, batch upserted",
  504. idx,
  505. len(categories),
  506. )
  507. rows = []
  508. if rows:
  509. upserted += stats_repo.upsert_rows(rows)
  510. result = {"words": len(categories), "upserted": upserted}
  511. logger.info("Popularity stats completed: %s", result)
  512. return result
  513. def _sync_pool_rows(partition_date: str) -> dict[str, Any]:
  514. """同步当天需求池主数据;内容哈希差分替代不可靠的行数门禁。"""
  515. return _sync_diff(partition_date)