diff --git a/pyrit/models/seeds/seed.py b/pyrit/models/seeds/seed.py index 3533846272..69dfeb0f64 100644 --- a/pyrit/models/seeds/seed.py +++ b/pyrit/models/seeds/seed.py @@ -9,13 +9,14 @@ from __future__ import annotations +import functools import logging import re import uuid from datetime import UTC, datetime -from typing import TYPE_CHECKING, Annotated, Any, TypeVar +from typing import TYPE_CHECKING, Annotated, Any, TypeVar, cast -from jinja2 import StrictUndefined, Undefined +from jinja2 import StrictUndefined, Undefined, meta from jinja2.sandbox import SandboxedEnvironment from pydantic import AwareDatetime, BaseModel, BeforeValidator, ConfigDict, Field @@ -23,7 +24,7 @@ from pyrit.models.seeds.seed_origin import SeedOrigin if TYPE_CHECKING: - from collections.abc import Iterator + from collections.abc import Callable, Iterator from pathlib import Path logger = logging.getLogger(__name__) @@ -102,6 +103,78 @@ def __bool__(self) -> bool: return True # Ensures it doesn't evaluate to False +class _DeferRenderError(Exception): + """Raised when an unresolved variable would decide a branch or a loop.""" + + +class _DeferringUndefined(PartialUndefined): + """ + Placeholder for the load-time render of trusted templates. + + It cannot decide a branch or a loop: answering at load would drop the {% if %} or {% for %} tags, + and the later render with the real parameters could not decide again. + """ + + def __iter__(self) -> Iterator[object]: + """ + Defer rendering instead of iterating over an unresolved variable. + + Raises: + _DeferRenderError: Always. + + """ + raise _DeferRenderError(self._undefined_name) + + def __bool__(self) -> bool: + """ + Defer rendering instead of testing an unresolved variable. + + Raises: + _DeferRenderError: Always. + + """ + raise _DeferRenderError(self._undefined_name) + + def __eq__(self, other: object) -> bool: + """ + Defer rendering instead of comparing an unresolved variable. + + Raises: + _DeferRenderError: Always. + + """ + raise _DeferRenderError(self._undefined_name) + + def __ne__(self, other: object) -> bool: + """ + Defer rendering instead of comparing an unresolved variable. + + Raises: + _DeferRenderError: Always. + + """ + raise _DeferRenderError(self._undefined_name) + + __hash__ = Undefined.__hash__ + + +_JinjaCallable = TypeVar("_JinjaCallable", bound="Callable[..., Any]") + + +def _deferring(function: _JinjaCallable) -> _JinjaCallable: + # Jinja tests such as `is defined` and the `default` filter check the value's type, not its truth. + # functools.wraps keeps Jinja's pass_environment marker, so the value may not be the first argument. + @functools.wraps(function) + def deferring(*args: Any, **kwargs: Any) -> Any: + for arg in args: + if isinstance(arg, _DeferringUndefined): + raise _DeferRenderError(arg._undefined_name) + return function(*args, **kwargs) + + # Same signature as the wrapped test or filter, so the environment's test and filter tables keep their types. + return cast("_JinjaCallable", deferring) + + class Seed(BaseModel): """Represents seed data with various attributes and metadata.""" @@ -232,6 +305,35 @@ def render_template_value_silent(self, **kwargs: Any) -> str: logger.error("Error rendering template: %s", e) return self.value + def _render_trusted_template_value(self, **kwargs: Any) -> str: + """ + Render a trusted template at load time, keeping the decisions its later render must make. + + Behaves like render_template_value_silent, except when a missing parameter would decide a branch, + a loop, a comparison, a Jinja test or the `default` filter. Then the template is kept as-is, with a + `{% set %}` in front for each supplied parameter it uses, so later renders and stored copies keep them. + + Args: + kwargs: Key-value pairs to replace in the SeedPrompt value. + + Returns: + The rendered value, or the unchanged value when rendering is deferred. + + """ + env = SandboxedEnvironment(undefined=_DeferringUndefined) + env.tests = {name: _deferring(test) for name, test in env.tests.items()} + env.filters["default"] = env.filters["d"] = _deferring(env.filters["default"]) + + try: + return env.from_string(self.value).render(**kwargs) + except _DeferRenderError: + used = meta.find_undeclared_variables(env.parse(self.value)) + bindings = "".join(f"{{% set {name} = {str(kwargs[name])!r} %}}" for name in sorted(used & kwargs.keys())) + return bindings + self.value + except Exception as e: + logger.error("Error rendering template: %s", e) + return self.value + @staticmethod def escape_for_jinja(value: str) -> str: """ diff --git a/pyrit/models/seeds/seed_objective.py b/pyrit/models/seeds/seed_objective.py index fb0d5ddd88..a6b858eba9 100644 --- a/pyrit/models/seeds/seed_objective.py +++ b/pyrit/models/seeds/seed_objective.py @@ -46,5 +46,5 @@ def _validate_and_render(self) -> SeedObjective: raise ValueError("SeedObjective cannot be a general technique.") # Only trusted templates are rendered through Jinja — see seed_prompt.py for details. if self.is_jinja_template: - self.value = self.render_template_value_silent(**PATHS_DICT) + self.value = self._render_trusted_template_value(**PATHS_DICT) return self diff --git a/pyrit/models/seeds/seed_prompt.py b/pyrit/models/seeds/seed_prompt.py index 171460ad01..3c10c44650 100644 --- a/pyrit/models/seeds/seed_prompt.py +++ b/pyrit/models/seeds/seed_prompt.py @@ -118,7 +118,7 @@ def _render_and_infer_data_type(self) -> SeedPrompt: # crafted payload containing "{% endraw %}" can escape the raw wrapper and execute # arbitrary Jinja expressions. See seed_objective.py for the same pattern. if self.is_jinja_template: - self.value = self.render_template_value_silent(**PATHS_DICT) + self.value = self._render_trusted_template_value(**PATHS_DICT) if not self.data_type: # If data_type is not provided, infer it from the value diff --git a/tests/unit/datasets/test_jailbreak_text.py b/tests/unit/datasets/test_jailbreak_text.py index 9deb982a2e..cf7ffeebac 100644 --- a/tests/unit/datasets/test_jailbreak_text.py +++ b/tests/unit/datasets/test_jailbreak_text.py @@ -308,3 +308,13 @@ def test_load_random_template_raises_when_no_prompt_templates(self) -> None: instance = TextJailBreak.__new__(TextJailBreak) with pytest.raises(ValueError, match="No jailbreak template with a single 'prompt' parameter"): instance._load_random_template() + + +def test_get_jailbreak_keeps_extra_kwargs_of_a_template_with_a_prompt_guard(): + """Extra kwargs rendered at construction survive when the template guards the prompt with an if.""" + jailbreak = TextJailBreak( + string_template="Style: {{ style }}. {% if prompt %}{{ prompt }}{% endif %}", + style="brief", + ) + + assert jailbreak.get_jailbreak("Explain rainbows") == "Style: brief. Explain rainbows" diff --git a/tests/unit/executor/attack/multi_turn/test_crescendo.py b/tests/unit/executor/attack/multi_turn/test_crescendo.py index 799708408c..744822f67a 100644 --- a/tests/unit/executor/attack/multi_turn/test_crescendo.py +++ b/tests/unit/executor/attack/multi_turn/test_crescendo.py @@ -710,6 +710,7 @@ async def test_setup_sets_adversarial_chat_system_prompt( call_args = mock_adversarial_chat.set_system_prompt_async.call_args assert "Test objective" in call_args.kwargs["system_prompt"] assert "15" in call_args.kwargs["system_prompt"] # Check for the max_turns value + assert "Prior Conversation Context" not in call_args.kwargs["system_prompt"] assert call_args.kwargs["conversation_id"] == basic_context.session.adversarial_chat_conversation_id async def test_setup_handles_prepended_conversation_with_refusal( diff --git a/tests/unit/models/test_seed.py b/tests/unit/models/test_seed.py index a977f35bf4..d09c300294 100644 --- a/tests/unit/models/test_seed.py +++ b/tests/unit/models/test_seed.py @@ -9,6 +9,9 @@ import numpy as np import pytest +import yaml +from jinja2 import StrictUndefined +from jinja2.sandbox import SandboxedEnvironment from PIL import Image from scipy.io import wavfile @@ -204,6 +207,103 @@ def test_render_template_value_silent_blocks_ssti_via_endraw_injection(): assert "__class__" not in result or result == raw_wrapped +@pytest.mark.parametrize( + ("template_value", "parameters", "expected"), + [ + ("{% if flag %}A{% else %}B{% endif %}", {"flag": True}, "A"), + ("{% if flag %}A{% else %}B{% endif %}", {"flag": False}, "B"), + ("{% if flag %}A{% else %}B{% endif %}", {"flag": None}, "B"), + ("{% if other %}A{% elif flag %}B{% else %}C{% endif %}", {"other": False, "flag": False}, "C"), + ('{{ "A" if flag else "B" }}', {"flag": False}, "B"), + ("{% set local = flag %}{% if local %}A{% else %}B{% endif %}", {"flag": False}, "B"), + ("{% if flag is defined %}A{% else %}B{% endif %}", {}, "B"), + ("{{ flag | default('B') }}", {}, "B"), + ("{% if flag == 'a' %}A{% else %}B{% endif %}", {"flag": "b"}, "B"), + ("{% if flag is filter %}A{% else %}B{% endif %}", {"flag": "upper"}, "A"), + ("{% if flag is test %}A{% else %}B{% endif %}", {"flag": "nope"}, "B"), + ("{% for item in ['A', 'B'] if flag %}{{ item }}{% endfor %}", {"flag": False}, ""), + ( + "{% macro show(rows) %}{% for row in rows %}[{{ row }}]{% endfor %}{% endmacro %}{{ show(items) }}", + {"items": ["x", "y"]}, + "[x][y]", + ), + ], +) +def test_seed_prompt_keeps_template_whose_missing_parameter_decides_a_branch(template_value, parameters, expected): + template = SeedPrompt(value=template_value, data_type="text", is_jinja_template=True) + + assert template.value == template_value + assert template.render_template_value(**parameters) == expected + + +@pytest.mark.parametrize( + ("conversation_context", "expected_tail"), + [(None, ""), ("two turns", "Context: two turns")], +) +def test_seed_prompt_keeps_path_resolved_at_load_when_its_condition_is_deferred(conversation_context, expected_tail): + template = SeedPrompt( + value="Path: {{ datasets_path }}. {% if conversation_context %}Context: {{ conversation_context }}{% endif %}", + data_type="text", + is_jinja_template=True, + ) + + # Memory rebuilds a stored prompt from its value alone, without is_jinja_template + reloaded = SeedPrompt(value=template.value, data_type="text") + + for seed in (template, reloaded): + rendered = seed.render_template_value(conversation_context=conversation_context) + assert rendered == f"Path: {DATASETS_PATH}. {expected_tail}" + + +def test_render_template_value_silent_decides_conditions_as_before(): + template = SeedPrompt( + value="{{ style }} {% if prompt %}{{ prompt }}{% endif %}", data_type="text", is_jinja_template=True + ) + + assert template.render_template_value_silent(style="brief") == "brief {{ prompt }}" + + +def test_render_template_value_silent_renders_condition_once_its_parameters_are_provided(): + template = SeedPrompt( + value="{% if flag %}{{ datasets_path }} {{ prompt }}{% endif %}", + data_type="text", + is_jinja_template=True, + ) + + assert template.render_template_value_silent(flag=True) == f"{DATASETS_PATH} {{{{ prompt }}}}" + + +def test_render_template_value_silent_renders_if_guard_on_loop_variable(): + seed = SeedPrompt( + value="{% for item in items %}{% if item %}[{{ item }}]{% endif %}{% endfor %}", + data_type="text", + is_jinja_template=True, + ) + + assert seed.render_template_value_silent(items=["a", "", "b"]) == "[a][b]" + + +_CONVERSATION_CONTEXT_TEMPLATES = sorted( + path + for path in pathlib.Path(DATASETS_PATH, "executors").rglob("*.yaml") + if "{% if conversation_context %}" in path.read_text(encoding="utf-8") +) + + +@pytest.mark.parametrize("template_path", _CONVERSATION_CONTEXT_TEMPLATES, ids=lambda path: path.stem) +@pytest.mark.parametrize("conversation_context", [None, ""]) +def test_loaded_template_renders_like_its_source(template_path, conversation_context): + seed_prompt = SeedPrompt.from_yaml_file(template_path) + parameters = {name: f"<{name}>" for name in seed_prompt.parameters or []} + parameters["conversation_context"] = conversation_context + source = yaml.safe_load(template_path.read_text(encoding="utf-8"))["value"] + + expected = SandboxedEnvironment(undefined=StrictUndefined).from_string(source).render(**parameters) + + assert seed_prompt.render_template_value(**parameters) == expected + assert ("" in expected) == (conversation_context is not None) + + def test_seed_group_untrusted_auto_escapes(): group = SeedGroup(seeds=[{"value": '{{ "".__class__ }}', "data_type": "text"}]) seed = group.prompts[0]