94 lines
3.7 KiB
Python
94 lines
3.7 KiB
Python
"""
|
|
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
|
|
|
|
|
|
class TestTaskPoolStageAndAsync:
|
|
"""Test TaskPool stage tracking and async execution (Task 3)"""
|
|
|
|
def test_task_has_stage_field(self):
|
|
"""Test Task has stage field with default empty string"""
|
|
from sanguo_orchestrator.task import Task
|
|
t = Task(task_id="t1", task_type="cta")
|
|
assert t.stage == ""
|
|
|
|
def test_pool_submit_work_returns_future(self):
|
|
"""Test submit_work() returns Future from executor"""
|
|
from unittest.mock import MagicMock
|
|
from sanguo_orchestrator.pool import TaskPool
|
|
pool = TaskPool(max_workers=2)
|
|
pool.executor = MagicMock() # mock executor to avoid spawning real processes
|
|
mock_future = MagicMock()
|
|
pool.executor.submit.return_value = mock_future
|
|
fut = pool.submit_work("t1", func=lambda: 1)
|
|
assert fut is mock_future
|
|
pool.executor.submit.assert_called_once()
|
|
|
|
def test_pool_update_and_get_stage(self):
|
|
"""Test update_stage() and get_stage() methods"""
|
|
from sanguo_orchestrator.pool import TaskPool
|
|
from sanguo_orchestrator.task import Task
|
|
pool = TaskPool(max_workers=2)
|
|
pool.submit("t1", "cta")
|
|
pool.update_stage("t1", "回测中")
|
|
assert pool.get_stage("t1") == "回测中"
|
|
assert pool.get_task("t1").stage == "回测中" |