diff --git a/burr/core/graph.py b/burr/core/graph.py index 5889c41f6..700ac4cad 100644 --- a/burr/core/graph.py +++ b/burr/core/graph.py @@ -248,8 +248,11 @@ def visualize( ), ) for g_key, g_value in engine_kwargs.items(): + if g_value is None and g_key.endswith("_attr"): + # Same as not passing it: keep Burr's defaults. + continue if isinstance(g_value, dict): - digraph_attr[g_key].update(**g_value) + digraph_attr.setdefault(g_key, {}).update(**g_value) else: digraph_attr[g_key] = g_value digraph = graphviz.Digraph(**digraph_attr) diff --git a/tests/core/test_graphviz_display.py b/tests/core/test_graphviz_display.py index 7eda47cba..fe0109221 100644 --- a/tests/core/test_graphviz_display.py +++ b/tests/core/test_graphviz_display.py @@ -19,6 +19,7 @@ import pytest +from burr.core import ApplicationBuilder from burr.core.graph import GraphBuilder from tests.core.test_graph import PassedInAction @@ -99,3 +100,38 @@ def test_visualize_include_state_multiline_label(reads: list, writes: list, expe digraph = graph.visualize(include_state=True) assert expected_label in digraph.source + + +def test_visualize_engine_kwargs_attr_dicts(graph): + """Attribute dicts passed through ``engine_kwargs`` reach the graphviz.Digraph, + including ones (like ``edge_attr``) that have no Burr default to merge into.""" + digraph = graph.visualize( + graph_attr={"rankdir": "LR"}, + node_attr={"fontname": "Courier"}, + edge_attr={"color": "red"}, + ) + + assert digraph.graph_attr["rankdir"] == "LR" + assert digraph.graph_attr["ranksep"] == "0.4" # Burr default is kept + assert digraph.node_attr["fontname"] == "Courier" + assert digraph.node_attr["fillcolor"] == "#b4d8e4" # Burr default is kept + assert digraph.edge_attr == {"color": "red"} + + +def test_application_visualize_edge_attr(graph): + """``Application.visualize`` is the path users call; it forwards ``edge_attr``.""" + app = ApplicationBuilder().with_graph(graph).with_entrypoint("counter").build() + + digraph = app.visualize(edge_attr={"color": "red"}) + + assert digraph.edge_attr == {"color": "red"} + assert digraph.graph_attr["rankdir"] == "TB" # Burr default is kept + + +def test_visualize_none_attr_keeps_defaults(graph): + """Passing ``graph_attr=None`` behaves like not passing it.""" + digraph = graph.visualize(graph_attr=None, node_attr=None, edge_attr=None) + + assert digraph.graph_attr["rankdir"] == "TB" + assert digraph.node_attr["fillcolor"] == "#b4d8e4" + assert digraph.edge_attr == {}