open-swe/agent/middleware/check_message_queue.py
Johannes du Plessis b19804536c
feat: handle images sent to non-vision models in Slack, Linear, and web UI (#1560)
* feat: handle images sent to non-vision models in Slack, Linear, and web UI

Add vision capability checks across all image input paths. When a user
sends images to a text-only model (e.g. GLM 5.2, DeepSeek V4 Pro), the
images are now skipped and a warning is injected into the prompt instead
of sending unsupported content to the model.

- Slack: resolve model at webhook time, skip image fetch + add warning
- Linear: same pattern as Slack
- Queued message middleware: read resolved model from thread metadata,
  strip images from queued payloads for text-only models
- Web UI: disable submit + show inline warning when images are attached
  to a non-vision model selection
- Shared: resolve_agent_model_id helper + vision_not_supported_warning

Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>

* test: mock resolve_agent_model_id in Slack mention test

The test_process_slack_mention_queues_active_thread_message test was
missing a mock for the new resolve_agent_model_id call added to the
Slack webhook handler, causing a TypeError when image URLs triggered
the model resolution path.

Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>

* fix: include vision warning in queued payload for text-only models

Update the prompt variable (not just content_blocks) before clearing
image_urls so the queued payload also carries the warning text when a
Slack/Linear follow-up arrives while the thread is busy.

Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>

---------

Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
2026-06-17 09:12:52 -07:00

197 lines
6.8 KiB
Python

"""Before-model middleware that injects queued messages into state.
Checks the LangGraph store for pending messages (e.g. follow-up Linear
comments that arrived while the agent was busy) and injects them as new
human messages before the next model call.
"""
from __future__ import annotations
import logging
from typing import Any
import httpx
from langchain.agents.middleware import AgentState, before_model
from langgraph.config import get_config, get_store
from langgraph.runtime import Runtime
from langgraph_sdk import get_client
from ..dashboard.options import model_supports_images
from ..utils.multimodal import fetch_image_block, vision_not_supported_warning
logger = logging.getLogger(__name__)
DASHBOARD_HANDOFF_MARKER = "[Open SWE Web handoff]"
DASHBOARD_HANDOFF_INSTRUCTION = (
f"{DASHBOARD_HANDOFF_MARKER} This follow-up was sent from Web. "
"The conversation has moved to Web, so answer in the dashboard stream with a normal "
"assistant message. Do not call slack_thread_reply unless a later Slack message explicitly "
"moves the conversation back to Slack."
)
class LinearNotifyState(AgentState):
"""Extended agent state for tracking Linear notifications."""
linear_messages_sent_count: int
async def _resolve_thread_model_id(thread_id: str) -> str | None:
"""Read the resolved model from thread metadata (set by ``get_agent``)."""
try:
client = get_client()
thread = await client.threads.get(thread_id)
metadata = thread.get("metadata") if isinstance(thread, dict) else None
if not isinstance(metadata, dict):
return None
model = metadata.get("model")
return model if isinstance(model, str) and model else None
except Exception:
logger.debug("Could not read thread metadata for model resolution", exc_info=True)
return None
async def _build_blocks_from_payload(
payload: dict[str, Any],
*,
model_id: str | None = None,
) -> list[dict[str, Any]]:
text = payload.get("text", "")
image_urls = payload.get("image_urls", []) or []
images = payload.get("images", []) or []
blocks: list[dict[str, Any]] = []
if text:
blocks.append({"type": "text", "text": text})
if isinstance(images, list):
blocks.extend(image for image in images if isinstance(image, dict))
if not image_urls:
return blocks
if model_id and not model_supports_images(model_id):
logger.warning(
"Skipping %d queued image(s): model %s does not support images",
len(image_urls),
model_id,
)
if text:
blocks[0] = {
"type": "text",
"text": text + vision_not_supported_warning(model_id, len(image_urls)),
}
return blocks
async with httpx.AsyncClient() as client:
for image_url in image_urls:
image_block = await fetch_image_block(image_url, client)
if image_block:
blocks.append(image_block)
return blocks
def _is_dashboard_queued_message(content: object) -> bool:
return isinstance(content, dict) and content.get("source") == "dashboard"
@before_model(state_schema=LinearNotifyState)
async def check_message_queue_before_model( # noqa: PLR0911
state: LinearNotifyState, # noqa: ARG001
runtime: Runtime, # noqa: ARG001
) -> dict[str, Any] | None:
"""Middleware that checks for queued messages before each model call.
If messages are found in the queue for this thread, it extracts all messages,
adds them to the conversation state as new human messages, and clears the queue.
Messages are processed in FIFO order (oldest first).
This enables handling of follow-up comments that arrive while the agent is busy.
The agent will see the new messages and can incorporate them into its response.
"""
try:
config = get_config()
configurable = config.get("configurable", {})
thread_id = configurable.get("thread_id")
if not thread_id:
return None
try:
store = get_store()
except Exception as e: # noqa: BLE001
logger.debug("Could not get store from context: %s", e)
return None
if store is None:
return None
namespace = ("queue", thread_id)
try:
queued_item = await store.aget(namespace, "pending_messages")
except Exception as e: # noqa: BLE001
logger.warning("Failed to get queued item: %s", e)
return None
if queued_item is None:
return None
queued_value = queued_item.value
queued_messages = queued_value.get("messages", [])
# Delete early to prevent duplicate processing if middleware runs again
await store.adelete(namespace, "pending_messages")
if not queued_messages:
return None
logger.info(
"Found %d queued message(s) for thread %s, injecting into state",
len(queued_messages),
thread_id,
)
has_images = any(
isinstance(msg.get("content"), dict)
and (msg["content"].get("image_urls") or msg["content"].get("images"))
for msg in queued_messages
)
resolved_model_id: str | None = None
if has_images:
resolved_model_id = await _resolve_thread_model_id(thread_id)
content_blocks: list[dict[str, Any]] = []
for msg in queued_messages:
content = msg.get("content")
if _is_dashboard_queued_message(content):
content_blocks.append({"type": "text", "text": DASHBOARD_HANDOFF_INSTRUCTION})
if isinstance(content, dict) and (
"text" in content or "image_urls" in content or "images" in content
):
logger.debug("Queued message contains text + image URLs")
blocks = await _build_blocks_from_payload(content, model_id=resolved_model_id)
content_blocks.extend(blocks)
continue
if isinstance(content, list):
logger.debug("Queued message contains %d content block(s)", len(content))
content_blocks.extend(content)
continue
if isinstance(content, str) and content:
logger.debug("Queued message contains text content")
content_blocks.append({"type": "text", "text": content})
if not content_blocks:
return None
new_message = {
"role": "user",
"content": content_blocks,
}
logger.info(
"Injected %d queued message(s) into state for thread %s",
len(content_blocks),
thread_id,
)
return {"messages": [new_message]} # noqa: TRY300
except Exception:
logger.exception("Error in check_message_queue_before_model")
return None