open-swe/agent/middleware/sanitize_fireworks_messages.py

80 lines
2.8 KiB
Python
Raw 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
"""Strip legacy ``function_call`` fields before Fireworks provider calls.
The LangSmith LLM Gateway enforces strict request-schema validation and rejects
messages carrying the legacy OpenAI ``function_call`` field
(``Extra inputs are not permitted, field: 'messages[N].function_call'``).
``langchain_fireworks`` forwards ``function_call`` from
``AIMessage.additional_kwargs`` whenever it is present — even when the message
also carries modern ``tool_calls`` — so a message that was originally produced
by (or deserialized from) a provider that populated the legacy field will break
every subsequent Fireworks turn once routed through the gateway.
This middleware drops the redundant legacy field from assistant messages before
they reach the Fireworks serializer, mirroring the thinking-block sanitizer.
"""
from __future__ import annotations
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_fireworks.chat_models import ChatFireworks
except ImportError: # pragma: no cover
ChatFireworks = None # type: ignore[assignment, misc]
def _is_chat_fireworks(model: object) -> bool:
if ChatFireworks is None:
return False
seen: set[int] = set()
current = model
for _ in range(10):
if isinstance(current, ChatFireworks):
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 _sanitize_messages(messages: list[Any]) -> None:
for message in messages:
if not isinstance(message, AIMessage):
continue
additional_kwargs = message.additional_kwargs
if not isinstance(additional_kwargs, dict) or "function_call" not in additional_kwargs:
continue
additional_kwargs.pop("function_call", None)
class SanitizeFireworksMessagesMiddleware(AgentMiddleware):
"""Drop legacy ``function_call`` fields before Fireworks provider calls."""
def wrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], ModelResponse],
) -> ModelCallResult:
if _is_chat_fireworks(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_fireworks(request.model):
_sanitize_messages(request.messages)
return await handler(request)