9201ef759e
Harness Compat / harness compat (push) Failing after 0s
CI / test on 3.12 (standard) (push) Has been cancelled
CI / test on 3.13 (standard) (push) Has been cancelled
CI / test on 3.14 (standard) (push) Has been cancelled
CI / test on 3.10 (all-extras) (push) Has been cancelled
CI / test on 3.11 (all-extras) (push) Has been cancelled
CI / test on 3.12 (all-extras) (push) Has been cancelled
CI / test on 3.14 (pydantic-ai-slim) (push) Has been cancelled
CI / test on 3.10 (pydantic-evals) (push) Has been cancelled
CI / test on 3.11 (pydantic-evals) (push) Has been cancelled
CI / test on 3.12 (pydantic-evals) (push) Has been cancelled
CI / deploy-docs-preview (push) Has been cancelled
CI / build release artifacts (push) Has been cancelled
CI / publish to PyPI (push) Has been cancelled
CI / Send tweet (push) Has been cancelled
CI / lint (push) Has been cancelled
CI / mypy (push) Has been cancelled
CI / docs (push) Has been cancelled
CI / test on 3.10 (standard) (push) Has been cancelled
CI / test on 3.11 (standard) (push) Has been cancelled
CI / test on 3.13 (all-extras) (push) Has been cancelled
CI / test on 3.14 (all-extras) (push) Has been cancelled
CI / test on 3.10 (pydantic-ai-slim) (push) Has been cancelled
CI / test on 3.11 (pydantic-ai-slim) (push) Has been cancelled
CI / test on 3.12 (pydantic-ai-slim) (push) Has been cancelled
CI / test on 3.13 (pydantic-ai-slim) (push) Has been cancelled
CI / test on 3.13 (pydantic-evals) (push) Has been cancelled
CI / test on 3.14 (pydantic-evals) (push) Has been cancelled
CI / test on 3.10 (lowest-versions) (push) Has been cancelled
CI / test on 3.11 (lowest-versions) (push) Has been cancelled
CI / test on 3.12 (lowest-versions) (push) Has been cancelled
CI / test on 3.13 (lowest-versions) (push) Has been cancelled
CI / test on 3.14 (lowest-versions) (push) Has been cancelled
CI / test examples on 3.11 (push) Has been cancelled
CI / test examples on 3.12 (push) Has been cancelled
CI / test examples on 3.13 (push) Has been cancelled
CI / test examples on 3.14 (push) Has been cancelled
CI / coverage (push) Has been cancelled
CI / check (push) Has been cancelled
CI / deploy-docs (push) Has been cancelled
386 lines
12 KiB
Python
386 lines
12 KiB
Python
"""Tests for broadcast (parallel) and map (fan-out) operations."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from dataclasses import dataclass, field
|
|
|
|
import anyio
|
|
import pytest
|
|
|
|
from pydantic_graph import GraphBuilder, StepContext
|
|
from pydantic_graph.join import reduce_list_append
|
|
|
|
pytestmark = pytest.mark.anyio
|
|
|
|
|
|
@dataclass
|
|
class CounterState:
|
|
values: list[int] = field(default_factory=list[int])
|
|
|
|
|
|
async def test_broadcast_to_multiple_steps():
|
|
"""Test broadcasting the same data to multiple parallel steps."""
|
|
g = GraphBuilder(state_type=CounterState, output_type=list[int])
|
|
|
|
@g.step
|
|
async def source(ctx: StepContext[CounterState, None, None]) -> int:
|
|
return 10
|
|
|
|
@g.step
|
|
async def add_one(ctx: StepContext[CounterState, None, int]) -> int:
|
|
return ctx.inputs + 1
|
|
|
|
@g.step
|
|
async def add_two(ctx: StepContext[CounterState, None, int]) -> int:
|
|
return ctx.inputs + 2
|
|
|
|
@g.step
|
|
async def add_three(ctx: StepContext[CounterState, None, int]) -> int:
|
|
return ctx.inputs + 3
|
|
|
|
collect = g.join(reduce_list_append, initial_factory=list[int])
|
|
|
|
g.add(
|
|
g.edge_from(g.start_node).to(source),
|
|
g.edge_from(source).to(add_one, add_two, add_three),
|
|
g.edge_from(add_one, add_two, add_three).to(collect),
|
|
g.edge_from(collect).to(g.end_node),
|
|
)
|
|
|
|
graph = g.build()
|
|
result = await graph.run(state=CounterState())
|
|
# Results can be in any order due to parallel execution
|
|
assert sorted(result) == [11, 12, 13]
|
|
|
|
|
|
async def test_map_over_list():
|
|
"""Test mapping a list to process items in parallel."""
|
|
g = GraphBuilder(state_type=CounterState, output_type=list[int])
|
|
|
|
@g.step
|
|
async def generate_list(ctx: StepContext[CounterState, None, None]) -> list[int]:
|
|
return [1, 2, 3, 4, 5]
|
|
|
|
@g.step
|
|
async def square(ctx: StepContext[CounterState, None, int]) -> int:
|
|
return ctx.inputs * ctx.inputs
|
|
|
|
collect = g.join(reduce_list_append, initial_factory=list[int])
|
|
|
|
g.add_mapping_edge(generate_list, square)
|
|
g.add(
|
|
g.edge_from(g.start_node).to(generate_list),
|
|
g.edge_from(square).to(collect),
|
|
g.edge_from(collect).to(g.end_node),
|
|
)
|
|
|
|
graph = g.build()
|
|
result = await graph.run(state=CounterState())
|
|
assert sorted(result) == [1, 4, 9, 16, 25]
|
|
|
|
|
|
async def test_map_with_labels():
|
|
"""Test map operation with labeled edges."""
|
|
g = GraphBuilder(state_type=CounterState, output_type=list[str])
|
|
|
|
@g.step
|
|
async def generate_numbers(ctx: StepContext[CounterState, None, None]) -> list[int]:
|
|
return [10, 20, 30]
|
|
|
|
@g.step
|
|
async def stringify(ctx: StepContext[CounterState, None, int]) -> str:
|
|
return f'Value: {ctx.inputs}'
|
|
|
|
collect = g.join(reduce_list_append, initial_factory=list[str])
|
|
|
|
g.add_mapping_edge(
|
|
generate_numbers,
|
|
stringify,
|
|
pre_map_label='before map',
|
|
post_map_label='after map',
|
|
)
|
|
g.add(
|
|
g.edge_from(g.start_node).to(generate_numbers),
|
|
g.edge_from(stringify).to(collect),
|
|
g.edge_from(collect).to(g.end_node),
|
|
)
|
|
|
|
graph = g.build()
|
|
result = await graph.run(state=CounterState())
|
|
assert sorted(result) == ['Value: 10', 'Value: 20', 'Value: 30']
|
|
|
|
|
|
async def test_map_empty_list():
|
|
"""Test mapping an empty list."""
|
|
g = GraphBuilder(state_type=CounterState, output_type=list[int])
|
|
|
|
@g.step
|
|
async def generate_empty(ctx: StepContext[CounterState, None, None]) -> list[int]:
|
|
return []
|
|
|
|
@g.step
|
|
async def double(ctx: StepContext[CounterState, None, int]) -> int:
|
|
return ctx.inputs * 2 # pragma: no cover
|
|
|
|
collect = g.join(reduce_list_append, initial_factory=list[int])
|
|
|
|
g.add_mapping_edge(generate_empty, double, downstream_join_id=collect.id)
|
|
g.add(
|
|
g.edge_from(g.start_node).to(generate_empty),
|
|
g.edge_from(double).to(collect),
|
|
g.edge_from(collect).to(g.end_node),
|
|
)
|
|
|
|
graph = g.build()
|
|
result = await graph.run(state=CounterState())
|
|
assert result == []
|
|
|
|
|
|
async def test_map_non_empty_list_with_downstream_join_id():
|
|
"""A non-empty map with `downstream_join_id` must fire the join once with the fully reduced value."""
|
|
g = GraphBuilder(state_type=CounterState, output_type=list[int])
|
|
|
|
@g.step
|
|
async def produce(ctx: StepContext[CounterState, None, None]) -> list[int]:
|
|
return [1, 2, 3]
|
|
|
|
@g.step
|
|
async def identity(ctx: StepContext[CounterState, None, int]) -> int:
|
|
return ctx.inputs
|
|
|
|
collect = g.join(reduce_list_append, initial_factory=list[int])
|
|
|
|
@g.step
|
|
async def after_join(ctx: StepContext[CounterState, None, list[int]]) -> list[int]:
|
|
ctx.state.values.append(len(ctx.inputs))
|
|
return ctx.inputs
|
|
|
|
g.add(g.edge_from(g.start_node).to(produce))
|
|
g.add_mapping_edge(produce, identity, downstream_join_id=collect.id)
|
|
g.add(
|
|
g.edge_from(identity).to(collect),
|
|
g.edge_from(collect).to(after_join),
|
|
g.edge_from(after_join).to(g.end_node),
|
|
)
|
|
|
|
graph = g.build()
|
|
state = CounterState()
|
|
result = await graph.run(state=state)
|
|
assert sorted(result) == [1, 2, 3]
|
|
# `after_join` fired exactly once, with the fully reduced list
|
|
assert state.values == [3]
|
|
|
|
|
|
async def test_parallel_maps_with_downstream_join_id():
|
|
"""Two independent map-and-join pipelines running in parallel, each with its own `downstream_join_id`.
|
|
|
|
Each join uses `preferred_parent_fork='closest'` so its reducer is keyed to its own map's fork run
|
|
rather than to the shared upstream broadcast; the two reducers therefore have distinct fork run ids and
|
|
can be active at the same time. When that happens, completing a task in one pipeline triggers a
|
|
completion check that encounters the *other* pipeline's reducer, whose parent fork run is absent from
|
|
the completing task's fork stack. That reducer must simply be skipped, and both pipelines must still
|
|
reduce correctly.
|
|
|
|
The two events force the completion order `a[0]` then `b` then `a[1]`, so pipeline A's reducer
|
|
is still active (awaiting its second item) when pipeline B's item completes — the only situation in
|
|
which the completion check sees a reducer keyed to a fork run outside the completing task's stack.
|
|
"""
|
|
g = GraphBuilder(state_type=CounterState, output_type=list[list[int]])
|
|
|
|
a_started = anyio.Event()
|
|
b_done = anyio.Event()
|
|
|
|
@g.step
|
|
async def seed(ctx: StepContext[CounterState, None, None]) -> int:
|
|
return 0
|
|
|
|
@g.step
|
|
async def produce_a(ctx: StepContext[CounterState, None, int]) -> list[int]:
|
|
return [1, 2]
|
|
|
|
@g.step
|
|
async def produce_b(ctx: StepContext[CounterState, None, int]) -> list[int]:
|
|
return [3]
|
|
|
|
@g.step
|
|
async def identity_a(ctx: StepContext[CounterState, None, int]) -> int:
|
|
if ctx.inputs == 1:
|
|
a_started.set() # pipeline A's reducer is about to be created
|
|
else:
|
|
await b_done.wait() # hold A's second item until pipeline B has completed
|
|
return ctx.inputs
|
|
|
|
@g.step
|
|
async def identity_b(ctx: StepContext[CounterState, None, int]) -> int:
|
|
await a_started.wait() # let pipeline A's reducer come into existence first
|
|
b_done.set()
|
|
return ctx.inputs
|
|
|
|
join_a = g.join(reduce_list_append, initial_factory=list[int], preferred_parent_fork='closest')
|
|
join_b = g.join(reduce_list_append, initial_factory=list[int], preferred_parent_fork='closest')
|
|
combine = g.join(reduce_list_append, initial_factory=list[list[int]])
|
|
|
|
@g.step
|
|
async def emit_a(ctx: StepContext[CounterState, None, list[int]]) -> list[int]:
|
|
return sorted(ctx.inputs)
|
|
|
|
@g.step
|
|
async def emit_b(ctx: StepContext[CounterState, None, list[int]]) -> list[int]:
|
|
return sorted(ctx.inputs)
|
|
|
|
g.add(
|
|
g.edge_from(g.start_node).to(seed),
|
|
g.edge_from(seed).to(produce_a, produce_b),
|
|
)
|
|
g.add_mapping_edge(produce_a, identity_a, downstream_join_id=join_a.id)
|
|
g.add_mapping_edge(produce_b, identity_b, downstream_join_id=join_b.id)
|
|
g.add(
|
|
g.edge_from(identity_a).to(join_a),
|
|
g.edge_from(identity_b).to(join_b),
|
|
g.edge_from(join_a).to(emit_a),
|
|
g.edge_from(join_b).to(emit_b),
|
|
g.edge_from(emit_a).to(combine),
|
|
g.edge_from(emit_b).to(combine),
|
|
g.edge_from(combine).to(g.end_node),
|
|
)
|
|
|
|
graph = g.build()
|
|
result = await graph.run(state=CounterState())
|
|
assert sorted(result) == [[1, 2], [3]]
|
|
|
|
|
|
async def test_nested_broadcasts():
|
|
"""Test nested broadcast operations."""
|
|
g = GraphBuilder(state_type=CounterState, output_type=list[int])
|
|
|
|
@g.step
|
|
async def start_value(ctx: StepContext[CounterState, None, None]) -> int:
|
|
return 5
|
|
|
|
@g.step
|
|
async def path_a1(ctx: StepContext[CounterState, None, int]) -> int:
|
|
return ctx.inputs + 1
|
|
|
|
@g.step
|
|
async def path_a2(ctx: StepContext[CounterState, None, int]) -> int:
|
|
return ctx.inputs + 10
|
|
|
|
@g.step
|
|
async def path_b1(ctx: StepContext[CounterState, None, int]) -> int:
|
|
return ctx.inputs * 2
|
|
|
|
@g.step
|
|
async def path_b2(ctx: StepContext[CounterState, None, int]) -> int:
|
|
return ctx.inputs * 3
|
|
|
|
collect = g.join(reduce_list_append, initial_factory=list[int])
|
|
|
|
g.add(
|
|
g.edge_from(g.start_node).to(start_value),
|
|
g.edge_from(start_value).to(path_a1, path_b1),
|
|
g.edge_from(path_a1).to(path_a2),
|
|
g.edge_from(path_b1).to(path_b2),
|
|
g.edge_from(path_a2, path_b2).to(collect),
|
|
g.edge_from(collect).to(g.end_node),
|
|
)
|
|
|
|
graph = g.build()
|
|
result = await graph.run(state=CounterState())
|
|
# path_a: 5 + 1 + 10 = 16
|
|
# path_b: 5 * 2 * 3 = 30
|
|
assert sorted(result) == [16, 30]
|
|
|
|
|
|
async def test_map_then_broadcast():
|
|
"""Test mapping followed by broadcasting from each map item."""
|
|
g = GraphBuilder(state_type=CounterState, output_type=list[int])
|
|
|
|
@g.step
|
|
async def generate_list(ctx: StepContext[CounterState, None, None]) -> list[int]:
|
|
return [10, 20]
|
|
|
|
@g.step
|
|
async def add_one(ctx: StepContext[CounterState, None, int]) -> int:
|
|
return ctx.inputs + 1
|
|
|
|
@g.step
|
|
async def add_two(ctx: StepContext[CounterState, None, int]) -> int:
|
|
return ctx.inputs + 2
|
|
|
|
collect = g.join(reduce_list_append, initial_factory=list[int])
|
|
|
|
g.add(
|
|
g.edge_from(g.start_node).to(generate_list),
|
|
g.edge_from(generate_list).map().to(add_one, add_two),
|
|
g.edge_from(add_one, add_two).to(collect),
|
|
g.edge_from(collect).to(g.end_node),
|
|
)
|
|
|
|
graph = g.build()
|
|
result = await graph.run(state=CounterState())
|
|
# From 10: 11, 12
|
|
# From 20: 21, 22
|
|
assert sorted(result) == [11, 12, 21, 22]
|
|
|
|
|
|
async def test_multiple_sequential_maps():
|
|
"""Test multiple sequential map operations."""
|
|
g = GraphBuilder(state_type=CounterState, output_type=list[str])
|
|
|
|
@g.step
|
|
async def generate_pairs(ctx: StepContext[CounterState, None, None]) -> list[tuple[int, int]]:
|
|
return [(1, 2), (3, 4)]
|
|
|
|
@g.step
|
|
async def unpack_pair(ctx: StepContext[CounterState, None, tuple[int, int]]) -> list[int]:
|
|
return [ctx.inputs[0], ctx.inputs[1]]
|
|
|
|
@g.step
|
|
async def stringify(ctx: StepContext[CounterState, None, int]) -> str:
|
|
return f'num:{ctx.inputs}'
|
|
|
|
collect = g.join(reduce_list_append, initial_factory=list[str])
|
|
|
|
g.add(
|
|
g.edge_from(g.start_node).to(generate_pairs),
|
|
g.edge_from(generate_pairs).map().to(unpack_pair),
|
|
g.edge_from(unpack_pair).map().to(stringify),
|
|
g.edge_from(stringify).to(collect),
|
|
g.edge_from(collect).to(g.end_node),
|
|
)
|
|
|
|
graph = g.build()
|
|
result = await graph.run(state=CounterState())
|
|
assert sorted(result) == ['num:1', 'num:2', 'num:3', 'num:4']
|
|
|
|
|
|
async def test_broadcast_with_different_outputs():
|
|
"""Test that broadcasts can produce different types of outputs."""
|
|
g = GraphBuilder(state_type=CounterState, output_type=list[int | str])
|
|
|
|
@g.step
|
|
async def source(ctx: StepContext[CounterState, None, None]) -> int:
|
|
return 42
|
|
|
|
@g.step
|
|
async def return_int(ctx: StepContext[CounterState, None, int]) -> int:
|
|
return ctx.inputs
|
|
|
|
@g.step
|
|
async def return_str(ctx: StepContext[CounterState, None, int]) -> str:
|
|
return str(ctx.inputs)
|
|
|
|
collect = g.join(reduce_list_append, initial_factory=list[int | str])
|
|
|
|
g.add(
|
|
g.edge_from(g.start_node).to(source),
|
|
g.edge_from(source).to(return_int, return_str),
|
|
g.edge_from(return_int, return_str).to(collect),
|
|
g.edge_from(collect).to(g.end_node),
|
|
)
|
|
|
|
graph = g.build()
|
|
result = await graph.run(state=CounterState())
|
|
# Order may vary
|
|
assert set(result) == {42, '42'}
|