| 1 | # -*- coding: utf-8 -*- |
| 2 | """ |
| 3 | 智能体基类 - 所有阶段智能体的抽象接口 |
| 4 | """ |
| 5 | |
| 6 | import logging |
| 7 | from abc import ABC, abstractmethod |
| 8 | from typing import Any, Optional, Dict, Callable |
| 9 | |
| 10 | logger = logging.getLogger(__name__) |
| 11 | |
| 12 | SESSION_PARAM_KEYS = [ |
| 13 | "idea", "user_textbox_input", "style", "video_ratio", "video_resolution", |
| 14 | "llm_model", "vlm_model", |
| 15 | "image_t2i_model", "image_it2i_model", "video_model", |
| 16 | "video_first_frame_model", "video_start_end_model", "video_reference_model", |
| 17 | "video_generation_mode", |
| 18 | "video_style", "expand_idea", "enable_concurrency", "web_search", "episodes" |
| 19 | ] |
| 20 | |
| 21 | |
| 22 | class AgentInterface(ABC): |
| 23 | """所有智能体必须实现的接口""" |
| 24 | |
| 25 | def __init__(self, name: str = ""): |
| 26 | self.name = name |
| 27 | self.cancellation_check: Optional[Callable] = None |
| 28 | self.progress_callback: Optional[Callable] = None |
| 29 | |
| 30 | def _merge_session_params(self, input_data: Any) -> Dict: |
| 31 | """从编排器注入的 session 快照补齐缺失参数。""" |
| 32 | if not isinstance(input_data, dict): |
| 33 | return {} |
| 34 | |
| 35 | session_meta = self._session_meta(input_data) |
| 36 | merged_data = input_data.copy() |
| 37 | for key in SESSION_PARAM_KEYS: |
| 38 | if key not in merged_data or not merged_data[key]: |
| 39 | if key in session_meta and session_meta[key] is not None: |
| 40 | merged_data[key] = session_meta[key] |
| 41 | return merged_data |
| 42 | |
| 43 | def _session_meta(self, input_data: Dict) -> Dict: |
| 44 | meta = input_data.get("_session_meta") if isinstance(input_data, dict) else {} |
| 45 | return meta if isinstance(meta, dict) else {} |
| 46 | |
| 47 | def _session_artifacts(self, input_data: Dict) -> Dict: |
| 48 | artifacts = input_data.get("_session_artifacts") if isinstance(input_data, dict) else {} |
| 49 | return artifacts if isinstance(artifacts, dict) else {} |
| 50 | |
| 51 | def _session_artifact(self, input_data: Dict, stage: str) -> Dict: |
| 52 | artifact = self._session_artifacts(input_data).get(stage, {}) |
| 53 | return artifact if isinstance(artifact, dict) else {} |
| 54 | |
| 55 | def set_cancellation_check(self, fn: Callable): |
| 56 | self.cancellation_check = fn |
| 57 | |
| 58 | def set_progress_callback(self, fn: Callable): |
| 59 | self.progress_callback = fn |
| 60 | |
| 61 | def _report_progress(self, phase: str, step_desc: str, percent: float, data: dict = None): |
| 62 | if self.progress_callback: |
| 63 | self.progress_callback(phase, step_desc, percent, data) |
| 64 | |
| 65 | def _check_cancel(self): |
| 66 | if self.cancellation_check and self.cancellation_check(): |
| 67 | raise RuntimeError(f"Agent [{self.name}] cancelled by user") |
| 68 | |
| 69 | def _require_input(self, input_data: Dict, key: str) -> str: |
| 70 | value = input_data.get(key) |
| 71 | if not value: |
| 72 | raise ValueError(f"Missing required model configuration: {key}") |
| 73 | return str(value) |
| 74 | |
| 75 | def _cancellable_query(self, llm, prompt: str, image_urls=[], model="gemini-3-flash-preview", safe_content=True, task_id=None, web_search=False): |
| 76 | """在 LLM 调用前后检查取消状态""" |
| 77 | self._check_cancel() |
| 78 | # 将位置参数映射给 llm.query |
| 79 | result = llm.query(prompt, image_urls, model, safe_content, task_id, web_search) |
| 80 | self._check_cancel() |
| 81 | return result |
| 82 | |
| 83 | def _get_style_prompt(self, style_name: str) -> str: |
| 84 | """从 prompts/style/{style_name}.txt 读取对应的视觉提示词""" |
| 85 | import os |
| 86 | style_file = os.path.join('prompts', 'style', f"{style_name}.txt") |
| 87 | if os.path.exists(style_file): |
| 88 | with open(style_file, 'r', encoding='utf-8') as f: |
| 89 | return f.read().strip() |
| 90 | # Fallback to English style name if file doesn't exist |
| 91 | return style_name + " style" |
| 92 | |
| 93 | # -------- 抽象方法 -------- |
| 94 | |
| 95 | @abstractmethod |
| 96 | async def process(self, input_data: Any, intervention: Optional[Dict] = None) -> Dict: |
| 97 | """ |
| 98 | 核心处理逻辑 |
| 99 | |
| 100 | Args: |
| 101 | input_data: 来自上一阶段的输入数据 |
| 102 | intervention: 用户介入修改内容 |
| 103 | |
| 104 | Returns: |
| 105 | dict: { "payload": ..., "requires_intervention": bool } |
| 106 | """ |
| 107 | pass |
| 108 |