diff --git a/diagrams/__init__.py b/diagrams/__init__.py index 1fa343b0..eef4edbf 100644 --- a/diagrams/__init__.py +++ b/diagrams/__init__.py @@ -196,11 +196,11 @@ class Diagram: if not self._validate_direction(direction): raise ValueError(f'"{direction}" is not a valid direction') - self.dot.graph_attr["rankdir"] = direction + self.dot.graph_attr["rankdir"] = direction.upper() if not self._validate_curvestyle(curvestyle): raise ValueError(f'"{curvestyle}" is not a valid curvestyle') - self.dot.graph_attr["splines"] = curvestyle + self.dot.graph_attr["splines"] = curvestyle.lower() if isinstance(outformat, list): for one_format in outformat: diff --git a/tests/test_diagram.py b/tests/test_diagram.py index 75fc02b8..32564555 100644 --- a/tests/test_diagram.py +++ b/tests/test_diagram.py @@ -43,6 +43,11 @@ class DiagramTest(unittest.TestCase): with self.assertRaises(ValueError): Diagram(curvestyle=cvs) + def test_case_insensitive_graph_attributes(self): + diagram = Diagram(direction="lr", curvestyle="CURVED") + self.assertEqual(diagram.dot.graph_attr["rankdir"], "LR") + self.assertEqual(diagram.dot.graph_attr["splines"], "curved") + def test_validate_theme(self): # Valid themes. for theme in ("neutral", "pastel", "blues", "greens", "orange"):