fix(vision): reserve the per-image embed slot atomically

repeat_refusal checked the counter and record_embed incremented it in separate
lock sections, so a concurrent tool batch on the same image (the incident issued
4 at once) all passed a check taken before any of them recorded and the cap of 3
let 6 embeds through. Reserve the slot inside the check's lock section and
release it when the embed then fails.
This commit is contained in:
teknium1
2026-09-16 23:48:05 -07:00
committed by Teknium
parent f4452169d1
commit 87f29fb1b9
3 changed files with 69 additions and 10 deletions

View File

@@ -74,6 +74,36 @@ class TestRepeatCap:
assert "already been loaded" in payload["error"] and "max_calls_per_image" in payload["error"]
assert _embedded(other), "a different image in the same session is not affected"
def test_parallel_batch_on_one_image_cannot_overshoot_the_cap(self, tmp_path):
"""The executor runs a tool batch concurrently: 6 simultaneous loads of one image with cap 3
must yield exactly 3 embeds — the slot is reserved atomically, not check-then-record."""
import contextvars
import threading
shot = _png(tmp_path / "shot.png")
results = [None] * 6
def one(i):
results[i] = _embedded(asyncio.new_event_loop().run_until_complete(_vision_analyze_native(shot, "q")))
with delegated_child_context("child-parallel"):
threads = [threading.Thread(target=contextvars.copy_context().run, args=(one, i)) for i in range(6)]
for th in threads:
th.start()
for th in threads:
th.join()
assert sum(results) == 3
assert budget._repeat_counts[("child-parallel", budget._image_key(shot))] == 3
def test_failed_embed_releases_its_reserved_slot(self, tmp_path):
"""A refused/failed load (missing file) must not burn one of the three slots."""
shot = _png(tmp_path / "shot.png")
with delegated_child_context("child-release"):
for _ in range(3):
assert json.loads(_load(str(tmp_path / "missing.png")))["success"] is False
assert all(_embedded(_load(shot)) for _ in range(3))
assert ("child-release", budget._image_key(str(tmp_path / "missing.png"))) not in budget._repeat_counts
def test_main_agent_is_unlimited_unless_configured(self, tmp_path):
shot = _png(tmp_path / "shot.png")
assert all(_embedded(_load(shot)) for _ in range(5))

View File

@@ -38,7 +38,9 @@ from tools.debug_helpers import DebugSession
from tools.website_policy import check_website_access
from tools.vision_tools_history_budget import (
record_embed as _record_embed,
release_embed as _release_embed,
repeat_refusal as _repeat_refusal,
resolve_repeat_cap as _resolve_repeat_cap,
resolve_embed_target_bytes as _resolve_embed_target_bytes,
)
from tools.vision_tools_image_prep import (
@@ -589,10 +591,13 @@ async def _vision_analyze_native(
or a JSON error string (the normal tool-result contract) on failure."""
if not isinstance(image_url, str) or not image_url.strip():
return tool_error("image_url is required", success=False)
# A cap > 0 RESERVES the slot here (atomic check-and-count); released below if no embed happens.
refusal = _repeat_refusal(image_url)
if refusal is not None:
return refusal
reserved = _resolve_repeat_cap() > 0
prepared: Optional[_PreparedImage] = None
embedded = False
try:
from tools.interrupt import is_interrupted
if is_interrupted():
@@ -619,7 +624,9 @@ async def _vision_analyze_native(
# Reject rather than embed a session-wedging payload.
if len(image_data_url) > _MAX_BASE64_BYTES:
return tool_error(_too_large_message(image_data_url), success=False)
_record_embed(image_url)
embedded = True
if not reserved:
_record_embed(image_url)
return _build_native_vision_tool_result(
image_url=image_url, question=question, image_data_url=image_data_url,
image_size_bytes=prepared.size_bytes,
@@ -628,6 +635,8 @@ async def _vision_analyze_native(
logger.warning("Native vision fast path failed: %s", exc)
return tool_error(f"Native vision failed: {exc}", success=False)
finally:
if reserved and not embedded:
_release_embed(image_url)
# Only delete temp files we created — never user-provided paths.
if prepared is not None:
_unlink_quietly(prepared.path)

View File

@@ -86,14 +86,19 @@ def _count_key(image_url: str) -> tuple[str, str]:
def repeat_refusal(image_url: str) -> Optional[str]:
"""Tool-error JSON when this image already hit its per-session embed cap, else ``None``."""
"""Reserve one native embed of ``image_url`` for the current session; tool-error JSON when the
per-session cap is already spent, else ``None``. Check and count are ONE lock section: a
parallel tool batch on the same image (the incident's 4 concurrent calls) must not all pass a
check taken before any of them recorded. Callers ``release_embed`` when the embed then fails."""
cap = resolve_repeat_cap()
if cap <= 0:
return None
key = _count_key(image_url)
with _repeat_lock:
count = _repeat_counts.get(_count_key(image_url), 0)
if count < cap:
return None
count = _repeat_counts.get(key, 0)
if count < cap:
_record_embed_locked(key)
return None
return tool_error(
f"vision_analyze refused: this image has already been loaded into context {count} time(s) "
"in this session (region crops of the same file count too), and every native load re-sends "
@@ -103,11 +108,26 @@ def repeat_refusal(image_url: str) -> Optional[str]:
)
def _record_embed_locked(key: tuple[str, str]) -> None:
_repeat_counts[key] = _repeat_counts.get(key, 0) + 1
# Bound long-lived gateway memory: evict the oldest (session, image) entries.
while len(_repeat_counts) > _REPEAT_COUNTS_MAX_KEYS:
_repeat_counts.pop(next(iter(_repeat_counts)))
def record_embed(image_url: str) -> None:
"""Count one successful native embed of ``image_url`` for the current session."""
"""Count one successful native embed of ``image_url`` for the current session (uncapped
sessions only — capped ones are counted by the reservation in :func:`repeat_refusal`)."""
with _repeat_lock:
_record_embed_locked(_count_key(image_url))
def release_embed(image_url: str) -> None:
"""Give back a slot reserved by :func:`repeat_refusal` when the embed did not happen."""
key = _count_key(image_url)
with _repeat_lock:
_repeat_counts[key] = _repeat_counts.get(key, 0) + 1
# Bound long-lived gateway memory: evict the oldest (session, image) entries.
while len(_repeat_counts) > _REPEAT_COUNTS_MAX_KEYS:
_repeat_counts.pop(next(iter(_repeat_counts)))
count = _repeat_counts.get(key, 0)
if count > 1:
_repeat_counts[key] = count - 1
else:
_repeat_counts.pop(key, None)