xueyiming 4 дней назад
Родитель
Сommit
87cad8bd27

+ 14 - 0
agents/find_agent/__init__.py

@@ -6,6 +6,7 @@ find_agent — 老年受众高潜视频发现 Agent
 from __future__ import annotations
 
 import logging
+from concurrent.futures import ThreadPoolExecutor, TimeoutError as FuturesTimeoutError
 
 from agents.find_agent.agent import create_find_agent
 from agents.find_agent.async_runner import _run_coroutine, arun_find_agent
@@ -17,6 +18,8 @@ __all__ = ["create_find_agent", "run_find_agent"]
 
 logger = logging.getLogger(__name__)
 
+_PUBLISH_TIMEOUT_SECONDS = 120.0
+
 
 def run_find_agent(
     user_input: str,
@@ -37,6 +40,17 @@ def run_find_agent(
             timeout_seconds=agent.settings.find_agent_timeout_seconds,
         )
     )
+    try:
+        with ThreadPoolExecutor(max_workers=1) as publish_executor:
+            publish_future = publish_executor.submit(agent._finish_run, result)
+            publish_future.result(timeout=_PUBLISH_TIMEOUT_SECONDS)
+    except FuturesTimeoutError:
+        logger.error(
+            "find_agent publish timed out after %.0fs; continuing without blocking worker",
+            _PUBLISH_TIMEOUT_SECONDS,
+        )
+    except Exception:
+        logger.exception("find_agent publish failed")
     try:
         persist_model_recommendations(result)
     except Exception:

+ 34 - 9
agents/find_agent/async_runner.py

@@ -10,7 +10,6 @@ from supply_agent.types import AgentResult
 
 logger = logging.getLogger(__name__)
 
-_thread_loop = threading.local()
 _ASYNC_CLEANUP_TIMEOUT_SECONDS = 5.0
 
 
@@ -71,10 +70,10 @@ async def arun_find_agent(
     *,
     timeout_seconds: float,
 ) -> AgentResult:
-    """在单个事件循环内运行 find_agent 并确保资源释放。"""
+    """在单个事件循环内运行 find_agent 核心循环并确保资源释放(不含 OSS 发布)。"""
     try:
         return await asyncio.wait_for(
-            agent.arun(user_input),
+            agent.arun_core(user_input),
             timeout=timeout_seconds,
         )
     except TimeoutError:
@@ -86,8 +85,30 @@ async def arun_find_agent(
         await _close_agent_async_resources(agent)
 
 
+def _shutdown_worker_loop(loop: asyncio.AbstractEventLoop) -> None:
+    """Release async generators/executor before closing a worker-thread loop."""
+    shutdown_timeout = 3.0
+    try:
+        loop.run_until_complete(
+            asyncio.wait_for(loop.shutdown_asyncgens(), timeout=shutdown_timeout)
+        )
+    except Exception:
+        logger.debug("shutdown_asyncgens failed", exc_info=True)
+    shutdown_executor = getattr(loop, "shutdown_default_executor", None)
+    if shutdown_executor is not None:
+        try:
+            loop.run_until_complete(
+                asyncio.wait_for(
+                    shutdown_executor(),
+                    timeout=shutdown_timeout,
+                )
+            )
+        except Exception:
+            logger.debug("shutdown_default_executor failed", exc_info=True)
+
+
 def _run_coroutine(coro) -> AgentResult:
-    """同步入口:主线程用 asyncio.run,worker 线程复用线程级 loop。"""
+    """同步入口:主线程用 asyncio.run;worker 线程每次新建并关闭独立 loop。"""
     try:
         asyncio.get_running_loop()
     except RuntimeError:
@@ -98,9 +119,13 @@ def _run_coroutine(coro) -> AgentResult:
     if threading.current_thread() is threading.main_thread():
         return asyncio.run(coro)
 
-    loop = getattr(_thread_loop, "loop", None)
-    if loop is None or loop.is_closed():
-        loop = asyncio.new_event_loop()
-        _thread_loop.loop = loop
+    loop = asyncio.new_event_loop()
     asyncio.set_event_loop(loop)
-    return loop.run_until_complete(coro)
+    try:
+        return loop.run_until_complete(coro)
+    finally:
+        try:
+            _shutdown_worker_loop(loop)
+        finally:
+            loop.close()
+            asyncio.set_event_loop(None)

+ 2 - 2
agents/find_agent/demand_run.py

@@ -145,7 +145,7 @@ def load_find_demand_contexts(
     grades: Iterable[str] = ("S", "A"),
     top_limit: int | None = None,
 ) -> list[FindDemandContext]:
-    """加载指定业务日 S/A 需求,每个需求词组装为一条完整上下文。"""
+    """加载指定业务日全部 S/A 需求(可选 top_limit 截断),每个需求词组装为一条完整上下文。"""
     grades_rows = DemandGradeRepository(session).list_by_biz_dt_and_grades(biz_dt, grades)
     if not grades_rows:
         return []
@@ -243,7 +243,7 @@ def list_find_demand_contexts(
     grades: Iterable[str] = ("S", "A"),
     top_limit: int | None = None,
 ) -> tuple[str, list[FindDemandContext]]:
-    """返回解析后的业务日与全部待执行上下文。"""
+    """返回解析后的业务日与待执行上下文(默认当日全部 S/A,可选 top_limit)。"""
     resolved_biz_dt = _resolve_biz_dt(biz_dt)
     with get_session() as session:
         contexts = load_find_demand_contexts(

+ 9 - 3
supply_agent/agent/core.py

@@ -159,15 +159,21 @@ class Agent:
         self._finish_run(result)
         return result
 
-    async def arun(
+    async def arun_core(
         self, user_input: str, *, history: list[Message] | None = None
     ) -> AgentResult:
-        """Run the agent asynchronously."""
+        """Run the agent loop without closing logs or publishing artifacts."""
         self.logger.start_run(user_input, model=self.model, agent_name=self.name)
         messages = list(history or [])
         messages.append(Message(role=Role.USER, content=user_input))
         loop = self._create_loop(messages)
-        result = await loop.arun()
+        return await loop.arun()
+
+    async def arun(
+        self, user_input: str, *, history: list[Message] | None = None
+    ) -> AgentResult:
+        """Run the agent asynchronously."""
+        result = await self.arun_core(user_input, history=history)
         self._finish_run(result)
         return result
 

+ 1 - 1
supply_infra/config.py

@@ -36,7 +36,7 @@ class InfraSettings(BaseSettings):
     )
     mysql_pool_size_api: int = Field(default=6, ge=1, alias="MYSQL_POOL_SIZE_API")
     mysql_pool_size_control: int = Field(
-        default=1,
+        default=2,
         ge=1,
         alias="MYSQL_POOL_SIZE_CONTROL",
     )

+ 18 - 0
supply_infra/db/session.py

@@ -55,6 +55,24 @@ def dispose_engine() -> None:
     _SessionLocal = None
 
 
+def ensure_mysql_pool_capacity(min_connections: int) -> None:
+    """Grow the process-local pool before parallel DB access in the same process."""
+    required = max(1, int(min_connections))
+    settings = get_infra_settings()
+    if settings.selected_mysql_pool_size >= required:
+        return
+
+    import os
+
+    role = settings.process_role
+    if role in {"scheduler", "worker", "reconciler", "step"}:
+        os.environ["MYSQL_POOL_SIZE_CONTROL"] = str(required)
+    else:
+        os.environ["MYSQL_POOL_SIZE"] = str(required)
+    get_infra_settings.cache_clear()
+    dispose_engine()
+
+
 def init_db() -> dict[str, list[str]]:
     """Create all tables (dev / first-run). Import models before calling."""
     import supply_infra.db.models  # noqa: F401 — register all models

+ 1 - 5
supply_infra/pipeline/registry.py

@@ -93,10 +93,7 @@ def _expand(context: StepContext) -> dict[str, Any]:
 
 
 def _discover(context: StepContext) -> dict[str, Any]:
-    from supply_infra.scheduler.constants import (
-        PIPELINE_FIND_AGENT_TOP_DEMANDS,
-        PIPELINE_FIND_AGENT_WORKERS,
-    )
+    from supply_infra.scheduler.constants import PIPELINE_FIND_AGENT_WORKERS
     from supply_infra.scheduler.jobs.discover_videos_from_demands import (
         discover_videos_from_demands,
     )
@@ -104,7 +101,6 @@ def _discover(context: StepContext) -> dict[str, Any]:
     return discover_videos_from_demands(
         context.biz_dt,
         workers=PIPELINE_FIND_AGENT_WORKERS,
-        top_limit=PIPELINE_FIND_AGENT_TOP_DEMANDS,
     )
 
 

+ 2 - 2
supply_infra/scheduler/constants.py

@@ -3,9 +3,9 @@
 SUPPLY_PIPELINE_JOB_ID = "run_supply_pipeline"
 SUPPLY_PIPELINE_JOB_NAME = "供给数据流水线"
 
-# find_agent:每日按 score 取 top N 需求,2 线程并行找片
-PIPELINE_FIND_AGENT_TOP_DEMANDS = 200
+# find_agent:当日全部 S/A 需求(有拓展点位),2 线程并行;有效视频满 200 提前结束
 PIPELINE_FIND_AGENT_WORKERS = 2
+# step 子进程内 MYSQL_POOL_SIZE_CONTROL 默认须 >= 并行 worker 数
 
 # 供给流水线每日触发时间(Asia/Shanghai)
 PIPELINE_CRON_HOUR = 14

+ 60 - 11
supply_infra/scheduler/jobs/discover_videos_from_demands.py

@@ -1,5 +1,5 @@
 """
-从 S/A 级需求及其拓展点位触发 find_agent 视频发现。
+从全部 S/A 级需求及其拓展点位触发 find_agent 视频发现;有效视频满 200 提前结束
 
 任务层负责查库与组装上下文;Agent 负责搜索、画像与分池落库。
 """
@@ -23,7 +23,7 @@ from agents.find_agent.demand_run import (
 from supply_infra.config import get_infra_settings
 from supply_infra.db.repositories.demand_grade_repo import DemandGradeRepository
 from supply_infra.db.repositories.video_discovery_repo import VideoDiscoveryRepository
-from supply_infra.db.session import get_session
+from supply_infra.db.session import ensure_mysql_pool_capacity, get_session
 
 logger = logging.getLogger(__name__)
 
@@ -122,14 +122,16 @@ def discover_videos_from_demands(
     force: bool = False,
 ) -> dict[str, Any]:
     """
-    对指定业务日的 S/A 需求拓展点位逐条调用 find_agent。
+    对指定业务日全部 S/A 级需求(有拓展点位)逐条调用 find_agent。
 
     每条记录对应一个 demand_grade + 其下全部视频与全部拓展点位。
     执行前会预写 video_discovery_run,并按 biz_dt + demand_grade_id 跳过已执行记录。
-    top_limit 按 score 取当日 top N 需求(S 优先于 A)。
+    默认处理全部 S/A(S 优先于 A、再按 score 排序);仅 CLI --top-limit 可人为截断。
+    当日 primary+backup 去重视频达到 200 时提前结束,否则跑完待处理队列。
     """
     started_at = datetime.now()
     batch_run_id = uuid.uuid4().hex
+    ensure_mysql_pool_capacity(max(1, int(workers)))
 
     try:
         resolved_biz_dt, contexts = list_find_demand_contexts(
@@ -220,11 +222,23 @@ def discover_videos_from_demands(
             )
             batch = contexts[next_context : next_context + batch_size]
             next_context += batch_size
-            futures = [
-                executor.submit(process_single_discover, ctx, force=force)
+            logger.info(
+                "discover videos batch start: biz_dt=%s batch_size=%d "
+                "completed_batches=%d/%d demands=%s",
+                resolved_biz_dt,
+                len(batch),
+                next_context // batch_size,
+                (len(contexts) + batch_size - 1) // batch_size,
+                [
+                    f"{ctx.demand_grade_id}:{ctx.demand_name}"
+                    for ctx in batch
+                ],
+            )
+            future_to_ctx = {
+                executor.submit(process_single_discover, ctx, force=force): ctx
                 for ctx in batch
-            ]
-            pending = set(futures)
+            }
+            pending = set(future_to_ctx.keys())
             batch_started = monotonic()
             while pending:
                 completed, pending = wait(
@@ -233,30 +247,65 @@ def discover_videos_from_demands(
                     return_when=FIRST_COMPLETED,
                 )
                 if not completed:
+                    pending_demands = [
+                        f"{future_to_ctx[future].demand_grade_id}:"
+                        f"{future_to_ctx[future].demand_name}"
+                        for future in pending
+                    ]
                     logger.info(
                         "discover videos still running: biz_dt=%s "
-                        "batch_pending=%d processed=%d elapsed_seconds=%d",
+                        "batch_pending=%d processed=%d elapsed_seconds=%d "
+                        "pending_demands=%s",
                         resolved_biz_dt,
                         len(pending),
                         result["processed"],
                         int(monotonic() - batch_started),
+                        pending_demands,
                     )
                     continue
 
                 for future in completed:
+                    ctx = future_to_ctx[future]
                     try:
                         item_result = future.result()
                     except Exception as exc:
                         logger.exception(
-                            "discover videos worker 出现未捕获错误: biz_dt=%s",
+                            "discover videos worker 出现未捕获错误: biz_dt=%s "
+                            "grade_id=%s demand=%s",
                             resolved_biz_dt,
+                            ctx.demand_grade_id,
+                            ctx.demand_name,
                         )
                         result["failed"] += 1
                         result["processed"] += 1
-                        result["errors"].append({"error": str(exc)})
+                        result["errors"].append(
+                            {
+                                "demand_grade_id": ctx.demand_grade_id,
+                                "demand_name": ctx.demand_name,
+                                "error": str(exc),
+                            }
+                        )
+                        logger.info(
+                            "discover videos batch item failed: grade_id=%s "
+                            "demand=%s batch_remaining=%d total_processed=%d",
+                            ctx.demand_grade_id,
+                            ctx.demand_name,
+                            len(pending),
+                            result["processed"],
+                        )
                         continue
 
                     result["processed"] += 1
+                    logger.info(
+                        "discover videos batch item finished: grade_id=%s demand=%s "
+                        "success=%s skipped=%s batch_remaining=%d total_processed=%d",
+                        ctx.demand_grade_id,
+                        ctx.demand_name,
+                        item_result.get("success"),
+                        item_result.get("skipped"),
+                        len(pending),
+                        result["processed"],
+                    )
                     if item_result.get("skipped"):
                         result["skipped"] += 1
                         continue

+ 1 - 1
tests/supply_infra/scheduler/test_discover_videos_from_demands.py

@@ -67,7 +67,7 @@ class _SlowAgent:
     def __init__(self) -> None:
         self.llm = type("LLM", (), {"_async_client": _AsyncClient()})()
 
-    async def arun(self, _user_input: str) -> None:
+    async def arun_core(self, _user_input: str) -> None:
         await asyncio.sleep(60)