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()