SamarpeetGarad commited on
Commit
27b8fe0
·
verified ·
1 Parent(s): a2a86b3

Upload orchestrator/workflow.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. orchestrator/workflow.py +301 -0
orchestrator/workflow.py ADDED
@@ -0,0 +1,301 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ RadioFlow Orchestrator
3
+ Coordinates the multi-agent workflow for radiology analysis
4
+ """
5
+
6
+ import time
7
+ from dataclasses import dataclass, field
8
+ from typing import Any, Dict, List, Optional, Callable
9
+ from datetime import datetime
10
+ from PIL import Image
11
+
12
+ from agents import (
13
+ CXRAnalyzerAgent,
14
+ FindingInterpreterAgent,
15
+ ReportGeneratorAgent,
16
+ PriorityRouterAgent,
17
+ BaseAgent,
18
+ AgentResult
19
+ )
20
+ from utils.metrics import MetricsTracker
21
+
22
+
23
+ @dataclass
24
+ class WorkflowResult:
25
+ """Complete result from the RadioFlow workflow"""
26
+ workflow_id: str
27
+ status: str # "success", "partial", "error"
28
+ start_time: str
29
+ end_time: str
30
+ total_duration_ms: float
31
+
32
+ # Agent results
33
+ cxr_analysis: Optional[AgentResult] = None
34
+ finding_interpretation: Optional[AgentResult] = None
35
+ report: Optional[AgentResult] = None
36
+ priority_routing: Optional[AgentResult] = None
37
+
38
+ # Aggregated outputs
39
+ final_report: str = ""
40
+ priority_level: str = "ROUTINE"
41
+ priority_score: float = 0.0
42
+ findings_count: int = 0
43
+ critical_findings: List[str] = field(default_factory=list)
44
+
45
+ # Errors
46
+ errors: List[str] = field(default_factory=list)
47
+
48
+ def to_dict(self) -> Dict:
49
+ return {
50
+ "workflow_id": self.workflow_id,
51
+ "status": self.status,
52
+ "start_time": self.start_time,
53
+ "end_time": self.end_time,
54
+ "total_duration_ms": self.total_duration_ms,
55
+ "final_report": self.final_report,
56
+ "priority_level": self.priority_level,
57
+ "priority_score": self.priority_score,
58
+ "findings_count": self.findings_count,
59
+ "critical_findings": self.critical_findings,
60
+ "agent_results": {
61
+ "cxr_analysis": self.cxr_analysis.to_dict() if self.cxr_analysis else None,
62
+ "finding_interpretation": self.finding_interpretation.to_dict() if self.finding_interpretation else None,
63
+ "report": self.report.to_dict() if self.report else None,
64
+ "priority_routing": self.priority_routing.to_dict() if self.priority_routing else None,
65
+ },
66
+ "errors": self.errors
67
+ }
68
+
69
+
70
+ class RadioFlowOrchestrator:
71
+ """
72
+ Main orchestrator for the RadioFlow multi-agent system.
73
+
74
+ Coordinates the sequential execution of:
75
+ 1. CXR Analyzer (Image Analysis)
76
+ 2. Finding Interpreter (Clinical Interpretation)
77
+ 3. Report Generator (Structured Report)
78
+ 4. Priority Router (Urgency Assessment)
79
+ """
80
+
81
+ def __init__(self, demo_mode: bool = True):
82
+ """
83
+ Initialize the orchestrator.
84
+
85
+ Args:
86
+ demo_mode: If True, agents use simulated outputs for faster demos
87
+ """
88
+ self.demo_mode = demo_mode
89
+ self.metrics = MetricsTracker()
90
+
91
+ # Initialize agents
92
+ self.agents: Dict[str, BaseAgent] = {
93
+ "cxr_analyzer": CXRAnalyzerAgent(demo_mode=demo_mode),
94
+ "finding_interpreter": FindingInterpreterAgent(demo_mode=demo_mode),
95
+ "report_generator": ReportGeneratorAgent(demo_mode=demo_mode),
96
+ "priority_router": PriorityRouterAgent(demo_mode=demo_mode)
97
+ }
98
+
99
+ # Workflow state
100
+ self._current_workflow_id: Optional[str] = None
101
+ self._workflow_callbacks: List[Callable] = []
102
+
103
+ # Agent order for pipeline
104
+ self._agent_order = [
105
+ "cxr_analyzer",
106
+ "finding_interpreter",
107
+ "report_generator",
108
+ "priority_router"
109
+ ]
110
+
111
+ def load_all_models(self) -> Dict[str, bool]:
112
+ """Load all agent models. Returns dict of agent_name -> success."""
113
+ results = {}
114
+ for name, agent in self.agents.items():
115
+ try:
116
+ results[name] = agent.load_model()
117
+ except Exception as e:
118
+ print(f"Failed to load {name}: {e}")
119
+ results[name] = False
120
+ return results
121
+
122
+ def add_callback(self, callback: Callable[[str, AgentResult], None]):
123
+ """Add a callback to be called after each agent completes."""
124
+ self._workflow_callbacks.append(callback)
125
+
126
+ def _notify_callbacks(self, agent_name: str, result: AgentResult):
127
+ """Notify all callbacks of agent completion."""
128
+ for callback in self._workflow_callbacks:
129
+ try:
130
+ callback(agent_name, result)
131
+ except Exception as e:
132
+ print(f"Callback error: {e}")
133
+
134
+ def process(
135
+ self,
136
+ image: Image.Image,
137
+ clinical_context: Optional[Dict] = None,
138
+ workflow_id: Optional[str] = None
139
+ ) -> WorkflowResult:
140
+ """
141
+ Run the complete RadioFlow workflow.
142
+
143
+ Args:
144
+ image: Chest X-ray image (PIL Image)
145
+ clinical_context: Optional clinical information
146
+ workflow_id: Optional ID for tracking
147
+
148
+ Returns:
149
+ WorkflowResult with complete analysis
150
+ """
151
+ # Initialize workflow
152
+ start_time = time.time()
153
+ start_timestamp = datetime.now().isoformat()
154
+
155
+ if workflow_id is None:
156
+ workflow_id = f"rf_{datetime.now().strftime('%Y%m%d_%H%M%S_%f')}"
157
+
158
+ self._current_workflow_id = workflow_id
159
+ self.metrics.start_workflow(workflow_id)
160
+
161
+ # Prepare context
162
+ context = clinical_context or {}
163
+
164
+ # Initialize result
165
+ result = WorkflowResult(
166
+ workflow_id=workflow_id,
167
+ status="processing",
168
+ start_time=start_timestamp,
169
+ end_time="",
170
+ total_duration_ms=0
171
+ )
172
+
173
+ errors = []
174
+
175
+ try:
176
+ # ============================================
177
+ # STAGE 1: CXR Analysis
178
+ # ============================================
179
+ print(f"[{workflow_id}] Stage 1: CXR Analysis...")
180
+ cxr_result = self.agents["cxr_analyzer"](image, context)
181
+ result.cxr_analysis = cxr_result
182
+ self.metrics.record_agent("CXR Analyzer", cxr_result.processing_time_ms, cxr_result.status == "success")
183
+ self._notify_callbacks("cxr_analyzer", cxr_result)
184
+
185
+ if cxr_result.status == "error":
186
+ errors.append(f"CXR Analyzer: {cxr_result.error_message}")
187
+
188
+ # ============================================
189
+ # STAGE 2: Finding Interpretation
190
+ # ============================================
191
+ print(f"[{workflow_id}] Stage 2: Finding Interpretation...")
192
+ interpretation_input = cxr_result.data if cxr_result.status == "success" else {}
193
+ interpretation_result = self.agents["finding_interpreter"](interpretation_input, context)
194
+ result.finding_interpretation = interpretation_result
195
+ self.metrics.record_agent("Finding Interpreter", interpretation_result.processing_time_ms, interpretation_result.status == "success")
196
+ self._notify_callbacks("finding_interpreter", interpretation_result)
197
+
198
+ if interpretation_result.status == "error":
199
+ errors.append(f"Finding Interpreter: {interpretation_result.error_message}")
200
+
201
+ # ============================================
202
+ # STAGE 3: Report Generation
203
+ # ============================================
204
+ print(f"[{workflow_id}] Stage 3: Report Generation...")
205
+ report_input = interpretation_result.data if interpretation_result.status == "success" else {}
206
+ report_result = self.agents["report_generator"](report_input, context)
207
+ result.report = report_result
208
+ self.metrics.record_agent("Report Generator", report_result.processing_time_ms, report_result.status == "success")
209
+ self._notify_callbacks("report_generator", report_result)
210
+
211
+ if report_result.status == "error":
212
+ errors.append(f"Report Generator: {report_result.error_message}")
213
+
214
+ # ============================================
215
+ # STAGE 4: Priority Routing
216
+ # ============================================
217
+ print(f"[{workflow_id}] Stage 4: Priority Routing...")
218
+ # Pass original findings through context for priority assessment
219
+ priority_context = {
220
+ **context,
221
+ "original_findings": cxr_result.data.get("findings", []) if cxr_result.data else []
222
+ }
223
+ priority_input = report_result.data if report_result.status == "success" else {}
224
+ priority_result = self.agents["priority_router"](priority_input, priority_context)
225
+ result.priority_routing = priority_result
226
+ self.metrics.record_agent("Priority Router", priority_result.processing_time_ms, priority_result.status == "success")
227
+ self._notify_callbacks("priority_router", priority_result)
228
+
229
+ if priority_result.status == "error":
230
+ errors.append(f"Priority Router: {priority_result.error_message}")
231
+
232
+ # ============================================
233
+ # Aggregate Results
234
+ # ============================================
235
+ result.final_report = report_result.data.get("full_report", "") if report_result.data else ""
236
+ result.priority_level = priority_result.data.get("priority_level", "ROUTINE") if priority_result.data else "ROUTINE"
237
+ result.priority_score = priority_result.data.get("priority_score", 0.0) if priority_result.data else 0.0
238
+ result.findings_count = len(cxr_result.data.get("findings", [])) if cxr_result.data else 0
239
+ result.critical_findings = priority_result.data.get("critical_findings_detected", []) if priority_result.data else []
240
+
241
+ # Determine overall status
242
+ if not errors:
243
+ result.status = "success"
244
+ elif len(errors) < 4:
245
+ result.status = "partial"
246
+ else:
247
+ result.status = "error"
248
+
249
+ result.errors = errors
250
+
251
+ except Exception as e:
252
+ result.status = "error"
253
+ result.errors = [str(e)]
254
+ print(f"[{workflow_id}] Workflow error: {e}")
255
+
256
+ finally:
257
+ # Finalize timing
258
+ end_time = time.time()
259
+ result.end_time = datetime.now().isoformat()
260
+ result.total_duration_ms = (end_time - start_time) * 1000
261
+
262
+ # Record metrics
263
+ self.metrics.end_workflow(
264
+ findings_count=result.findings_count,
265
+ priority_score=result.priority_score,
266
+ status=result.status
267
+ )
268
+
269
+ print(f"[{workflow_id}] Workflow complete in {result.total_duration_ms:.0f}ms")
270
+
271
+ return result
272
+
273
+ def get_agent_statuses(self) -> Dict[str, Dict]:
274
+ """Get status of all agents."""
275
+ return {
276
+ name: {
277
+ "name": agent.name,
278
+ "model": agent.model_name,
279
+ "loaded": agent.is_loaded,
280
+ "metrics": agent.get_metrics()
281
+ }
282
+ for name, agent in self.agents.items()
283
+ }
284
+
285
+ def get_workflow_metrics(self) -> str:
286
+ """Get formatted workflow metrics."""
287
+ return self.metrics.format_for_display()
288
+
289
+ def reset(self):
290
+ """Reset orchestrator state."""
291
+ self._current_workflow_id = None
292
+ for agent in self.agents.values():
293
+ agent.reset_metrics()
294
+ self.metrics = MetricsTracker()
295
+
296
+
297
+ def create_orchestrator(demo_mode: bool = True) -> RadioFlowOrchestrator:
298
+ """Factory function to create an orchestrator instance."""
299
+ orchestrator = RadioFlowOrchestrator(demo_mode=demo_mode)
300
+ orchestrator.load_all_models()
301
+ return orchestrator