open-swe/agent/middleware/sanitize_openai_responses.py

86 lines
2.6 KiB
Python
Raw Permalink Normal View History

feat: port LangSmith LLM Gateway routing from upstream (#1671, #1673, #1674, #1678) (#155) * feat: port LangSmith LLM Gateway routing from upstream (#1671, #1673, #1674, #1678) Ports four upstream commits that add opt-in LLM call routing through the LangSmith Gateway, preserving fork conventions (Bedrock/Fireworks model IDs, no-agent-attribution, bun toolchain). - #1671 (e9dc6e01): opt-in gateway routing — new gateway.py, team-settings toggle, admin UI section, wired into make_model for all graph entrypoints - #1673 (702ef908): dedicated LANGSMITH_GATEWAY_API_KEY precedence over platform LANGSMITH_API_KEY - #1674 (5f7c2f46): fix Fireworks gateway base URL to /fireworks (bare host, SDK appends /v1/chat/completions) + SanitizeFireworksMessagesMiddleware - #1678 (73b7d1c0): fix OpenAI Responses reasoning replay — SanitizeOpenAIResponsesMiddleware, store/include config for encrypted reasoning content, reasoning_effort coercion for Chat Completions fallback Refs #134 * fix: downgrade gateway not-routed log to debug, add Bedrock UI note, add sanitizer parity - Downgrade logger.warning to logger.debug in gateway_overrides for not-routed providers and missing API key (Bedrock is the default provider in this fork, so these are expected steady states) - Add Bedrock to the LLMGatewaySection route-toggle description so admins know it is not routed through the gateway - Add SanitizeOpenAIResponsesMiddleware to chat.py for parity with server.py and reviewer.py - Restore the Bedrock region comment in model.py that explains the AWS_REGION / AWS_DEFAULT_REGION precedence Refs #138 --------- Co-authored-by: amoussa1229 <166072409+amoussa1229@users.noreply.github.com>
2026-07-09 14:44:15 -04:00
"""Middleware that removes stale OpenAI Responses reasoning references."""
from __future__ import annotations
import logging
from collections.abc import Awaitable, Callable
from typing import Any
from langchain.agents.middleware import AgentMiddleware
from langchain.agents.middleware.types import ModelCallResult, ModelRequest, ModelResponse
from langchain_core.messages import AIMessage
try:
from langchain_openai import ChatOpenAI
except ImportError: # pragma: no cover
ChatOpenAI = None # type: ignore[assignment, misc]
logger = logging.getLogger(__name__)
def _is_chat_openai(model: object) -> bool:
if ChatOpenAI is None:
return False
seen: set[int] = set()
current = model
for _ in range(10):
if isinstance(current, ChatOpenAI):
return True
current_id = id(current)
if current_id in seen:
return False
seen.add(current_id)
bound = getattr(current, "bound", None)
if bound is None or bound is current:
return False
current = bound
return False
def _is_stale_reasoning_reference(block: dict[str, Any]) -> bool:
if block.get("type") != "reasoning":
return False
if block.get("encrypted_content"):
return False
block_id = block.get("id")
return isinstance(block_id, str) and block_id.startswith("rs_")
def _sanitize_messages(messages: list[Any]) -> None:
removed = 0
for message in messages:
if not isinstance(message, AIMessage) or not isinstance(message.content, list):
continue
content = []
for block in message.content:
if isinstance(block, dict) and _is_stale_reasoning_reference(block):
removed += 1
continue
content.append(block)
if len(content) != len(message.content):
message.content = content
if removed:
logger.warning("Removed %d stale OpenAI Responses reasoning reference(s)", removed)
class SanitizeOpenAIResponsesMiddleware(AgentMiddleware):
"""Drop non-replayable OpenAI Responses reasoning item references."""
def wrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], ModelResponse],
) -> ModelCallResult:
if _is_chat_openai(request.model):
_sanitize_messages(request.messages)
return handler(request)
async def awrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
) -> Any:
if _is_chat_openai(request.model):
_sanitize_messages(request.messages)
return await handler(request)