fix(llm): cap the tool calls one assistant response may queue (#977)

* fix(llm): cap the tool calls one assistant response may queue

* fix(llm): cap the subscription backend's responses too

---------

Co-authored-by: Ahmed Allam <ahmed39652003@gmail.com>
This commit is contained in:
devin-ai-integration[bot]
2026-08-06 00:06:59 +03:00
committed by GitHub
co-authored by Ahmed Allam
parent 68ea6fca65
commit 8bd6c8e87a
7 changed files with 302 additions and 28 deletions
+51 -20
View File
@@ -4,6 +4,7 @@ from __future__ import annotations
import contextlib
import inspect
import logging
import os
import time
from typing import TYPE_CHECKING, Any, cast
@@ -36,6 +37,7 @@ from openai.types.shared import Reasoning
from strix.config import codex
from strix.config.loader import load_settings
from strix.config.tool_call_ids import TurnCallIdRewriter, dedupe_input
from strix.config.tool_call_limits import TurnToolCallLimiter
if TYPE_CHECKING:
@@ -54,6 +56,9 @@ if TYPE_CHECKING:
from strix.config.settings import LlmSettings, ReasoningEffort, Settings
logger = logging.getLogger(__name__)
def request_timeout_extra_args(timeout_s: float | None) -> dict[str, float] | None:
"""Per-request model timeout; a plain float so ``ModelSettings.to_json_dict()`` stays serializable.""" # noqa: E501
if not timeout_s or timeout_s <= 0:
@@ -235,18 +240,34 @@ class _NonStreamingModel(Model):
yield _completed_stream_event(response, getattr(self._inner, "model", None))
class _UniqueToolCallIdModel(Model):
"""Keep tool-call ids unique so a recycled id can't invalidate the history.
class _TurnGuardModel(Model):
"""Keep one turn from corrupting the conversation or running away.
Providers that number tool calls per turn (``exec_command:0``, ...) restart
the counter each turn, so the same id eventually appears twice in one
conversation and strict providers reject every subsequent request. Ids that
collide with the history are rewritten before the turn is recorded, and
already-corrupted histories are repaired on the way out.
Tool-call ids: providers that number calls per turn (``exec_command:0``,
...) restart the counter each turn, so the same id eventually appears twice
in one conversation and strict providers reject every subsequent request.
Ids that collide with the history are rewritten before the turn is
recorded, and already-corrupted histories are repaired on the way out.
Tool-call volume: a degenerate response can queue hundreds of calls that
the run loop then honours one by one. Only the first
``LLM_MAX_TOOL_CALLS_PER_TURN`` calls of a response are kept.
"""
def __init__(self, inner: Model) -> None:
def __init__(self, inner: Model, *, max_tool_calls_per_turn: int = 0) -> None:
self._inner = inner
self._max_tool_calls_per_turn = max_tool_calls_per_turn
def _limiter(self) -> TurnToolCallLimiter:
return TurnToolCallLimiter(self._max_tool_calls_per_turn)
def _log_dropped(self, limiter: TurnToolCallLimiter) -> None:
if limiter.dropped:
logger.warning(
"dropped %d tool call(s) past the per-response limit of %d",
limiter.dropped,
self._max_tool_calls_per_turn,
)
async def close(self) -> None:
await self._inner.close()
@@ -282,7 +303,9 @@ class _UniqueToolCallIdModel(Model):
conversation_id=conversation_id,
prompt=prompt,
)
response.output = rewriter.rewrite_items(list(response.output))
limiter = self._limiter()
response.output = limiter.filter_items(rewriter.rewrite_items(list(response.output)))
self._log_dropped(limiter)
return response
async def stream_response(
@@ -301,6 +324,7 @@ class _UniqueToolCallIdModel(Model):
) -> AsyncIterator[TResponseStreamEvent]:
sanitized = dedupe_input(input)
rewriter = TurnCallIdRewriter(sanitized)
limiter = self._limiter()
stream = self._inner.stream_response(
system_instructions,
cast("str | list[TResponseInputItem]", sanitized),
@@ -314,20 +338,26 @@ class _UniqueToolCallIdModel(Model):
prompt=prompt,
)
async for event in stream:
yield _rewrite_event_call_ids(event, rewriter)
guarded = _guard_event(event, rewriter, limiter)
if guarded is not None:
yield guarded
self._log_dropped(limiter)
def _rewrite_event_call_ids(
event: TResponseStreamEvent, rewriter: TurnCallIdRewriter
) -> TResponseStreamEvent:
def _guard_event(
event: TResponseStreamEvent, rewriter: TurnCallIdRewriter, limiter: TurnToolCallLimiter
) -> TResponseStreamEvent | None:
if isinstance(event, ResponseOutputItemAddedEvent | ResponseOutputItemDoneEvent):
rewritten = rewriter.rewrite_item(event.item)
if not limiter.allow(rewritten):
return None
if rewritten is not event.item:
return event.model_copy(update={"item": rewritten})
return event
if isinstance(event, ResponseCompletedEvent):
output = rewriter.rewrite_items(list(event.response.output))
if output != list(event.response.output):
original = list(event.response.output)
output = limiter.filter_items(rewriter.rewrite_items(original))
if output != original:
return event.model_copy(
update={"response": event.response.model_copy(update={"output": output})}
)
@@ -403,15 +433,16 @@ class StrixProvider(MultiProvider):
# The ChatGPT subscription backend is always streamed; it has no
# non-streaming mode to fall back to, so LLM_DISABLE_STREAMING
# does not apply here.
return _CodexResponsesModel(
model: Model = _CodexResponsesModel(
slug,
codex.get_subscription_client(),
reasoning_effort=llm.reasoning_effort,
)
model = super().get_model(model_name)
if llm.disable_streaming:
model = _NonStreamingModel(model)
return _UniqueToolCallIdModel(model)
else:
model = super().get_model(model_name)
if llm.disable_streaming:
model = _NonStreamingModel(model)
return _TurnGuardModel(model, max_tool_calls_per_turn=llm.max_tool_calls_per_turn)
DEFAULT_MODEL_RETRY = ModelRetrySettings(