Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@

from __future__ import annotations

import re
from enum import Enum
from fnmatch import fnmatch
from string import Formatter
Expand Down Expand Up @@ -450,9 +451,20 @@ def validate_skip_references(
return violations


_FORMAT_FIELD_REFERENCE_PATTERN = re.compile(r"(?<!\{)\{\s*(\w+)\s*(?:![sra])?(?::[^{}]*(?:\{[^{}]*\}[^{}]*)*)?\}(?!\})")


def _get_string_formatter_references(template: str, allowed_references: list[str]) -> list[str]:
return [
k[1].strip()
for k in Formatter().parse(template)
if len(k) > 1 and k[1] is not None and k[1].strip() in allowed_references
]
try:
return [
k[1].strip()
for k in Formatter().parse(template)
if len(k) > 1 and k[1] is not None and k[1].strip() in allowed_references
]
except ValueError:
# Unmatched literal braces (e.g. JSON examples like 'output format: }') are
# invalid f-string syntax but valid Jinja text, so ``Formatter().parse`` can
# raise ``ValueError``. Fall back to a tolerant regex scan that skips Jinja
# ``{{ ... }}`` expressions so the advisory check still detects ``{column}``
# references without crashing validation.
return [m for m in _FORMAT_FIELD_REFERENCE_PATTERN.findall(template) if m in allowed_references]
85 changes: 85 additions & 0 deletions packages/data-designer-engine/tests/engine/test_validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -236,6 +236,91 @@ def test_validate_detect_f_string_syntax():
assert violations[0].level == ViolationLevel.WARNING


def test_validate_prompt_templates_with_literal_braces() -> None:
"""Literal braces that are valid Jinja text must not crash the f-string advisory check."""
columns = [
SamplerColumnConfig(
name="random_number",
sampler_type="uniform",
params={"low": 0, "high": 10},
),
LLMTextColumnConfig(
name="literal_close_brace",
prompt="Why is {{ random_number }} your favorite number? End with a literal } brace.",
model_alias=STUB_MODEL_ALIAS,
),
LLMTextColumnConfig(
name="literal_open_brace",
system_prompt="Prefix every answer with an opening brace { like this.",
prompt="Describe {{ random_number }}.",
model_alias=STUB_MODEL_ALIAS,
),
]
violations = validate_prompt_templates(columns, [c.name for c in columns])
assert len(violations) == 0


def test_validate_detect_f_string_syntax_with_literal_braces() -> None:
"""f-string references are still detected when the prompt also contains literal braces."""
columns = [
SamplerColumnConfig(
name="random_number",
sampler_type="uniform",
params={"low": 0, "high": 10},
),
LLMTextColumnConfig(
name="f_string_ref_after_literal_brace",
prompt="End with a literal } brace. Why is {random_number} and {{ random_number }} your favorite number?",
model_alias=STUB_MODEL_ALIAS,
),
]
violations = validate_prompt_templates(columns, [c.name for c in columns])
assert len(violations) == 1
assert violations[0].type == ViolationType.F_STRING_SYNTAX
assert violations[0].column == "f_string_ref_after_literal_brace"
assert violations[0].level == ViolationLevel.WARNING


def test_validate_detect_nested_format_spec_ref_with_literal_braces() -> None:
"""Fallback also detects references whose format spec nests another field."""
columns = [
SamplerColumnConfig(
name="random_number",
sampler_type="uniform",
params={"low": 0, "high": 10},
),
LLMTextColumnConfig(
name="nested_spec_ref_after_literal_brace",
prompt="Literal } value {random_number:{width}} jinja {{ random_number }}.",
model_alias=STUB_MODEL_ALIAS,
),
]
violations = validate_prompt_templates(columns, [c.name for c in columns])
assert len(violations) == 1
assert violations[0].type == ViolationType.F_STRING_SYNTAX
assert violations[0].column == "nested_spec_ref_after_literal_brace"


def test_validate_detect_formatted_f_string_refs_with_literal_braces() -> None:
"""Fallback also catches references carrying a conversion or format spec."""
columns = [
SamplerColumnConfig(
name="random_number",
sampler_type="uniform",
params={"low": 0, "high": 10},
),
LLMTextColumnConfig(
name="formatted_ref_after_literal_brace",
prompt="End with a literal } brace. Padded {random_number:03d}, repr {random_number!r}, jinja {{ random_number }}.",
model_alias=STUB_MODEL_ALIAS,
),
]
violations = validate_prompt_templates(columns, [c.name for c in columns])
assert len(violations) == 1
assert violations[0].type == ViolationType.F_STRING_SYNTAX
assert violations[0].column == "formatted_ref_after_literal_brace"


def test_validate_column_config_with_multi_modal_context():
column = LLMTextColumnConfig(
name="image_description",
Expand Down
39 changes: 39 additions & 0 deletions packages/data-designer/tests/interface/test_data_designer.py
Original file line number Diff line number Diff line change
Expand Up @@ -1555,6 +1555,45 @@ def test_validate_raises_error_when_seed_collides(
data_designer.validate(config_builder)


def test_validate_with_literal_braces_in_prompts(
stub_artifact_path: Path,
stub_model_providers: list[ModelProvider],
stub_check_models_model_configs: list[ModelConfig],
stub_managed_assets_path: Path,
) -> None:
"""Literal braces that are valid Jinja text must not crash validation.

Regression test for #904: an unmatched ``}`` in a prompt or an unmatched ``{``
in a system prompt is invalid f-string syntax but valid Jinja text, so the
f-string advisory check must not leak ``ValueError`` from ``string.Formatter``.
"""
config_builder = DataDesignerConfigBuilder(model_configs=stub_check_models_model_configs)
config_builder.add_column(
SamplerColumnConfig(
name="topic",
sampler_type=SamplerType.CATEGORY,
params=CategorySamplerParams(values=["science"]),
)
)
config_builder.add_column(
LLMTextColumnConfig(
name="story",
model_alias="stub-model",
prompt="Write about {{ topic }}. End with a literal } brace.",
system_prompt="Open every answer with a literal { brace.",
)
)

data_designer = DataDesigner(
artifact_path=stub_artifact_path,
model_providers=stub_model_providers,
secret_resolver=PlaintextResolver(),
managed_assets_path=stub_managed_assets_path,
)

assert data_designer.validate(config_builder) is None


def test_init_auto_configures_logging_by_default(
stub_artifact_path: Path,
stub_model_providers: list[ModelProvider],
Expand Down
Loading