backfill_demand_video_expansion_point_desc.py 5.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162
  1. #!/usr/bin/env python3
  2. """补全 demand_video_expansion 表中缺失的 point_desc。
  3. 从 multi_demand_video_point 按 (video_id, point_type, expanded_text=point_data) 匹配;
  4. 查不到则保持空值。
  5. 用法:
  6. .venv/bin/python scripts/backfill_demand_video_expansion_point_desc.py
  7. .venv/bin/python scripts/backfill_demand_video_expansion_point_desc.py --biz-dt 20260721
  8. .venv/bin/python scripts/backfill_demand_video_expansion_point_desc.py --biz-dt 20260721 --dry-run
  9. """
  10. from __future__ import annotations
  11. import argparse
  12. import json
  13. import logging
  14. import sys
  15. from pathlib import Path
  16. from typing import Any
  17. from sqlalchemy import or_, select, update
  18. _ROOT = Path(__file__).resolve().parents[1]
  19. if str(_ROOT) not in sys.path:
  20. sys.path.insert(0, str(_ROOT))
  21. from agents.demand_video_expand_agent.tools.batch_save_demand_expansions import (
  22. _fill_missing_point_descs,
  23. )
  24. from supply_infra.db.models.demand_video_expansion import DemandVideoExpansion
  25. from supply_infra.db.session import get_session
  26. logger = logging.getLogger(__name__)
  27. _BATCH_SIZE = 500
  28. def _list_rows_missing_point_desc(biz_dt: str | None) -> list[dict[str, Any]]:
  29. stmt = select(DemandVideoExpansion).where(
  30. DemandVideoExpansion.is_delete == 0,
  31. or_(
  32. DemandVideoExpansion.point_desc.is_(None),
  33. DemandVideoExpansion.point_desc == "",
  34. ),
  35. )
  36. if biz_dt:
  37. stmt = stmt.where(DemandVideoExpansion.biz_dt == biz_dt)
  38. stmt = stmt.order_by(DemandVideoExpansion.id)
  39. with get_session() as session:
  40. rows = session.scalars(stmt).all()
  41. return [
  42. {
  43. "id": int(row.id),
  44. "biz_dt": str(row.biz_dt),
  45. "video_id": str(row.video_id),
  46. "point_type": str(row.point_type),
  47. "expanded_text": str(row.expanded_text),
  48. "point_desc": row.point_desc,
  49. }
  50. for row in rows
  51. ]
  52. def backfill_missing_point_descs(
  53. biz_dt: str | None = None,
  54. *,
  55. dry_run: bool = False,
  56. ) -> dict[str, Any]:
  57. rows = _list_rows_missing_point_desc(biz_dt)
  58. result: dict[str, Any] = {
  59. "biz_dt": biz_dt,
  60. "dry_run": dry_run,
  61. "missing_total": len(rows),
  62. "filled": 0,
  63. "still_empty": 0,
  64. "updated": 0,
  65. "samples": [],
  66. }
  67. if not rows:
  68. return result
  69. with get_session() as session:
  70. _fill_missing_point_descs(rows, session)
  71. to_update: list[dict[str, Any]] = []
  72. for row in rows:
  73. if row.get("point_desc"):
  74. to_update.append(row)
  75. result["filled"] += 1
  76. if len(result["samples"]) < 10:
  77. result["samples"].append(
  78. {
  79. "id": row["id"],
  80. "video_id": row["video_id"],
  81. "point_type": row["point_type"],
  82. "expanded_text": row["expanded_text"],
  83. "point_desc": row["point_desc"][:80]
  84. if len(str(row["point_desc"])) > 80
  85. else row["point_desc"],
  86. }
  87. )
  88. else:
  89. result["still_empty"] += 1
  90. if dry_run or not to_update:
  91. result["updated"] = 0
  92. return result
  93. with get_session() as session:
  94. for i in range(0, len(to_update), _BATCH_SIZE):
  95. batch = to_update[i : i + _BATCH_SIZE]
  96. for row in batch:
  97. session.execute(
  98. update(DemandVideoExpansion)
  99. .where(DemandVideoExpansion.id == int(row["id"]))
  100. .values(point_desc=row["point_desc"])
  101. )
  102. result["updated"] += len(batch)
  103. return result
  104. def main(argv: list[str] | None = None) -> int:
  105. parser = argparse.ArgumentParser(
  106. description="补全 demand_video_expansion 缺失的 point_desc",
  107. )
  108. parser.add_argument("--biz-dt", default=None, help="业务日期 YYYYMMDD,默认全表")
  109. parser.add_argument("--dry-run", action="store_true", help="仅统计,不写库")
  110. parser.add_argument("--json", action="store_true", help="以 JSON 输出结果")
  111. args = parser.parse_args(argv)
  112. logging.basicConfig(
  113. level=logging.INFO,
  114. format="%(asctime)s %(levelname)s %(name)s: %(message)s",
  115. )
  116. result = backfill_missing_point_descs(args.biz_dt, dry_run=bool(args.dry_run))
  117. if args.json:
  118. print(json.dumps(result, ensure_ascii=False, indent=2, default=str))
  119. else:
  120. print("\n=== point_desc 补全 ===")
  121. print(f"biz_dt={result.get('biz_dt') or '全部'}")
  122. print(f"dry_run={result.get('dry_run')}")
  123. print(f"缺失记录={result.get('missing_total')}")
  124. print(f"可补全={result.get('filled')}")
  125. print(f"仍为空={result.get('still_empty')}")
  126. print(f"已更新={result.get('updated')}")
  127. if result.get("samples"):
  128. print("\n示例:")
  129. for item in result["samples"]:
  130. print(
  131. f" id={item['id']} video={item['video_id']} "
  132. f"type={item['point_type']} text={item['expanded_text']!r} "
  133. f"desc={item['point_desc']!r}"
  134. )
  135. return 0
  136. if __name__ == "__main__":
  137. raise SystemExit(main())