feat(orchestrator): task 状态机 + pool 封装
This commit is contained in:
@@ -0,0 +1,12 @@
|
|||||||
|
"""
|
||||||
|
Sanguo Orchestrator Module
|
||||||
|
Task coordination and execution management for Sanguo VeighNa platform
|
||||||
|
"""
|
||||||
|
|
||||||
|
__version__ = "0.1.0"
|
||||||
|
|
||||||
|
from .task import TaskState, Task
|
||||||
|
from .pool import TaskPool
|
||||||
|
from .runner import Orchestrator
|
||||||
|
|
||||||
|
__all__ = ["TaskState", "Task", "TaskPool", "Orchestrator"]
|
||||||
@@ -0,0 +1,29 @@
|
|||||||
|
"""
|
||||||
|
Task pool for managing multiple tasks in memory
|
||||||
|
Provides task storage and status tracking (not actual multiprocessing)
|
||||||
|
"""
|
||||||
|
from .task import Task, TaskState
|
||||||
|
|
||||||
|
|
||||||
|
class TaskPool:
|
||||||
|
"""Manages task storage and status tracking"""
|
||||||
|
|
||||||
|
def __init__(self, max_workers: int = 2):
|
||||||
|
"""Initialize task pool with maximum workers"""
|
||||||
|
self.max_workers = max_workers
|
||||||
|
self._tasks: dict[str, Task] = {}
|
||||||
|
|
||||||
|
def submit(self, task_id: str, task_type: str) -> Task:
|
||||||
|
"""Submit a new task to the pool"""
|
||||||
|
task = Task(task_id=task_id, task_type=task_type)
|
||||||
|
self._tasks[task_id] = task
|
||||||
|
return task
|
||||||
|
|
||||||
|
def get_status(self, task_id: str) -> TaskState | None:
|
||||||
|
"""Get task status by ID"""
|
||||||
|
task = self._tasks.get(task_id)
|
||||||
|
return task.status if task else None
|
||||||
|
|
||||||
|
def get_task(self, task_id: str) -> Task | None:
|
||||||
|
"""Get task object by ID"""
|
||||||
|
return self._tasks.get(task_id)
|
||||||
@@ -0,0 +1,40 @@
|
|||||||
|
"""
|
||||||
|
Task state management for Sanguo Orchestrator
|
||||||
|
Defines Task state machine and transitions
|
||||||
|
"""
|
||||||
|
import enum
|
||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
|
||||||
|
class TaskState(enum.Enum):
|
||||||
|
"""Task execution states"""
|
||||||
|
PENDING = "pending"
|
||||||
|
RUNNING = "running"
|
||||||
|
DONE = "done"
|
||||||
|
FAILED = "failed"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class Task:
|
||||||
|
"""Represents a task in the orchestrator"""
|
||||||
|
task_id: str
|
||||||
|
task_type: str
|
||||||
|
status: TaskState = TaskState.PENDING
|
||||||
|
result_id: int | None = None
|
||||||
|
error_msg: str | None = None
|
||||||
|
|
||||||
|
def start(self):
|
||||||
|
"""Transition from PENDING to RUNNING"""
|
||||||
|
if self.status != TaskState.PENDING:
|
||||||
|
raise ValueError(f"不能从 {self.status.name} 启动")
|
||||||
|
self.status = TaskState.RUNNING
|
||||||
|
|
||||||
|
def complete(self, result_id: int):
|
||||||
|
"""Transition to DONE with result"""
|
||||||
|
self.status = TaskState.DONE
|
||||||
|
self.result_id = result_id
|
||||||
|
|
||||||
|
def fail(self, error_msg: str):
|
||||||
|
"""Transition to FAILED with error message"""
|
||||||
|
self.status = TaskState.FAILED
|
||||||
|
self.error_msg = error_msg
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
# Tests for sanguo_orchestrator module
|
||||||
@@ -0,0 +1,62 @@
|
|||||||
|
"""
|
||||||
|
Tests for sanguo_orchestrator.pool module
|
||||||
|
Tests TaskPool task management
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
from sanguo_orchestrator.pool import TaskPool
|
||||||
|
from sanguo_orchestrator.task import TaskState
|
||||||
|
|
||||||
|
|
||||||
|
class TestTaskPool:
|
||||||
|
"""Test TaskPool initialization and task management"""
|
||||||
|
|
||||||
|
def test_task_pool_initialization(self):
|
||||||
|
"""Test TaskPool initializes with max_workers"""
|
||||||
|
pool = TaskPool(max_workers=2)
|
||||||
|
assert pool.max_workers == 2
|
||||||
|
assert pool._tasks == {}
|
||||||
|
|
||||||
|
def test_task_pool_submit_creates_task(self):
|
||||||
|
"""Test submit() creates and stores a new Task"""
|
||||||
|
pool = TaskPool(max_workers=2)
|
||||||
|
task = pool.submit(task_id="test_1", task_type="test")
|
||||||
|
assert task.task_id == "test_1"
|
||||||
|
assert task.task_type == "test"
|
||||||
|
assert task.status == TaskState.PENDING
|
||||||
|
assert "test_1" in pool._tasks
|
||||||
|
|
||||||
|
def test_task_pool_get_status_existing(self):
|
||||||
|
"""Test get_status() returns TaskState for existing task"""
|
||||||
|
pool = TaskPool(max_workers=2)
|
||||||
|
pool.submit(task_id="test_1", task_type="test")
|
||||||
|
status = pool.get_status("test_1")
|
||||||
|
assert status == TaskState.PENDING
|
||||||
|
|
||||||
|
def test_task_pool_get_status_nonexistent(self):
|
||||||
|
"""Test get_status() returns None for non-existent task"""
|
||||||
|
pool = TaskPool(max_workers=2)
|
||||||
|
status = pool.get_status("nonexistent")
|
||||||
|
assert status is None
|
||||||
|
|
||||||
|
def test_task_pool_get_task_existing(self):
|
||||||
|
"""Test get_task() returns Task object for existing task"""
|
||||||
|
pool = TaskPool(max_workers=2)
|
||||||
|
submitted_task = pool.submit(task_id="test_1", task_type="test")
|
||||||
|
retrieved_task = pool.get_task("test_1")
|
||||||
|
assert retrieved_task is submitted_task
|
||||||
|
assert retrieved_task.task_id == "test_1"
|
||||||
|
|
||||||
|
def test_task_pool_get_task_nonexistent(self):
|
||||||
|
"""Test get_task() returns None for non-existent task"""
|
||||||
|
pool = TaskPool(max_workers=2)
|
||||||
|
task = pool.get_task("nonexistent")
|
||||||
|
assert task is None
|
||||||
|
|
||||||
|
def test_task_pool_multiple_tasks(self):
|
||||||
|
"""Test TaskPool can handle multiple tasks"""
|
||||||
|
pool = TaskPool(max_workers=2)
|
||||||
|
task1 = pool.submit(task_id="test_1", task_type="test")
|
||||||
|
task2 = pool.submit(task_id="test_2", task_type="test")
|
||||||
|
assert len(pool._tasks) == 2
|
||||||
|
assert pool.get_status("test_1") == TaskState.PENDING
|
||||||
|
assert pool.get_status("test_2") == TaskState.PENDING
|
||||||
@@ -0,0 +1,72 @@
|
|||||||
|
"""
|
||||||
|
Tests for sanguo_orchestrator.task module
|
||||||
|
Tests Task state transitions and validation
|
||||||
|
"""
|
||||||
|
import pytest
|
||||||
|
from sanguo_orchestrator.task import TaskState, Task
|
||||||
|
|
||||||
|
|
||||||
|
class TestTaskState:
|
||||||
|
"""Test TaskState enum values"""
|
||||||
|
|
||||||
|
def test_task_state_enum_values(self):
|
||||||
|
"""Test TaskState has correct enum values"""
|
||||||
|
assert TaskState.PENDING.value == "pending"
|
||||||
|
assert TaskState.RUNNING.value == "running"
|
||||||
|
assert TaskState.DONE.value == "done"
|
||||||
|
assert TaskState.FAILED.value == "failed"
|
||||||
|
|
||||||
|
|
||||||
|
class TestTask:
|
||||||
|
"""Test Task dataclass and state transitions"""
|
||||||
|
|
||||||
|
def test_task_initial_state(self):
|
||||||
|
"""Test Task starts with PENDING state"""
|
||||||
|
task = Task(task_id="test_1", task_type="test")
|
||||||
|
assert task.task_id == "test_1"
|
||||||
|
assert task.task_type == "test"
|
||||||
|
assert task.status == TaskState.PENDING
|
||||||
|
assert task.result_id is None
|
||||||
|
assert task.error_msg is None
|
||||||
|
|
||||||
|
def test_task_start_from_pending(self):
|
||||||
|
"""Test start() transitions PENDING to RUNNING"""
|
||||||
|
task = Task(task_id="test_1", task_type="test")
|
||||||
|
task.start()
|
||||||
|
assert task.status == TaskState.RUNNING
|
||||||
|
|
||||||
|
def test_task_start_from_running_raises_error(self):
|
||||||
|
"""Test start() from RUNNING raises ValueError"""
|
||||||
|
task = Task(task_id="test_1", task_type="test")
|
||||||
|
task.start()
|
||||||
|
with pytest.raises(ValueError, match="不能从 RUNNING 启动"):
|
||||||
|
task.start()
|
||||||
|
|
||||||
|
def test_task_start_from_done_raises_error(self):
|
||||||
|
"""Test start() from DONE raises ValueError"""
|
||||||
|
task = Task(task_id="test_1", task_type="test")
|
||||||
|
task.status = TaskState.DONE
|
||||||
|
with pytest.raises(ValueError, match="不能从 DONE 启动"):
|
||||||
|
task.start()
|
||||||
|
|
||||||
|
def test_task_start_from_failed_raises_error(self):
|
||||||
|
"""Test start() from FAILED raises ValueError"""
|
||||||
|
task = Task(task_id="test_1", task_type="test")
|
||||||
|
task.status = TaskState.FAILED
|
||||||
|
with pytest.raises(ValueError, match="不能从 FAILED 启动"):
|
||||||
|
task.start()
|
||||||
|
|
||||||
|
def test_task_complete_transitions_to_done(self):
|
||||||
|
"""Test complete() transitions to DONE and sets result_id"""
|
||||||
|
task = Task(task_id="test_1", task_type="test")
|
||||||
|
task.complete(result_id=12345)
|
||||||
|
assert task.status == TaskState.DONE
|
||||||
|
assert task.result_id == 12345
|
||||||
|
|
||||||
|
def test_task_fail_transitions_to_failed(self):
|
||||||
|
"""Test fail() transitions to FAILED and sets error_msg"""
|
||||||
|
task = Task(task_id="test_1", task_type="test")
|
||||||
|
task.fail(error_msg="Test error")
|
||||||
|
assert task.status == TaskState.FAILED
|
||||||
|
assert task.error_msg == "Test error"
|
||||||
|
assert task.result_id is None
|
||||||
Reference in New Issue
Block a user