fix(router): harden streaming fallback wrapper for bridge iterators

- FallbackResponsesStreamWrapper now uses getattr fallbacks when copying
  attributes from the source iterator. The bridge path
  (LiteLLMCompletionStreamingIterator used by Anthropic/Bedrock/Vertex)
  does not call super().__init__ and is missing response, logging_obj
  (it uses litellm_logging_obj), responses_api_provider_config,
  start_time, request_data, call_type, and _hidden_params. Previously,
  wrapper construction raised AttributeError for any streaming fallback
  on the bridge path.
- _aresponses_with_streaming_fallbacks now deep-copies the
  litellm_metadata (and metadata) dicts into fallback_kwargs. The
  primary attempt mutates this dict in place via
  _update_kwargs_with_deployment, so a shallow copy of kwargs was
  leaking primary-deployment fields (deployment, model_info, api_base)
  into the mid-stream fallback request.

Co-authored-by: Yassin Kortam <yassin@berri.ai>
This commit is contained in:
Cursor Agent 2026-05-20 17:24:14 +00:00
parent 3e18cbb8c8
commit f73a9e3b63
No known key found for this signature in database

View file

@ -2481,6 +2481,7 @@ class Router:
def __init__(self, async_generator: AsyncGenerator):
import time
from datetime import datetime
self._async_generator = async_generator
# Mirror every attribute BaseResponsesAPIStreamingIterator.__init__
@ -2489,28 +2490,44 @@ class Router:
# we copy from source_iterator where applicable and use safe
# defaults elsewhere. This keeps inherited methods (e.g.
# _check_max_streaming_duration, _handle_failure) safe to call.
self.response = source_iterator.response
self.model = source_iterator.model
self.logging_obj = source_iterator.logging_obj
#
# The bridge path (LiteLLMCompletionStreamingIterator used by
# Anthropic/Bedrock/Vertex) does not call super().__init__ and
# is missing many of these attributes — use getattr fallbacks
# so wrapper construction never raises AttributeError. The
# bridge stores the logging object as `litellm_logging_obj`.
self.response = getattr(source_iterator, "response", None)
self.model = getattr(source_iterator, "model", None)
self.logging_obj = getattr(
source_iterator,
"logging_obj",
getattr(source_iterator, "litellm_logging_obj", None),
)
self.finished = False
self.responses_api_provider_config = (
source_iterator.responses_api_provider_config
self.responses_api_provider_config = getattr(
source_iterator, "responses_api_provider_config", None
)
self.completed_response = None
self.start_time = source_iterator.start_time
self.start_time = getattr(source_iterator, "start_time", datetime.now())
self._failure_handled = False
self._completed_response_cached = False
self._completed_response_logged = False
self._completed_response_cache_hit = None
self._persist_completed_response_before_logging = True
self._stream_created_time = time.time()
self.litellm_metadata = source_iterator.litellm_metadata
self.custom_llm_provider = source_iterator.custom_llm_provider
self.request_data = source_iterator.request_data
self.call_type = source_iterator.call_type
self.litellm_metadata = getattr(
source_iterator, "litellm_metadata", None
)
self.custom_llm_provider = getattr(
source_iterator, "custom_llm_provider", None
)
self.request_data = getattr(source_iterator, "request_data", {}) or {}
self.call_type = getattr(source_iterator, "call_type", None)
# Preserve hidden params so response headers (model_id,
# api_base, additional_headers) keep flowing.
self._hidden_params = dict(source_iterator._hidden_params or {})
self._hidden_params = dict(
getattr(source_iterator, "_hidden_params", None) or {}
)
def __aiter__(self):
return self
@ -4683,9 +4700,25 @@ class Router:
)
# Snapshot the request kwargs before _ageneric_api_call_with_fallbacks
# mutates them. The original_generic_function is preserved so the
# per-attempt helper knows which underlying API to call on fallback.
# mutates them. A shallow copy alone is not enough: the primary
# attempt mutates nested dicts in place — notably `litellm_metadata`,
# which `_update_kwargs_with_deployment` populates with
# deployment-specific fields (`deployment`, `model_info`, `api_base`,
# tags, etc.). Without an explicit copy of that dict, the shallow
# copy would still share its reference, leaking primary-deployment
# metadata into the mid-stream fallback request.
#
# We avoid `copy.deepcopy` on the full kwargs because it can contain
# non-deepcopyable objects (logging handles, async clients, etc.).
# The original_generic_function is preserved so the per-attempt
# helper knows which underlying API to call on fallback.
fallback_kwargs: Dict[str, Any] = kwargs.copy()
if isinstance(fallback_kwargs.get("litellm_metadata"), dict):
fallback_kwargs["litellm_metadata"] = copy.deepcopy(
fallback_kwargs["litellm_metadata"]
)
if isinstance(fallback_kwargs.get("metadata"), dict):
fallback_kwargs["metadata"] = copy.deepcopy(fallback_kwargs["metadata"])
fallback_kwargs["original_generic_function"] = original_function
response = await self._ageneric_api_call_with_fallbacks(