From cb24bc26990bab6eb683779b6925e0dc6841e3b2 Mon Sep 17 00:00:00 2001 From: Hyeonsang Cho Date: Thu, 27 Aug 2026 10:12:37 +0900 Subject: [PATCH] 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. --- backend/AGENTS.md | 2 +- backend/docs/middleware-execution-flow.md | 53 +-- .../deerflow/agents/middlewares/AGENTS.md | 2 +- .../middlewares/view_image_middleware.py | 152 +++---- .../harness/deerflow/runtime/serialization.py | 8 +- backend/tests/test_view_image_middleware.py | 412 +++++++++--------- 6 files changed, 294 insertions(+), 335 deletions(-) diff --git a/backend/AGENTS.md b/backend/AGENTS.md index 7b823a704..2a247457f 100644 --- a/backend/AGENTS.md +++ b/backend/AGENTS.md @@ -289,7 +289,7 @@ See [docs/summarization.md](docs/summarization.md) for details. For models with `supports_vision: true`: - `ViewImageMiddleware` processes images in conversation - `view_image_tool` added to agent's toolset -- Images are converted to base64 and injected into a hidden message carrying both a reserved ID prefix and a server-owned metadata marker for the model call; Gateway strips that marker from untrusted input, and the middleware requires both identifiers before removing the message. The `before_model` and `model` node checkpoints for that call still contain the payload; after `after_model` cleanup, subsequent checkpoints retain only lightweight `viewed_images` metadata, while client-chosen IDs survive +- Images are converted to base64 and appended to the model request as a hidden message carrying both a reserved ID prefix and a server-owned metadata marker; Gateway strips that marker from untrusted input, and the middleware requires both identifiers to recognize its own message. The middleware injects inside `wrap_model_call`, so the payload never enters graph state: checkpoints retain only lightweight `viewed_images` metadata, while client-chosen IDs survive. It also sweeps its own message out of every request before rebuilding it, so a payload stranded in an older checkpoint by an interrupted run stops being resent ## Code Style diff --git a/backend/docs/middleware-execution-flow.md b/backend/docs/middleware-execution-flow.md index 99d638938..6fd5ecec2 100644 --- a/backend/docs/middleware-execution-flow.md +++ b/backend/docs/middleware-execution-flow.md @@ -16,7 +16,7 @@ | 7 | TodoMiddleware | | ✓ | ✓ | | ✓ | | ✓ | ✗ | `plan_mode` 参数 | | 8 | TitleMiddleware | | | ✓ | | | | ✓ | ✗ | `auto_title` | | 9 | MemoryMiddleware | | | | ✓ | | | ✓ | ✗ | `memory` | -| 10 | ViewImageMiddleware | | ✓ | | | | | ✓ | ✗ | `vision` | +| 10 | ViewImageMiddleware | | | | | ✓ | | ✓ | ✗ | `vision` | | 11 | SubagentLimitMiddleware | | | ✓ | | | | ✓ | ✗ | `subagent` | | 12 | LoopDetectionMiddleware | ✓ | | ✓ | ✓ | ✓ | | ✓ | ✗ | 始终开启 | | 13 | ClarificationMiddleware | | | | | | ✓ | ✓ | ✗ | 始终最后 | @@ -38,18 +38,12 @@ graph TB TD["[0] ThreadData
创建线程目录"] --> UL["[1] Uploads
扫描上传文件"] --> SB["[2] Sandbox
获取沙箱"] --> LD_BA["[12] LoopDetection
清理 stale warning"] end - subgraph BM ["before_model 正序 0→N"] + subgraph WM ["wrap_model_call 外→内 0→N"] direction TB - VI["[10] ViewImage
注入图片 base64"] + DTC_WM["[3] DanglingToolCall
补悬空 ToolMessage"] --> VI["[10] ViewImage
注入图片 base64(仅请求内)"] --> LD_WM["[12] LoopDetection
注入当前 run warning"] end - subgraph WM ["wrap_model_call"] - direction TB - DTC_WM["[3] DanglingToolCall
补悬空 ToolMessage"] --> LD_WM["[12] LoopDetection
注入当前 run warning"] - end - - LD_BA --> VI - VI --> DTC_WM + LD_BA --> DTC_WM LD_WM --> M["MODEL"] subgraph AM ["after_model 反序 N→0"] @@ -74,8 +68,8 @@ graph TB classDef afterAgentNode fill:#a0b5a8,stroke:#637a6b,color:#2d3239 classDef terminalNode fill:#a8b5a0,stroke:#6b7a63,color:#2d3239 - class TD,UL,SB,LD_BA,VI beforeNode - class DTC_WM,LD_WM wrapModelNode + class TD,UL,SB,LD_BA beforeNode + class DTC_WM,VI,LD_WM wrapModelNode class M modelNode class LD,SL,TI afterModelNode class LD_CLEAN,SBR,MEM afterAgentNode @@ -91,8 +85,8 @@ sequenceDiagram participant UL as UploadsMiddleware participant SB as SandboxMiddleware participant LD as LoopDetectionMiddleware - participant VI as ViewImageMiddleware participant DTC as DanglingToolCallMiddleware + participant VI as ViewImageMiddleware participant M as MODEL participant SL as SubagentLimitMiddleware participant TI as TitleMiddleware @@ -113,14 +107,14 @@ sequenceDiagram SB ->> LD: before_agent activate LD Note right of LD: before_agent 清理同 thread 旧 run 的 pending warning - LD ->> VI: before_model - activate VI - Note right of VI: before_model 注入图片 base64 - - VI ->> DTC: wrap_model_call + LD ->> DTC: wrap_model_call activate DTC Note right of DTC: wrap_model_call 补悬空 ToolMessage - DTC ->> LD: wrap_model_call + + DTC ->> VI: wrap_model_call + activate VI + Note right of VI: wrap_model_call 把图片 base64 追加到请求(不写入 state) + VI ->> LD: wrap_model_call Note right of LD: wrap_model_call drain 当前 run warning 并追加到末尾 LD ->> M: messages + tools activate M @@ -138,14 +132,14 @@ sequenceDiagram activate TI Note right of TI: after_model 生成标题 - TI -->> DTC: done + TI -->> VI: done deactivate TI - deactivate DTC - - VI -->> SB: done deactivate VI + DTC -->> SB: done + deactivate DTC + Note right of LD: after_agent 清理当前 run 未消费 warning Note right of MEM: after_agent 入队记忆 @@ -166,7 +160,8 @@ sequenceDiagram 列表位置决定在洋葱中的层级 — 位置 0 最外层,位置 N 最内层: ``` -进入 before_*: [0] → [1] → [2] → ... → [10] → MODEL +进入 before_*: [0] → [1] → [2] → ... → [7] → MODEL +进入 wrap_model_call: [3] → [10] → [12] → MODEL(外→内,同样正序) 退出 after_*: MODEL → [13] → [11] → ... → [6] → [3] → [2] → [0] ↑ 最内层最先执行 ``` @@ -233,8 +228,8 @@ sequenceDiagram participant UL as Uploads participant SB as Sandbox participant LD as LoopDetection - participant VI as ViewImage participant DTC as DanglingToolCall + participant VI as ViewImage participant M as MODEL participant SL as SubagentLimit participant TI as Title @@ -250,11 +245,11 @@ sequenceDiagram Note right of LD: before_agent 清理 stale pending warning loop 每轮对话(tool call 循环) - SB ->> VI: . - Note right of VI: before_model 注入图片 - VI ->> DTC: . + SB ->> DTC: . Note right of DTC: wrap_model_call 补悬空工具结果 - DTC ->> LD: . + DTC ->> VI: . + Note right of VI: wrap_model_call 把图片追加到请求 + VI ->> LD: . Note right of LD: wrap_model_call 注入当前 run warning LD ->> M: messages + tools M -->> LD: AI response diff --git a/backend/packages/harness/deerflow/agents/middlewares/AGENTS.md b/backend/packages/harness/deerflow/agents/middlewares/AGENTS.md index 79d8b26ca..6870deef3 100644 --- a/backend/packages/harness/deerflow/agents/middlewares/AGENTS.md +++ b/backend/packages/harness/deerflow/agents/middlewares/AGENTS.md @@ -84,7 +84,7 @@ Before changing a later authorization phase, read the [authorization RFC](../../ 20. **TokenUsageMiddleware** - *(optional, if `token_usage.enabled`)* Records token usage metrics; subagent usage is read from terminal `ToolMessage.additional_kwargs` in the current run and merged back into the dispatching AIMessage by message position. The same state update marks the ToolMessage with `subagent_token_usage_attributed=true`, so checkpoint replay or middleware re-entry cannot add the cumulative snapshot twice; missing/malformed usage or a result with no matching dispatch remains unmarked and retryable. 21. **TitleMiddleware** - Auto-generates the thread title after the first complete exchange and normalizes structured message content before prompting the title model. If a first-turn run is interrupted before this middleware can write a title, `runtime/runs/worker.py` keeps the run in a finalizing state, persists a local fallback title from the latest checkpoint or original run input, and then syncs it to `threads_meta.display_name`. Replacement runs admitted by `multitask_strategy="interrupt"` / `"rollback"` wait for older same-thread finalization before entering the graph; the interrupted run only skips the fallback title write once a later run has started and may have advanced the checkpoint. 22. **MemoryMiddleware** - Queues conversations for async memory update (filters to user + final AI responses); captures the runtime-resolved user so standalone LangGraph Server reads and writes stay in the same bucket -23. **ViewImageMiddleware** - *(optional, if the model supports vision)* Injects a hidden HumanMessage with base64 image data, identified by a reserved ID prefix plus a server-owned metadata marker, before the LLM call. Because `before_model`, `model`, and `after_model` are separate graph nodes, the `before_model` and `model` node checkpoints for that call still contain the payload; `after_model` / `aafter_model` then emits `RemoveMessage`, so subsequent checkpoints do not retain it +23. **ViewImageMiddleware** - *(optional, if the model supports vision)* Appends a hidden HumanMessage with base64 image data, identified by a reserved ID prefix plus a server-owned metadata marker, to `ModelRequest.messages` in `wrap_model_call` / `awrap_model_call`. The payload lives only in that request and is never returned as a state update, so no checkpoint carries it and an interrupted run cannot strand it in history; state keeps only the lightweight `viewed_images` metadata. It owns that context and rebuilds it per call: its own message is swept out of the request first — a thread checkpointed by the earlier `before_model`/`after_model` pair (which wrote the payload into state and took it back out with `RemoveMessage`) can carry one that reached state but was never removed, and leaving it in would resend that base64 in every later request for the life of the thread — then a freshly built one is appended when warranted. The sweep requires both the reserved ID prefix and the server-owned marker, so a client cannot get its own message dropped; unmarked leftovers predating the marker are left in place and merely not duplicated 24. **McpRoutingMiddleware** - *(optional, if `tool_search.enabled` and PR1 MCP routing metadata produce a routing index)* Auto-promotes matching deferred MCP tool schemas before the model call by writing a minimal `promoted` state update. It matches only the latest real `HumanMessage`, uses the global `tool_search.auto_promote_top_k` limit (default 3, clamped to 1..5), never executes tools, and must be installed before `DeferredToolFilterMiddleware` 25. **DeferredToolFilterMiddleware** - *(optional, if `tool_search.enabled`)* Hides deferred (MCP) tool schemas from the bound model until `tool_search` or `McpRoutingMiddleware` promotes them (reads per-thread promotions from `ThreadState.promoted`, hash-scoped) 26. **SystemMessageCoalescingMiddleware** - Merges every SystemMessage into a single leading SystemMessage per request; provider-agnostic fix for strict backends (vLLM/SGLang/Qwen/Anthropic) that reject non-leading system messages. Touches the per-request payload only (checkpoint state unchanged); on midnight crossings only the latest `dynamic_context_reminder` SystemMessage survives. The subagent builder places its date-only context middleware immediately before this coalescer, so the built-in subagent prompt and hidden date reminder still reach providers as one leading system block diff --git a/backend/packages/harness/deerflow/agents/middlewares/view_image_middleware.py b/backend/packages/harness/deerflow/agents/middlewares/view_image_middleware.py index d1f1a4225..3f148cd54 100644 --- a/backend/packages/harness/deerflow/agents/middlewares/view_image_middleware.py +++ b/backend/packages/harness/deerflow/agents/middlewares/view_image_middleware.py @@ -1,16 +1,17 @@ -"""Middleware for injecting image details into conversation before LLM call.""" +"""Middleware for injecting image details into the model request.""" import asyncio import base64 import logging +from collections.abc import Awaitable, Callable from pathlib import Path from typing import override from uuid import uuid4 from deerflow_extension_api import ContentKind, provenance_kwargs from langchain.agents.middleware import AgentMiddleware -from langchain_core.messages import AIMessage, HumanMessage, RemoveMessage, ToolMessage -from langgraph.runtime import Runtime +from langchain.agents.middleware.types import ModelCallResult, ModelRequest, ModelResponse +from langchain_core.messages import AIMessage, AnyMessage, HumanMessage, ToolMessage from deerflow.agents.thread_state import ThreadState @@ -29,18 +30,24 @@ class ViewImageMiddlewareState(ThreadState): class ViewImageMiddleware(AgentMiddleware[ViewImageMiddlewareState]): - """Injects image details as a human message before LLM calls when view_image tools have completed. + """Injects image details into the model request when view_image tool calls have completed. This middleware: - 1. Runs before each LLM call + 1. Wraps each LLM call 2. Checks if the last assistant message contains view_image tool calls 3. Verifies all tool calls in that message have been completed (have corresponding ToolMessages) - 4. If conditions are met, creates a human message with all viewed image details (including base64 data) - 5. Adds the message to state so the LLM can see and analyze the images - 6. Removes the transient message after the LLM call so later checkpoints do not retain its base64 data + 4. If conditions are met, appends a human message with all viewed image details (including base64 data) + 5. Hands the augmented request to the model so it can see and analyze the images This enables the LLM to automatically receive and analyze images that were loaded via view_image tool, without requiring explicit user prompts to describe the images. + + Injection happens in ``wrap_model_call`` on purpose: the message exists only + in ``ModelRequest.messages`` and is never returned as a state update, so no + checkpoint carries the base64 payload and an interrupted run cannot strand it + in history. Do not move this back to a ``before_model``/``after_model`` pair + -- that writes the payload into state and can only take it out again + afterwards (see #4267). """ state_schema = ViewImageMiddlewareState @@ -183,16 +190,15 @@ class ViewImageMiddleware(AgentMiddleware[ViewImageMiddlewareState]): return content_blocks - def _should_inject_image_message(self, state: ViewImageMiddlewareState) -> bool: - """Determine if we should inject an image details message. + def _should_inject_image_message(self, messages: list[AnyMessage]) -> bool: + """Determine if we should append an image details message. Args: - state: Current state + messages: Messages about to be sent to the model Returns: - True if we should inject the message + True if we should append the message """ - messages = state.get("messages", []) if not messages: return False @@ -209,13 +215,14 @@ class ViewImageMiddleware(AgentMiddleware[ViewImageMiddlewareState]): if not self._all_tools_completed(messages, last_assistant_msg): return False - # Check if we've already added an image details message - # Look for a human message after the last assistant message that contains image details + # Skip when image details are already present. ``_inject`` has stripped + # this middleware's own messages by now, so what remains are unmarked + # ones from checkpoints written before the marker existed -- those cannot + # be told apart from user-authored text with certainty, so they are left + # in place and simply not duplicated. assistant_idx = messages.index(last_assistant_msg) for msg in messages[assistant_idx + 1 :]: if isinstance(msg, HumanMessage): - if self._is_image_context_message(msg): - return False content_str = str(msg.content) if "Here are the images you've viewed" in content_str or "Here are the details of the images you've viewed" in content_str: # Already added, don't add again @@ -236,86 +243,55 @@ class ViewImageMiddleware(AgentMiddleware[ViewImageMiddlewareState]): }, ) - @staticmethod - def _remove_image_context_messages(state: ViewImageMiddlewareState) -> dict | None: - """Remove transient image context messages after the model consumed them.""" - removals = [RemoveMessage(id=msg.id) for msg in state.get("messages", []) if ViewImageMiddleware._is_image_context_message(msg)] - if not removals: - return None - return {"messages": removals} - - def _inject_image_message(self, state: ViewImageMiddlewareState) -> dict | None: - """Internal helper to inject image details message. + def _inject(self, request: ModelRequest) -> ModelRequest: + """Rebuild the request's image context from ``viewed_images``. Args: - state: Current state + request: The pending model request Returns: - State update with additional human message, or None if no update needed + A request whose messages carry exactly the image context this call + warrants -- one freshly built message, or none -- leaving the + original request untouched when there is nothing to change """ - if not self._should_inject_image_message(state): - return None + # This middleware owns the image context and rebuilds it from + # ``viewed_images`` on every call, so drop any copy already in the list. + # A thread checkpointed by the earlier before_model/after_model pair can + # carry one that reached state but was never removed (the run died during + # the model call); left in place it 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 this can never drop a user-authored message. + messages = [message for message in request.messages if not self._is_image_context_message(message)] + dropped_stranded = len(messages) != len(request.messages) + if dropped_stranded: + logger.debug("Dropping %d stranded image context message(s) from the model request", len(request.messages) - len(messages)) - # Create the image details message with text and image content - image_content = self._create_image_details_message(state) + if not self._should_inject_image_message(messages): + return request.override(messages=messages) if dropped_stranded else request - # Create a new human message with mixed content (text + images). This is - # internal context for the model only, so hide it from the chat UI and IM - # channels (matches the other middleware-injected context messages). - human_msg = self._create_image_context_message(image_content) + # Mixed content (text + images) for the model only, so hide it from the + # chat UI and IM channels (matches the other middleware-injected context + # messages) even though it never leaves this request. + image_content = self._create_image_details_message(request.state or {}) + logger.debug("Injecting image details message with images into the model request") - logger.debug("Injecting image details message with images before LLM call") - - # Return state update with the new message - return {"messages": [human_msg]} + return request.override(messages=[*messages, self._create_image_context_message(image_content)]) @override - def before_model(self, state: ViewImageMiddlewareState, runtime: Runtime) -> dict | None: - """Inject image details message before LLM call if view_image tools have completed (sync version). - - This runs before each LLM call, checking if the previous turn included view_image - tool calls that have all completed. If so, it injects a human message with the image - details so the LLM can see and analyze the images. - - Args: - state: Current state - runtime: Runtime context (unused but required by interface) - - Returns: - State update with additional human message, or None if no update needed - """ - return self._inject_image_message(state) + def wrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], ModelResponse], + ) -> ModelCallResult: + return handler(self._inject(request)) @override - async def abefore_model(self, state: ViewImageMiddlewareState, runtime: Runtime) -> dict | None: - """Inject image details message before LLM call if view_image tools have completed (async version). - - This runs before each LLM call, checking if the previous turn included view_image - tool calls that have all completed. If so, it injects a human message with the image - details so the LLM can see and analyze the images. - - Args: - state: Current state - runtime: Runtime context (unused but required by interface) - - Returns: - State update with additional human message, or None if no update needed - """ - if not self._should_inject_image_message(state): - return None - # Image reads + base64 encoding can be slow (up to 20MB), so offload - # the blocking work to a thread rather than stalling the event loop. - image_content = await asyncio.to_thread(self._create_image_details_message, state) - human_msg = self._create_image_context_message(image_content) - logger.debug("Injecting image details message with images before LLM call") - return {"messages": [human_msg]} - - @override - def after_model(self, state: ViewImageMiddlewareState, runtime: Runtime) -> dict | None: - """Remove model-only image data before subsequent checkpoints (sync version).""" - return self._remove_image_context_messages(state) - - @override - async def aafter_model(self, state: ViewImageMiddlewareState, runtime: Runtime) -> dict | None: - """Remove model-only image data before subsequent checkpoints (async version).""" - return self._remove_image_context_messages(state) + async def awrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], Awaitable[ModelResponse]], + ) -> ModelCallResult: + # Image reads + base64 encoding can be slow (up to 20MB), so offload the + # blocking work to a thread rather than stalling the event loop. + return await handler(await asyncio.to_thread(self._inject, request)) diff --git a/backend/packages/harness/deerflow/runtime/serialization.py b/backend/packages/harness/deerflow/runtime/serialization.py index a9b5cea2f..6d0c84587 100644 --- a/backend/packages/harness/deerflow/runtime/serialization.py +++ b/backend/packages/harness/deerflow/runtime/serialization.py @@ -75,9 +75,11 @@ def strip_data_url_image_blocks(messages: list[dict[str, Any]]) -> list[dict[str """Remove ``data:``-scheme ``image_url`` blocks from *hide_from_ui* messages. The history and run-wait endpoints return checkpoint-persisted messages to - the frontend. ``ViewImageMiddleware`` stores full base64 image payloads in - ``hide_from_ui`` human messages — these are internal model context and must - not be sent over the wire (huge response bodies, no UI value). + the frontend. ``ViewImageMiddleware`` now keeps its base64 image payloads + inside the model request, but threads checkpointed by earlier versions still + hold them in ``hide_from_ui`` human messages — these are internal model + context and must not be sent over the wire (huge response bodies, no UI + value). Only content blocks of type ``image_url`` whose URL starts with ``data:`` are stripped. Text blocks, ``https://`` image URLs, and non-hidden diff --git a/backend/tests/test_view_image_middleware.py b/backend/tests/test_view_image_middleware.py index cf58ab46f..4018441d0 100644 --- a/backend/tests/test_view_image_middleware.py +++ b/backend/tests/test_view_image_middleware.py @@ -1,9 +1,9 @@ """Unit tests for ViewImageMiddleware. Tests cover the middleware's ability to inject image details (including base64 -payloads) as a HumanMessage before the next LLM call, triggered only when the -previous assistant turn contained `view_image` tool calls that have all been -completed with corresponding ToolMessages. +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). @@ -12,12 +12,13 @@ Covered behavior: - `_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 was already added. -- `_inject_image_message` returns a state update with a HumanMessage, or None - when injection is not warranted. -- `before_model` and `abefore_model` expose the same behavior sync/async. -- `after_model` and `aafter_model` remove only the transient image message so - later checkpoints do not retain its base64 payload. + 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 @@ -26,10 +27,10 @@ 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, HumanMessage, RemoveMessage, SystemMessage, ToolMessage -from langgraph.graph.message import add_messages +from langchain_core.messages import AIMessage, AnyMessage, HumanMessage, SystemMessage, ToolMessage from deerflow.agents.middlewares.view_image_middleware import ( _IMAGE_CONTEXT_MESSAGE_MARKER_KEY, @@ -45,10 +46,19 @@ def _other_tool_call(call_id: str = "call_other", name: str = "bash") -> dict: return {"name": name, "id": call_id, "args": {"command": "ls"}} -def _runtime() -> MagicMock: - """Minimal Runtime stub. The middleware doesn't use it today, but the - interface requires it.""" - return MagicMock() +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): @@ -59,6 +69,10 @@ class _CaptureChatMessages(BaseCallbackHandler): 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 @@ -310,126 +324,98 @@ class TestCreateImageDetailsMessage: class TestShouldInjectImageMessage: def test_false_when_no_messages(self): mw = ViewImageMiddleware() - assert mw._should_inject_image_message({"messages": []}) is False - - def test_false_when_messages_key_missing(self): - mw = ViewImageMiddleware() - assert mw._should_inject_image_message({}) is False + assert mw._should_inject_image_message([]) is False def test_false_when_no_assistant_message(self): mw = ViewImageMiddleware() - state = {"messages": [HumanMessage(content="hello")]} - assert mw._should_inject_image_message(state) is False + 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()]) - state = { - "messages": [assistant, ToolMessage(content="ok", tool_call_id="call_other")], - } - assert mw._should_inject_image_message(state) is False + 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")]) - state = {"messages": [assistant]} # no ToolMessage yet - assert mw._should_inject_image_message(state) is False + assert mw._should_inject_image_message([assistant]) is False # no ToolMessage yet - def test_true_when_all_preconditions_met(self, tmp_path): + def test_true_when_all_preconditions_met(self): mw = ViewImageMiddleware() assistant = AIMessage(content="", tool_calls=[_view_image_call("c1")]) - img_meta = _make_viewed_image(tmp_path) - state = { - "messages": [assistant, ToolMessage(content="ok", tool_call_id="c1")], - "viewed_images": {"/img.png": img_meta}, - } - assert mw._should_inject_image_message(state) is True + 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, tmp_path): - """If a HumanMessage with the recognized header is already present after - the assistant turn, we must not inject a duplicate.""" + 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") - img_meta = _make_viewed_image(tmp_path) - state = { - "messages": [ - assistant, - ToolMessage(content="ok", tool_call_id="c1"), - already_injected, - ], - "viewed_images": {"/img.png": img_meta}, - } - assert mw._should_inject_image_message(state) is False + 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. - The middleware's own `_inject_image_message` creates a HumanMessage - whose `.content` is 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 `_should_inject_image_message` - still detects the marker via `str(msg.content)`. + 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")]) - img_meta = _make_viewed_image(tmp_path) - viewed_images = {"/img.png": img_meta} + 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) - already_injected = HumanMessage(content=real_injected_content) - state = { - "messages": [ - assistant, - ToolMessage(content="ok", tool_call_id="c1"), - already_injected, - ], - "viewed_images": viewed_images, - } - assert mw._should_inject_image_message(state) is False + 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, tmp_path): + 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: ...") - img_meta = _make_viewed_image(tmp_path) - state = { - "messages": [ - assistant, - ToolMessage(content="ok", tool_call_id="c1"), - legacy, - ], - "viewed_images": {"/img.png": img_meta}, - } - assert mw._should_inject_image_message(state) is False + messages = [ + assistant, + ToolMessage(content="ok", tool_call_id="c1"), + legacy, + ] + assert mw._should_inject_image_message(messages) is False -class TestInjectImageMessage: - def test_returns_none_when_should_not_inject(self): +class TestInject: + def test_returns_request_unchanged_when_should_not_inject(self): mw = ViewImageMiddleware() - state = {"messages": []} - assert mw._inject_image_message(state) is None + request = _model_request([HumanMessage(content="hi")]) + assert mw._inject(request) is request - def test_returns_state_update_with_human_message(self, tmp_path): + def test_appends_image_context_message_to_request(self, tmp_path): mw = ViewImageMiddleware() assistant = AIMessage(content="", tool_calls=[_view_image_call("c1")]) - img_meta = _make_viewed_image(tmp_path) - state = { - "messages": [assistant, ToolMessage(content="ok", tool_call_id="c1")], - "viewed_images": {"/img.png": img_meta}, - } + original = [assistant, ToolMessage(content="ok", tool_call_id="c1")] + request = _model_request(original, {"/img.png": _make_viewed_image(tmp_path)}) - result = mw._inject_image_message(state) + injected_request = mw._inject(request) - assert isinstance(result, dict) - assert "messages" in result - assert len(result["messages"]) == 1 - injected = result["messages"][0] + 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) @@ -441,110 +427,146 @@ class TestInjectImageMessage: assert injected.id is not None assert injected.id.startswith("view-image-context:") - -class TestBeforeModel: - def test_before_model_returns_none_when_preconditions_not_met(self): - mw = ViewImageMiddleware() - state = {"messages": [HumanMessage(content="hi")]} - assert mw.before_model(state, _runtime()) is None - - def test_before_model_returns_injection_when_ready(self, tmp_path): + 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")]) - img_meta = _make_viewed_image(tmp_path) - state = { - "messages": [assistant, ToolMessage(content="ok", tool_call_id="c1")], - "viewed_images": {"/img.png": img_meta}, - } - result = mw.before_model(state, _runtime()) - assert result is not None - assert isinstance(result["messages"][0], HumanMessage) + 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)}, + ) - @pytest.mark.anyio - async def test_abefore_model_matches_sync_behavior(self, 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")]) - img_meta = _make_viewed_image(tmp_path) - state = { - "messages": [assistant, ToolMessage(content="ok", tool_call_id="c1")], - "viewed_images": {"/img.png": img_meta}, - } - result = await mw.abefore_model(state, _runtime()) - assert result is not None - assert isinstance(result["messages"][0], HumanMessage) + 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_abefore_model_returns_none_when_no_injection(self): + async def test_awrap_model_call_matches_sync_behavior(self, tmp_path): mw = ViewImageMiddleware() - state = {"messages": []} - assert await mw.abefore_model(state, _runtime()) is None + 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 TestAfterModel: - def test_graph_exposes_image_context_only_during_model_call(self, tmp_path): +class TestGraphIntegration: + def _graph_and_capture(self): capture = _CaptureChatMessages() model = FakeMessagesListChatModel( responses=[AIMessage(content="I can see the image.")], callbacks=[capture], ) - graph = create_agent(model=model, tools=[], middleware=[ViewImageMiddleware()]) - assistant = AIMessage(content="", tool_calls=[_view_image_call("c1")]) - img_meta = _make_viewed_image(tmp_path) + return create_agent(model=model, tools=[], middleware=[ViewImageMiddleware()]), capture - result = graph.invoke( - { - "messages": [assistant, ToolMessage(content="ok", tool_call_id="c1")], - "viewed_images": {"/img.png": img_meta}, - } - ) + 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)}, + } - model_image_messages = [message for message in capture.messages if isinstance(message, HumanMessage) and message.id and message.id.startswith("view-image-context:")] + 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) - assert all(not (isinstance(message, HumanMessage) and message.id and message.id.startswith("view-image-context:")) for message in result["messages"]) + # 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"]) == [] - def test_removes_transient_image_message_from_later_state(self, tmp_path): - mw = ViewImageMiddleware() - assistant = AIMessage(content="", tool_calls=[_view_image_call("c1")]) - tool_result = ToolMessage(content="ok", tool_call_id="c1") - img_meta = _make_viewed_image(tmp_path) - before_state = { - "messages": [assistant, tool_result], - "viewed_images": {"/img.png": img_meta}, - } - injected = mw.before_model(before_state, _runtime())["messages"][0] - model_response = AIMessage(content="I can see the image.") - model_state = { - **before_state, - "messages": [assistant, tool_result, injected, model_response], - } + @pytest.mark.anyio + async def test_async_graph_matches_sync_behavior(self, tmp_path): + graph, capture = self._graph_and_capture() - result = mw.after_model(model_state, _runtime()) + result = await graph.ainvoke(self._input(tmp_path)) - assert result is not None - assert len(result["messages"]) == 1 - removal = result["messages"][0] - assert isinstance(removal, RemoveMessage) - assert removal.id == injected.id - - checkpoint_messages = add_messages(model_state["messages"], result["messages"]) - assert injected.id not in {message.id for message in checkpoint_messages} - assert model_response in checkpoint_messages - - def test_does_not_remove_unmarked_human_messages(self): - mw = ViewImageMiddleware() - state = { - "messages": [ - HumanMessage( - id="view-image-context:client-supplied", - content="Here are the images you've viewed: user-authored text", - additional_kwargs={"hide_from_ui": True}, - ), - AIMessage(content="response"), - ] - } - - assert mw.after_model(state, _runtime()) is None + 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 @@ -568,50 +590,14 @@ class TestAfterModel: client_message = normalized["messages"][0] assert _IMAGE_CONTEXT_MESSAGE_MARKER_KEY not in client_message.additional_kwargs - capture = _CaptureChatMessages() - model = FakeMessagesListChatModel( - responses=[AIMessage(content="I can see the image.")], - callbacks=[capture], - ) - graph = create_agent(model=model, tools=[], middleware=[ViewImageMiddleware()]) - assistant = AIMessage(content="", tool_calls=[_view_image_call("c1")]) + graph, capture = self._graph_and_capture() + graph_input = self._input(tmp_path) - result = graph.invoke( - { - "messages": [ - client_message, - assistant, - ToolMessage(content="ok", tool_call_id="c1"), - ], - "viewed_images": {"/img.png": _make_viewed_image(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(isinstance(message, HumanMessage) and message.id != client_id and message.additional_kwargs.get(_IMAGE_CONTEXT_MESSAGE_MARKER_KEY) is True 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)) - - @pytest.mark.anyio - async def test_aafter_model_matches_sync_cleanup(self, tmp_path): - mw = ViewImageMiddleware() - assistant = AIMessage(content="", tool_calls=[_view_image_call("c1")]) - tool_result = ToolMessage(content="ok", tool_call_id="c1") - img_meta = _make_viewed_image(tmp_path) - before_state = { - "messages": [assistant, tool_result], - "viewed_images": {"/img.png": img_meta}, - } - injected = (await mw.abefore_model(before_state, _runtime()))["messages"][0] - model_state = { - **before_state, - "messages": [assistant, tool_result, injected, AIMessage(content="response")], - } - - result = await mw.aafter_model(model_state, _runtime()) - - assert result is not None - assert isinstance(result["messages"][0], RemoveMessage) - assert result["messages"][0].id == injected.id