Enhance multimodal handling in LLM model (#256)

* Enhance multimodal handling in LLM model

- Updated `_flatten_message_content` to preserve media blocks (images, files) while flattening text content.
- Introduced `_sanitize_messages` to manage media hoisting for tool messages, ensuring compatibility with OpenAI APIs.
- Modified `_patch_openai_compat_content` to accommodate new media handling logic, including retry mechanisms for media errors.
- Added comprehensive tests for media preservation, including various scenarios with images, files, and unsupported media types.

* fix: preserve order of text and media blocks in message flattening

* test: add tests for _strip_media_types to ensure position preservation and deduplication
This commit is contained in:
Xi Zhang
2026-06-03 01:06:46 +01:00
committed by GitHub
parent d348076f40
commit 9cffe9d457
3 changed files with 1152 additions and 29 deletions
+319 -28
View File
@@ -194,62 +194,260 @@ def _is_ccproxy_codex() -> bool:
# ---------------------------------------------------------------------------
# Utility + Patch: Flatten list content to strings for OpenAI-compatible APIs.
# DeepSeek, SiliconFlow, etc. reject assistant messages whose content is a
# list rather than a string.
# list rather than a string. Image and file (PDF/document) blocks are
# preserved, not flattened away.
# ---------------------------------------------------------------------------
_SKIP_CONTENT_TYPES = frozenset({"thinking", "reasoning", "reasoning_content"})
# Media block types preserved when flattening (positive allowlist;
# thinking/reasoning still dropped). Images + files (PDF/documents): both
# serialize on OpenAI-compatible APIs and capable models read them. `video` is
# deliberately EXCLUDED (langchain-openai raises ValueError on it); `audio` is
# omitted (almost no model support). is_data_content_block is unreliable
# (False for OpenAI image_url / Anthropic image-source). Kept as separate
# image/file sets so the no-media fallback can gate per modality.
_IMAGE_CONTENT_TYPES = frozenset({"image", "image_url", "input_image"})
_FILE_CONTENT_TYPES = frozenset({"file", "input_file", "document"})
_MEDIA_CONTENT_TYPES = _IMAGE_CONTENT_TYPES | _FILE_CONTENT_TYPES
def _flatten_message_content(content: Any) -> str | Any:
"""Convert list-of-blocks content to a plain string.
def _flatten_message_content(content: Any) -> str | list[Any] | Any:
"""Convert list-of-blocks content to a string, preserving media blocks.
Thinking/reasoning blocks are dropped. When a media block (image or file)
is present, returns a list in the ORIGINAL order — consecutive text is
joined into a single text block, media blocks kept as-is — so captions stay
next to the attachment they describe. Otherwise returns a plain string.
Non-list input is returned unchanged.
Args:
content: Message content — either a string, a list of content blocks
(dicts with ``type`` and ``text`` keys), or another type.
content: Message content — a string, a list of content blocks, or
another type.
Returns:
A plain string with text blocks joined by double newlines.
Thinking/reasoning blocks are skipped. Non-list input is
returned unchanged.
A plain string for text-only content, a list of blocks when media is
present, or the input unchanged for non-list input.
"""
if isinstance(content, str):
return content
if not isinstance(content, list):
return content
parts: list[str] = []
ordered_blocks: list[Any] = []
saw_media = False
def _flush_text() -> None:
if parts:
ordered_blocks.append({"type": "text", "text": "\n\n".join(parts)})
parts.clear()
for block in content:
if isinstance(block, dict):
if block.get("type") in _SKIP_CONTENT_TYPES:
btype = block.get("type")
if btype in _MEDIA_CONTENT_TYPES:
# Keep media as-is (never mutate; upstream copy.copy is shallow)
# and preserve its position relative to surrounding text.
_flush_text()
ordered_blocks.append(block)
saw_media = True
continue
if btype in _SKIP_CONTENT_TYPES:
continue
text = block.get("text")
if text:
parts.append(text)
elif isinstance(block, str):
parts.append(block)
if saw_media:
_flush_text()
return ordered_blocks
return "\n\n".join(parts) if parts else ""
def _patch_openai_compat_content(model: Any) -> None:
def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> list[Any]:
"""Flatten list content for OpenAI-compatible APIs, preserving media.
Text/reasoning content is flattened to a string; image blocks are
preserved. Tool messages cannot carry media
on OpenAI-compatible APIs (content must be a string), so when
``hoist_tool_media`` is set the media in a tool result is moved into a
HumanMessage emitted after that turn's (possibly parallel) tool messages.
Anthropic-routed providers accept tool-result media natively and pass
``hoist_tool_media=False`` to keep it inline.
"""
import copy
from langchain_core.messages import HumanMessage
out: list[Any] = []
pending_media: list[Any] = [] # media hoisted out of a run of tool messages
def _flush() -> None:
if pending_media:
out.append(HumanMessage(content=list(pending_media)))
pending_media.clear()
for msg in messages:
is_tool = getattr(msg, "type", None) == "tool"
if not is_tool:
_flush() # emit hoisted media before any non-tool message
if not isinstance(msg.content, list):
out.append(msg)
continue
flat = _flatten_message_content(msg.content)
if hoist_tool_media and is_tool and isinstance(flat, list):
text_blocks = [
b for b in flat if isinstance(b, dict) and b.get("type") == "text"
]
media_blocks = [
b for b in flat if not (isinstance(b, dict) and b.get("type") == "text")
]
tool_msg = copy.copy(msg)
# Join ALL text runs (interleaved content can yield more than one)
# so no text is lost; tool content must be a string on OpenAI-compat.
tool_msg.content = (
"\n\n".join(b["text"] for b in text_blocks)
if text_blocks
else "[media content provided in the following message]"
)
out.append(tool_msg)
pending_media.extend(media_blocks)
else:
msg = copy.copy(msg)
msg.content = flat
out.append(msg)
_flush() # conversation may end with tool messages
return out
# Fallback for models that reject some media input (no vision / no file
# support): replace blocks of the unsupported types with a placeholder so the
# conversation can keep going instead of erroring every turn that re-sends them.
_UNSUPPORTED_MEDIA_PLACEHOLDER = (
"[attachment omitted: this model does not support this input type]"
)
# Markers that identify WHICH modality an error rejects, so only the implicated
# type is remembered (not all media). Generic markers map to all media.
_IMAGE_ERROR_MARKERS = ("image_url", "image input", "support image")
_FILE_ERROR_MARKERS = ("file input", "pdf input", "document input")
# `expected \`text\`` keeps the backticks — the bare "expected text" phrase is
# too broad (matches non-media schema errors like "... expected text").
_GENERIC_MEDIA_MARKERS = ("multimodal", "expected `text`")
def _media_types_in(messages: list[Any]) -> set[str]:
"""Set of preserved-media block types present across the messages."""
found: set[str] = set()
for msg in messages:
content = getattr(msg, "content", None)
if isinstance(content, list):
for b in content:
if isinstance(b, dict) and b.get("type") in _MEDIA_CONTENT_TYPES:
found.add(b["type"])
return found
def _media_error_types(exc: Exception) -> set[str]:
"""Media modalities the error text implicates (empty if not media-specific).
Used to remember ONLY the rejected modality, so e.g. a PDF rejection never
disables images.
"""
text = str(exc).lower()
types: set[str] = set()
if any(m in text for m in _IMAGE_ERROR_MARKERS):
types |= _IMAGE_CONTENT_TYPES
if any(m in text for m in _FILE_ERROR_MARKERS):
types |= _FILE_CONTENT_TYPES
if any(m in text for m in _GENERIC_MEDIA_MARKERS):
types |= _MEDIA_CONTENT_TYPES
return types
def _is_http_400(exc: Exception) -> bool:
"""True if the error carries an HTTP 400 status (bad request)."""
status = getattr(exc, "status_code", None)
if status is None:
status = getattr(getattr(exc, "response", None), "status_code", None)
return status == 400
def _strip_media_types(messages: list[Any], types: set[str]) -> list[Any]:
"""Replace blocks of the given types with a placeholder text block.
Each stripped block is replaced IN PLACE (consecutive ones collapse into one
placeholder), so surrounding text/media keep their original positions and a
model that rejects only one modality still receives the others in order.
"""
import copy
out: list[Any] = []
for msg in messages:
content = getattr(msg, "content", None)
if not isinstance(content, list) or not any(
isinstance(b, dict) and b.get("type") in types for b in content
):
out.append(msg)
continue
kept: list[Any] = []
last_was_placeholder = False
for b in content:
if isinstance(b, dict) and b.get("type") in types:
if not last_was_placeholder:
kept.append(
{"type": "text", "text": _UNSUPPORTED_MEDIA_PLACEHOLDER}
)
last_was_placeholder = True
continue
kept.append(b)
last_was_placeholder = False
msg = copy.copy(msg)
msg.content = kept
out.append(msg)
return out
def _patch_openai_compat_content(model: Any, hoist_tool_media: bool = True) -> None:
"""Flatten list content to strings before OpenAI-compatible API calls.
Wraps ``_generate`` / ``_agenerate`` to prevent "invalid type: sequence,
expected a string" errors from strict APIs like DeepSeek.
Wraps ``_generate`` / ``_agenerate`` / ``_stream`` / ``_astream`` to prevent
"invalid type: sequence, expected a string" errors from strict APIs like
DeepSeek, while preserving image and file (PDF/document) blocks. Tracks a
per-modality ``blocked`` set of media types the model rejects: when a request
fails with a media-looking error it is retried with the offending modalities
stripped to a placeholder, and those types are remembered ONLY after the
stripped retry succeeds — so an unrelated 400 never permanently degrades the
model, and a PDF rejection never disables images. The profile pre-seeds the
set per modality (``image_inputs``/``pdf_inputs`` ``is False``).
Args:
model: A LangChain chat model instance to patch in-place.
hoist_tool_media: When True (OpenAI-compatible), media in a tool result
is hoisted into a following HumanMessage (tool content must be a
string). Anthropic-routed providers pass False (native support).
"""
import copy
import functools
from langchain_core.messages import BaseMessage
def _sanitize_messages(messages: list[BaseMessage]) -> list[BaseMessage]:
out: list[BaseMessage] = []
for msg in messages:
if isinstance(msg.content, list):
msg = copy.copy(msg)
msg.content = _flatten_message_content(msg.content)
out.append(msg)
return out
# Media block types this model rejects. Pre-seeded per modality from the
# profile, then grown reactively — but only after a stripped retry succeeds.
profile = getattr(model, "profile", None)
blocked: set[str] = set()
if isinstance(profile, dict):
if profile.get("image_inputs") is False:
blocked |= _IMAGE_CONTENT_TYPES
if profile.get("pdf_inputs") is False:
blocked |= _FILE_CONTENT_TYPES
def _prepare(messages: list[BaseMessage]) -> list[BaseMessage]:
msgs = _strip_media_types(messages, blocked) if blocked else messages
return _sanitize_messages(msgs, hoist_tool_media)
def _stripped(messages: list[BaseMessage], suspects: set[str]) -> list[BaseMessage]:
return _sanitize_messages(
_strip_media_types(messages, blocked | suspects), hoist_tool_media
)
orig_generate = getattr(model, "_generate", None)
if orig_generate is None:
@@ -259,7 +457,23 @@ def _patch_openai_compat_content(model: Any) -> None:
def _patched_generate(
messages: list[BaseMessage], *args: Any, **kwargs: Any
) -> Any:
return orig_generate(_sanitize_messages(messages), *args, **kwargs)
prepared = _prepare(messages)
suspects = _media_types_in(messages) - blocked
if not suspects:
return orig_generate(prepared, *args, **kwargs)
try:
return orig_generate(prepared, *args, **kwargs)
except Exception as exc:
culprit = _media_error_types(exc) & suspects
if not culprit and not _is_http_400(exc):
raise
try:
result = orig_generate(_stripped(messages, suspects), *args, **kwargs)
except Exception:
# stripping didn't help — surface original, don't cache
raise exc from None
blocked.update(culprit) # cache only the marker-identified modality
return result
model._generate = _patched_generate
@@ -270,7 +484,24 @@ def _patch_openai_compat_content(model: Any) -> None:
async def _patched_agenerate(
messages: list[BaseMessage], *args: Any, **kwargs: Any
) -> Any:
return await orig_agenerate(_sanitize_messages(messages), *args, **kwargs)
prepared = _prepare(messages)
suspects = _media_types_in(messages) - blocked
if not suspects:
return await orig_agenerate(prepared, *args, **kwargs)
try:
return await orig_agenerate(prepared, *args, **kwargs)
except Exception as exc:
culprit = _media_error_types(exc) & suspects
if not culprit and not _is_http_400(exc):
raise
try:
result = await orig_agenerate(
_stripped(messages, suspects), *args, **kwargs
)
except Exception:
raise exc from None
blocked.update(culprit)
return result
model._agenerate = _patched_agenerate
@@ -283,7 +514,38 @@ def _patch_openai_compat_content(model: Any) -> None:
def _patched_stream(
messages: list[BaseMessage], *args: Any, **kwargs: Any
) -> Any:
return orig_stream(_sanitize_messages(messages), *args, **kwargs)
prepared = _prepare(messages)
suspects = _media_types_in(messages) - blocked
if not suspects:
yield from orig_stream(prepared, *args, **kwargs)
return
started = False
try:
for chunk in orig_stream(prepared, *args, **kwargs):
started = True
yield chunk
return
except Exception as exc:
culprit = _media_error_types(exc) & suspects
if started or (not culprit and not _is_http_400(exc)):
raise
media_exc = exc
media_culprit = culprit
retry_started = False
try:
for chunk in orig_stream(
_stripped(messages, suspects), *args, **kwargs
):
if not retry_started:
retry_started = True
blocked.update(media_culprit)
yield chunk
except Exception:
if retry_started:
raise
raise media_exc from None
if not retry_started:
raise media_exc from None # stripped retry produced nothing
model._stream = _patched_stream
@@ -294,10 +556,39 @@ def _patch_openai_compat_content(model: Any) -> None:
async def _patched_astream(
messages: list[BaseMessage], *args: Any, **kwargs: Any
) -> Any:
async for chunk in orig_astream(
_sanitize_messages(messages), *args, **kwargs
):
yield chunk
prepared = _prepare(messages)
suspects = _media_types_in(messages) - blocked
if not suspects:
async for chunk in orig_astream(prepared, *args, **kwargs):
yield chunk
return
started = False
try:
async for chunk in orig_astream(prepared, *args, **kwargs):
started = True
yield chunk
return
except Exception as exc:
culprit = _media_error_types(exc) & suspects
if started or (not culprit and not _is_http_400(exc)):
raise
media_exc = exc
media_culprit = culprit
retry_started = False
try:
async for chunk in orig_astream(
_stripped(messages, suspects), *args, **kwargs
):
if not retry_started:
retry_started = True
blocked.update(media_culprit)
yield chunk
except Exception:
if retry_started:
raise
raise media_exc from None
if not retry_started:
raise media_exc from None # stripped retry produced nothing
model._astream = _patched_astream