fix(agent): track image rejections per model across a fallback chain

The first head stored a single rejecting (provider, model) and kept the
turn-global `_vision_supported` as the recovery guard. In a fallback
chain that fails: model A rejects images and retries text-only, a later
error activates model B, the restart rebuilds api_messages from history
so B receives the images, and when B rejects them too the branch is
skipped because `_vision_supported` is already False — the request
falls through to generic error handling. Recording B also overwrote A,
so A was no longer treated as text-only on later turns.

`_image_rejecting_models` is now a set of every rejecting model, and it
is also the guard: each model's first rejection runs the recovery and a
repeat rejection from the same model still falls through, so the retry
cannot loop. image_model_key() names the key in one place.

Adds a test for the two-model sequence (fails on the previous head) and
one pinning that a repeat rejection from the same model does not retry.

Thanks to @ehz0ah for the review.

(cherry picked from commit 225fd76ccf90cd3ec509a9327f725ac02d996f24)
This commit is contained in:
rodricksz4h5
2026-09-21 17:43:29 +05:30
committed by kshitij
parent 1ef306f68d
commit 7299015092
4 changed files with 61 additions and 11 deletions

View File

@@ -2403,9 +2403,9 @@ def init_agent(
agent.request_overrides = dict(request_overrides or {})
agent.prefill_messages = prefill_messages or [] # Prefilled conversation turns
agent._force_ascii_payload = False
# (provider, model) that rejected image content this session; build_api_request strips images
# from requests to that model only, so history keeps them for any model that can see.
agent._image_rejecting_model = None
# Every (provider, model) that rejected image content this session. build_api_request strips
# images from requests to those models only, so history keeps them for any model that can see.
agent._image_rejecting_models = set()
_init_prompt_cache_config(agent)
_init_turn_state(agent, run_budget_seconds)

View File

@@ -400,17 +400,23 @@ _IMAGE_REJECTION_PHRASES = (
)
def image_model_key(agent: Any) -> tuple:
"""``(provider, model)`` the agent is currently talking to — the key image rejections are
tracked by, so each model in a fallback chain is judged on its own."""
return (getattr(agent, "provider", None), getattr(agent, "model", None))
def strip_images_for_rejecting_model(agent: Any, api_messages: Any) -> bool:
"""Send-path image strip for a model that rejected image content (see turn_recovery).
Runs on the per-call ``api_messages`` copy in Hermes's own message format, BEFORE the
provider-specific conversion: the part types this stripper knows are that format's, and a
converted payload (Bedrock Converse ``{"image": ...}`` blocks carry no ``type``) would slip
past it. History is never touched. Keyed on the rejecting (provider, model), so a model
past it. History is never touched. Keyed on each rejecting (provider, model), so a model
that accepts images gets them again.
"""
rejecting = getattr(agent, "_image_rejecting_model", None)
if rejecting is None or rejecting != (getattr(agent, "provider", None), getattr(agent, "model", None)):
rejected = getattr(agent, "_image_rejecting_models", None)
if not isinstance(rejected, set) or image_model_key(agent) not in rejected:
return False
return isinstance(api_messages, list) and _strip_images_from_messages(api_messages)
@@ -428,7 +434,7 @@ __all__ = [
"_escape_invalid_chars_in_json_strings", "_repair_tool_call_arguments",
"_strip_non_ascii", "_sanitize_messages_non_ascii", "_sanitize_tools_non_ascii",
"_strip_images_from_messages", "_sanitize_structure_non_ascii", "sanitize_outbound_kwargs",
"strip_images_for_rejecting_model",
"strip_images_for_rejecting_model", "image_model_key",
# call_id policy owners
"deterministic_call_id", "coalesce_tool_call_id", "tool_call_id_variants",
"tool_result_id_variants", "uniquify_tool_call_ids",

View File

@@ -24,7 +24,7 @@ from agent.error_classifier import FailoverReason, classify_api_error
from agent.message_sanitization import (
_looks_like_image_content_rejection, _sanitize_messages_non_ascii,
_sanitize_messages_surrogates, _sanitize_structure_non_ascii, _sanitize_structure_surrogates,
_strip_images_from_messages, _strip_non_ascii,
_strip_images_from_messages, _strip_non_ascii, image_model_key,
close_interrupted_tool_sequence,
)
from agent.thinking_timeout_guidance import build_thinking_timeout_guidance, is_thinking_timeout
@@ -244,14 +244,20 @@ def recover_before_classification(
_err_status = getattr(api_error, "status_code", None)
# 4xx-only gate: 5xx/timeouts are transient and take the retry path.
_status_ok = _err_status is None or (400 <= int(_err_status) < 500)
if getattr(agent, "_vision_supported", True) and _looks_like_image_content_rejection(_err_body) and _status_ok:
# Guarded PER MODEL, not by the turn-global ``_vision_supported``: in a fallback chain the next
# model can reject images too, and a turn-wide flag would skip its recovery and fail the turn.
_model_key = image_model_key(agent)
_rejected = getattr(agent, "_image_rejecting_models", None)
if not isinstance(_rejected, set):
_rejected = agent._image_rejecting_models = set()
if _model_key not in _rejected and _looks_like_image_content_rejection(_err_body) and _status_ok:
agent._vision_supported = False
# Send-path only. A rejection says what THIS model accepts, not what the conversation
# holds: stripping ``messages`` (canonical history) and forcing a flush deleted every
# image — and every image-only message — from state.db for good, so a later switch to a
# vision model found them gone. Same failure as the ASCII strip in #117802. Record the
# model; build_api_request strips images from each request to it instead.
agent._image_rejecting_model = (getattr(agent, "provider", None), getattr(agent, "model", None))
_rejected.add(_model_key)
if isinstance(api_messages, list):
_strip_images_from_messages(api_messages)
_vlines(

View File

@@ -318,7 +318,7 @@ class TestRejectionNeverReachesPersistedHistory:
return SimpleNamespace(
provider=provider, model=model, _vision_supported=True, _force_ascii_payload=False,
_image_rejecting_model=None, _db_flush_scan_prefix=7, log_prefix="",
_image_rejecting_models=set(), _db_flush_scan_prefix=7, log_prefix="",
_vprint=lambda *a, **k: None,
)
@@ -379,3 +379,41 @@ class TestRejectionNeverReachesPersistedHistory:
api_messages = self._history()
assert strip_images_for_rejecting_model(agent, api_messages) is False
assert str(api_messages).count("data:image/png") == 2
def test_every_model_in_a_fallback_chain_is_tracked(self):
"""Two models reject images in the same turn (fallback A -> B). A turn-global guard
skipped B's recovery once A had tripped it, failing the turn; recording only one model
also forgot A on later turns. Each model is now judged and remembered on its own."""
from agent.message_sanitization import strip_images_for_rejecting_model
agent = self._agent(provider="p", model="model-a")
retry_a, _ = self._recover(agent, self._history(), [])
# The fallback restart rebuilds api_messages from history, images included, for B.
agent.model = "model-b"
rebuilt = self._history()
assert strip_images_for_rejecting_model(agent, rebuilt) is False
assert "image_url" in str(rebuilt)
# B rejects too, in the same turn: its recovery must still run.
retry_b, _ = self._recover(agent, self._history(), [])
assert retry_a is True and retry_b is True
assert agent._image_rejecting_models == {("p", "model-a"), ("p", "model-b")}
for model in ("model-a", "model-b"):
agent.model = model
api_messages = self._history()
assert strip_images_for_rejecting_model(agent, api_messages) is True, model
assert "image_url" not in str(api_messages)
agent.model = "model-c"
api_messages = self._history()
assert strip_images_for_rejecting_model(agent, api_messages) is False
assert str(api_messages).count("data:image/png") == 2
def test_a_repeat_rejection_from_the_same_model_does_not_loop(self):
"""The per-model guard still stops re-entry: a second rejection from a model already
known to reject images falls through to normal error handling instead of retrying forever."""
agent = self._agent()
assert self._recover(agent, self._history(), [])[0] is True
assert self._recover(agent, self._history(), [])[0] is False