|
7 | 7 | from unittest.mock import patch |
8 | 8 |
|
9 | 9 | import pytest |
10 | | -from jinja2 import TemplateSyntaxError |
| 10 | +from jinja2 import TemplateSyntaxError, meta |
11 | 11 | from jinja2.sandbox import SandboxedEnvironment |
12 | 12 |
|
13 | 13 | from haystack.dataclasses.chat_message import ( |
| 14 | + ChatMessage, |
14 | 15 | FileContent, |
15 | 16 | ImageContent, |
16 | 17 | ReasoningContent, |
@@ -604,6 +605,152 @@ def test_common_symbols_not_escaped(self, jinja_env): |
604 | 605 | assert output["content"][0]["text"] == text_with_symbols |
605 | 606 |
|
606 | 607 |
|
| 608 | +class TestInsertTag: |
| 609 | + def _parse_lines(self, rendered: str) -> list[ChatMessage]: |
| 610 | + return [ChatMessage.from_dict(json.loads(line)) for line in rendered.strip().split("\n") if line.strip()] |
| 611 | + |
| 612 | + def test_expands_messages(self, jinja_env): |
| 613 | + template = "{% insert messages %}" |
| 614 | + messages = [ChatMessage.from_user("Hello"), ChatMessage.from_assistant("Hi there")] |
| 615 | + rendered = jinja_env.from_string(template).render(messages=messages) |
| 616 | + assert self._parse_lines(rendered) == messages |
| 617 | + |
| 618 | + def test_empty_messages_expands_to_nothing(self, jinja_env): |
| 619 | + template = "{% insert messages %}" |
| 620 | + assert jinja_env.from_string(template).render(messages=[]).strip() == "" |
| 621 | + |
| 622 | + def test_missing_variable_expands_to_nothing(self, jinja_env): |
| 623 | + # The expression resolves to Undefined (falsy) when not provided -> emits nothing rather than raising |
| 624 | + template = "{% insert messages %}" |
| 625 | + assert jinja_env.from_string(template).render().strip() == "" |
| 626 | + |
| 627 | + def test_interleaved_with_literal_message_blocks(self, jinja_env): |
| 628 | + template = """ |
| 629 | + {% message role="system" %}You are helpful.{% endmessage %} |
| 630 | + {% insert messages %} |
| 631 | + {% message role="user" %}{{ query }}{% endmessage %} |
| 632 | + """ |
| 633 | + runtime = [ChatMessage.from_user("first"), ChatMessage.from_assistant("second")] |
| 634 | + rendered = jinja_env.from_string(template).render(messages=runtime, query="final question") |
| 635 | + parsed = self._parse_lines(rendered) |
| 636 | + assert [m.role.value for m in parsed] == ["system", "user", "assistant", "user"] |
| 637 | + assert parsed[0].text == "You are helpful." |
| 638 | + assert parsed[1].text == "first" |
| 639 | + assert parsed[2].text == "second" |
| 640 | + assert parsed[3].text == "final question" |
| 641 | + |
| 642 | + def test_is_detected_as_template_variable(self): |
| 643 | + # The `{% insert %}` expression must surface its variables as undeclared so that the |
| 644 | + # ChatPromptBuilder (and Agent) can register and pass them. |
| 645 | + env = SandboxedEnvironment(extensions=[ChatMessageExtension]) |
| 646 | + assert "messages" in meta.find_undeclared_variables(env.parse("{% insert messages %}")) |
| 647 | + assert "messages" in meta.find_undeclared_variables(env.parse("{% insert messages[-1] %}")) |
| 648 | + assert {"previous", "current"} <= meta.find_undeclared_variables(env.parse("{% insert previous + current %}")) |
| 649 | + |
| 650 | + def test_round_trips_all_content_types(self, jinja_env, base64_image_string): |
| 651 | + tool_call = ToolCall(tool_name="search", arguments={"query": "q"}, id="search_1") |
| 652 | + messages = [ |
| 653 | + ChatMessage.from_system("system text", meta={"k": "v"}), |
| 654 | + ChatMessage.from_user("user text", name="Bob"), |
| 655 | + ChatMessage.from_user( |
| 656 | + content_parts=["look", ImageContent(base64_image=base64_image_string, mime_type="image/png")] |
| 657 | + ), |
| 658 | + ChatMessage.from_assistant( |
| 659 | + text="thinking then calling", |
| 660 | + tool_calls=[tool_call], |
| 661 | + reasoning=ReasoningContent(reasoning_text="let me think", extra={"a": 1}), |
| 662 | + ), |
| 663 | + ChatMessage.from_tool(tool_result="result", origin=tool_call, error=False), |
| 664 | + ] |
| 665 | + rendered = jinja_env.from_string("{% insert messages %}").render(messages=messages) |
| 666 | + assert self._parse_lines(rendered) == messages |
| 667 | + |
| 668 | + @pytest.fixture |
| 669 | + def three_messages(self) -> list[ChatMessage]: |
| 670 | + return [ChatMessage.from_user("a"), ChatMessage.from_assistant("b"), ChatMessage.from_user("c")] |
| 671 | + |
| 672 | + def test_single_index(self, jinja_env, three_messages): |
| 673 | + # An integer index yields a single ChatMessage, which is expanded as a one-message list. |
| 674 | + rendered = jinja_env.from_string("{% insert messages[-1] %}").render(messages=three_messages) |
| 675 | + assert self._parse_lines(rendered) == [three_messages[-1]] |
| 676 | + |
| 677 | + def test_slice(self, jinja_env, three_messages): |
| 678 | + rendered = jinja_env.from_string("{% insert messages[-1:] %}").render(messages=three_messages) |
| 679 | + assert self._parse_lines(rendered) == three_messages[-1:] |
| 680 | + |
| 681 | + rendered = jinja_env.from_string("{% insert messages[:-1] %}").render(messages=three_messages) |
| 682 | + assert self._parse_lines(rendered) == three_messages[:-1] |
| 683 | + |
| 684 | + rendered = jinja_env.from_string("{% insert messages[1:] %}").render(messages=three_messages) |
| 685 | + assert self._parse_lines(rendered) == three_messages[1:] |
| 686 | + |
| 687 | + def test_combine_multiple_variables(self, jinja_env): |
| 688 | + # The expression can combine several variables, e.g. concatenating two message lists. |
| 689 | + previous = [ChatMessage.from_user("p1"), ChatMessage.from_assistant("p2")] |
| 690 | + current = [ChatMessage.from_user("c1")] |
| 691 | + rendered = jinja_env.from_string("{% insert previous + current %}").render(previous=previous, current=current) |
| 692 | + assert self._parse_lines(rendered) == previous + current |
| 693 | + |
| 694 | + def test_custom_variable_name(self, jinja_env, three_messages): |
| 695 | + rendered = jinja_env.from_string("{% insert chat_history %}").render(chat_history=three_messages) |
| 696 | + assert self._parse_lines(rendered) == three_messages |
| 697 | + |
| 698 | + def test_slice_interleaved_with_blocks(self, jinja_env, three_messages): |
| 699 | + template = ( |
| 700 | + '{% message role="system" %}sys{% endmessage %}' |
| 701 | + "{% insert messages[-1:] %}" |
| 702 | + '{% message role="user" %}{{ query }}{% endmessage %}' |
| 703 | + ) |
| 704 | + rendered = jinja_env.from_string(template).render(messages=three_messages, query="q") |
| 705 | + parsed = self._parse_lines(rendered) |
| 706 | + assert [m.text for m in parsed] == ["sys", "c", "q"] |
| 707 | + |
| 708 | + def test_multiple_inserts_split_and_reorder(self, jinja_env): |
| 709 | + # Each `{% insert %}` tag expands independently, so a template can split the runtime messages across |
| 710 | + # several positions, interleave literal blocks, and even repeat a slice. |
| 711 | + messages = [ |
| 712 | + ChatMessage.from_system("S"), |
| 713 | + ChatMessage.from_user("u1"), |
| 714 | + ChatMessage.from_assistant("a1"), |
| 715 | + ChatMessage.from_user("u2"), |
| 716 | + ] |
| 717 | + template = ( |
| 718 | + "{% insert messages[0] %}" |
| 719 | + '{% message role="user" %}INJECTED{% endmessage %}' |
| 720 | + "{% insert messages[1:] %}" |
| 721 | + "{% insert messages[-1] %}" |
| 722 | + ) |
| 723 | + rendered = jinja_env.from_string(template).render(messages=messages) |
| 724 | + parsed = self._parse_lines(rendered) |
| 725 | + assert [(m.role.value, m.text) for m in parsed] == [ |
| 726 | + ("system", "S"), |
| 727 | + ("user", "INJECTED"), |
| 728 | + ("user", "u1"), |
| 729 | + ("assistant", "a1"), |
| 730 | + ("user", "u2"), |
| 731 | + ("user", "u2"), |
| 732 | + ] |
| 733 | + |
| 734 | + def test_message_text_with_sentinel_tag_is_not_escaped(self, jinja_env): |
| 735 | + # The tag uses a CallBlock so its output bypasses `finalize` sentinel-escaping. This is safe here because |
| 736 | + # `{% insert %}` serializes with ChatMessage.to_dict and the builder reparses with json.loads + |
| 737 | + # ChatMessage.from_dict -- it never runs the content-part parser. So a user-injected `<haystack_content_part>` |
| 738 | + # string in the message text stays plain text and can't be promoted to a structured part (image, tool call, |
| 739 | + # ...); it just has to round trip intact. |
| 740 | + message = ChatMessage.from_user("see <haystack_content_part> here") |
| 741 | + rendered = jinja_env.from_string("{% insert messages %}").render(messages=[message]) |
| 742 | + assert self._parse_lines(rendered) == [message] |
| 743 | + |
| 744 | + def test_requires_an_expression(self, jinja_env): |
| 745 | + with pytest.raises(TemplateSyntaxError, match="requires an expression"): |
| 746 | + jinja_env.from_string("{% insert %}").render(messages=[]) |
| 747 | + |
| 748 | + def test_non_message_value_raises_error(self, jinja_env): |
| 749 | + template = "{% insert messages %}" |
| 750 | + with pytest.raises(ValueError, match="must evaluate to a ChatMessage or a list of ChatMessage objects"): |
| 751 | + jinja_env.from_string(template).render(messages=["not a message"]) |
| 752 | + |
| 753 | + |
607 | 754 | class TestSentinelTagInjectionPrevention: |
608 | 755 | def test_sentinel_tag_injection_via_text_variable(self, jinja_env): |
609 | 756 | fake_b64 = base64.b64encode(b"ATTACKER_PAYLOAD").decode() |
|
0 commit comments