registry.py 6.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198
  1. from __future__ import annotations
  2. import hashlib
  3. import json
  4. from collections.abc import Callable
  5. from pathlib import Path
  6. from typing import Any
  7. from supply_infra.config import get_infra_settings
  8. from supply_infra.db.repositories.pipeline_outbox_repo import PipelineOutboxRepository
  9. from supply_infra.db.session import get_session
  10. from supply_infra.pipeline.contracts import StepContext
  11. StepHandler = Callable[[StepContext], dict[str, Any]]
  12. _REPO_ROOT = Path(__file__).resolve().parents[2]
  13. _MAX_INLINE_EFFECT_BYTES = 512_000
  14. def _global_tree(context: StepContext) -> dict[str, Any]:
  15. from supply_infra.scheduler.jobs.sync_global_tree_odps_to_mysql import (
  16. sync_global_tree_odps_to_mysql,
  17. )
  18. return sync_global_tree_odps_to_mysql(
  19. partition_date=str(context.date_snapshot["global_tree_partition"])
  20. )
  21. def _demand_source(context: StepContext) -> dict[str, Any]:
  22. from supply_infra.scheduler.jobs.demand_pool.sync import _sync_pool_rows
  23. return _sync_pool_rows(str(context.date_snapshot["demand_pool_partition"]))
  24. def _demand_classify(context: StepContext) -> dict[str, Any]:
  25. from supply_infra.scheduler.jobs.demand_pool.sync import _classify_words
  26. return _classify_words(context.biz_dt)
  27. def _demand_rel(_context: StepContext) -> dict[str, Any]:
  28. from supply_infra.scheduler.jobs.demand_pool.belong_rel import (
  29. sync_demand_belong_pool_rel,
  30. )
  31. return sync_demand_belong_pool_rel()
  32. def _real_metrics(context: StepContext) -> dict[str, Any]:
  33. from supply_infra.scheduler.jobs.demand_pool.sync import enrich_real_rov_vov_7d
  34. return enrich_real_rov_vov_7d(context.biz_dt)
  35. def _popularity(context: StepContext) -> dict[str, Any]:
  36. from supply_infra.scheduler.jobs.demand_pool.sync import compute_popularity_stats
  37. return compute_popularity_stats(context.biz_dt)
  38. def _tree_weight(context: StepContext) -> dict[str, Any]:
  39. from supply_infra.scheduler.jobs.demand_pool.tree_weight import (
  40. compute_category_tree_weight,
  41. )
  42. return compute_category_tree_weight(context.biz_dt)
  43. def _source_videos(context: StepContext) -> dict[str, Any]:
  44. from supply_infra.scheduler.jobs.demand_pool.videos import sync_multi_demand_videos
  45. return sync_multi_demand_videos(
  46. limit=None,
  47. decode_dt=str(context.date_snapshot["video_decode_partition"]),
  48. )
  49. def _grade(context: StepContext) -> dict[str, Any]:
  50. from supply_infra.scheduler.jobs.grade_demand_pool import grade_demand_pool
  51. return grade_demand_pool(context.biz_dt)
  52. def _expand(context: StepContext) -> dict[str, Any]:
  53. from supply_infra.scheduler.jobs.expand_demand_from_video_points import (
  54. expand_demand_from_video_points,
  55. )
  56. return expand_demand_from_video_points(
  57. context.biz_dt,
  58. workers=5,
  59. )
  60. def _discover(context: StepContext) -> dict[str, Any]:
  61. from supply_infra.scheduler.constants import PIPELINE_FIND_AGENT_WORKERS
  62. from supply_infra.scheduler.jobs.discover_videos_from_demands import (
  63. discover_videos_from_demands,
  64. )
  65. return discover_videos_from_demands(
  66. context.biz_dt,
  67. workers=PIPELINE_FIND_AGENT_WORKERS,
  68. )
  69. def _aigc_write_record(context: StepContext) -> dict[str, Any]:
  70. from supply_infra.scheduler.jobs.publish_videos_from_discovery import (
  71. publish_videos_from_discovery,
  72. )
  73. settings = get_infra_settings()
  74. payload = publish_videos_from_discovery(biz_dt=context.biz_dt)
  75. publish_success = bool(payload.get("success", True))
  76. canonical = json.dumps(
  77. payload,
  78. ensure_ascii=False,
  79. sort_keys=True,
  80. separators=(",", ":"),
  81. default=str,
  82. )
  83. encoded = canonical.encode("utf-8")
  84. payload_hash = hashlib.sha256(encoded).hexdigest()
  85. idempotency_key = f"aigc_write_record:{context.biz_dt}:{payload_hash}"
  86. payload_uri: str | None = None
  87. stored_payload: dict[str, Any] | None = payload
  88. if len(encoded) > _MAX_INLINE_EFFECT_BYTES:
  89. log_dir = Path(settings.pipeline_log_dir)
  90. if not log_dir.is_absolute():
  91. log_dir = _REPO_ROOT / log_dir
  92. effect_dir = log_dir / context.run_id / "effects"
  93. effect_dir.mkdir(parents=True, exist_ok=True)
  94. effect_path = effect_dir / f"{context.step_run_id}-{payload_hash}.json"
  95. temporary_path = effect_path.with_suffix(".json.tmp")
  96. temporary_path.write_bytes(encoded)
  97. temporary_path.replace(effect_path)
  98. payload_uri = str(effect_path)
  99. stored_payload = {
  100. "externalized": True,
  101. "payload_bytes": len(encoded),
  102. "payload_hash": payload_hash,
  103. }
  104. with get_session() as session:
  105. record = PipelineOutboxRepository(session).record_effect(
  106. run_id=context.run_id,
  107. step_run_id=context.step_run_id,
  108. effect_type="aigc_write_record",
  109. idempotency_key=idempotency_key,
  110. payload_hash=payload_hash,
  111. payload=stored_payload,
  112. payload_uri=payload_uri,
  113. )
  114. outbox_id = record.outbox_id
  115. result: dict[str, Any] = {
  116. "success": publish_success,
  117. "effect_recorded": True,
  118. "external_request_made": int(payload.get("candidate_count", 0) or 0) > 0,
  119. "outbox_id": outbox_id,
  120. "payload_hash": payload_hash,
  121. "payload_uri": payload_uri,
  122. "candidate_count": payload.get("candidate_count", 0),
  123. "batch_count": payload.get("batch_count", 0),
  124. "failed_batch_count": payload.get("failed_batch_count", 0),
  125. }
  126. if not publish_success:
  127. result["error"] = str(
  128. payload.get("error")
  129. or f"AIGC publish failed ({result['failed_batch_count']} batch(es))"
  130. )
  131. result["error_code"] = "aigc_publish_failed"
  132. return result
  133. STEP_REGISTRY: dict[str, StepHandler] = {
  134. "global_tree_sync": _global_tree,
  135. "demand_pool_source_sync": _demand_source,
  136. "demand_classify": _demand_classify,
  137. "demand_belong_rel_sync": _demand_rel,
  138. "real_metrics_sync": _real_metrics,
  139. "popularity_stats": _popularity,
  140. "category_tree_weight": _tree_weight,
  141. "source_video_sync": _source_videos,
  142. "demand_grade": _grade,
  143. "demand_expand": _expand,
  144. "video_discovery": _discover,
  145. "aigc_write_record": _aigc_write_record,
  146. }
  147. def execute_registered_step(context: StepContext) -> dict[str, Any]:
  148. handler = STEP_REGISTRY.get(context.step_key)
  149. if handler is None:
  150. raise KeyError(f"Unknown pipeline step: {context.step_key}")
  151. payload = handler(context)
  152. if not isinstance(payload, dict):
  153. return {"success": True, "result": payload}
  154. return payload