diff --git a/src/comfydv/format_string.py b/src/comfydv/format_string.py index 49eea2a..2ba0f40 100644 --- a/src/comfydv/format_string.py +++ b/src/comfydv/format_string.py @@ -21,7 +21,7 @@ import sys from typing import Any, Dict, List from aiohttp import web -from jinja2 import exceptions, sandbox +from jinja2 import exceptions, meta, sandbox # Set up logger for this module logger = logging.getLogger(__name__) @@ -252,38 +252,48 @@ class FormatString: >>> # Test with additional context (should be excluded) >>> keys = FormatString._extract_keys("Time: {{ datetime.now() }}") >>> assert keys == [] + >>> # Test {% for %} control structures: the loop variable is bound by the + >>> # template itself and must not be treated as a required input, while the + >>> # iterable it draws from must be. + >>> keys = FormatString._extract_keys( + ... "{% for hint in extraction_hints %}{{ hint }}{% endfor %}" + ... ) + >>> assert keys == ['extraction_hints'] --> """ variables = [] seen = set() def add_var(var): - var = var.split("|")[0].split(".")[0].strip() if var not in seen and var not in FormatString.additional_context: seen.add(var) variables.append(var) - # Extract variables from Jinja2 expressions {{ }} - for match in re.finditer( - r"\{\{\s*([\w.]+)(?:\s*\|[\w\s]+)?(?:\.[^\(\)]+\(\))?\s*\}\}", template - ): - add_var(match.group(1)) - - # Extract variables from f-string style { } + # Extract variables from Python str.format() style { } for match in re.finditer(r"\{(\w+)\}", template): add_var(match.group(1)) - # Extract variables from Jinja2 control structures {% %} - for structure in re.finditer(r"\{%.*?%\}", template): - for var in re.findall(r"\b(\w+)\|\b", structure.group(0)): - if not var.startswith("end") and var not in { - "if", - "else", - "elif", - "for", - "in", - }: - add_var(var) + # Extract variables referenced anywhere in Jinja2 syntax ({{ }} expressions + # and {% %} control structures) by parsing the template with Jinja2 itself + # rather than approximating it with regexes. This is what correctly excludes + # names bound within the template (e.g. the `hint` loop variable in + # `{% for hint in extraction_hints %}`) while still surfacing names the + # template expects the caller to supply (e.g. `extraction_hints`). + try: + template_ast = FormatString.jinja_env.parse(template) + except exceptions.TemplateSyntaxError: + pass + else: + # find_undeclared_variables returns an unordered set; sort by first + # textual occurrence so extraction order is deterministic and matches + # the order the template reads left to right (callers rely on this + # for positional outputs, e.g. two {{ }} variables in sequence). + undeclared = sorted( + meta.find_undeclared_variables(template_ast), + key=lambda name: template.find(name), + ) + for var in undeclared: + add_var(var) return variables diff --git a/tests/conftest.py b/tests/conftest.py index 1128edb..a170dfa 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -60,6 +60,15 @@ def pytest_configure(config): sys.modules["folder_paths"] = MockFolderPaths + # Force "comfydv" to resolve to src/comfydv and get cached in sys.modules now, + # while our sys.path.insert(0, ...) above is still the definitive answer. The + # repo root's own __init__.py (ComfyUI's custom-node entry point) is also a + # valid "comfydv" package from certain sys.path states pytest transiently + # constructs during fixture setup; without this, a later bare `import comfydv` + # (e.g. in the _clear_ollama_caches fixture) can resolve to that root package + # instead, which lacks the _llm submodule and fails with ModuleNotFoundError. + import comfydv # noqa: F401 + # --------------------------------------------------------------------------- # Ollama fixtures (used by @pytest.mark.integration tests) diff --git a/tests/test_format_string.py b/tests/test_format_string.py index fdf86c3..e56a186 100644 --- a/tests/test_format_string.py +++ b/tests/test_format_string.py @@ -35,9 +35,32 @@ class TestVariableExtraction: def test_extract_jinja2_with_multiple_filters(self, format_string_class): """Test extraction of variables with multiple Jinja2 filters.""" keys = format_string_class._extract_keys("{{ name | upper | trim }}") - # Multiple chained filters may not extract - that's a limitation of the regex - # Just test that it doesn't crash - assert isinstance(keys, list) + assert keys == ["name"] + + def test_extract_jinja2_for_loop_excludes_loop_variable(self, format_string_class): + """The for-loop target (e.g. `hint`) is bound by the template and must + not be treated as a required input; the iterable it draws from must be.""" + keys = format_string_class._extract_keys( + "{% for hint in extraction_hints %}- {{ hint }}\n{% endfor %}" + ) + assert keys == ["extraction_hints"] + + def test_extract_jinja2_if_condition_variable(self, format_string_class): + """A variable referenced only in an {% if %} condition must still be + detected, even without a matching {{ }} expression elsewhere.""" + keys = format_string_class._extract_keys( + "{% if extraction_hints is defined and extraction_hints %}yes{% endif %}" + ) + assert keys == ["extraction_hints"] + + def test_extract_jinja2_filter_with_arguments(self, format_string_class): + """A filter called with arguments (e.g. tojson(indent=2)) has parens + in the way of the old regex's anchor to the closing }} — the variable + must still be detected.""" + keys = format_string_class._extract_keys( + "{{ scene_manifest | tojson(indent=2) }}" + ) + assert keys == ["scene_manifest"] def test_extract_jinja2_multiple_variables(self, format_string_class): """Test extraction of multiple variables from Jinja2 template.""" @@ -199,10 +222,11 @@ class TestJinja2Formatting: unique_id="test9", value=sample_data["value"], )["result"] - # value is not extracted as a variable because it's used in an expression - assert len(result) == 2 # Just formatted_string, saved_file_path + # value is extracted even though it's used in an expression + assert len(result) == 3 # formatted_string, saved_file_path, value assert result[0] == "Result: 10" assert result[1] == "" + assert result[2] == str(sample_data["value"]) class TestInlineDisplay: