""" 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 == "回测中"