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