File size: 10,853 Bytes
d82bbe4 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 | """
轨迹构建器 - 将收集的原始数据转换为标准 TRJ 格式
支持:
1. 从 State 和 Collector 构建完整轨迹
2. 自动检测 ReAct/Workflow 模式
3. 提取关键信息和统计数据
"""
from datetime import datetime
from typing import Any, Dict, List, Optional
from dataflow_agent.trajectory.models import (
Trajectory,
TrajectoryStep,
TrajectoryMode,
StepRole,
)
from dataflow_agent.trajectory.collector import TrajectoryCollector
from dataflow_agent.state import MainState
from dataflow_agent.logger import get_logger
log = get_logger(__name__)
class TrajectoryBuilder:
"""
轨迹构建器
将 TrajectoryCollector 收集的原始步骤数据和 State 对象
转换为标准的 Trajectory 对象
"""
def __init__(self):
pass
def build_from_state(self,
state: MainState,
collector: TrajectoryCollector,
workflow_name: str,
user_id: str = None,
session_id: str = None) -> Trajectory:
"""
从 State 和 Collector 构建完整轨迹
Args:
state: Workflow 执行后的最终状态
collector: 轨迹收集器
workflow_name: Workflow 名称
user_id: 用户 ID
session_id: 会话 ID
Returns:
完整的 Trajectory 对象
"""
log.info(f"[TrajectoryBuilder] 开始构建轨迹: {workflow_name}")
# 1. 生成 trace_id
trace_id = Trajectory.generate_trace_id()
# 2. 获取步骤
steps = collector.finish()
# 3. 检测模式
mode = self._detect_mode(state, steps)
# 4. 提取输入
inputs = self._extract_inputs(state, collector)
# 5. 提取输出
final_output = self._extract_final_output(state)
# 6. 判断状态
status = self._determine_status(state, steps)
# 7. 计算统计信息
total_duration_ms = self._calculate_total_duration(steps)
total_tokens = self._calculate_total_tokens(steps)
# 8. 构建 Trajectory
trajectory = Trajectory(
trace_id=trace_id,
workflow_name=workflow_name,
timestamp=datetime.now().isoformat(),
status=status,
mode=mode,
user_id=user_id,
session_id=session_id or getattr(state.request, 'session_id', None),
inputs=inputs,
steps=steps,
final_output=final_output,
total_duration_ms=total_duration_ms,
total_tokens=total_tokens,
metadata=collector.get_metadata()
)
# 更新统计
trajectory.total_llm_calls = sum(len(step.llm_calls) for step in steps)
trajectory.total_tool_calls = sum(len(step.tool_calls) for step in steps)
log.info(f"[TrajectoryBuilder] 轨迹构建完成: {trace_id}, "
f"模式={mode}, 步骤数={len(steps)}, 状态={status}")
return trajectory
def build_from_steps(self,
steps: List[TrajectoryStep],
workflow_name: str,
inputs: Dict[str, Any] = None,
final_output: Any = None,
**kwargs) -> Trajectory:
"""
直接从步骤列表构建轨迹(不依赖 State)
Args:
steps: 步骤列表
workflow_name: Workflow 名称
inputs: 输入数据
final_output: 最终输出
**kwargs: 其他参数
Returns:
Trajectory 对象
"""
trace_id = Trajectory.generate_trace_id()
mode = self._detect_mode_from_steps(steps)
status = "success" if not any(step.error for step in steps) else "failed"
trajectory = Trajectory(
trace_id=trace_id,
workflow_name=workflow_name,
timestamp=datetime.now().isoformat(),
status=status,
mode=mode,
inputs=inputs or {},
steps=steps,
final_output=final_output,
**kwargs
)
# 更新统计
trajectory.total_llm_calls = sum(len(step.llm_calls) for step in steps)
trajectory.total_tool_calls = sum(len(step.tool_calls) for step in steps)
trajectory.total_duration_ms = self._calculate_total_duration(steps)
trajectory.total_tokens = self._calculate_total_tokens(steps)
return trajectory
def _detect_mode(self, state: MainState, steps: List[TrajectoryStep]) -> str:
"""
检测轨迹模式
通过分析 State 和 Steps 判断是 ReAct 还是 Workflow 模式
"""
# 检查是否有 thought 字段(ReAct 特征)
has_thoughts = any(step.thought for step in steps)
# 检查是否有 observation 字段(ReAct 特征)
has_observations = any(step.observation for step in steps)
# 检查消息历史(ReAct 通常有更多的对话轮次)
messages = getattr(state, 'messages', [])
has_many_messages = len(messages) > 5
# 检查是否有明确的 agent 角色步骤
has_agent_steps = any(step.role == StepRole.AGENT.value for step in steps)
if has_thoughts or has_observations:
return TrajectoryMode.REACT.value
elif has_agent_steps and has_many_messages:
return TrajectoryMode.HYBRID.value
else:
return TrajectoryMode.WORKFLOW.value
def _detect_mode_from_steps(self, steps: List[TrajectoryStep]) -> str:
"""仅从步骤检测模式"""
has_thoughts = any(step.thought for step in steps)
has_observations = any(step.observation for step in steps)
if has_thoughts or has_observations:
return TrajectoryMode.REACT.value
else:
return TrajectoryMode.WORKFLOW.value
def _extract_inputs(self, state: MainState, collector: TrajectoryCollector) -> Dict[str, Any]:
"""
提取输入数据
优先级:
1. Collector 记录的初始输入
2. State.request 中的字段
3. State 的其他相关字段
"""
inputs = {}
# 从 collector 获取
collector_inputs = collector.get_initial_inputs()
if collector_inputs:
inputs.update(collector_inputs)
# 从 state.request 提取
if hasattr(state, 'request'):
request = state.request
# 提取常见字段
if hasattr(request, 'target') and request.target:
inputs['query'] = request.target
if hasattr(request, 'model'):
inputs['model'] = request.model
if hasattr(request, 'language'):
inputs['language'] = request.language
# 提取其他可能的输入字段
for field in ['json_file', 'python_file_path', 'keywords', 'style']:
if hasattr(request, field):
value = getattr(request, field)
if value:
inputs[field] = value
return inputs
def _extract_final_output(self, state: MainState) -> Any:
"""
提取最终输出
从 State 的不同字段中提取最终结果
"""
# 尝试从 agent_results 获取
if hasattr(state, 'agent_results') and state.agent_results:
# 获取最后一个 agent 的结果
last_agent_result = None
for agent_name, result in state.agent_results.items():
if isinstance(result, dict) and 'results' in result:
last_agent_result = result['results']
if last_agent_result:
return last_agent_result
# 尝试从特定字段获取
for field in ['final_output', 'execution_result', 'pipeline_structure_code',
'recommendation', 'icon_prompt', 'research_summary']:
if hasattr(state, field):
value = getattr(state, field)
if value:
return value
# 如果都没有,返回整个 state 的字典表示(简化版)
return {"status": "completed"}
def _determine_status(self, state: MainState, steps: List[TrajectoryStep]) -> str:
"""
判断执行状态
Returns:
"success" | "failed" | "partial"
"""
# 检查是否有错误步骤
has_errors = any(step.error for step in steps)
# 检查 execution_result
if hasattr(state, 'execution_result'):
exec_result = state.execution_result
if isinstance(exec_result, dict):
if exec_result.get('success') is False:
return "failed"
elif exec_result.get('success') is True:
return "success"
# 根据错误情况判断
if has_errors:
# 如果所有步骤都有错误,则失败
if all(step.error for step in steps):
return "failed"
else:
return "partial"
return "success"
def _calculate_total_duration(self, steps: List[TrajectoryStep]) -> Optional[float]:
"""计算总耗时"""
if not steps:
return None
total = sum(step.duration_ms for step in steps if step.duration_ms)
return total if total > 0 else None
def _calculate_total_tokens(self, steps: List[TrajectoryStep]) -> Optional[Dict[str, int]]:
"""计算总 token 使用量"""
total_prompt = 0
total_completion = 0
for step in steps:
for llm_call in step.llm_calls:
if llm_call.token_usage:
total_prompt += llm_call.token_usage.get('prompt', 0)
total_completion += llm_call.token_usage.get('completion', 0)
if total_prompt > 0 or total_completion > 0:
return {
'prompt': total_prompt,
'completion': total_completion,
'total': total_prompt + total_completion
}
return None
# ==================== 便捷函数 ====================
def create_builder() -> TrajectoryBuilder:
"""创建轨迹构建器"""
return TrajectoryBuilder()
|