| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357 |
- """
- 从 S/A 级需求关联视频中挖掘拓展需求。
- 任务层负责查库与组装上下文;Agent 仅做语义判断与落库。
- """
- from __future__ import annotations
- import json
- import logging
- import uuid
- from concurrent.futures import ThreadPoolExecutor, as_completed
- from datetime import datetime
- from typing import Any
- from zoneinfo import ZoneInfo
- from agents.demand_video_expand_agent.run import (
- DemandExpandContext,
- VideoPoint,
- extract_saved_count,
- judge_demand_expansion,
- )
- from supply_infra.config import get_infra_settings
- from supply_infra.db.models.demand_grade import DemandGrade
- from supply_infra.db.repositories.demand_grade_repo import DemandGradeRepository
- from supply_infra.db.repositories.demand_video_expansion_repo import (
- DemandVideoExpansionRunRepository,
- )
- from supply_infra.db.repositories.multi_demand_video_point_repo import (
- MultiDemandVideoPointRepository,
- )
- from supply_infra.db.session import get_session
- logger = logging.getLogger(__name__)
- _DEFAULT_WORKERS = 5
- def _resolve_biz_dt(biz_dt: str | None) -> str:
- if biz_dt:
- text = str(biz_dt).strip()
- if len(text) == 8 and text.isdigit():
- return text
- raise ValueError(f"biz_dt 格式无效,应为 YYYYMMDD: {biz_dt!r}")
- with get_session() as session:
- latest = DemandGradeRepository(session).get_latest_biz_dt()
- if latest:
- return str(latest)
- timezone = ZoneInfo(get_infra_settings().scheduler_timezone)
- return datetime.now(timezone).strftime("%Y%m%d")
- def _parse_video_ids(raw: Any) -> list[str]:
- if raw is None:
- return []
- items: list[Any]
- if isinstance(raw, str):
- text = raw.strip()
- if not text:
- return []
- try:
- parsed = json.loads(text)
- items = list(parsed) if isinstance(parsed, list) else [text]
- except (ValueError, TypeError):
- items = [part.strip() for part in text.split(",") if part.strip()]
- elif isinstance(raw, (list, tuple)):
- items = list(raw)
- else:
- return []
- out: list[str] = []
- seen: set[str] = set()
- for item in items:
- vid = str(item).strip() if item is not None else ""
- if not vid or vid in seen:
- continue
- seen.add(vid)
- out.append(vid)
- return out
- def _to_float(value: Any) -> float | None:
- if value is None:
- return None
- return float(value)
- def load_expand_contexts(
- session,
- biz_dt: str,
- *,
- run_id: str,
- skip_finished: bool = True,
- ) -> tuple[list[DemandExpandContext], dict[str, int]]:
- """加载待拓展判断的需求上下文,并返回跳过统计。"""
- stats = {
- "total_sa": 0,
- "skipped_no_video": 0,
- "skipped_no_points": 0,
- "skipped_already_done": 0,
- }
- grades = DemandGradeRepository(session).list_by_biz_dt_and_grades(biz_dt, ("S", "A"))
- stats["total_sa"] = len(grades)
- finished_ids: set[int] = set()
- if skip_finished:
- finished_ids = DemandVideoExpansionRunRepository(session).list_finished_grade_ids(
- biz_dt
- )
- rows_with_video: list[tuple[DemandGrade, list[str]]] = []
- all_video_ids: set[str] = set()
- for row in grades:
- if int(row.id) in finished_ids:
- stats["skipped_already_done"] += 1
- continue
- video_ids = _parse_video_ids(row.video_list)
- if not video_ids:
- stats["skipped_no_video"] += 1
- continue
- rows_with_video.append((row, video_ids))
- all_video_ids.update(video_ids)
- points_by_vid = MultiDemandVideoPointRepository(session).list_by_video_ids(all_video_ids)
- contexts: list[DemandExpandContext] = []
- for row, video_ids in rows_with_video:
- points: list[VideoPoint] = []
- for vid in video_ids:
- for point in points_by_vid.get(vid, []):
- points.append(
- VideoPoint(
- video_id=vid,
- point_type=str(point.get("point_type") or ""),
- point_data=point.get("point_data"),
- point_desc=point.get("point_desc"),
- )
- )
- if not points:
- stats["skipped_no_points"] += 1
- continue
- contexts.append(
- DemandExpandContext(
- biz_dt=biz_dt,
- run_id=run_id,
- demand_grade_id=int(row.id),
- demand_name=str(row.demand_name),
- grade=str(row.grade),
- score=_to_float(row.score),
- video_ids=video_ids,
- points=points,
- )
- )
- return contexts, stats
- def _record_run(
- *,
- biz_dt: str,
- run_id: str,
- source_demand_grade_id: int,
- saved_count: int,
- status: str,
- error_message: str | None = None,
- ) -> None:
- with get_session() as session:
- DemandVideoExpansionRunRepository(session).upsert_run(
- biz_dt=biz_dt,
- run_id=run_id,
- source_demand_grade_id=source_demand_grade_id,
- saved_count=saved_count,
- status=status,
- error_message=error_message,
- )
- def _process_single_expand(
- ctx: DemandExpandContext,
- *,
- biz_dt: str,
- run_id: str,
- ) -> dict[str, Any]:
- """并发 worker:对单个需求执行拓展判断并落库执行记录。"""
- try:
- agent_result = judge_demand_expansion(ctx)
- saved_count = extract_saved_count(agent_result)
- _record_run(
- biz_dt=biz_dt,
- run_id=run_id,
- source_demand_grade_id=ctx.demand_grade_id,
- saved_count=saved_count,
- status="finished",
- )
- logger.info(
- "expand demand done: grade_id=%s demand=%s saved=%d iterations=%d",
- ctx.demand_grade_id,
- ctx.demand_name,
- saved_count,
- agent_result.iterations,
- )
- return {
- "success": True,
- "demand_grade_id": ctx.demand_grade_id,
- "demand_name": ctx.demand_name,
- "saved_count": saved_count,
- "iterations": agent_result.iterations,
- }
- except Exception as exc:
- error_text = str(exc)
- _record_run(
- biz_dt=biz_dt,
- run_id=run_id,
- source_demand_grade_id=ctx.demand_grade_id,
- saved_count=0,
- status="failed",
- error_message=error_text,
- )
- logger.exception(
- "expand demand failed: grade_id=%s demand=%s",
- ctx.demand_grade_id,
- ctx.demand_name,
- )
- return {
- "success": False,
- "demand_grade_id": ctx.demand_grade_id,
- "demand_name": ctx.demand_name,
- "error": error_text,
- }
- def expand_demand_from_video_points(
- biz_dt: str | None = None,
- *,
- skip_finished: bool = True,
- workers: int = _DEFAULT_WORKERS,
- ) -> dict[str, Any]:
- """
- 对指定业务日的 S/A 需求执行视频点位拓展判断。
- 程序负责查需求与点位;无视频或无点位则跳过;有数据则并发调用 Agent。
- """
- started_at = datetime.now()
- run_id = uuid.uuid4().hex
- try:
- resolved_biz_dt = _resolve_biz_dt(biz_dt)
- except Exception as exc:
- logger.exception("expand_demand_from_video_points preflight failed")
- return {
- "success": False,
- "error": str(exc),
- "started_at": started_at.isoformat(),
- "finished_at": datetime.now().isoformat(),
- }
- with get_session() as session:
- contexts, preload_stats = load_expand_contexts(
- session,
- resolved_biz_dt,
- run_id=run_id,
- skip_finished=skip_finished,
- )
- logger.info(
- "expand_demand_from_video_points start: biz_dt=%s run_id=%s workers=%s pending=%s",
- resolved_biz_dt,
- run_id,
- workers,
- len(contexts),
- )
- result: dict[str, Any] = {
- "success": True,
- "run_id": run_id,
- "biz_dt": resolved_biz_dt,
- "started_at": started_at.isoformat(),
- "workers": 0,
- **preload_stats,
- "processed": 0,
- "saved_total": 0,
- "failed": 0,
- "errors": [],
- }
- if not contexts:
- finished_at = datetime.now()
- result["finished_at"] = finished_at.isoformat()
- result["duration_seconds"] = round((finished_at - started_at).total_seconds(), 2)
- logger.info("expand_demand_from_video_points finished: %s", result)
- return result
- worker_count = max(1, min(int(workers), len(contexts)))
- result["workers"] = worker_count
- with ThreadPoolExecutor(max_workers=worker_count) as executor:
- futures = [
- executor.submit(_process_single_expand, ctx, biz_dt=resolved_biz_dt, run_id=run_id)
- for ctx in contexts
- ]
- for future in as_completed(futures):
- try:
- item_result = future.result()
- except Exception as exc:
- logger.exception("expand demand worker 出现未捕获错误: biz_dt=%s", resolved_biz_dt)
- result["failed"] += 1
- result["errors"].append({"error": str(exc)})
- continue
- result["processed"] += 1
- if item_result.get("success"):
- result["saved_total"] += int(item_result.get("saved_count") or 0)
- continue
- result["failed"] += 1
- result["errors"].append(
- {
- "demand_grade_id": item_result.get("demand_grade_id"),
- "demand_name": item_result.get("demand_name"),
- "error": item_result.get("error"),
- }
- )
- finished_at = datetime.now()
- result["finished_at"] = finished_at.isoformat()
- result["duration_seconds"] = round((finished_at - started_at).total_seconds(), 2)
- result["success"] = result["failed"] == 0
- logger.info("expand_demand_from_video_points finished: %s", result)
- return result
- if __name__ == "__main__":
- import sys
- from supply_infra.scheduler.cli_result import run_cli
- _args = sys.argv[1:]
- _biz_dt = _args[0] if _args and not _args[0].startswith("-") else None
- _workers_arg = None
- if _biz_dt and len(_args) > 1 and not _args[1].startswith("-"):
- _workers_arg = _args[1]
- _skip_finished = "--force" not in _args
- run_cli(
- lambda: expand_demand_from_video_points(
- _biz_dt,
- skip_finished=_skip_finished,
- workers=int(_workers_arg) if _workers_arg else 5,
- ),
- label="expand_demand_from_video_points",
- )
|