Przeglądaj źródła

增加同步任务和日志上传

xueyiming 2 tygodni temu
rodzic
commit
dd5aef898a

+ 9 - 1
.env.example

@@ -36,4 +36,12 @@ ODPS_ENDPOINT=https://service.cn.maxcompute.aliyun.com/api
 
 # Scheduler
 SCHEDULER_ENABLED=true
-SCHEDULER_TIMEZONE=Asia/Shanghai
+SCHEDULER_TIMEZONE=Asia/Shanghai
+# Aliyun OSS (agent 运行日志可视化上传)
+ALIYUN_OSS_ACCESS_KEY_ID=
+ALIYUN_OSS_ACCESS_KEY_SECRET=
+ALIYUN_OSS_REGION=cn-hangzhou
+ALIYUN_OSS_BUCKET=art-pubbucket
+ALIYUN_OSS_ROOT_PREFIX=supply_agent
+ALIYUN_OSS_PUBLIC_BASE_URL=http://rescdn.yishihui.com
+LOG_OSS_UPLOAD_ENABLED=true

+ 7 - 3
README.md

@@ -161,12 +161,16 @@ async for event in agent.astream("Your question"):
 
 | 文件 | 说明 |
 |------|------|
-| `run_<id>.log` | 人类可读的完整日志 |
-| `run_<id>.jsonl` | 结构化事件流(推荐用于可视化) |
+| `run_<agent>_<id>.log` | 人类可读的完整日志 |
+| `run_<agent>_<id>.jsonl` | 结构化事件流(推荐用于可视化) |
+
+运行结束后会自动:生成 `.html` 可视化页 → 上传 `.log` / `.jsonl` / `.html` 到 OSS(`supply_agent/<agent_name>/`)→ 写入 MySQL `oss_logs`。
+
+可用 `LOG_OSS_UPLOAD_ENABLED=false` 关闭上传。
 
 事件类型:`run_start` → `llm_input` → `llm_output`(含 reasoning)→ `tool_call`(完整入参/返回)→ … → `run_end`。
 
-生成可视化页面:
+手动生成可视化页面:
 
 ```bash
 # 无参数:为 logs/ 下全部运行生成可视化页面

+ 2 - 1
agents/demand_belong_category_agent/agent.py

@@ -24,9 +24,10 @@ def create_demand_belong_category_agent(
     """创建 find_agent 实例,注册所有相关工具。"""
     agent = Agent(
         settings=settings,
+        name="demand_belong_category_agent",
         model=model,
         system_prompt=DEMAND_BELONG_CATEGORY_AGENT_SYSTEM_PROMPT,
-        max_iterations = 30,
+        max_iterations=30,
     )
 
     # 本 Agent 专属工具

+ 5 - 90
agents/demand_belong_category_agent/run.py

@@ -4,103 +4,18 @@
 from agents.demand_belong_category_agent import create_demand_belong_category_agent
 
 
-def main() -> None:
+def main(word_list: list[str]) -> None:
     agent = create_demand_belong_category_agent()
     print(f"demand_belong_category_agent ready | model={agent.model}")
     print(f"tools: {agent.tools.list_tools()}")
     print()
-    user_input = '''
-    以下是待分类的词语
-    
-时间界定
-暴雨预警
-防汛避险
-明确时间界定
-暴雨预警背景
-发布防汛避险通告
-健康防护
-健康防护知识
-地域风险预警
-广西洪灾
-洪涝灾害救援
-科普
-河南救援队广西洪灾健康防护
-科普洪涝灾害救援健康防护知识
-中毒
-症状
-食品安全
-食用安全警示
-剧烈中毒症状
-凉拌黄瓜食用安全警示
-科普食品安全知识
-全民健身
-寿命数据
-政策
-案例
-河南
-长寿公式
-防大于治
-寿命数据对比
-长寿公式引用
-边角地改造案例
-河南本地数据
-防大于治逻辑
-国家5年健身计划
-解读全民健身政策
-危险场景
-台风天
-危险场景具象化
-原理向警示逻辑转化
-官方荣誉
-紧急救援
-见义勇为
-官方荣誉背书
-紧急救援场景
-曾凡林见义勇为
-报道见义勇为事迹
-主任 
-医院放射科主任
-受贿
-受贿案
-通过APP答题受贿
-18套房产受贿
-揭露医院放射科主任受贿案
-饭局
-饭局现场
-饭局现场叙事
-公共饭局场景
-探讨儿童餐桌礼仪
-央视
-新闻联播
-权威背书
-央视权威背书
-于东来登上新闻联播
-发布会
-宕昌县山体滑坡救援事件
-官方
-权威
-灾害救援
-官方发布会通报
-权威信源引用
-通报灾害救援进展
-防汛安全
-防汛提示
-沈阳非必要不出门防汛提示
-发布防汛安全提示
-社区化场景落地
-解读国务院全民健身政策
-谣言
-辟谣
-事实与谣言对比
-蒲家人中奖谣言辟谣
-辟谣发票中奖谣言
-地缘政治
-政治博弈
-地缘政治绑定
+    words_str = ",".join(word_list)
+    user_input = f'''
+    以下是待分类的词语 : {words_str}
     '''
     result = agent.run(user_input)
     print(result.content)
     print(f"\n[iterations={result.iterations}, tool_calls={result.tool_calls_made}]")
 
 if __name__ == "__main__":
-    main()
+    main(['1','2'])

+ 1 - 0
agents/find_agent/agent.py

@@ -31,6 +31,7 @@ def create_find_agent(
     """创建 find_agent 实例,注册所有相关工具。"""
     agent = Agent(
         settings=settings,
+        name="find_agent",
         model=model,
         system_prompt=FIND_AGENT_SYSTEM_PROMPT,
     )

+ 1 - 0
pyproject.toml

@@ -18,6 +18,7 @@ dependencies = [
     "pymysql>=1.1",
     "apscheduler>=3.10",
     "python-dotenv>=1.0",
+    "oss2>=2.18.0",
 ]
 
 [project.optional-dependencies]

+ 16 - 8
supply_agent/agent/core.py

@@ -6,6 +6,7 @@ from supply_agent.agent.loop import AgentLoop
 from supply_agent.config import Settings, get_settings
 from supply_agent.llm.client import LLMClient
 from supply_agent.logging.logger import AgentLogger
+from supply_agent.logging.publish import publish_run_artifacts
 from supply_agent.skills.registry import SkillRegistry
 from supply_agent.tools.registry import ToolRegistry
 from supply_agent.types import AgentEvent, AgentEventType, AgentResult, Message, Role
@@ -50,6 +51,7 @@ class Agent:
         self,
         settings: Settings | None = None,
         *,
+        name: str | None = None,
         model: str | None = None,
         system_prompt: str | None = None,
         tools: ToolRegistry | None = None,
@@ -59,6 +61,7 @@ class Agent:
         reasoning_effort: str | None = "medium",
         logger: AgentLogger | None = None,
     ) -> None:
+        self.name = name
         self.settings = settings or get_settings()
         self.logger = logger or AgentLogger(
             self.settings.logs_dir,
@@ -132,33 +135,38 @@ class Agent:
             active_skills=self._active_skills,
         )
 
+    def _finish_run(self, result: AgentResult) -> None:
+        """Close run logs and publish visualization artifacts."""
+        self.logger.end_run(result)
+        publish_run_artifacts(self.logger)
+
     def run(self, user_input: str, *, history: list[Message] | None = None) -> AgentResult:
         """Run the agent synchronously with a user message."""
-        self.logger.start_run(user_input, model=self.model)
+        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 = loop.run()
-        self.logger.end_run(result)
+        self._finish_run(result)
         return result
 
     async def arun(
         self, user_input: str, *, history: list[Message] | None = None
     ) -> AgentResult:
         """Run the agent asynchronously."""
-        self.logger.start_run(user_input, model=self.model)
+        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()
-        self.logger.end_run(result)
+        self._finish_run(result)
         return result
 
     def stream(
         self, user_input: str, *, history: list[Message] | None = None
     ) -> Iterator[AgentEvent]:
         """Stream agent events during execution."""
-        self.logger.start_run(user_input, model=self.model)
+        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)
@@ -175,13 +183,13 @@ class Agent:
                 )
             yield event
         if final_result:
-            self.logger.end_run(final_result)
+            self._finish_run(final_result)
 
     async def astream(
         self, user_input: str, *, history: list[Message] | None = None
     ) -> AsyncIterator[AgentEvent]:
         """Async stream agent events."""
-        self.logger.start_run(user_input, model=self.model)
+        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)
@@ -198,7 +206,7 @@ class Agent:
                 )
             yield event
         if final_result:
-            self.logger.end_run(final_result)
+            self._finish_run(final_result)
 
     def reset(self) -> None:
         """Reset active skills state."""

+ 4 - 1
supply_agent/logging/__init__.py

@@ -1,14 +1,17 @@
 """Logging module for SupplyAgent."""
 
-from supply_agent.logging.logger import AgentLogger, get_agent_logger
+from supply_agent.logging.logger import AgentLogger, get_agent_logger, slugify_agent_name
 from supply_agent.logging.parser import load_run_events, summarize_run
+from supply_agent.logging.publish import publish_run_artifacts
 from supply_agent.logging.visualize import generate_visualization, render_html
 
 __all__ = [
     "AgentLogger",
     "get_agent_logger",
+    "slugify_agent_name",
     "load_run_events",
     "summarize_run",
+    "publish_run_artifacts",
     "generate_visualization",
     "render_html",
 ]

+ 35 - 2
supply_agent/logging/logger.py

@@ -2,6 +2,7 @@ from __future__ import annotations
 
 import json
 import logging
+import re
 import uuid
 from datetime import datetime, timezone
 from pathlib import Path
@@ -53,6 +54,19 @@ def _extract_skill_name(skill_name_or_args: str) -> str:
     return skill_name_or_args
 
 
+def slugify_agent_name(name: str | None) -> str | None:
+    """Make agent name safe for use in log filenames / OSS paths."""
+    if not name:
+        return None
+    slug = re.sub(r"[^\w\-]+", "_", name.strip(), flags=re.UNICODE)
+    slug = re.sub(r"_+", "_", slug).strip("_.")
+    return slug[:64] or None
+
+
+# Backwards-compatible alias
+_slugify_agent_name = slugify_agent_name
+
+
 class _FullContentFormatter(logging.Formatter):
     """Formatter that never truncates message content."""
 
@@ -79,11 +93,16 @@ class AgentLogger:
         self._logger: logging.Logger | None = None
         self._seq: int = 0
         self._jsonl_fh: Any | None = None
+        self._agent_name: str | None = None
 
     @property
     def run_id(self) -> str | None:
         return self._run_id
 
+    @property
+    def agent_name(self) -> str | None:
+        return self._agent_name
+
     @property
     def log_file(self) -> Path | None:
         return self._log_file
@@ -92,14 +111,26 @@ class AgentLogger:
     def jsonl_file(self) -> Path | None:
         return self._jsonl_file
 
-    def start_run(self, user_input: str, *, model: str) -> str:
+    def start_run(
+        self,
+        user_input: str,
+        *,
+        model: str,
+        agent_name: str | None = None,
+    ) -> str:
         """Start a new run log file. Returns the run id."""
         if not self.enabled:
             return ""
 
         self.logs_dir.mkdir(parents=True, exist_ok=True)
         timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
-        self._run_id = f"{timestamp}_{uuid.uuid4().hex[:8]}"
+        short_id = uuid.uuid4().hex[:8]
+        slug = _slugify_agent_name(agent_name)
+        self._agent_name = agent_name
+        if slug:
+            self._run_id = f"{slug}_{timestamp}_{short_id}"
+        else:
+            self._run_id = f"{timestamp}_{short_id}"
         self._log_file = self.logs_dir / f"run_{self._run_id}.log"
         self._jsonl_file = self.logs_dir / f"run_{self._run_id}.jsonl"
         self._seq = 0
@@ -125,6 +156,7 @@ class AgentLogger:
             "run_start",
             {
                 "run_id": self._run_id,
+                "agent_name": agent_name,
                 "model": model,
                 "user_input": user_input,
                 "log_file": str(self._log_file),
@@ -256,6 +288,7 @@ class AgentLogger:
             "run_end",
             {
                 "run_id": self._run_id,
+                "agent_name": self._agent_name,
                 "iterations": result.iterations,
                 "tool_calls_made": result.tool_calls_made,
                 "skills_used": result.skills_used,

+ 1 - 0
supply_agent/logging/parser.py

@@ -183,6 +183,7 @@ def summarize_run(events: list[dict[str, Any]]) -> dict[str, Any]:
 
     return {
         "run_id": start_data.get("run_id") or end_data.get("run_id") or (start or {}).get("run_id"),
+        "agent_name": start_data.get("agent_name") or end_data.get("agent_name"),
         "model": start_data.get("model"),
         "user_input": start_data.get("user_input"),
         "iterations": end_data.get("iterations") or (max(iterations) if iterations else 0),

+ 77 - 0
supply_agent/logging/publish.py

@@ -0,0 +1,77 @@
+"""Publish agent run logs: HTML visualize → OSS upload → MySQL record."""
+
+from __future__ import annotations
+
+import logging
+from pathlib import Path
+from typing import TYPE_CHECKING
+
+from supply_agent.logging.logger import slugify_agent_name
+from supply_agent.logging.visualize import generate_visualization
+from supply_infra.config import get_infra_settings
+from supply_infra.db.repositories.oss_log_repo import OssLogRepository
+from supply_infra.db.session import get_session
+from supply_infra.oss.client import OssClient
+
+if TYPE_CHECKING:
+    from supply_agent.logging.logger import AgentLogger
+
+_log = logging.getLogger("supply_agent.publish")
+
+
+def publish_run_artifacts(agent_logger: AgentLogger) -> str | None:
+    """
+    After a run finishes: generate HTML, upload .log/.jsonl/.html to OSS,
+    and insert a row into ``oss_logs``.
+
+    Returns the public HTML URL, or None if skipped / failed.
+    Failures are logged and do not raise (so agent runs are not broken).
+    """
+    settings = get_infra_settings()
+    if not settings.log_oss_upload_enabled:
+        return None
+    if not settings.aliyun_oss_access_key_id or not settings.aliyun_oss_access_key_secret:
+        _log.debug("OSS credentials missing; skip log publish")
+        return None
+
+    log_file = agent_logger.log_file
+    jsonl_file = agent_logger.jsonl_file
+    if not log_file or not log_file.exists():
+        _log.warning("No log file to publish")
+        return None
+
+    try:
+        source = jsonl_file if jsonl_file and jsonl_file.exists() else log_file
+        html_path = generate_visualization(source)
+
+        agent_slug = slugify_agent_name(agent_logger.agent_name) or "agent"
+        log_name = html_path.stem
+
+        client = OssClient(settings)
+        files = [p for p in (log_file, jsonl_file, html_path) if p and Path(p).exists()]
+        html_url: str | None = None
+
+        for path in files:
+            key = client.object_key(agent_slug, Path(path).name)
+            url = client.upload_file(path, key)
+            if Path(path).suffix.lower() == ".html":
+                html_url = url
+            _log.info("Uploaded %s → %s", path.name, url)
+
+        if not html_url:
+            raise RuntimeError("HTML was not uploaded")
+
+        with get_session() as session:
+            OssLogRepository(session).create(
+                log_name=log_name,
+                agent_name=agent_logger.agent_name or agent_slug,
+                oss_path=html_url,
+            )
+
+        _log.info("oss_logs saved: log_name=%s url=%s", log_name, html_url)
+        print(f"[publish] 可视化已上传: {html_url}")
+        return html_url
+    except Exception:
+        _log.exception("Failed to publish run artifacts")
+        print("[publish] 日志上传/入库失败,详见日志;不影响本次 Agent 结果")
+        return None

+ 54 - 31
supply_agent/logging/visualize.py

@@ -28,9 +28,10 @@ def _pretty(obj: Any) -> str:
 
 
 def _pre(text: Any, *, klass: str = "code") -> str:
-    content = _pretty(text) if not isinstance(text, str) else text
-    if content is None:
+    if text is None:
         content = ""
+    else:
+        content = _pretty(text)
     return f'<pre class="{klass}">{_esc(content)}</pre>'
 
 
@@ -116,8 +117,8 @@ def _render_messages(
             body_parts.append(_pre(msg["content"], klass="code prose"))
         if msg.get("tool_calls"):
             body_parts.append(
-                '<div class="sublabel">tool_calls</div>'
-                + _pre(msg["tool_calls"])
+                '<div class="sublabel">决定调用的工具</div>'
+                + _render_tool_call_cards(msg["tool_calls"])
             )
         if msg.get("tool_call_id"):
             body_parts.append(
@@ -127,15 +128,24 @@ def _render_messages(
             body_parts.append(
                 f'<div class="meta-line">name: <code>{_esc(msg["name"])}</code></div>'
             )
-        if default_open is None:
+        # Tool result content: try pretty JSON (already handled by _pre via content)        if default_open is None:
             open_attr = "open" if role != "system" else ""
         else:
             open_attr = "open" if default_open else ""
+        preview = _preview(msg.get("content"))
+        if not preview and msg.get("tool_calls"):
+            names = []
+            for tc in msg["tool_calls"]:
+                name, _, _ = _normalize_tool_call(tc)
+                names.append(name)
+            preview = "调用: " + ", ".join(names)
+        elif not preview and msg.get("name"):
+            preview = f"tool → {msg['name']}"
         parts.append(
             f"""
             <details class="msg" {open_attr}>
               <summary>{_role_badge(role)} <span class="msg-idx">#{index_offset + i + 1}</span>
-                <span class="msg-preview">{_esc(_preview(msg.get("content")))}</span>
+                <span class="msg-preview">{_esc(preview)}</span>
               </summary>
               <div class="msg-body">{"".join(body_parts) or '<p class="muted">(empty)</p>'}</div>
             </details>
@@ -217,12 +227,7 @@ def _render_llm_input(
 
 def _render_reasoning(reasoning: Any) -> str:
     if not reasoning:
-        return """
-        <div class="reasoning empty">
-          <div class="sublabel">思考过程</div>
-          <p class="muted">本次回复未返回 reasoning 字段</p>
-        </div>
-        """
+        return ""
     return f"""
     <div class="reasoning">
       <div class="sublabel">思考过程</div>
@@ -240,19 +245,23 @@ def _parse_tool_arguments(arguments: Any) -> Any:
     return arguments
 
 
-def _render_planned_tool_calls(tool_calls: list[dict[str, Any]]) -> str:
-    """Show each planned tool as name + args; keep raw JSON in a collapsible."""
+def _normalize_tool_call(tc: dict[str, Any]) -> tuple[str, Any, str]:
+    """Return (name, parsed_args, id) from flat or OpenAI function-wrapped tool_call."""
+    name = tc.get("name") or (tc.get("function") or {}).get("name") or "?"
+    raw_args = tc.get("arguments")
+    if raw_args is None and isinstance(tc.get("function"), dict):
+        raw_args = tc["function"].get("arguments")
+    return name, _parse_tool_arguments(raw_args), str(tc.get("id") or "")
+
+
+def _render_tool_call_cards(tool_calls: list[dict[str, Any]]) -> str:
+    """Render tool calls as name + args cards; raw JSON in a collapsible."""
     if not tool_calls:
         return ""
 
     cards: list[str] = []
     for i, tc in enumerate(tool_calls, 1):
-        name = tc.get("name") or (tc.get("function") or {}).get("name") or "?"
-        raw_args = tc.get("arguments")
-        if raw_args is None and isinstance(tc.get("function"), dict):
-            raw_args = tc["function"].get("arguments")
-        args = _parse_tool_arguments(raw_args)
-        tc_id = tc.get("id") or ""
+        name, args, tc_id = _normalize_tool_call(tc)
         cards.append(
             f"""
             <div class="planned-tool">
@@ -268,7 +277,6 @@ def _render_planned_tool_calls(tool_calls: list[dict[str, Any]]) -> str:
         )
 
     return f"""
-    <h4>模型决定调用的工具</h4>
     <div class="planned-tool-list">{"".join(cards)}</div>
     <details class="raw-block">
       <summary>查看原始 tool_calls JSON</summary>
@@ -277,9 +285,22 @@ def _render_planned_tool_calls(tool_calls: list[dict[str, Any]]) -> str:
     """
 
 
+def _render_planned_tool_calls(tool_calls: list[dict[str, Any]]) -> str:
+    """Show each planned tool as name + args; omit entirely when empty."""
+    if not tool_calls:
+        return ""
+    return f"""
+    <h4>模型决定调用的工具</h4>
+    {_render_tool_call_cards(tool_calls)}
+    """
+
+
 def _render_llm_output(data: dict[str, Any], seq: int, iteration: Any) -> str:
     tool_calls = data.get("tool_calls") or []
+    reasoning = data.get("reasoning")
+    content = data.get("content")
     usage = data.get("usage")
+
     usage_html = ""
     if usage:
         usage_html = f"""
@@ -291,14 +312,16 @@ def _render_llm_output(data: dict[str, Any], seq: int, iteration: Any) -> str:
         </div>
         """
 
-    tool_calls_html = _render_planned_tool_calls(tool_calls)
-    content = data.get("content")
     content_html = (
-        f"<h4>模型输出文本</h4>{_pre(content, klass='code prose')}"
-        if content
-        else '<h4>模型输出文本</h4><p class="muted">(空,可能仅有 tool_calls)</p>'
+        f"<h4>模型输出文本</h4>{_pre(content, klass='code prose')}" if content else ""
     )
 
+    tags: list[str] = [f'<span class="tag">iteration {iteration}</span>']
+    if reasoning:
+        tags.append('<span class="tag">有思考</span>')
+    if tool_calls:
+        tags.append(f'<span class="tag">{len(tool_calls)} tool calls</span>')
+
     raw_html = ""
     if data.get("raw"):
         raw_html = f"""
@@ -314,15 +337,13 @@ def _render_llm_output(data: dict[str, Any], seq: int, iteration: Any) -> str:
         <div class="step-num">Step {seq}</div>
         <div class="card-title">LLM 输出</div>
         <div class="card-tags">
-          <span class="tag">iteration {iteration}</span>
-          <span class="tag">{'有思考' if data.get('reasoning') or data.get('has_reasoning') else '无思考'}</span>
-          <span class="tag">{len(tool_calls)} tool calls</span>
+          {"".join(tags)}
         </div>
       </header>
       <div class="card-body">
-        {_render_reasoning(data.get("reasoning"))}
+        {_render_reasoning(reasoning)}
         {content_html}
-        {tool_calls_html}
+        {_render_planned_tool_calls(tool_calls)}
         {usage_html}
         {raw_html}
       </div>
@@ -793,6 +814,7 @@ def render_html(events: list[dict[str, Any]]) -> str:
   <div class="layout">
     <aside class="sidebar">
       <h1>Agent Trace</h1>
+      <div class="run-id">{_esc(meta.get("agent_name") or "agent")}</div>
       <div class="run-id">{_esc(meta.get("run_id"))}</div>
       <nav>
         {_nav_items(events)}
@@ -802,6 +824,7 @@ def render_html(events: list[dict[str, Any]]) -> str:
       <section class="hero">
         <h2>运行概览</h2>
         <div class="stats">
+          <div class="stat"><div class="label">Agent</div><div class="value">{_esc(meta.get("agent_name") or "—")}</div></div>
           <div class="stat"><div class="label">Model</div><div class="value">{_esc(meta.get("model") or "—")}</div></div>
           <div class="stat"><div class="label">Iterations</div><div class="value">{_esc(meta.get("iterations"))}</div></div>
           <div class="stat"><div class="label">Tool Calls</div><div class="value">{_esc(meta.get("tool_calls_made"))}</div></div>

+ 12 - 0
supply_infra/config.py

@@ -43,6 +43,18 @@ class InfraSettings(BaseSettings):
     scheduler_timezone: str = Field(default="Asia/Shanghai", alias="SCHEDULER_TIMEZONE")
     scheduler_enabled: bool = Field(default=True, alias="SCHEDULER_ENABLED")
 
+    # Aliyun OSS (agent run log publishing)
+    aliyun_oss_access_key_id: str = Field(default="", alias="ALIYUN_OSS_ACCESS_KEY_ID")
+    aliyun_oss_access_key_secret: str = Field(default="", alias="ALIYUN_OSS_ACCESS_KEY_SECRET")
+    aliyun_oss_region: str = Field(default="cn-hangzhou", alias="ALIYUN_OSS_REGION")
+    aliyun_oss_bucket: str = Field(default="art-pubbucket", alias="ALIYUN_OSS_BUCKET")
+    aliyun_oss_root_prefix: str = Field(default="supply_agent", alias="ALIYUN_OSS_ROOT_PREFIX")
+    aliyun_oss_public_base_url: str = Field(
+        default="http://rescdn.yishihui.com",
+        alias="ALIYUN_OSS_PUBLIC_BASE_URL",
+    )
+    log_oss_upload_enabled: bool = Field(default=True, alias="LOG_OSS_UPLOAD_ENABLED")
+
     @property
     def mysql_url(self) -> str:
         user = quote_plus(self.mysql_user)

+ 9 - 1
supply_infra/db/models/__init__.py

@@ -3,5 +3,13 @@
 from supply_infra.db.models.demand_belong_category import DemandBelongCategory
 from supply_infra.db.models.global_tree_category import GlobalTreeCategory
 from supply_infra.db.models.global_tree_element import GlobalTreeElement
+from supply_infra.db.models.multi_demand_pool_di import MultiDemandPoolDi
+from supply_infra.db.models.oss_log import OssLog
 
-__all__ = ["DemandBelongCategory", "GlobalTreeCategory", "GlobalTreeElement"]
+__all__ = [
+    "DemandBelongCategory",
+    "GlobalTreeCategory",
+    "GlobalTreeElement",
+    "MultiDemandPoolDi",
+    "OssLog",
+]

+ 40 - 0
supply_infra/db/models/multi_demand_pool_di.py

@@ -0,0 +1,40 @@
+from __future__ import annotations
+
+from datetime import datetime
+
+from sqlalchemy import BigInteger, Float, Index, String, Text, func
+from sqlalchemy.orm import Mapped, mapped_column
+
+from supply_infra.db.base import Base
+
+
+class MultiDemandPoolDi(Base):
+    """策略需求天级表 — 从 ODPS dwd_multi_demand_pool_di 同步。"""
+
+    __tablename__ = "multi_demand_pool_di"
+    __table_args__ = (
+        Index("idx_biz_dt", "biz_dt"),
+        Index("idx_strategy_demand", "strategy", "demand_id"),
+    )
+
+    id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
+    strategy: Mapped[str] = mapped_column(String(128), nullable=False, comment="策略")
+    demand_id: Mapped[str] = mapped_column(String(128), nullable=False, comment="需求id")
+    demand_name: Mapped[str] = mapped_column(String(256), nullable=False, comment="需求名称")
+    weight: Mapped[float | None] = mapped_column(Float, nullable=True, comment="权重")
+    type: Mapped[str | None] = mapped_column(String(64), nullable=True, comment="需求类型")
+    video_count: Mapped[int | None] = mapped_column(BigInteger, nullable=True, comment="视频数量")
+    video_list: Mapped[str | None] = mapped_column(Text, nullable=True, comment="视频列表")
+    extend: Mapped[str | None] = mapped_column(Text, nullable=True, comment="拓展字段")
+    biz_dt: Mapped[str] = mapped_column(String(32), nullable=False, comment="业务日期")
+    create_time: Mapped[datetime] = mapped_column(
+        nullable=False,
+        server_default=func.now(),
+        comment="创建时间",
+    )
+    update_time: Mapped[datetime] = mapped_column(
+        nullable=False,
+        server_default=func.now(),
+        onupdate=func.now(),
+        comment="更新时间",
+    )

+ 34 - 0
supply_infra/db/models/oss_log.py

@@ -0,0 +1,34 @@
+from __future__ import annotations
+
+from datetime import datetime
+
+from sqlalchemy import BigInteger, Integer, String, func
+from sqlalchemy.orm import Mapped, mapped_column
+
+from supply_infra.db.base import Base
+
+
+class OssLog(Base):
+    """Agent 运行日志在 OSS 上的记录。"""
+
+    __tablename__ = "oss_logs"
+
+    id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
+    log_name: Mapped[str | None] = mapped_column(String(256), nullable=True, comment="日志名称")
+    # DDL 示例里误标为 int;按业务存 agent 名称字符串
+    agent_name: Mapped[str | None] = mapped_column(String(256), nullable=True, comment="agent名称")
+    oss_path: Mapped[str | None] = mapped_column(String(512), nullable=True, comment="oss路径")
+    is_delete: Mapped[int] = mapped_column(
+        Integer, default=0, nullable=False, comment="是否删除 0-正常 1-删除"
+    )
+    create_time: Mapped[datetime] = mapped_column(
+        nullable=False,
+        server_default=func.now(),
+        comment="创建时间",
+    )
+    update_time: Mapped[datetime] = mapped_column(
+        nullable=False,
+        server_default=func.now(),
+        onupdate=func.now(),
+        comment="更新时间",
+    )

+ 4 - 0
supply_infra/db/repositories/__init__.py

@@ -6,10 +6,14 @@ from supply_infra.db.repositories.demand_belong_category_repo import (
 )
 from supply_infra.db.repositories.global_tree_category_repo import GlobalTreeCategoryRepository
 from supply_infra.db.repositories.global_tree_element_repo import GlobalTreeElementRepository
+from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPoolDiRepository
+from supply_infra.db.repositories.oss_log_repo import OssLogRepository
 
 __all__ = [
     "BaseRepository",
     "DemandBelongCategoryRepository",
     "GlobalTreeCategoryRepository",
     "GlobalTreeElementRepository",
+    "MultiDemandPoolDiRepository",
+    "OssLogRepository",
 ]

+ 18 - 0
supply_infra/db/repositories/demand_belong_category_repo.py

@@ -1,5 +1,8 @@
 from __future__ import annotations
 
+from collections.abc import Iterable
+
+from sqlalchemy import select
 from sqlalchemy.dialects.mysql import insert
 
 from supply_infra.db.models.demand_belong_category import DemandBelongCategory
@@ -13,6 +16,21 @@ class DemandBelongCategoryRepository(BaseRepository[DemandBelongCategory]):
 
     model = DemandBelongCategory
 
+    def get_existing_names(self, names: Iterable[str]) -> set[str]:
+        """返回 names 中已存在于表内的名称(含软删除)。"""
+        name_list = [n for n in names if n]
+        if not name_list:
+            return set()
+
+        existing: set[str] = set()
+        for i in range(0, len(name_list), _BATCH_SIZE):
+            batch = name_list[i : i + _BATCH_SIZE]
+            stmt = select(DemandBelongCategory.name).where(
+                DemandBelongCategory.name.in_(batch)
+            )
+            existing.update(n for n in self.session.scalars(stmt).all() if n)
+        return existing
+
     def bulk_insert_ignore(self, rows: list[dict]) -> int:
         """批量插入,MySQL 按 name 唯一索引忽略已存在行。"""
         if not rows:

+ 79 - 0
supply_infra/db/repositories/multi_demand_pool_di_repo.py

@@ -0,0 +1,79 @@
+from __future__ import annotations
+
+from sqlalchemy import delete, func, select, tuple_
+from sqlalchemy.dialects.mysql import insert
+
+from supply_infra.db.models.multi_demand_pool_di import MultiDemandPoolDi
+from supply_infra.db.repositories.base import BaseRepository
+
+_BATCH_SIZE = 1000
+
+RowKey = tuple[str, str]
+
+
+class MultiDemandPoolDiRepository(BaseRepository[MultiDemandPoolDi]):
+    """策略需求天级表 repository。"""
+
+    model = MultiDemandPoolDi
+
+    def count_by_biz_dt(self, biz_dt: str) -> int:
+        """统计指定业务日期去重行数(strategy + demand_id)。"""
+        stmt = (
+            select(func.count())
+            .select_from(
+                select(MultiDemandPoolDi.strategy, MultiDemandPoolDi.demand_id)
+                .where(MultiDemandPoolDi.biz_dt == biz_dt)
+                .distinct()
+                .subquery()
+            )
+        )
+        return int(self.session.scalar(stmt) or 0)
+
+    def list_keys_by_biz_dt(self, biz_dt: str) -> set[RowKey]:
+        """查询指定业务日期的 (strategy, demand_id) 键集合。"""
+        stmt = select(MultiDemandPoolDi.strategy, MultiDemandPoolDi.demand_id).where(
+            MultiDemandPoolDi.biz_dt == biz_dt
+        )
+        return {(str(s), str(d)) for s, d in self.session.execute(stmt).all()}
+
+    def delete_by_biz_dt(self, biz_dt: str) -> int:
+        """删除指定业务日期的全部数据,便于按日重跑。"""
+        stmt = delete(MultiDemandPoolDi).where(MultiDemandPoolDi.biz_dt == biz_dt)
+        result = self.session.execute(stmt)
+        return result.rowcount or 0
+
+    def delete_by_keys(self, biz_dt: str, keys: list[RowKey]) -> int:
+        """按 (strategy, demand_id) 批量删除指定业务日期数据。"""
+        if not keys:
+            return 0
+
+        deleted = 0
+        for i in range(0, len(keys), _BATCH_SIZE):
+            batch = keys[i : i + _BATCH_SIZE]
+            stmt = delete(MultiDemandPoolDi).where(
+                MultiDemandPoolDi.biz_dt == biz_dt,
+                tuple_(MultiDemandPoolDi.strategy, MultiDemandPoolDi.demand_id).in_(batch),
+            )
+            result = self.session.execute(stmt)
+            deleted += result.rowcount or 0
+        return deleted
+
+    def list_demand_names_by_biz_dt(self, biz_dt: str) -> list[str]:
+        """查询指定业务日期的全部 demand_name。"""
+        stmt = select(MultiDemandPoolDi.demand_name).where(
+            MultiDemandPoolDi.biz_dt == biz_dt
+        )
+        return [name for name in self.session.scalars(stmt).all() if name]
+
+    def bulk_insert(self, rows: list[dict]) -> int:
+        """批量插入。"""
+        if not rows:
+            return 0
+
+        inserted = 0
+        for i in range(0, len(rows), _BATCH_SIZE):
+            batch = rows[i : i + _BATCH_SIZE]
+            stmt = insert(MultiDemandPoolDi).values(batch)
+            result = self.session.execute(stmt)
+            inserted += result.rowcount
+        return inserted

+ 23 - 0
supply_infra/db/repositories/oss_log_repo.py

@@ -0,0 +1,23 @@
+from __future__ import annotations
+
+from supply_infra.db.models.oss_log import OssLog
+from supply_infra.db.repositories.base import BaseRepository
+
+
+class OssLogRepository(BaseRepository[OssLog]):
+    model = OssLog
+
+    def create(
+        self,
+        *,
+        log_name: str,
+        agent_name: str | None,
+        oss_path: str,
+    ) -> OssLog:
+        entity = OssLog(
+            log_name=log_name,
+            agent_name=agent_name,
+            oss_path=oss_path,
+            is_delete=0,
+        )
+        return self.add(entity)

+ 33 - 0
supply_infra/odps/client.py

@@ -84,6 +84,39 @@ class ODPSClient:
         """
         return self.execute_sql(sql)
 
+    def fetch_multi_demand_pool(self, bizdate: str) -> list[dict[str, Any]]:
+        """拉取 dwd_multi_demand_pool_di 策略需求天级数据(不同步 video_list)。"""
+        sql = f"""
+        SELECT  strategy
+                ,demand_id
+                ,demand_name
+                ,weight
+                ,`type`
+                ,video_count
+                ,extend
+        FROM    loghubods.dwd_multi_demand_pool_di
+        WHERE   dt = '{bizdate}'
+        """
+        return self.execute_sql(sql)
+
+    def count_multi_demand_pool(self, bizdate: str) -> int:
+        """统计 dwd_multi_demand_pool_di 分区去重后行数(strategy + demand_id)。"""
+        sql = f"""
+        SELECT  COUNT(1) AS cnt
+        FROM    (
+                    SELECT  strategy
+                            ,demand_id
+                    FROM    loghubods.dwd_multi_demand_pool_di
+                    WHERE   dt = '{bizdate}'
+                    GROUP BY strategy
+                             ,demand_id
+                ) t
+        """
+        rows = self.execute_sql(sql)
+        if not rows:
+            return 0
+        return int(rows[0].get("cnt") or 0)
+
 
 @lru_cache
 def get_odps_client() -> ODPSClient:

+ 5 - 0
supply_infra/oss/__init__.py

@@ -0,0 +1,5 @@
+"""Aliyun OSS helpers."""
+
+from supply_infra.oss.client import OssClient
+
+__all__ = ["OssClient"]

+ 49 - 0
supply_infra/oss/client.py

@@ -0,0 +1,49 @@
+"""Aliyun OSS client for uploading agent run artifacts."""
+
+from __future__ import annotations
+
+from pathlib import Path
+
+import oss2
+
+from supply_infra.config import InfraSettings, get_infra_settings
+
+
+class OssClient:
+    """Thin wrapper around oss2 Bucket."""
+
+    def __init__(self, settings: InfraSettings | None = None) -> None:
+        self.settings = settings or get_infra_settings()
+        if not self.settings.aliyun_oss_access_key_id:
+            raise ValueError("ALIYUN_OSS_ACCESS_KEY_ID is not configured")
+        if not self.settings.aliyun_oss_access_key_secret:
+            raise ValueError("ALIYUN_OSS_ACCESS_KEY_SECRET is not configured")
+
+        endpoint = f"https://oss-{self.settings.aliyun_oss_region}.aliyuncs.com"
+        auth = oss2.Auth(
+            self.settings.aliyun_oss_access_key_id,
+            self.settings.aliyun_oss_access_key_secret,
+        )
+        self._bucket = oss2.Bucket(auth, endpoint, self.settings.aliyun_oss_bucket)
+        self._root = self.settings.aliyun_oss_root_prefix.strip("/")
+        self._public_base = self.settings.aliyun_oss_public_base_url.rstrip("/")
+
+    def object_key(self, *parts: str) -> str:
+        cleaned = [self._root] if self._root else []
+        for part in parts:
+            p = str(part).strip("/")
+            if p:
+                cleaned.append(p)
+        return "/".join(cleaned)
+
+    def public_url(self, object_key: str) -> str:
+        key = object_key.lstrip("/")
+        return f"{self._public_base}/{key}"
+
+    def upload_file(self, local_path: Path | str, object_key: str) -> str:
+        """Upload a local file and return its public CDN URL."""
+        path = Path(local_path)
+        if not path.is_file():
+            raise FileNotFoundError(f"File not found: {path}")
+        self._bucket.put_object_from_file(object_key, str(path))
+        return self.public_url(object_key)

+ 12 - 0
supply_infra/scheduler/app.py

@@ -7,6 +7,9 @@ from apscheduler.triggers.cron import CronTrigger
 
 from supply_infra.config import get_infra_settings
 from supply_infra.scheduler.jobs.sync_global_tree_odps_to_mysql import sync_global_tree_odps_to_mysql
+from supply_infra.scheduler.jobs.sync_multi_demand_pool_odps_to_mysql import (
+    sync_multi_demand_pool_odps_to_mysql,
+)
 
 logger = logging.getLogger(__name__)
 
@@ -25,6 +28,15 @@ def create_scheduler() -> BlockingScheduler:
         replace_existing=True,
     )
 
+    # 每天 12:00 从 ODPS 同步策略需求天级表到 MySQL(当天 dt)
+    scheduler.add_job(
+        sync_multi_demand_pool_odps_to_mysql,
+        trigger=CronTrigger(hour=12, minute=0),
+        id="sync_multi_demand_pool_odps_to_mysql",
+        name="ODPS → MySQL 策略需求池同步",
+        replace_existing=True,
+    )
+
     logger.info("Scheduler configured with %d job(s)", len(scheduler.get_jobs()))
     return scheduler
 

+ 211 - 0
supply_infra/scheduler/jobs/sync_multi_demand_pool_odps_to_mysql.py

@@ -0,0 +1,211 @@
+"""
+定时任务:从 ODPS 同步策略需求天级表到 MySQL,并对新词做归属分类。
+
+流程:
+1. 比对当天 ODPS / MySQL 行数,相同则跳过写入
+2. 有差异时拉取 ODPS,按 (strategy, demand_id) 只插入缺失、删除多余
+3. 查询当天 demand_name,按空格分词并去重
+4. 过滤 demand_belong_category 中已存在的词
+5. 剩余词按 100 个一批,调用 demand_belong_category_agent
+   (测试阶段仅调用 1 批)
+"""
+from __future__ import annotations
+
+import logging
+from datetime import datetime
+from typing import Any
+
+from agents.demand_belong_category_agent.run import main as classify_demand_words
+from supply_infra.db.repositories.demand_belong_category_repo import (
+    DemandBelongCategoryRepository,
+)
+from supply_infra.db.repositories.multi_demand_pool_di_repo import MultiDemandPoolDiRepository
+from supply_infra.db.session import get_session
+from supply_infra.odps.client import get_odps_client
+
+logger = logging.getLogger(__name__)
+
+_WORD_BATCH_SIZE = 100
+# 测试阶段只跑 1 批;上线后改为 None 表示跑完全部
+_MAX_CLASSIFY_BATCHES = 1
+
+RowKey = tuple[str, str]
+
+
+def _row_key(row: dict[str, Any]) -> RowKey:
+    return (row["strategy"], row["demand_id"])
+
+
+def _to_mysql_rows(raw_rows: list[dict[str, Any]], biz_dt: str) -> list[dict[str, Any]]:
+    """转换为 MySQL 行,按 (strategy, demand_id) 去重(保留最后一条)。"""
+    by_key: dict[RowKey, dict[str, Any]] = {}
+    for row in raw_rows:
+        strategy = row.get("strategy")
+        demand_id = row.get("demand_id")
+        demand_name = row.get("demand_name")
+        if strategy is None or demand_id is None or demand_name is None:
+            logger.warning("Skip row with missing required fields: %s", row)
+            continue
+
+        mapped = {
+            "strategy": str(strategy),
+            "demand_id": str(demand_id),
+            "demand_name": str(demand_name),
+            "weight": row.get("weight"),
+            "type": str(row["type"]) if row.get("type") is not None else None,
+            "video_count": row.get("video_count"),
+            "video_list": None,
+            "extend": str(row["extend"]) if row.get("extend") is not None else None,
+            "biz_dt": biz_dt,
+        }
+        by_key[_row_key(mapped)] = mapped
+    return list(by_key.values())
+
+
+def _tokenize_demand_names(demand_names: list[str]) -> list[str]:
+    """按空格分词,去空、去重(保持首次出现顺序)。"""
+    seen: set[str] = set()
+    tokens: list[str] = []
+    for name in demand_names:
+        for token in str(name).split():
+            word = token.strip()
+            if not word or word in seen:
+                continue
+            seen.add(word)
+            tokens.append(word)
+    return tokens
+
+
+def _chunked(items: list[str], size: int) -> list[list[str]]:
+    return [items[i : i + size] for i in range(0, len(items), size)]
+
+
+def _classify_new_words(biz_dt: str, *, max_batches: int | None = _MAX_CLASSIFY_BATCHES) -> dict:
+    """分词 → 过滤已存在词 → 分批调用归属分类 agent。"""
+    with get_session() as session:
+        pool_repo = MultiDemandPoolDiRepository(session)
+        category_repo = DemandBelongCategoryRepository(session)
+
+        demand_names = pool_repo.list_demand_names_by_biz_dt(biz_dt)
+        tokens = _tokenize_demand_names(demand_names)
+        existing = category_repo.get_existing_names(tokens)
+        new_words = [w for w in tokens if w not in existing]
+
+    batches = _chunked(new_words, _WORD_BATCH_SIZE)
+    if max_batches is not None:
+        batches = batches[:max_batches]
+
+    logger.info(
+        "Classify prepare: demand_names=%d tokens=%d existing=%d new=%d batches=%d (max=%s)",
+        len(demand_names),
+        len(tokens),
+        len(existing),
+        len(new_words),
+        len(batches),
+        max_batches,
+    )
+
+    classified_batches = 0
+    for idx, batch in enumerate(batches, start=1):
+        logger.info("Classifying batch %d/%d (%d words)", idx, len(batches), len(batch))
+        classify_demand_words(batch)
+        classified_batches += 1
+
+    return {
+        "demand_names": len(demand_names),
+        "tokens": len(tokens),
+        "existing_filtered": len(existing),
+        "new_words": len(new_words),
+        "batches_total": (len(new_words) + _WORD_BATCH_SIZE - 1) // _WORD_BATCH_SIZE
+        if new_words
+        else 0,
+        "batches_ran": classified_batches,
+    }
+
+
+def _sync_diff(partition_date: str) -> dict[str, Any]:
+    """行数不同时拉取 ODPS,只同步 (strategy, demand_id) 差异。"""
+    odps = get_odps_client()
+    raw_rows = odps.fetch_multi_demand_pool(partition_date)
+    mysql_rows = _to_mysql_rows(raw_rows, partition_date)
+    odps_keys = {_row_key(r) for r in mysql_rows}
+    odps_by_key = {_row_key(r): r for r in mysql_rows}
+
+    with get_session() as session:
+        repo = MultiDemandPoolDiRepository(session)
+        mysql_keys = repo.list_keys_by_biz_dt(partition_date)
+
+        to_insert_keys = odps_keys - mysql_keys
+        to_delete_keys = mysql_keys - odps_keys
+
+        insert_rows = [odps_by_key[k] for k in to_insert_keys]
+        deleted = repo.delete_by_keys(partition_date, list(to_delete_keys))
+        inserted = repo.bulk_insert(insert_rows)
+
+    logger.info(
+        "Diff sync: odps=%d mysql_before=%d insert=%d delete=%d",
+        len(odps_keys),
+        len(mysql_keys),
+        inserted,
+        deleted,
+    )
+    return {
+        "fetched": len(raw_rows),
+        "odps_unique": len(odps_keys),
+        "mysql_before": len(mysql_keys),
+        "inserted": inserted,
+        "deleted": deleted,
+        "skipped_invalid": len(raw_rows) - len(mysql_rows),
+    }
+
+
+def sync_multi_demand_pool_odps_to_mysql(partition_date: str | None = None) -> dict:
+    """
+    从 ODPS 增量同步策略需求天级数据到 MySQL,并对新词做归属分类。
+
+    Args:
+        partition_date: 分区日期 (YYYYMMDD),默认当天
+    """
+    if partition_date is None:
+        partition_date = datetime.now().strftime("%Y%m%d")
+
+    logger.info(
+        "Starting multi demand pool ODPS → MySQL sync for partition: %s",
+        partition_date,
+    )
+
+    odps = get_odps_client()
+    odps_count = odps.count_multi_demand_pool(partition_date)
+    with get_session() as session:
+        mysql_count = MultiDemandPoolDiRepository(session).count_by_biz_dt(partition_date)
+
+    logger.info("Count check: odps=%d mysql=%d", odps_count, mysql_count)
+
+    if odps_count == mysql_count:
+        sync_stats: dict[str, Any] = {
+            "skipped_same_count": True,
+            "odps_count": odps_count,
+            "mysql_count": mysql_count,
+            "fetched": 0,
+            "inserted": 0,
+            "deleted": 0,
+        }
+        logger.info("Same count, skip ODPS data sync")
+    else:
+        sync_stats = {
+            "skipped_same_count": False,
+            "odps_count": odps_count,
+            "mysql_count": mysql_count,
+            **_sync_diff(partition_date),
+        }
+
+    classify_stats = _classify_new_words(partition_date)
+
+    result = {
+        "partition_date": partition_date,
+        **sync_stats,
+        "classify": classify_stats,
+        "synced_at": datetime.now().isoformat(),
+    }
+    logger.info("Multi demand pool sync completed: %s", result)
+    return result