deer-flow/backend/tests/test_view_image_middleware.py
Hyeonsang Cho cb24bc2699
perf(middleware): stop checkpointing view_image base64 payloads (#5014)
ViewImageMiddleware injected the viewed-image message from before_model and
removed it again from after_model. before_model, model, and after_model are
separate graph nodes, so every view_image turn cost two extra nodes and two
state writes, and up to 20MB of base64 sat in two checkpoints for the duration
of the model call. A run interrupted in that window (user cancel, restart)
stranded the payload in history for good.

Inject from wrap_model_call instead, so the message lives only in
ModelRequest.messages and is never returned as a state update:

- before_model/after_model (and the async pair) are replaced by
  wrap_model_call/awrap_model_call; _remove_image_context_messages and its
  RemoveMessage bookkeeping go with them. The async hook keeps the existing
  asyncio.to_thread offload for the file read and base64 encode.
- _should_inject_image_message gates on request.messages rather than state, so
  the decision is made against what the model will actually see.
- _inject sweeps this middleware's own message out of the request before
  rebuilding it. Dropping after_model also drops the cleanup it did on every
  call, so without the sweep a payload stranded by an older interrupted run
  would ride along in every later request for the life of the thread. Matching
  requires both the reserved id prefix and the server-owned marker, and Gateway
  strips that marker from client input, so a user message is never dropped.

Chain position is unchanged, and wrap_model_call nests first-registered
outermost, so TokenBudgetMiddleware still sees the image message and enforces
the input budget against it.

Checkpoint rows that already hold a stranded payload keep it on disk. It is
inert -- never sent to a provider, and strip_data_url_image_blocks keeps it off
the wire -- and reclaiming it would mean keeping the node this change removes.

tests/test_view_image_middleware.py is rewritten around the new hook (43
tests): sync/async at unit and graph level, the stranded sweep, and the
client-message protection. Docs: middleware chain entry 23, Vision Support, the
middleware-execution-flow hook matrix and diagrams, and the
strip_data_url_image_blocks docstring.
2026-08-27 09:12:37 +08:00

604 lines
26 KiB
Python

"""Unit tests for ViewImageMiddleware.
Tests cover the middleware's ability to inject image details (including base64
payloads) into the model request, triggered only when the previous assistant
turn contained `view_image` tool calls that have all been completed with
corresponding ToolMessages.
Covered behavior:
- `_get_last_assistant_message` returns the most recent AIMessage (or None).
- `_has_view_image_tool` only matches assistant messages with `view_image` tool calls.
- `_all_tools_completed` verifies every tool call id has a matching ToolMessage.
- `_create_image_details_message` produces correctly structured content blocks,
reading image files on-demand from disk (no base64 stored in state).
- `_should_inject_image_message` gates injection on all preconditions, including
deduplication when an image-details message is already in the request.
- `_inject` rebuilds the request's image context: it sweeps out any copy left
in the message list by an interrupted run before deciding whether to append a
freshly built one.
- `wrap_model_call` and `awrap_model_call` expose the same behavior sync/async,
handing the payload to the model without ever writing it to state — so no
checkpoint retains it, even if the run is interrupted mid-call.
"""
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from langchain.agents import create_agent
from langchain.agents.middleware.types import ModelRequest
from langchain_core.callbacks import BaseCallbackHandler
from langchain_core.language_models.fake_chat_models import FakeMessagesListChatModel
from langchain_core.messages import AIMessage, AnyMessage, HumanMessage, SystemMessage, ToolMessage
from deerflow.agents.middlewares.view_image_middleware import (
_IMAGE_CONTEXT_MESSAGE_MARKER_KEY,
ViewImageMiddleware,
)
def _view_image_call(call_id: str = "call_1", path: str = "/mnt/user-data/uploads/img.png") -> dict:
return {"name": "view_image", "id": call_id, "args": {"image_path": path}}
def _other_tool_call(call_id: str = "call_other", name: str = "bash") -> dict:
return {"name": name, "id": call_id, "args": {"command": "ls"}}
def _model_request(messages: list[AnyMessage], viewed_images: dict | None = None) -> ModelRequest:
"""Build a real ModelRequest so `.override()` behaves as it does in the graph."""
return ModelRequest(
model=FakeMessagesListChatModel(responses=[AIMessage(content="ok")]),
messages=list(messages),
system_message=None,
tool_choice=None,
tools=[],
response_format=None,
state={"messages": list(messages), "viewed_images": viewed_images or {}},
runtime=MagicMock(),
model_settings={},
)
class _CaptureChatMessages(BaseCallbackHandler):
def __init__(self):
self.messages = []
def on_chat_model_start(self, serialized, messages, **kwargs):
self.messages = messages[0]
def _image_context_messages(messages: list[AnyMessage]) -> list[HumanMessage]:
return [message for message in messages if isinstance(message, HumanMessage) and message.id and message.id.startswith("view-image-context:")]
def _make_viewed_image(tmp_path, filename="img.png", mime_type="image/png", data=b"\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR"):
"""Create a real image file and return viewed_images metadata dict."""
img_path = tmp_path / filename
img_path.write_bytes(data)
return {
"mime_type": mime_type,
"size": len(data),
"actual_path": str(img_path),
}
class TestGetLastAssistantMessage:
def test_returns_none_on_empty_list(self):
mw = ViewImageMiddleware()
assert mw._get_last_assistant_message([]) is None
def test_returns_none_when_no_ai_message(self):
mw = ViewImageMiddleware()
messages = [
SystemMessage(content="sys"),
HumanMessage(content="hi"),
]
assert mw._get_last_assistant_message(messages) is None
def test_returns_most_recent_ai_message(self):
mw = ViewImageMiddleware()
older = AIMessage(content="older")
newer = AIMessage(content="newer")
messages = [HumanMessage(content="q"), older, HumanMessage(content="q2"), newer]
assert mw._get_last_assistant_message(messages) is newer
class TestHasViewImageTool:
def test_returns_false_when_tool_calls_attr_missing(self):
"""Exercise the `not hasattr(message, "tool_calls")` guard.
AIMessage always has a `tool_calls` attribute, so we use a plain
object that truly lacks the attribute to cover this branch.
"""
mw = ViewImageMiddleware()
msg = SimpleNamespace(content="just text") # no tool_calls attribute
assert not hasattr(msg, "tool_calls") # precondition
assert mw._has_view_image_tool(msg) is False
def test_returns_false_when_ai_message_has_no_tool_calls(self):
"""AIMessage without tool_calls kwarg defaults to an empty list."""
mw = ViewImageMiddleware()
msg = AIMessage(content="just text")
assert mw._has_view_image_tool(msg) is False
def test_returns_false_when_tool_calls_empty(self):
mw = ViewImageMiddleware()
msg = AIMessage(content="", tool_calls=[])
assert mw._has_view_image_tool(msg) is False
def test_returns_true_when_view_image_present(self):
mw = ViewImageMiddleware()
msg = AIMessage(content="", tool_calls=[_view_image_call()])
assert mw._has_view_image_tool(msg) is True
def test_returns_true_when_view_image_mixed_with_others(self):
mw = ViewImageMiddleware()
msg = AIMessage(
content="",
tool_calls=[_other_tool_call(), _view_image_call(call_id="call_vi")],
)
assert mw._has_view_image_tool(msg) is True
def test_returns_false_when_only_other_tools(self):
mw = ViewImageMiddleware()
msg = AIMessage(content="", tool_calls=[_other_tool_call()])
assert mw._has_view_image_tool(msg) is False
class TestAllToolsCompleted:
def test_returns_false_when_no_tool_calls(self):
mw = ViewImageMiddleware()
assistant = AIMessage(content="", tool_calls=[])
assert mw._all_tools_completed([assistant], assistant) is False
def test_returns_true_when_all_completed(self):
mw = ViewImageMiddleware()
assistant = AIMessage(
content="",
tool_calls=[_view_image_call("c1"), _view_image_call("c2", "/p2.png")],
)
messages = [
assistant,
ToolMessage(content="ok", tool_call_id="c1"),
ToolMessage(content="ok", tool_call_id="c2"),
]
assert mw._all_tools_completed(messages, assistant) is True
def test_returns_false_when_some_tool_call_unanswered(self):
mw = ViewImageMiddleware()
assistant = AIMessage(
content="",
tool_calls=[_view_image_call("c1"), _view_image_call("c2", "/p2.png")],
)
messages = [assistant, ToolMessage(content="ok", tool_call_id="c1")]
assert mw._all_tools_completed(messages, assistant) is False
def test_returns_false_when_assistant_not_in_messages(self):
mw = ViewImageMiddleware()
assistant = AIMessage(content="", tool_calls=[_view_image_call("c1")])
# assistant is not part of the list, so messages.index() will raise and be caught
messages = [HumanMessage(content="hi")]
assert mw._all_tools_completed(messages, assistant) is False
def test_ignores_tool_messages_before_assistant(self):
mw = ViewImageMiddleware()
assistant = AIMessage(content="", tool_calls=[_view_image_call("c1")])
# A stale ToolMessage with matching id appears BEFORE the assistant turn.
# It should not count — only ToolMessages after the assistant close the call.
messages = [
ToolMessage(content="stale", tool_call_id="c1"),
assistant,
]
assert mw._all_tools_completed(messages, assistant) is False
class TestCreateImageDetailsMessage:
def test_returns_placeholder_when_no_images(self):
mw = ViewImageMiddleware()
state = {"viewed_images": {}}
blocks = mw._create_image_details_message(state)
assert blocks == [{"type": "text", "text": "No images have been viewed."}]
def test_returns_placeholder_when_state_missing_key(self):
mw = ViewImageMiddleware()
blocks = mw._create_image_details_message({})
assert blocks == [{"type": "text", "text": "No images have been viewed."}]
def test_builds_blocks_for_single_image(self, tmp_path):
mw = ViewImageMiddleware()
img_meta = _make_viewed_image(tmp_path, "cat.png")
state = {
"viewed_images": {
"/path/to/cat.png": img_meta,
}
}
blocks = mw._create_image_details_message(state)
# header text + per-image description text + per-image image_url block
assert len(blocks) == 3
assert blocks[0] == {"type": "text", "text": "Here are the images you've viewed:"}
assert blocks[1]["type"] == "text"
assert "/path/to/cat.png" in blocks[1]["text"]
assert "image/png" in blocks[1]["text"]
assert blocks[2]["type"] == "image_url"
assert blocks[2]["image_url"]["url"].startswith("data:image/png;base64,")
def test_builds_blocks_for_multiple_images(self, tmp_path):
mw = ViewImageMiddleware()
img1 = _make_viewed_image(tmp_path, "a.png", data=b"\x89PNG\r\n\x1a\nfake-png")
img2 = _make_viewed_image(tmp_path, "b.jpg", mime_type="image/jpeg", data=b"\xff\xd8\xff\xe0fake-jpeg")
state = {
"viewed_images": {
"/a.png": img1,
"/b.jpg": img2,
}
}
blocks = mw._create_image_details_message(state)
# 1 header + (1 description + 1 image_url) per image = 5 blocks
assert len(blocks) == 5
image_url_blocks = [b for b in blocks if isinstance(b, dict) and b.get("type") == "image_url"]
assert len(image_url_blocks) == 2
urls = {b["image_url"]["url"] for b in image_url_blocks}
assert any(u.startswith("data:image/png;base64,") for u in urls)
assert any(u.startswith("data:image/jpeg;base64,") for u in urls)
def test_omits_image_url_block_when_file_missing(self, tmp_path):
mw = ViewImageMiddleware()
state = {
"viewed_images": {
"/broken.png": {
"mime_type": "image/png",
"size": 0,
"actual_path": str(tmp_path / "nonexistent.png"),
},
}
}
blocks = mw._create_image_details_message(state)
# header + description + error text (file no longer available)
assert len(blocks) == 3
assert all(not (isinstance(b, dict) and b.get("type") == "image_url") for b in blocks)
def test_uses_unknown_mime_type_when_missing(self, tmp_path):
mw = ViewImageMiddleware()
img_meta = _make_viewed_image(tmp_path, "mystery.bin", mime_type="unknown")
state = {
"viewed_images": {
"/mystery.bin": img_meta,
}
}
blocks = mw._create_image_details_message(state)
# The description block should mention unknown
description_blocks = [b for b in blocks if b.get("type") == "text" and "/mystery.bin" in b.get("text", "")]
assert len(description_blocks) == 1
assert "unknown" in description_blocks[0]["text"]
def test_omits_image_url_when_read_raises_oserror(self, tmp_path, monkeypatch):
"""A failure during on-demand read must not crash the middleware."""
img_meta = _make_viewed_image(tmp_path, "ok.png")
state = {
"viewed_images": {
"/ok.png": img_meta,
}
}
def _raise(*args, **kwargs):
raise OSError("disk error")
monkeypatch.setattr("builtins.open", _raise)
mw = ViewImageMiddleware()
blocks = mw._create_image_details_message(state)
# header + description + 'unavailable' text, no image_url block
assert all(not (isinstance(b, dict) and b.get("type") == "image_url") for b in blocks)
unavailable = [b for b in blocks if isinstance(b, dict) and b.get("type") == "text" and "unavailable" in b.get("text", "")]
assert len(unavailable) == 1
def test_omits_image_url_when_size_changes_between_view_and_inject(self, tmp_path):
"""Defense against TOCTOU growth: skip if current size differs from recorded size."""
img_meta = _make_viewed_image(tmp_path, "shrinking.png", data=b"original-larger-content")
# Grow the file after the metadata was written
img_meta_path = Path(img_meta["actual_path"])
img_meta_path.write_bytes(b"much-much-much-larger-content-bytes")
state = {"viewed_images": {"/shrinking.png": img_meta}}
mw = ViewImageMiddleware()
blocks = mw._create_image_details_message(state)
assert all(not (isinstance(b, dict) and b.get("type") == "image_url") for b in blocks)
def test_omits_image_url_when_size_exceeds_cap(self, tmp_path):
"""Records a small size but the actual file is large - the cap kicks in regardless."""
img_meta = _make_viewed_image(tmp_path, "huge.png", data=b"x" * 100)
img_meta_path = Path(img_meta["actual_path"])
# Grow past the cap (20 MB)
img_meta_path.write_bytes(b"y" * (21 * 1024 * 1024))
state = {"viewed_images": {"/huge.png": img_meta}}
mw = ViewImageMiddleware()
blocks = mw._create_image_details_message(state)
assert all(not (isinstance(b, dict) and b.get("type") == "image_url") for b in blocks)
class TestShouldInjectImageMessage:
def test_false_when_no_messages(self):
mw = ViewImageMiddleware()
assert mw._should_inject_image_message([]) is False
def test_false_when_no_assistant_message(self):
mw = ViewImageMiddleware()
assert mw._should_inject_image_message([HumanMessage(content="hello")]) is False
def test_false_when_no_view_image_tool_call(self):
mw = ViewImageMiddleware()
assistant = AIMessage(content="", tool_calls=[_other_tool_call()])
messages = [assistant, ToolMessage(content="ok", tool_call_id="call_other")]
assert mw._should_inject_image_message(messages) is False
def test_false_when_tool_not_completed(self):
mw = ViewImageMiddleware()
assistant = AIMessage(content="", tool_calls=[_view_image_call("c1")])
assert mw._should_inject_image_message([assistant]) is False # no ToolMessage yet
def test_true_when_all_preconditions_met(self):
mw = ViewImageMiddleware()
assistant = AIMessage(content="", tool_calls=[_view_image_call("c1")])
messages = [assistant, ToolMessage(content="ok", tool_call_id="c1")]
assert mw._should_inject_image_message(messages) is True
def test_false_when_already_injected(self):
"""A checkpoint written by an older version (or by a run that died before
its `RemoveMessage` cleanup landed) can still carry image details. Do not
add a duplicate on top of one."""
mw = ViewImageMiddleware()
assistant = AIMessage(content="", tool_calls=[_view_image_call("c1")])
already_injected = HumanMessage(content="Here are the images you've viewed: /img.png")
messages = [
assistant,
ToolMessage(content="ok", tool_call_id="c1"),
already_injected,
]
assert mw._should_inject_image_message(messages) is False
def test_false_when_already_injected_with_list_content(self, tmp_path):
"""Deduplication must recognize the real injected payload shape.
An unmarked leftover carries `.content` as a *list* of dicts (text +
image_url blocks), not a plain string. This test reuses
`_create_image_details_message` output to reproduce the realistic shape
and confirms the marker is still detected via `str(msg.content)`.
"""
mw = ViewImageMiddleware()
assistant = AIMessage(content="", tool_calls=[_view_image_call("c1")])
viewed_images = {"/img.png": _make_viewed_image(tmp_path)}
# Build content the same way the middleware would.
real_injected_content = mw._create_image_details_message({"viewed_images": viewed_images})
# Sanity: this is a list of blocks, not a plain string.
assert isinstance(real_injected_content, list)
messages = [
assistant,
ToolMessage(content="ok", tool_call_id="c1"),
HumanMessage(content=real_injected_content),
]
assert mw._should_inject_image_message(messages) is False
def test_false_when_legacy_details_marker_present(self):
"""The middleware also recognizes the legacy 'Here are the details of the
images you've viewed' marker as an already-injected signal."""
mw = ViewImageMiddleware()
assistant = AIMessage(content="", tool_calls=[_view_image_call("c1")])
legacy = HumanMessage(content="Here are the details of the images you've viewed: ...")
messages = [
assistant,
ToolMessage(content="ok", tool_call_id="c1"),
legacy,
]
assert mw._should_inject_image_message(messages) is False
class TestInject:
def test_returns_request_unchanged_when_should_not_inject(self):
mw = ViewImageMiddleware()
request = _model_request([HumanMessage(content="hi")])
assert mw._inject(request) is request
def test_appends_image_context_message_to_request(self, tmp_path):
mw = ViewImageMiddleware()
assistant = AIMessage(content="", tool_calls=[_view_image_call("c1")])
original = [assistant, ToolMessage(content="ok", tool_call_id="c1")]
request = _model_request(original, {"/img.png": _make_viewed_image(tmp_path)})
injected_request = mw._inject(request)
assert injected_request is not request
# The payload is appended last, so it directly follows the tool results.
assert injected_request.messages[:-1] == original
injected = injected_request.messages[-1]
assert isinstance(injected, HumanMessage)
# Mixed-content payload: list of text + image_url blocks
assert isinstance(injected.content, list)
assert any(isinstance(b, dict) and b.get("type") == "image_url" for b in injected.content)
# Internal injection: must be hidden from the chat UI (and IM channels),
# like the other middleware-injected context messages.
assert injected.additional_kwargs.get("hide_from_ui") is True
assert injected.additional_kwargs.get(_IMAGE_CONTEXT_MESSAGE_MARKER_KEY) is True
assert injected.id is not None
assert injected.id.startswith("view-image-context:")
def test_replaces_a_stranded_payload_instead_of_stacking_a_second_one(self, tmp_path):
"""A run that died during the model call can leave the old
before_model/after_model pair's message checkpointed. Rebuild it rather
than adding a second copy on top."""
mw = ViewImageMiddleware()
assistant = AIMessage(content="", tool_calls=[_view_image_call("c1")])
stranded = ViewImageMiddleware._create_image_context_message([{"type": "text", "text": "stale"}])
request = _model_request(
[assistant, ToolMessage(content="ok", tool_call_id="c1"), stranded],
{"/img.png": _make_viewed_image(tmp_path)},
)
injected = _image_context_messages(mw._inject(request).messages)
assert len(injected) == 1
assert injected[0].id != stranded.id
assert any(isinstance(b, dict) and b.get("type") == "image_url" for b in injected[0].content)
def test_drops_a_stranded_payload_even_when_no_injection_is_warranted(self):
"""Otherwise the stale base64 would ride along in every later request for
the life of the thread -- the old `after_model` swept it, so dropping the
hook must not lose that."""
mw = ViewImageMiddleware()
stranded = ViewImageMiddleware._create_image_context_message([{"type": "text", "text": "stale"}])
request = _model_request([HumanMessage(content="hi"), stranded, AIMessage(content="done")])
prepared = mw._inject(request)
assert prepared is not request
assert _image_context_messages(prepared.messages) == []
assert [type(m) for m in prepared.messages] == [HumanMessage, AIMessage]
def test_never_drops_a_client_message_wearing_the_reserved_prefix(self):
"""The prefix alone is not enough — Gateway strips the server-owned
marker from client input, and both are required to match."""
mw = ViewImageMiddleware()
client_message = HumanMessage(id="view-image-context:client-supplied", content="client-authored", additional_kwargs={"hide_from_ui": True})
request = _model_request([client_message, AIMessage(content="done")])
assert mw._inject(request) is request
def test_does_not_mutate_the_incoming_request(self, tmp_path):
mw = ViewImageMiddleware()
assistant = AIMessage(content="", tool_calls=[_view_image_call("c1")])
request = _model_request(
[assistant, ToolMessage(content="ok", tool_call_id="c1")],
{"/img.png": _make_viewed_image(tmp_path)},
)
mw._inject(request)
assert _image_context_messages(request.messages) == []
class TestWrapModelCall:
def test_handler_receives_the_image_context_message(self, tmp_path):
mw = ViewImageMiddleware()
assistant = AIMessage(content="", tool_calls=[_view_image_call("c1")])
request = _model_request(
[assistant, ToolMessage(content="ok", tool_call_id="c1")],
{"/img.png": _make_viewed_image(tmp_path)},
)
seen: list[ModelRequest] = []
def handler(prepared: ModelRequest) -> AIMessage:
seen.append(prepared)
return AIMessage(content="I can see the image.")
result = mw.wrap_model_call(request, handler)
assert result.content == "I can see the image."
assert len(_image_context_messages(seen[0].messages)) == 1
def test_handler_receives_request_unchanged_when_not_warranted(self):
mw = ViewImageMiddleware()
request = _model_request([HumanMessage(content="hi")])
seen: list[ModelRequest] = []
mw.wrap_model_call(request, lambda prepared: seen.append(prepared) or AIMessage(content="ok"))
assert seen[0] is request
@pytest.mark.anyio
async def test_awrap_model_call_matches_sync_behavior(self, tmp_path):
mw = ViewImageMiddleware()
assistant = AIMessage(content="", tool_calls=[_view_image_call("c1")])
request = _model_request(
[assistant, ToolMessage(content="ok", tool_call_id="c1")],
{"/img.png": _make_viewed_image(tmp_path)},
)
seen: list[ModelRequest] = []
async def handler(prepared: ModelRequest) -> AIMessage:
seen.append(prepared)
return AIMessage(content="I can see the image.")
result = await mw.awrap_model_call(request, handler)
assert result.content == "I can see the image."
assert len(_image_context_messages(seen[0].messages)) == 1
class TestGraphIntegration:
def _graph_and_capture(self):
capture = _CaptureChatMessages()
model = FakeMessagesListChatModel(
responses=[AIMessage(content="I can see the image.")],
callbacks=[capture],
)
return create_agent(model=model, tools=[], middleware=[ViewImageMiddleware()]), capture
def _input(self, tmp_path):
return {
"messages": [
AIMessage(content="", tool_calls=[_view_image_call("c1")]),
ToolMessage(content="ok", tool_call_id="c1"),
],
"viewed_images": {"/img.png": _make_viewed_image(tmp_path)},
}
def test_image_context_reaches_the_model_but_never_the_state(self, tmp_path):
graph, capture = self._graph_and_capture()
result = graph.invoke(self._input(tmp_path))
model_image_messages = _image_context_messages(capture.messages)
assert len(model_image_messages) == 1
assert any(block.get("type") == "image_url" for block in model_image_messages[0].content)
# Nothing is written back, so the payload is absent from every checkpoint
# rather than being added and then removed again.
assert _image_context_messages(result["messages"]) == []
@pytest.mark.anyio
async def test_async_graph_matches_sync_behavior(self, tmp_path):
graph, capture = self._graph_and_capture()
result = await graph.ainvoke(self._input(tmp_path))
assert len(_image_context_messages(capture.messages)) == 1
assert _image_context_messages(result["messages"]) == []
def test_graph_preserves_normalized_client_message_with_reserved_prefix(self, tmp_path):
from app.gateway.services import normalize_input
client_id = "view-image-context:client-supplied"
normalized = normalize_input(
{
"messages": [
{
"role": "user",
"id": client_id,
"content": "client-authored message",
"additional_kwargs": {
_IMAGE_CONTEXT_MESSAGE_MARKER_KEY: True,
"custom": "keep-me",
},
}
]
}
)
client_message = normalized["messages"][0]
assert _IMAGE_CONTEXT_MESSAGE_MARKER_KEY not in client_message.additional_kwargs
graph, capture = self._graph_and_capture()
graph_input = self._input(tmp_path)
result = graph.invoke({**graph_input, "messages": [client_message, *graph_input["messages"]]})
assert any(message.id == client_id for message in capture.messages)
assert any(message.id != client_id and message.additional_kwargs.get(_IMAGE_CONTEXT_MESSAGE_MARKER_KEY) is True for message in _image_context_messages(capture.messages))
persisted_client = next(message for message in result["messages"] if message.id == client_id)
assert persisted_client.content == "client-authored message"
assert persisted_client.additional_kwargs == {"custom": "keep-me"}
assert all(message.additional_kwargs.get(_IMAGE_CONTEXT_MESSAGE_MARKER_KEY) is not True for message in result["messages"] if isinstance(message, HumanMessage))