返回 VideoClaw
base_agent.py
根目录 / video-claw / video-claw / backend / core / agents / base_agent.py
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
108 lines PYTHON