Skip to content

Commit 3b1558f

Browse files
authored
feat: Extend ChatPromptBuilder with messages tag (#11486)
1 parent 2b09dbb commit 3b1558f

4 files changed

Lines changed: 318 additions & 3 deletions

File tree

haystack/utils/jinja2_chat_extension.py

Lines changed: 92 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -72,6 +72,20 @@ class ChatMessageExtension(Extension):
7272
{% endmessage %}
7373
```
7474
75+
This extension also provides an `{% insert %}` placeholder tag that evaluates an expression to a `ChatMessage`
76+
or a list of `ChatMessage` objects and expands it into the prompt, so a runtime conversation can be interleaved
77+
with literal `{% message %}` blocks:
78+
79+
```
80+
{% message role="system" %}You are a helpful assistant.{% endmessage %}
81+
{% insert messages %}
82+
{% message role="user" %}{{ query }}{% endmessage %}
83+
```
84+
85+
The expression can be a plain variable (`{% insert messages %}`), a slice or index
86+
(`{% insert messages[-1:] %}`, `{% insert messages[-1] %}`), or a combination of variables
87+
(`{% insert previous + current %}`).
88+
7589
### How it works
7690
1. The `{% message %}` tag is used to define a chat message.
7791
2. The message can contain text and other structured content parts.
@@ -85,24 +99,67 @@ class ChatMessageExtension(Extension):
8599

86100
SUPPORTED_ROLES = [role.value for role in ChatRole]
87101

88-
tags = {"message"}
102+
tags = {"message", "insert"}
89103

90104
def __init__(self, environment: Any) -> None:
91105
super().__init__(environment)
92106
environment.finalize = _escape_sentinel_tags
93107
environment.filters["templatize_part"] = templatize_part
94108

95109
def parse(self, parser: Any) -> nodes.Node | list[nodes.Node]:
110+
"""
111+
Dispatch parsing based on the tag that triggered the extension.
112+
113+
Handles both the single `{% message %}` block tag and the `{% insert %}` placeholder tag.
114+
115+
:param parser: The Jinja2 parser instance
116+
:return: A CallBlock node containing the parsed configuration
117+
"""
118+
tag = next(parser.stream)
119+
if tag.value == "insert":
120+
return self._parse_insert_tag(parser, tag.lineno)
121+
return self._parse_message_tag(parser, tag.lineno)
122+
123+
def _parse_insert_tag(self, parser: Any, lineno: int) -> nodes.Node:
124+
"""
125+
Parse the `{% insert %}` placeholder tag.
126+
127+
This bodyless tag evaluates an expression to a `ChatMessage` or a list of `ChatMessage` objects and expands
128+
it into the same JSON-line format produced by `{% message %}` blocks, so messages provided at runtime can be
129+
interleaved with literal message blocks (for example a system message above and a user message below).
130+
131+
The expression can be a plain variable (`{% insert messages %}`), a slice or index
132+
(`{% insert messages[-1:] %}`, `{% insert messages[-1] %}`), or a combination of variables
133+
(`{% insert previous + current %}`).
134+
135+
:param parser: The Jinja2 parser instance
136+
:param lineno: The line number of the tag, used for error reporting.
137+
:return: A CallBlock node that expands the evaluated expression.
138+
:raises TemplateSyntaxError: If the tag is not given an expression.
139+
"""
140+
if parser.stream.current.test("block_end"):
141+
raise TemplateSyntaxError(
142+
"The 'insert' tag requires an expression that evaluates to a ChatMessage or a list of ChatMessage "
143+
"objects, for example '{% insert messages %}' or '{% insert messages[-1:] %}'.",
144+
lineno,
145+
)
146+
expr = parser.parse_expression()
147+
# Bodyless tag: empty body, no matching end tag required.
148+
return nodes.CallBlock(
149+
self.call_method(name="_build_inserted_messages_json", args=[expr]), [], [], []
150+
).set_lineno(lineno)
151+
152+
def _parse_message_tag(self, parser: Any, lineno: int) -> nodes.Node | list[nodes.Node]:
96153
"""
97154
Parse the message tag and its attributes in the Jinja2 template.
98155
99156
This method handles the parsing of role (mandatory), name (optional), meta (optional) and message body content.
100157
101158
:param parser: The Jinja2 parser instance
159+
:param lineno: The line number of the tag, used for error reporting.
102160
:return: A CallBlock node containing the parsed message configuration
103161
:raises TemplateSyntaxError: If an invalid role is provided
104162
"""
105-
lineno = next(parser.stream).lineno
106163

107164
# Parse role attribute (mandatory)
108165
parser.stream.expect("name:role")
@@ -173,6 +230,39 @@ def _build_chat_message_json(self, role: str, name: str | None, meta: dict, call
173230

174231
return json.dumps(chat_message.to_dict()) + "\n"
175232

233+
def _build_inserted_messages_json(
234+
self,
235+
messages: list[ChatMessage] | ChatMessage,
236+
caller: Callable[[], str], # noqa: ARG002
237+
) -> str:
238+
"""
239+
Expand a list of ChatMessage objects into newline-separated JSON, one message per line.
240+
241+
This method is called by Jinja2 when processing an `{% insert %}` tag. It produces the same JSON-line format
242+
as `_build_chat_message_json`, so the messages are parsed back into ChatMessage objects by the
243+
ChatPromptBuilder alongside any literal `{% message %}` blocks. The full `ChatMessage.to_dict()` payload is
244+
serialized so that all content types (tool calls, tool call results, images, reasoning, name and meta) round
245+
trip without loss.
246+
247+
:param messages: The value the `{% insert %}` expression evaluated to. A missing or empty value expands to
248+
nothing. A single ChatMessage is also accepted, since indexing with an integer (for example
249+
`{% insert messages[-1] %}`) yields one message rather than a list. The value is validated at render time
250+
because it comes from untrusted template input.
251+
:param caller: Callable that returns the (empty) rendered body. Unused.
252+
:return: Newline-terminated JSON lines, one per message, or an empty string if there are no messages.
253+
:raises ValueError: If the value is not a ChatMessage or a list of ChatMessage objects.
254+
"""
255+
if isinstance(messages, ChatMessage):
256+
messages = [messages]
257+
if not messages:
258+
return ""
259+
if not isinstance(messages, (list, tuple)) or not all(isinstance(m, ChatMessage) for m in messages):
260+
raise ValueError(
261+
"The '{% insert %}' expression must evaluate to a ChatMessage or a list of ChatMessage objects. "
262+
f"Got: {type(messages).__name__}."
263+
)
264+
return "".join(json.dumps(message.to_dict()) + "\n" for message in messages)
265+
176266
@staticmethod
177267
def _parse_content_parts(content: str) -> list[ChatMessageContentT]:
178268
"""
Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,34 @@
1+
---
2+
features:
3+
- |
4+
Added a new ``{% insert %}`` tag to the Jinja2 string templates used by ``ChatPromptBuilder``.
5+
The tag is a placeholder that evaluates an expression to a ``ChatMessage`` or a list of ``ChatMessage``
6+
objects and expands it into the prompt, so messages provided at runtime can be interleaved with literal
7+
``{% message %}`` blocks. For example, you can wrap the runtime messages with a system message above
8+
and a templated user message below, then pass the messages (and any template variables) at run time:
9+
10+
.. code-block:: python
11+
12+
from haystack.components.builders import ChatPromptBuilder
13+
from haystack.dataclasses import ChatMessage
14+
15+
template = """
16+
{% message role="system" %}You are a helpful assistant.{% endmessage %}
17+
{% insert messages %}
18+
{% message role="user" %}{{ query }}{% endmessage %}
19+
"""
20+
21+
builder = ChatPromptBuilder(template=template)
22+
result = builder.run(
23+
messages=[ChatMessage.from_user("Hi"), ChatMessage.from_assistant("Hello!")],
24+
query="What's the weather?",
25+
)
26+
# result["prompt"] -> [system, user "Hi", assistant "Hello!", user "What's the weather?"]
27+
28+
All content types (tool calls, tool call results, images, reasoning, ``name`` and ``meta``) round trip
29+
without loss. A missing or empty value expands to nothing.
30+
31+
The expression can be a plain variable (``{% insert messages %}``), a slice or index
32+
(``{% insert messages[-1:] %}``, ``{% insert messages[-1] %}``), or a combination of variables
33+
(``{% insert previous + current %}``). Multiple ``{% insert %}`` tags can be used in a single template,
34+
so the runtime messages can be split, reordered, or repeated across different positions.

test/components/builders/test_chat_prompt_builder.py

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1076,3 +1076,47 @@ def test_poisoned_document_does_not_inject_image(self, in_memory_doc_store):
10761076

10771077
images = [p for p in msg._content if isinstance(p, ImageContent)]
10781078
assert len(images) == 0
1079+
1080+
def test_insert_placeholder(self):
1081+
builder = ChatPromptBuilder(template="{% insert messages %}")
1082+
assert builder.variables == ["messages"]
1083+
runtime = [ChatMessage.from_user("Hello"), ChatMessage.from_assistant("Hi")]
1084+
result = builder.run(messages=runtime)
1085+
assert result["prompt"] == runtime
1086+
1087+
def test_insert_placeholder_with_subscript(self):
1088+
builder = ChatPromptBuilder(template="{% insert messages[-1:] %}")
1089+
assert builder.variables == ["messages"]
1090+
runtime = [ChatMessage.from_user("first"), ChatMessage.from_assistant("last")]
1091+
result = builder.run(messages=runtime)
1092+
assert result["prompt"] == [ChatMessage.from_assistant("last")]
1093+
1094+
def test_insert_placeholder_custom_variable_name(self):
1095+
builder = ChatPromptBuilder(template="{% insert chat_history %}")
1096+
assert builder.variables == ["chat_history"]
1097+
runtime = [ChatMessage.from_user("Hello"), ChatMessage.from_assistant("Hi")]
1098+
result = builder.run(chat_history=runtime)
1099+
assert result["prompt"] == runtime
1100+
1101+
def test_insert_placeholder_combines_variables(self):
1102+
builder = ChatPromptBuilder(template="{% insert previous + current %}")
1103+
assert set(builder.variables) == {"previous", "current"}
1104+
previous = [ChatMessage.from_user("p1"), ChatMessage.from_assistant("p2")]
1105+
current = [ChatMessage.from_user("c1")]
1106+
result = builder.run(previous=previous, current=current)
1107+
assert result["prompt"] == previous + current
1108+
1109+
def test_insert_placeholder_interleaved_with_blocks(self):
1110+
template = (
1111+
'{% message role="system" %}You are helpful.{% endmessage %}'
1112+
"{% insert messages %}"
1113+
'{% message role="user" %}{{ query }}{% endmessage %}'
1114+
)
1115+
builder = ChatPromptBuilder(template=template)
1116+
assert set(builder.variables) == {"messages", "query"}
1117+
result = builder.run(messages=[ChatMessage.from_user("earlier")], query="now")
1118+
assert result["prompt"] == [
1119+
ChatMessage.from_system("You are helpful."),
1120+
ChatMessage.from_user("earlier"),
1121+
ChatMessage.from_user("now"),
1122+
]

test/utils/test_jinja2_chat_extension.py

Lines changed: 148 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,10 +7,11 @@
77
from unittest.mock import patch
88

99
import pytest
10-
from jinja2 import TemplateSyntaxError
10+
from jinja2 import TemplateSyntaxError, meta
1111
from jinja2.sandbox import SandboxedEnvironment
1212

1313
from haystack.dataclasses.chat_message import (
14+
ChatMessage,
1415
FileContent,
1516
ImageContent,
1617
ReasoningContent,
@@ -604,6 +605,152 @@ def test_common_symbols_not_escaped(self, jinja_env):
604605
assert output["content"][0]["text"] == text_with_symbols
605606

606607

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+
607754
class TestSentinelTagInjectionPrevention:
608755
def test_sentinel_tag_injection_via_text_variable(self, jinja_env):
609756
fake_b64 = base64.b64encode(b"ATTACKER_PAYLOAD").decode()

0 commit comments

Comments
 (0)