217 lines
8.1 KiB
Plaintext
217 lines
8.1 KiB
Plaintext
class TaskProcedureExecutor:
|
|
"""
|
|
Task 실행 흐름을 DAG(Directed Acyclic Graph) 기반으로 관리하고 실행합니다.
|
|
|
|
task_procedure 예시:
|
|
{
|
|
"IN": {"nexts": ["task_1"], "wait_until": []},
|
|
"task_1": {"nexts": ["task_2", "task_4"], "wait_until": []},
|
|
"task_2": {"nexts": ["task_3"], "wait_until": []},
|
|
"task_3": {"nexts": ["task_5"], "wait_until": []},
|
|
"task_4": {"nexts": ["task_5"], "wait_until": []},
|
|
"task_5": {"nexts": ["OUT"], "wait_until": ["task_3", "task_4"]}
|
|
}
|
|
"""
|
|
|
|
def __init__(self, tasks: Dict[str, Task], procedure: Dict[str, Dict], system_prompt: str):
|
|
"""
|
|
Args:
|
|
tasks: task_name -> Task 객체 매핑
|
|
procedure: task_procedure 정의 (IN, OUT 포함)
|
|
system_prompt: 모든 Task에 공통으로 적용할 시스템 프롬프트
|
|
"""
|
|
self.tasks = tasks
|
|
self.procedure = procedure
|
|
self.system_prompt = system_prompt
|
|
self.completed_tasks: Set[str] = set()
|
|
self.running_tasks: Set[str] = set()
|
|
self.task_results: Dict[str, Any] = {}
|
|
self.task_events: Dict[str, asyncio.Event] = {}
|
|
|
|
# 각 Task에 대한 완료 이벤트 생성
|
|
for task_name in tasks:
|
|
self.task_events[task_name] = asyncio.Event()
|
|
|
|
# wait_until 기반 의존성 로깅
|
|
wait_until_deps = {}
|
|
for task_name, config in procedure.items():
|
|
if task_name in ("IN", "OUT"):
|
|
continue
|
|
wait_until = config.get("wait_until", [])
|
|
if wait_until:
|
|
wait_until_deps[task_name] = wait_until
|
|
logger.info(f"[EXECUTOR] Wait-until dependencies: {wait_until_deps}")
|
|
|
|
def _can_start_task(self, task_name: str) -> bool:
|
|
"""Task가 시작 가능한지 확인 (wait_until 조건 체크)"""
|
|
if task_name in self.completed_tasks or task_name in self.running_tasks:
|
|
return False
|
|
|
|
task_config = self.procedure.get(task_name, {})
|
|
wait_until = task_config.get("wait_until", [])
|
|
|
|
# wait_until에 있는 모든 Task가 완료되어야 함
|
|
for dep_task in wait_until:
|
|
if dep_task not in self.completed_tasks:
|
|
return False
|
|
|
|
return True
|
|
|
|
def _get_dependencies(self, task_name: str) -> Set[str]:
|
|
"""Task의 의존성을 반환 (wait_until만 사용)"""
|
|
task_config = self.procedure.get(task_name, {})
|
|
wait_until = task_config.get("wait_until", [])
|
|
return set(wait_until)
|
|
|
|
async def _wait_for_dependencies(self, task_name: str):
|
|
"""Task의 의존성이 완료될 때까지 대기"""
|
|
dependencies = self._get_dependencies(task_name)
|
|
|
|
if not dependencies:
|
|
return
|
|
|
|
logger.info(f"[EXECUTOR] Task '{task_name}' waiting for: {dependencies}")
|
|
|
|
# 모든 의존 Task의 완료 이벤트를 기다림
|
|
await asyncio.gather(*[
|
|
self.task_events[dep].wait()
|
|
for dep in dependencies
|
|
if dep in self.task_events
|
|
])
|
|
|
|
logger.info(f"[EXECUTOR] Task '{task_name}' dependencies satisfied")
|
|
|
|
async def _run_single_task(self, task_name: str, on_task_event=None) -> Any:
|
|
"""단일 Task 실행"""
|
|
if task_name not in self.tasks:
|
|
logger.error(f"[EXECUTOR] Task '{task_name}' not found")
|
|
return None
|
|
|
|
task = self.tasks[task_name]
|
|
self.running_tasks.add(task_name)
|
|
|
|
# 의존성 확인 (자동 계산 + 명시적 wait_until)
|
|
dependencies = self._get_dependencies(task_name)
|
|
|
|
if dependencies:
|
|
# 의존성이 있으면 waiting 상태로 시작
|
|
task.status = TaskStatus.WAITING
|
|
if on_task_event:
|
|
await on_task_event({
|
|
"type": "task_start",
|
|
"task_name": task_name,
|
|
"status": "waiting"
|
|
})
|
|
# 의존성 대기
|
|
await self._wait_for_dependencies(task_name)
|
|
|
|
# 의존성 완료 후 running 상태로 변경
|
|
task.status = TaskStatus.RUNNING
|
|
|
|
if on_task_event:
|
|
await on_task_event({
|
|
"type": "task_start",
|
|
"task_name": task_name,
|
|
"status": "running"
|
|
})
|
|
|
|
try:
|
|
# Task 실행
|
|
async def task_on_iteration(iteration_num, content, is_complete):
|
|
if on_task_event:
|
|
await on_task_event({
|
|
"type": "task_iteration",
|
|
"task_name": task_name,
|
|
"iteration": iteration_num,
|
|
"content": content,
|
|
"is_complete": is_complete
|
|
})
|
|
|
|
result = await task.run(self.system_prompt, on_iteration=task_on_iteration)
|
|
|
|
self.task_results[task_name] = result
|
|
self.completed_tasks.add(task_name)
|
|
self.running_tasks.discard(task_name)
|
|
self.task_events[task_name].set() # 완료 신호
|
|
|
|
if on_task_event:
|
|
await on_task_event({
|
|
"type": "task_complete",
|
|
"task_name": task_name,
|
|
"status": "completed"
|
|
})
|
|
|
|
logger.info(f"[EXECUTOR] Task '{task_name}' completed successfully")
|
|
return result
|
|
|
|
except Exception as e:
|
|
task.status = TaskStatus.FAILED
|
|
task.error = str(e)
|
|
self.running_tasks.discard(task_name)
|
|
self.task_events[task_name].set() # 실패해도 이벤트 설정 (다른 Task 대기 해제)
|
|
|
|
if on_task_event:
|
|
await on_task_event({
|
|
"type": "task_error",
|
|
"task_name": task_name,
|
|
"error": str(e),
|
|
"status": "failed"
|
|
})
|
|
|
|
logger.error(f"[EXECUTOR] Task '{task_name}' failed: {e}")
|
|
raise
|
|
|
|
async def execute(self, on_task_event=None) -> Dict[str, Any]:
|
|
"""
|
|
전체 Task Procedure 실행
|
|
|
|
DAG 구조에 따라 병렬 실행 가능한 Task들은 동시에 실행합니다.
|
|
각 Task는 내부적으로 wait_until 의존성을 기다린 후 실행됩니다.
|
|
|
|
Args:
|
|
on_task_event: Task 이벤트 콜백 (task_start, task_iteration, task_complete, task_error)
|
|
|
|
Returns:
|
|
모든 Task의 결과를 담은 딕셔너리
|
|
"""
|
|
logger.info("[EXECUTOR] Starting task procedure execution")
|
|
|
|
# 모든 실행할 Task 수집 (BFS)
|
|
all_tasks_to_run = set()
|
|
in_config = self.procedure.get("IN", {})
|
|
initial_tasks = in_config.get("nexts", [])
|
|
|
|
if not initial_tasks:
|
|
logger.warning("[EXECUTOR] No initial tasks found in procedure")
|
|
return {}
|
|
|
|
queue = list(initial_tasks)
|
|
while queue:
|
|
task_name = queue.pop(0)
|
|
if task_name == "OUT" or task_name in all_tasks_to_run:
|
|
continue
|
|
all_tasks_to_run.add(task_name)
|
|
task_config = self.procedure.get(task_name, {})
|
|
for next_task in task_config.get("nexts", []):
|
|
if next_task != "OUT":
|
|
queue.append(next_task)
|
|
|
|
logger.info(f"[EXECUTOR] Tasks to execute: {all_tasks_to_run}")
|
|
|
|
# 모든 Task를 동시에 시작 (각 Task는 내부에서 의존성을 기다림)
|
|
async def run_task_with_dependencies(task_name: str):
|
|
"""Task를 의존성 대기 후 실행"""
|
|
try:
|
|
await self._run_single_task(task_name, on_task_event)
|
|
except Exception as e:
|
|
logger.error(f"[EXECUTOR] Error running task '{task_name}': {e}")
|
|
|
|
# 모든 Task를 병렬로 시작 - 각 Task는 _wait_for_dependencies에서 의존성 완료를 기다림
|
|
task_coroutines = [run_task_with_dependencies(task) for task in all_tasks_to_run]
|
|
await asyncio.gather(*task_coroutines, return_exceptions=True)
|
|
|
|
logger.info(f"[EXECUTOR] Task procedure completed. Completed tasks: {self.completed_tasks}")
|
|
|
|
return self.task_results
|
|
|