task_manager.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317
  1. """Task management service for Web-Local system dispatch (P4-M2).
  2. Handles task creation, dispatch, status tracking, progress updates,
  3. and result reception from local simulation executor.
  4. """
  5. import json
  6. import os
  7. import uuid
  8. from datetime import datetime
  9. from typing import Dict, List, Optional, Any
  10. from pathlib import Path
  11. from ..database import SessionLocal, Task
  12. class TaskManager:
  13. """Manages simulation tasks between Web and Local systems."""
  14. TASK_STATUSES = [
  15. "pending", # Created, waiting for dispatch
  16. "dispatched", # Sent to local executor
  17. "running", # Local executor is running
  18. "completed", # All points completed successfully
  19. "failed", # Task failed
  20. "cancelled", # User cancelled
  21. ]
  22. def __init__(self, output_dir: Optional[str] = None):
  23. self.output_dir = output_dir or os.path.join(
  24. os.path.dirname(os.path.dirname(os.path.dirname(__file__))),
  25. "output", "tasks"
  26. )
  27. os.makedirs(self.output_dir, exist_ok=True)
  28. def create_task(
  29. self,
  30. plan_id: Optional[int],
  31. plan_data: Dict[str, Any],
  32. parameters: List[Dict[str, Any]],
  33. task_name: Optional[str] = None,
  34. priority: int = 5,
  35. created_by: str = "web",
  36. ) -> Dict[str, Any]:
  37. """Create a new simulation task.
  38. Args:
  39. plan_id: Associated plan ID (optional)
  40. plan_data: Full plan data (boundary conditions, topology, etc.)
  41. parameters: List of parameter sets to simulate
  42. task_name: Optional task name
  43. priority: Task priority (1-10, higher = more urgent)
  44. created_by: Creator identifier
  45. Returns:
  46. Created task dict
  47. """
  48. task_uuid = str(uuid.uuid4())[:8]
  49. task_name = task_name or f"task_{task_uuid}"
  50. task_dir = os.path.join(self.output_dir, f"{datetime.now().strftime('%Y%m%d_%H%M%S')}_{task_name}")
  51. os.makedirs(task_dir, exist_ok=True)
  52. # Write task file for local executor
  53. task_file = os.path.join(task_dir, "task.json")
  54. task_payload = {
  55. "task_id": task_uuid,
  56. "task_name": task_name,
  57. "plan_id": plan_id,
  58. "plan_data": plan_data,
  59. "parameters": parameters,
  60. "priority": priority,
  61. "created_at": datetime.now().isoformat(),
  62. "total_points": len(parameters),
  63. }
  64. with open(task_file, "w", encoding="utf-8") as f:
  65. json.dump(task_payload, f, ensure_ascii=False, indent=2)
  66. # Save to database
  67. with SessionLocal() as db:
  68. db_task = Task(
  69. task_id=task_uuid,
  70. task_name=task_name,
  71. plan_id=plan_id,
  72. status="pending",
  73. priority=priority,
  74. total_points=len(parameters),
  75. completed_points=0,
  76. task_dir=task_dir,
  77. task_file=task_file,
  78. created_by=created_by,
  79. created_at=datetime.now(),
  80. )
  81. db.add(db_task)
  82. db.commit()
  83. db.refresh(db_task)
  84. return self._task_to_dict(db_task)
  85. def dispatch_task(self, task_id: str) -> Dict[str, Any]:
  86. """Mark task as dispatched and ready for local executor.
  87. Args:
  88. task_id: Task UUID
  89. Returns:
  90. Updated task dict
  91. """
  92. with SessionLocal() as db:
  93. task = db.query(Task).filter(Task.task_id == task_id).first()
  94. if not task:
  95. raise ValueError(f"Task {task_id} not found")
  96. if task.status != "pending":
  97. raise ValueError(f"Task {task_id} is not pending (status: {task.status})")
  98. task.status = "dispatched"
  99. task.dispatched_at = datetime.now()
  100. db.commit()
  101. db.refresh(task)
  102. return self._task_to_dict(task)
  103. def update_progress(
  104. self,
  105. task_id: str,
  106. current_point: int,
  107. total_points: Optional[int] = None,
  108. current_params: Optional[Dict[str, Any]] = None,
  109. elapsed_time: Optional[float] = None,
  110. ) -> Dict[str, Any]:
  111. """Update task progress from local executor.
  112. Args:
  113. task_id: Task UUID
  114. current_point: Current point index (0-based)
  115. total_points: Total points (optional, will use stored value)
  116. current_params: Current parameter values being simulated
  117. elapsed_time: Elapsed time in seconds
  118. Returns:
  119. Updated task dict
  120. """
  121. with SessionLocal() as db:
  122. task = db.query(Task).filter(Task.task_id == task_id).first()
  123. if not task:
  124. raise ValueError(f"Task {task_id} not found")
  125. if task.status in ("dispatched", "running"):
  126. task.status = "running"
  127. task.started_at = task.started_at or datetime.now()
  128. task.completed_points = current_point
  129. if total_points:
  130. task.total_points = total_points
  131. # Update progress metadata
  132. progress_data = {
  133. "current_point": current_point,
  134. "current_params": current_params,
  135. "elapsed_time": elapsed_time,
  136. "updated_at": datetime.now().isoformat(),
  137. }
  138. existing_progress = json.loads(task.progress_data or "{}")
  139. existing_progress.update(progress_data)
  140. task.progress_data = json.dumps(existing_progress, ensure_ascii=False)
  141. db.commit()
  142. db.refresh(task)
  143. return self._task_to_dict(task)
  144. def report_results(
  145. self,
  146. task_id: str,
  147. results: List[Dict[str, Any]],
  148. metrics: Optional[Dict[str, Any]] = None,
  149. logs: Optional[str] = None,
  150. duration: Optional[float] = None,
  151. status: str = "completed",
  152. ) -> Dict[str, Any]:
  153. """Report final results from local executor.
  154. Args:
  155. task_id: Task UUID
  156. results: List of simulation result dicts
  157. metrics: Aggregated metrics
  158. logs: Program logs
  159. duration: Total duration in seconds
  160. status: Final status (completed/failed)
  161. Returns:
  162. Updated task dict
  163. """
  164. with SessionLocal() as db:
  165. task = db.query(Task).filter(Task.task_id == task_id).first()
  166. if not task:
  167. raise ValueError(f"Task {task_id} not found")
  168. task.status = status
  169. task.completed_at = datetime.now()
  170. task.completed_points = len(results)
  171. if duration:
  172. task.duration = duration
  173. # Save results to file
  174. results_file = os.path.join(task.task_dir, "results.json")
  175. with open(results_file, "w", encoding="utf-8") as f:
  176. json.dump({"results": results, "metrics": metrics, "logs": logs}, f, ensure_ascii=False, indent=2)
  177. task.results_file = results_file
  178. if metrics:
  179. task.result_metrics = json.dumps(metrics, ensure_ascii=False)
  180. db.commit()
  181. db.refresh(task)
  182. return self._task_to_dict(task)
  183. def get_task(self, task_id: str) -> Optional[Dict[str, Any]]:
  184. """Get task by ID."""
  185. with SessionLocal() as db:
  186. task = db.query(Task).filter(Task.task_id == task_id).first()
  187. if not task:
  188. return None
  189. return self._task_to_dict(task)
  190. def list_tasks(
  191. self,
  192. status: Optional[str] = None,
  193. plan_id: Optional[int] = None,
  194. limit: int = 50,
  195. offset: int = 0,
  196. ) -> Dict[str, Any]:
  197. """List tasks with filters."""
  198. with SessionLocal() as db:
  199. query = db.query(Task)
  200. if status:
  201. query = query.filter(Task.status == status)
  202. if plan_id:
  203. query = query.filter(Task.plan_id == plan_id)
  204. total = query.count()
  205. tasks = query.order_by(Task.created_at.desc()).offset(offset).limit(limit).all()
  206. return {
  207. "total": total,
  208. "limit": limit,
  209. "offset": offset,
  210. "tasks": [self._task_to_dict(t) for t in tasks],
  211. }
  212. def cancel_task(self, task_id: str) -> Dict[str, Any]:
  213. """Cancel a pending or running task."""
  214. with SessionLocal() as db:
  215. task = db.query(Task).filter(Task.task_id == task_id).first()
  216. if not task:
  217. raise ValueError(f"Task {task_id} not found")
  218. if task.status in ("completed", "failed", "cancelled"):
  219. raise ValueError(f"Task {task_id} already finished (status: {task.status})")
  220. task.status = "cancelled"
  221. task.completed_at = datetime.now()
  222. db.commit()
  223. db.refresh(task)
  224. return self._task_to_dict(task)
  225. def get_task_results(self, task_id: str) -> Optional[Dict[str, Any]]:
  226. """Get task results file content."""
  227. task = self.get_task(task_id)
  228. if not task or not task.get("results_file"):
  229. return None
  230. if not os.path.exists(task["results_file"]):
  231. return None
  232. with open(task["results_file"], "r", encoding="utf-8") as f:
  233. return json.load(f)
  234. def _task_to_dict(self, task: Task) -> Dict[str, Any]:
  235. """Convert Task ORM object to dict."""
  236. result = {
  237. "id": task.id,
  238. "task_id": task.task_id,
  239. "task_name": task.task_name,
  240. "plan_id": task.plan_id,
  241. "status": task.status,
  242. "priority": task.priority,
  243. "total_points": task.total_points,
  244. "completed_points": task.completed_points,
  245. "progress": round((task.completed_points / task.total_points * 100), 1) if task.total_points else 0,
  246. "task_dir": task.task_dir,
  247. "task_file": task.task_file,
  248. "results_file": task.results_file,
  249. "created_by": task.created_by,
  250. "created_at": task.created_at.isoformat() if task.created_at else None,
  251. "dispatched_at": task.dispatched_at.isoformat() if task.dispatched_at else None,
  252. "started_at": task.started_at.isoformat() if task.started_at else None,
  253. "completed_at": task.completed_at.isoformat() if task.completed_at else None,
  254. "duration": task.duration,
  255. }
  256. if task.progress_data:
  257. try:
  258. result["progress_data"] = json.loads(task.progress_data)
  259. except (json.JSONDecodeError, Exception):
  260. result["progress_data"] = None
  261. if task.result_metrics:
  262. try:
  263. result["result_metrics"] = json.loads(task.result_metrics)
  264. except (json.JSONDecodeError, Exception):
  265. result["result_metrics"] = None
  266. return result
  267. # Global singleton
  268. _task_manager: Optional[TaskManager] = None
  269. def get_task_manager() -> TaskManager:
  270. """Get or create global TaskManager singleton."""
  271. global _task_manager
  272. if _task_manager is None:
  273. _task_manager = TaskManager()
  274. return _task_manager