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:
+319
-28
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user