Files
Liti-agent-Development/parallel_processing_definition.txt
T

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