"""Unit tests for task graph module.""" import pytest from abacus_common_logic.concurrent import Task, TaskGraph, TaskGraphError, TaskOutput class TestTask: """Tests for Task dataclass.""" def test_task_creation(self): """Test Task can be created with minimal parameters.""" def handler(x: int) -> int: return x * 2 task = Task(id='test', handler=handler) assert task.id == 'test' assert task.handler == handler assert task.params == {} assert task.depends_on == set() assert task.pool is None def test_task_with_dependencies(self): """Test Task with dependencies.""" def handler(): pass task = Task(id='task2', handler=handler, depends_on={'task1'}) assert task.depends_on == {'task1'} def test_task_with_params(self): """Test Task with parameters.""" def handler(x: int, y: int) -> int: return x + y task = Task(id='add', handler=handler, params={'x': 5, 'y': 10}) assert task.params == {'x': 5, 'y': 10} def test_task_with_pool(self): """Test Task with execution pool.""" def handler(): pass task = Task(id='task1', handler=handler, pool='io_pool') assert task.pool == 'io_pool' def test_task_frozen(self): """Test Task is frozen (immutable).""" def handler(): pass task = Task(id='test', handler=handler) with pytest.raises(AttributeError): task.id = 'new_id' class TestTaskOutput: """Tests for TaskOutput placeholder.""" def test_task_output_creation(self): """Test TaskOutput creation.""" output = TaskOutput(task_id='fetch') assert output.task_id == 'fetch' assert output.path is None def test_task_output_with_path(self): """Test TaskOutput with path.""" output = TaskOutput(task_id='fetch', path='results.data') assert output.task_id == 'fetch' assert output.path == 'results.data' class TestTaskGraph: """Tests for TaskGraph validation.""" def test_simple_linear_graph(self): """Test simple linear dependency chain.""" def handler(): pass tasks = [ Task(id='task1', handler=handler), Task(id='task2', handler=handler, depends_on={'task1'}), Task(id='task3', handler=handler, depends_on={'task2'}), ] graph = TaskGraph(tasks) assert len(graph.tasks) == 3 assert graph.in_degrees == {'task1': 0, 'task2': 1, 'task3': 1} def test_parallel_tasks(self): """Test parallel independent tasks.""" def handler(): pass tasks = [ Task(id='task1', handler=handler), Task(id='task2', handler=handler), Task(id='task3', handler=handler), ] graph = TaskGraph(tasks) assert all(deg == 0 for deg in graph.in_degrees.values()) def test_diamond_dependency(self): """Test diamond-shaped dependency graph.""" def handler(): pass tasks = [ Task(id='start', handler=handler), Task(id='left', handler=handler, depends_on={'start'}), Task(id='right', handler=handler, depends_on={'start'}), Task(id='end', handler=handler, depends_on={'left', 'right'}), ] graph = TaskGraph(tasks) assert graph.in_degrees == {'start': 0, 'left': 1, 'right': 1, 'end': 2} assert set(graph.adjacency_list['start']) == {'left', 'right'} assert set(graph.adjacency_list['left']) == {'end'} assert set(graph.adjacency_list['right']) == {'end'} def test_duplicate_task_ids_raises_error(self): """Test duplicate task IDs raise ValueError.""" def handler(): pass tasks = [ Task(id='task1', handler=handler), Task(id='task1', handler=handler), ] with pytest.raises(ValueError, match='Duplicate task IDs found: task1'): TaskGraph(tasks) def test_multiple_duplicate_task_ids(self): """Test multiple duplicate task IDs are all reported.""" def handler(): pass tasks = [ Task(id='task1', handler=handler), Task(id='task1', handler=handler), Task(id='task2', handler=handler), Task(id='task2', handler=handler), ] with pytest.raises(ValueError, match='Duplicate task IDs found'): TaskGraph(tasks) def test_unknown_dependency_raises_error(self): """Test depending on unknown task raises ValueError.""" def handler(): pass tasks = [ Task(id='task1', handler=handler), Task(id='task2', handler=handler, depends_on={'nonexistent'}), ] with pytest.raises( ValueError, match="Task 'task2' depends on unknown task 'nonexistent'" ): TaskGraph(tasks) def test_simple_cycle_detection(self): """Test simple two-task cycle is detected.""" def handler(): pass tasks = [ Task(id='task1', handler=handler, depends_on={'task2'}), Task(id='task2', handler=handler, depends_on={'task1'}), ] with pytest.raises(ValueError, match='Graph cycle\\(s\\) detected'): TaskGraph(tasks) def test_three_task_cycle_detection(self): """Test three-task cycle is detected.""" def handler(): pass tasks = [ Task(id='task1', handler=handler, depends_on={'task3'}), Task(id='task2', handler=handler, depends_on={'task1'}), Task(id='task3', handler=handler, depends_on={'task2'}), ] with pytest.raises(ValueError, match='Graph cycle\\(s\\) detected'): TaskGraph(tasks) def test_self_cycle_detection(self): """Test self-referencing task is detected.""" def handler(): pass tasks = [Task(id='task1', handler=handler, depends_on={'task1'})] with pytest.raises(ValueError, match='Graph cycle\\(s\\) detected'): TaskGraph(tasks) def test_complex_graph_with_multiple_paths(self): """Test complex graph with multiple convergent paths.""" def handler(): pass tasks = [ Task(id='a', handler=handler), Task(id='b', handler=handler, depends_on={'a'}), Task(id='c', handler=handler, depends_on={'a'}), Task(id='d', handler=handler, depends_on={'b', 'c'}), Task(id='e', handler=handler, depends_on={'b'}), Task(id='f', handler=handler, depends_on={'d', 'e'}), ] graph = TaskGraph(tasks) assert graph.in_degrees == {'a': 0, 'b': 1, 'c': 1, 'd': 2, 'e': 1, 'f': 2} def test_is_ancestor_direct_parent(self): """Test is_ancestor detects direct parent.""" def handler(): pass tasks = [ Task(id='parent', handler=handler), Task(id='child', handler=handler, depends_on={'parent'}), ] graph = TaskGraph(tasks) assert graph.is_ancestor('parent', 'child') is True assert graph.is_ancestor('child', 'parent') is False def test_is_ancestor_grandparent(self): """Test is_ancestor detects grandparent relationship.""" def handler(): pass tasks = [ Task(id='grandparent', handler=handler), Task(id='parent', handler=handler, depends_on={'grandparent'}), Task(id='child', handler=handler, depends_on={'parent'}), ] graph = TaskGraph(tasks) assert graph.is_ancestor('grandparent', 'child') is True assert graph.is_ancestor('parent', 'child') is True assert graph.is_ancestor('child', 'grandparent') is False def test_is_ancestor_with_siblings(self): """Test is_ancestor with sibling tasks.""" def handler(): pass tasks = [ Task(id='parent', handler=handler), Task(id='child1', handler=handler, depends_on={'parent'}), Task(id='child2', handler=handler, depends_on={'parent'}), ] graph = TaskGraph(tasks) assert graph.is_ancestor('parent', 'child1') is True assert graph.is_ancestor('parent', 'child2') is True assert graph.is_ancestor('child1', 'child2') is False assert graph.is_ancestor('child2', 'child1') is False def test_is_ancestor_diamond_pattern(self): """Test is_ancestor in diamond dependency pattern.""" def handler(): pass tasks = [ Task(id='start', handler=handler), Task(id='left', handler=handler, depends_on={'start'}), Task(id='right', handler=handler, depends_on={'start'}), Task(id='end', handler=handler, depends_on={'left', 'right'}), ] graph = TaskGraph(tasks) assert graph.is_ancestor('start', 'end') is True assert graph.is_ancestor('left', 'end') is True assert graph.is_ancestor('right', 'end') is True def test_task_output_valid_data_flow(self): """Test valid TaskOutput reference to ancestor.""" def handler(): pass tasks = [ Task(id='fetch', handler=handler), Task( id='process', handler=handler, depends_on={'fetch'}, params={'data': TaskOutput('fetch')}, ), ] # Should not raise TaskGraph(tasks) def test_task_output_invalid_non_ancestor(self): """Test TaskOutput reference to non-ancestor raises error.""" def handler(): pass tasks = [ Task(id='task1', handler=handler), Task( id='task2', handler=handler, params={'data': TaskOutput('task1')}, # No dependency! ), ] with pytest.raises( ValueError, match="Task 'task2' requests output from non-ancestor task 'task1'", ): TaskGraph(tasks) def test_task_output_nonexistent_task(self): """Test TaskOutput reference to nonexistent task raises error.""" def handler(): pass tasks = [ Task( id='task1', handler=handler, params={'data': TaskOutput('nonexistent')}, ), ] with pytest.raises( ValueError, match="Task 'task1' requests output from non-existent task 'nonexistent'", ): TaskGraph(tasks) def test_task_output_with_sibling(self): """Test TaskOutput cannot reference sibling task.""" def handler(): pass tasks = [ Task(id='parent', handler=handler), Task(id='child1', handler=handler, depends_on={'parent'}), Task( id='child2', handler=handler, depends_on={'parent'}, params={'data': TaskOutput('child1')}, ), ] with pytest.raises(ValueError, match='non-ancestor task'): TaskGraph(tasks) def test_task_output_valid_in_chain(self): """Test TaskOutput valid in a chain of dependencies.""" def handler(): pass tasks = [ Task(id='task1', handler=handler), Task( id='task2', handler=handler, depends_on={'task1'}, params={'input': TaskOutput('task1')}, ), Task( id='task3', handler=handler, depends_on={'task2'}, params={'input': TaskOutput('task2')}, ), ] # Should not raise graph = TaskGraph(tasks) assert len(graph.tasks) == 3 def test_task_output_can_reference_distant_ancestor(self): """Test TaskOutput can reference any ancestor, not just parent.""" def handler(): pass tasks = [ Task(id='task1', handler=handler), Task(id='task2', handler=handler, depends_on={'task1'}), Task( id='task3', handler=handler, depends_on={'task2'}, params={'data': TaskOutput('task1')}, # Skip task2 ), ] # Should not raise graph = TaskGraph(tasks) assert graph.is_ancestor('task1', 'task3') def test_empty_graph(self): """Test empty task list.""" graph = TaskGraph([]) assert len(graph.tasks) == 0 assert len(graph.adjacency_list) == 0 assert len(graph.in_degrees) == 0 def test_single_task(self): """Test graph with single task.""" def handler(): pass tasks = [Task(id='only', handler=handler)] graph = TaskGraph(tasks) assert len(graph.tasks) == 1 assert graph.in_degrees == {'only': 0} class TestTaskGraphError: """Tests for TaskGraphError exception.""" def test_task_graph_error_creation(self): """Test TaskGraphError creation.""" original = ValueError('Something went wrong') error = TaskGraphError('task123', original) assert error.task_id == 'task123' assert str(error) == "Task 'task123' failed: Something went wrong" def test_task_graph_error_preserves_original(self): """Test TaskGraphError preserves original exception.""" original = RuntimeError('Database connection failed') error = TaskGraphError('db_task', original) assert 'db_task' in str(error) assert 'Database connection failed' in str(error)