diff --git a/litellm/a2a_protocol/litellm_completion_bridge/handler.py b/litellm/a2a_protocol/litellm_completion_bridge/handler.py index a5c4463da8d..71993135037 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/handler.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/handler.py @@ -154,7 +154,7 @@ class A2ACompletionBridgeHandler: if a2a_provider_config is not None: verbose_logger.info("A2A: Using provider config for %s", custom_llm_provider) - provider_params: Final = {key: value for key, value in params.items() if key != "messages"} + provider_params: Final = dict(params) provider_kwargs: Final[dict[str, Any]] = { "request_id": request_id, "params": provider_params, @@ -227,7 +227,7 @@ class A2ACompletionBridgeHandler: if a2a_provider_config is not None: verbose_logger.info("A2A: Using provider config for %s (streaming)", custom_llm_provider) - provider_params: Final = {key: value for key, value in params.items() if key != "messages"} + provider_params: Final = dict(params) provider_kwargs: Final[dict[str, Any]] = { "request_id": request_id, "params": provider_params, @@ -237,8 +237,14 @@ class A2ACompletionBridgeHandler: } if litellm_params.get("timeout") is not None: provider_kwargs["timeout"] = litellm_params["timeout"] - async for chunk in a2a_provider_config.handle_streaming(**provider_kwargs): - yield chunk + provider_stream: Final = a2a_provider_config.handle_streaming(**provider_kwargs) + try: + async for chunk in provider_stream: + yield chunk + finally: + close_provider_stream = getattr(provider_stream, "aclose", None) + if close_provider_stream is not None: + await close_provider_stream() return @@ -274,26 +280,52 @@ class A2ACompletionBridgeHandler: # 3. Forward content as artifact updates accumulated_tool_calls: Final[list[object]] = [] # mutable-ok: collect streaming tool-call deltas + stream_usage: object | None = None + stream_finish_reason: str | None = None chunk_count = 0 - async for chunk in response: - chunk_count += 1 + try: + async for chunk in response: + chunk_count += 1 - # Extract delta content - content = "" - if chunk is not None and hasattr(chunk, "choices") and chunk.choices: - choice = chunk.choices[0] - if hasattr(choice, "delta") and choice.delta: - content = choice.delta.content or "" - tool_calls = getattr(choice.delta, "tool_calls", None) - if isinstance(tool_calls, (list, tuple)): - accumulated_tool_calls.extend(tool_calls) + raw_usage = getattr(chunk, "usage", None) + if isinstance(raw_usage, Mapping): + stream_usage = raw_usage + else: + dump_usage = getattr(raw_usage, "model_dump", None) + if callable(dump_usage): + dumped_usage = dump_usage(exclude_none=True) + if isinstance(dumped_usage, Mapping): + stream_usage = dumped_usage + else: + dict_usage = getattr(raw_usage, "dict", None) + if callable(dict_usage): + dumped_usage = dict_usage(exclude_none=True) + if isinstance(dumped_usage, Mapping): + stream_usage = dumped_usage - if content: - artifact_event: Final = A2ACompletionBridgeTransformation.create_artifact_update_event( - ctx=ctx, - text=content, - ) - yield artifact_event + # Extract delta content + content = "" + if chunk is not None and hasattr(chunk, "choices") and chunk.choices: + choice = chunk.choices[0] + raw_finish_reason = getattr(choice, "finish_reason", None) + if isinstance(raw_finish_reason, str) and raw_finish_reason: + stream_finish_reason = raw_finish_reason + if hasattr(choice, "delta") and choice.delta: + content = choice.delta.content or "" + tool_calls = getattr(choice.delta, "tool_calls", None) + if isinstance(tool_calls, (list, tuple)): + accumulated_tool_calls.extend(tool_calls) + + if content: + artifact_event: Final = A2ACompletionBridgeTransformation.create_artifact_update_event( + ctx=ctx, + text=content, + ) + yield artifact_event + finally: + close_response = getattr(response, "aclose", None) + if close_response is not None: + await close_response() # 4. Emit final status update (kind: "status-update", status: "completed", final: true) completed_event: Final = A2ACompletionBridgeTransformation.create_status_update_event( @@ -303,6 +335,10 @@ class A2ACompletionBridgeHandler: ) if accumulated_tool_calls: completed_event["result"]["tool_calls"] = accumulated_tool_calls + if stream_finish_reason: + completed_event["result"]["finish_reason"] = stream_finish_reason + if stream_usage is not None: + completed_event["usage"] = stream_usage yield completed_event verbose_logger.info( diff --git a/litellm/a2a_protocol/litellm_completion_bridge/transformation.py b/litellm/a2a_protocol/litellm_completion_bridge/transformation.py index c24af243c83..02006e98fa0 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/transformation.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/transformation.py @@ -151,6 +151,22 @@ class A2ACompletionBridgeTransformation: return [openai_message] + @staticmethod + def _model_dump(value: Any) -> dict[str, Any]: + if isinstance(value, dict): + return value + dump = getattr(value, "model_dump", None) + if callable(dump): + dumped = dump(exclude_none=True) + if isinstance(dumped, dict): + return dumped + dump = getattr(value, "dict", None) + if callable(dump): + dumped = dump(exclude_none=True) + if isinstance(dumped, dict): + return dumped + return {} + @staticmethod def openai_response_to_a2a_response( response: Any, @@ -170,16 +186,19 @@ class A2ACompletionBridgeTransformation: raw_choices: Final = getattr(response, "choices", None) if raw_choices: for choice in raw_choices: - content: Final = ( - getattr(getattr(choice, "message", None), "content", None) or "" - ) + raw_message = getattr(choice, "message", None) + message_fields: Final = A2ACompletionBridgeTransformation._model_dump(raw_message) + raw_content = message_fields.get("content") + if raw_content is None: + raw_content = getattr(raw_message, "content", None) + content: Final = raw_content if isinstance(raw_content, str) else "" message: Final = { "kind": "message", "role": "agent", "parts": [{"kind": "text", "text": content}], "messageId": uuid4().hex, } - raw_tool_calls = getattr(getattr(choice, "message", None), "tool_calls", None) + raw_tool_calls = message_fields.get("tool_calls") if raw_tool_calls: message["tool_calls"] = [ call.model_dump(exclude_none=True) @@ -189,10 +208,38 @@ class A2ACompletionBridgeTransformation: else call for call in raw_tool_calls ] - finish_reason: Final = getattr(choice, "finish_reason", None) + for field in ( + "annotations", + "audio", + "function_call", + "images", + "provider_specific_fields", + "reasoning_content", + "reasoning_items", + "thinking_blocks", + ): + value = message_fields.get(field) + if value is not None: + message[field] = value + choice_fields: Final = A2ACompletionBridgeTransformation._model_dump(choice) + finish_reason: Final = choice_fields.get("finish_reason") + if finish_reason is None: + raw_finish_reason = getattr(choice, "finish_reason", None) + finish_reason = raw_finish_reason if isinstance(raw_finish_reason, str) else None if finish_reason: message["finish_reason"] = finish_reason - serialized_choices.append({"index": len(serialized_choices), "message": message}) + choice_payload: Final[dict[str, Any]] = { + "index": len(serialized_choices), + "message": message, + } + logprobs = choice_fields.get("logprobs") + if logprobs is None: + raw_logprobs = getattr(choice, "logprobs", None) + logprobs = raw_logprobs if isinstance(raw_logprobs, dict) else None + if logprobs is not None: + choice_payload["logprobs"] = logprobs + message["logprobs"] = logprobs + serialized_choices.append(choice_payload) a2a_message: Final = ( serialized_choices[0]["message"] diff --git a/litellm/llms/a2a/chat/streaming_iterator.py b/litellm/llms/a2a/chat/streaming_iterator.py index 58b3b396a0e..636aef0e724 100644 --- a/litellm/llms/a2a/chat/streaming_iterator.py +++ b/litellm/llms/a2a/chat/streaming_iterator.py @@ -2,8 +2,10 @@ A2A Streaming Response Iterator """ +from collections.abc import Mapping from typing import Final +import litellm from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator from litellm.types.llms.openai import ChatCompletionToolCallChunk from litellm.types.utils import GenericStreamingChunk, ModelResponseStream @@ -73,13 +75,14 @@ class A2AModelResponseIterator(BaseModelResponseIterator): # Determine finish reason finish_reason: Final = self._get_finish_reason(chunk) tool_calls: Final = self._get_tool_calls(chunk) + usage: Final = self._get_usage(chunk) # Return generic streaming chunk return GenericStreamingChunk( text=text, is_finished=bool(finish_reason or tool_calls), finish_reason=finish_reason or ("tool_calls" if tool_calls else ""), - usage=None, + usage=usage, index=0, tool_use=tool_calls, ) @@ -105,6 +108,14 @@ class A2AModelResponseIterator(BaseModelResponseIterator): # Check for task completion if isinstance(result, dict): + explicit_finish_reason: Final = result.get("finish_reason") + if isinstance(explicit_finish_reason, str) and explicit_finish_reason: + return explicit_finish_reason + message: Final = result.get("message") + if isinstance(message, dict): + message_finish_reason: Final = message.get("finish_reason") + if isinstance(message_finish_reason, str) and message_finish_reason: + return message_finish_reason status: Final = result.get("status", {}) if isinstance(status, dict): state: Final = status.get("state") @@ -119,16 +130,53 @@ class A2AModelResponseIterator(BaseModelResponseIterator): return None + def _get_usage(self, chunk: dict) -> object | None: + raw_usage: object | None = chunk.get("usage") + result: Final = chunk.get("result", {}) + if raw_usage is None and isinstance(result, dict): + raw_usage = result.get("usage") + if raw_usage is None: + return None + if isinstance(raw_usage, Mapping): + try: + return litellm.Usage(**raw_usage) + except Exception: + return raw_usage + if hasattr(raw_usage, "model_dump"): + try: + return litellm.Usage(**raw_usage.model_dump(exclude_none=True)) + except Exception: + return raw_usage + return raw_usage + def _get_tool_calls(self, chunk: dict) -> ChatCompletionToolCallChunk | None: result: Final = chunk.get("result", {}) if not isinstance(result, dict): return None tool_calls = result.get("tool_calls") if isinstance(tool_calls, list) and tool_calls: - first_tool_call: Final = tool_calls[0] - return first_tool_call if isinstance(first_tool_call, dict) else None + return self._serialize_tool_call(tool_calls[0]) message = result.get("message") if isinstance(message, dict) and isinstance(message.get("tool_calls"), list) and message["tool_calls"]: - first_tool_call = message["tool_calls"][0] - return first_tool_call if isinstance(first_tool_call, dict) else None + return self._serialize_tool_call(message["tool_calls"][0]) return None + + @staticmethod + def _serialize_tool_call(tool_call: object) -> ChatCompletionToolCallChunk | None: + if isinstance(tool_call, dict): + return tool_call + if hasattr(tool_call, "model_dump"): + return tool_call.model_dump(exclude_none=True) + if hasattr(tool_call, "dict"): + return tool_call.dict(exclude_none=True) + return None + + async def aclose(self) -> None: + streaming_response = self.streaming_response + self.streaming_response = None + try: + await super().aclose() + finally: + close_stream = getattr(streaming_response, "aclose", None) + if close_stream is not None: + await close_stream() diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index 5826442865d..34e945f1415 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -52,6 +52,7 @@ _FORWARDED_REQUEST_PARAMS: Final = frozenset( "response_format", "seed", "service_tier", + "safety_identifier", "stop", "store", "temperature", @@ -64,6 +65,8 @@ _FORWARDED_REQUEST_PARAMS: Final = frozenset( "user", "verbosity", "web_search_options", + "output_config", + "prompt_cache_key", } ) _A2A_PRICING_PARAMS: Final = frozenset({"cost_per_query", "response_cost"}) | frozenset( @@ -159,6 +162,7 @@ async def _route_registered_provider( logging_obj: Final = data.get("litellm_logging_obj") if isinstance(logging_obj, Logging): + provider_params["no-log"] = True pricing_params = { key: litellm_params[key] for key in _A2A_PRICING_PARAMS @@ -168,7 +172,6 @@ async def _route_registered_provider( logging_obj.litellm_params.update(pricing_params) logging_obj.model_call_details["litellm_params"].update(pricing_params) logging_obj.custom_pricing = True - provider_params["no-log"] = True if stream: streaming_response: Final = A2ACompletionBridgeHandler.handle_streaming( @@ -214,64 +217,80 @@ async def _route_registered_provider( nested_message: Final = result_dict.get("message") response_message: Final = nested_message if isinstance(nested_message, Mapping) else result_dict response_choices: Final = response.get("choices") - choice_payloads: Final = ( - response_choices - if isinstance(response_choices, list) - else result_dict.get("choices") - ) + choice_payloads: Final = response_choices if isinstance(response_choices, list) else result_dict.get("choices") + + def _serialize_value(value: object) -> object: + if hasattr(value, "model_dump"): + return value.model_dump(exclude_none=True) + if hasattr(value, "dict"): + return value.dict(exclude_none=True) + return value + + def _build_message(message_payload: Mapping[str, object], content: str) -> Message: + message_kwargs: dict[str, object] = { + "content": content, + "role": "assistant", + } + raw_tool_calls = message_payload.get("tool_calls") + if isinstance(raw_tool_calls, list): + message_kwargs["tool_calls"] = raw_tool_calls + for field in ( + "audio", + "annotations", + "function_call", + "images", + "provider_specific_fields", + "reasoning_content", + "reasoning_items", + "thinking_blocks", + ): + value = message_payload.get(field) + if value is not None: + message_kwargs[field] = _serialize_value(value) + return Message(**message_kwargs) + if isinstance(choice_payloads, list) and choice_payloads: - model_choices = [ - Choices( - finish_reason=( - choice.get("finish_reason") - if isinstance(choice, Mapping) and isinstance(choice.get("finish_reason"), str) - else choice.get("message", {}).get("finish_reason") - if isinstance(choice, Mapping) - and isinstance(choice.get("message"), Mapping) - and isinstance(choice.get("message", {}).get("finish_reason"), str) + model_choices = [] + for choice_index, choice in enumerate(choice_payloads): + choice_mapping: Mapping[str, object] = choice if isinstance(choice, Mapping) else {} + raw_message = choice_mapping.get("message") + message_payload: Mapping[str, object] = raw_message if isinstance(raw_message, Mapping) else choice_mapping + choice_kwargs: dict[str, object] = { + "finish_reason": ( + choice_mapping.get("finish_reason") + if isinstance(choice_mapping.get("finish_reason"), str) + else message_payload.get("finish_reason") + if isinstance(message_payload.get("finish_reason"), str) else "stop" ), - index=choice.get("index", choice_index) - if isinstance(choice, Mapping) and isinstance(choice.get("index", choice_index), int) + "index": choice_mapping.get("index", choice_index) + if isinstance(choice_mapping.get("index", choice_index), int) else choice_index, - message=Message( - content=extract_text_from_a2a_response( - {"result": choice.get("message", choice)} - if isinstance(choice, Mapping) - else {"result": {}} - ), - role="assistant", - tool_calls=( - choice.get("message", {}).get("tool_calls") - if isinstance(choice, Mapping) - and isinstance(choice.get("message"), Mapping) - and isinstance(choice.get("message", {}).get("tool_calls"), list) - else choice.get("tool_calls") - if isinstance(choice, Mapping) and isinstance(choice.get("tool_calls"), list) - else None - ), + "message": _build_message( + message_payload, + extract_text_from_a2a_response({"result": message_payload}), ), - ) - for choice_index, choice in enumerate(choice_payloads) - ] + } + raw_logprobs = choice_mapping.get("logprobs", message_payload.get("logprobs")) + if raw_logprobs is not None: + choice_kwargs["logprobs"] = _serialize_value(raw_logprobs) + model_choices.append(Choices(**choice_kwargs)) else: tool_calls: Final = response_message.get("tool_calls") normalized_tool_calls: Final = tool_calls if isinstance(tool_calls, list) else None finish_reason: Final = response_message.get("finish_reason") text: Final = extract_text_from_a2a_response(response) - model_choices = [ - Choices( - finish_reason=( - finish_reason - if isinstance(finish_reason, str) - else "tool_calls" - if normalized_tool_calls - else "stop" - ), - index=0, - message=Message(content=text, role="assistant", tool_calls=normalized_tool_calls), - ) - ] + choice_kwargs = { + "finish_reason": ( + finish_reason if isinstance(finish_reason, str) else "tool_calls" if normalized_tool_calls else "stop" + ), + "index": 0, + "message": _build_message(response_message, text), + } + raw_logprobs = response_message.get("logprobs") + if raw_logprobs is not None: + choice_kwargs["logprobs"] = _serialize_value(raw_logprobs) + model_choices = [Choices(**choice_kwargs)] model_response: Final = ModelResponse( id=str(response.get("id") or request_id), model=model_name, diff --git a/tests/test_litellm/a2a_protocol/test_completion_bridge_streaming.py b/tests/test_litellm/a2a_protocol/test_completion_bridge_streaming.py index 23495d06629..d2b74d416a0 100644 --- a/tests/test_litellm/a2a_protocol/test_completion_bridge_streaming.py +++ b/tests/test_litellm/a2a_protocol/test_completion_bridge_streaming.py @@ -12,6 +12,8 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest +from litellm.types.utils import Choices, Message, ModelResponse + class TestA2AStreamingTransformation: """Test the A2A streaming transformation creates proper events.""" @@ -26,9 +28,7 @@ class TestA2AStreamingTransformation: "parts": [{"text": "Reply to ticket #4823"}], "metadata": {"skillId": "draft_reply"}, } - openai_messages = ( - A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message) - ) + openai_messages = A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message) # Metadata is forwarded on the run payload only, not duplicated on messages. assert "metadata" not in openai_messages[0] @@ -174,10 +174,7 @@ class TestA2AStreamingTransformation: assert "artifactId" in event["result"]["artifact"] assert event["result"]["artifact"]["name"] == "response" assert event["result"]["artifact"]["parts"][0]["kind"] == "text" - assert ( - event["result"]["artifact"]["parts"][0]["text"] - == "Hello, I am an AI assistant." - ) + assert event["result"]["artifact"]["parts"][0]["text"] == "Hello, I am an AI assistant." @pytest.mark.asyncio @@ -197,6 +194,8 @@ async def test_handle_streaming_emits_proper_events(): mock_chunk2.choices = [MagicMock()] mock_chunk2.choices[0].delta = MagicMock() mock_chunk2.choices[0].delta.content = " world" + mock_chunk2.choices[0].finish_reason = "length" + mock_chunk2.usage = {"prompt_tokens": 2, "completion_tokens": 3, "total_tokens": 5} async def mock_streaming_response(): yield mock_chunk1 @@ -246,6 +245,66 @@ async def test_handle_streaming_emits_proper_events(): assert events[4]["result"]["kind"] == "status-update" assert events[4]["result"]["status"]["state"] == "completed" assert events[4]["result"]["final"] is True + assert events[4]["result"]["finish_reason"] == "length" + assert events[4]["usage"]["total_tokens"] == 5 + + +@pytest.mark.asyncio +async def test_provider_config_receives_full_message_history(): + from litellm.a2a_protocol.litellm_completion_bridge.handler import ( + A2ACompletionBridgeHandler, + ) + + provider_config = MagicMock() + provider_config.handle_non_streaming = AsyncMock(return_value={"result": {}}) + messages = [ + {"role": "system", "content": "Be concise"}, + {"role": "user", "content": "Hello"}, + ] + params = { + "message": {"role": "user", "parts": []}, + "messages": messages, + } + + with patch( + "litellm.a2a_protocol.litellm_completion_bridge.handler.A2AProviderConfigManager.get_provider_config", + return_value=provider_config, + ): + await A2ACompletionBridgeHandler.handle_non_streaming( + request_id="req-1", + params=params, + litellm_params={"custom_llm_provider": "langflow", "model": "flow"}, + ) + + assert provider_config.handle_non_streaming.await_args.kwargs["params"]["messages"] == messages + + +def test_response_transform_preserves_audio_and_logprobs(): + from litellm.a2a_protocol.litellm_completion_bridge.transformation import ( + A2ACompletionBridgeTransformation, + ) + + response = ModelResponse( + id="resp-1", + model="test-model", + choices=[ + Choices( + finish_reason="stop", + index=0, + message=Message( + content="hello", + role="assistant", + audio={"data": "abc", "expires_at": 1, "transcript": "hello"}, + ), + logprobs={"content": []}, + ) + ], + ) + + transformed = A2ACompletionBridgeTransformation.openai_response_to_a2a_response(response) + + assert transformed["result"]["audio"]["data"] == "abc" + assert transformed["result"]["logprobs"] == {"content": []} @pytest.mark.asyncio diff --git a/tests/test_litellm/llms/a2a/chat/test_a2a_streaming_iterator.py b/tests/test_litellm/llms/a2a/chat/test_a2a_streaming_iterator.py index dafa2839ace..db303719864 100644 --- a/tests/test_litellm/llms/a2a/chat/test_a2a_streaming_iterator.py +++ b/tests/test_litellm/llms/a2a/chat/test_a2a_streaming_iterator.py @@ -4,6 +4,7 @@ import pytest from litellm.llms.a2a.chat.streaming_iterator import A2AModelResponseIterator from litellm.llms.a2a.common_utils import A2AError +from litellm.types.utils import Delta @pytest.mark.asyncio @@ -48,6 +49,55 @@ async def test_async_iterator_preserves_tool_calls(): assert chunk["finish_reason"] == "tool_calls" +@pytest.mark.asyncio +async def test_async_iterator_serializes_delta_tool_calls_and_usage(): + delta = Delta( + tool_calls=[ + { + "id": "call-1", + "type": "function", + "function": {"name": "lookup", "arguments": "{}"}, + } + ] + ) + + async def _events(): + yield { + "jsonrpc": "2.0", + "usage": {"prompt_tokens": 2, "completion_tokens": 3, "total_tokens": 5}, + "result": { + "tool_calls": [delta.tool_calls[0]], + "finish_reason": "length", + }, + } + + iterator = A2AModelResponseIterator(streaming_response=_events(), sync_stream=False) + chunk = await iterator.__aiter__().__anext__() + + assert chunk["tool_use"]["id"] == "call-1" + assert chunk["finish_reason"] == "length" + assert chunk["usage"].total_tokens == 5 + + +@pytest.mark.asyncio +async def test_async_iterator_closes_nested_stream(): + closed = False + + async def _events(): + nonlocal closed + try: + yield {"jsonrpc": "2.0", "result": {"kind": "artifact-update"}} + raise AssertionError("stream should be closed before a second event") + finally: + closed = True + + iterator = A2AModelResponseIterator(streaming_response=_events(), sync_stream=False) + await iterator.__aiter__().__anext__() + await iterator.aclose() + + assert closed is True + + @pytest.mark.asyncio async def test_async_iterator_propagates_jsonrpc_errors(): async def _events(): diff --git a/tests/test_litellm/proxy/test_route_a2a_models.py b/tests/test_litellm/proxy/test_route_a2a_models.py index 766745425c3..ca8a7ff5b2a 100644 --- a/tests/test_litellm/proxy/test_route_a2a_models.py +++ b/tests/test_litellm/proxy/test_route_a2a_models.py @@ -98,6 +98,9 @@ async def test_route_a2a_model_uses_registered_provider(): "temperature": 0.2, "timeout": 12.0, "tools": [{"type": "function", "function": {"name": "lookup"}}], + "output_config": {"format": "json"}, + "prompt_cache_key": "cache-key", + "safety_identifier": "safety-id", "proxy_server_request": { "headers": { "x-tenant": "tenant-1", @@ -144,6 +147,9 @@ async def test_route_a2a_model_uses_registered_provider(): assert bridge_kwargs["litellm_params"]["temperature"] == 0.2 assert bridge_kwargs["litellm_params"]["timeout"] == 12.0 assert bridge_kwargs["litellm_params"]["tools"] == data["tools"] + assert bridge_kwargs["litellm_params"]["output_config"] == data["output_config"] + assert bridge_kwargs["litellm_params"]["prompt_cache_key"] == data["prompt_cache_key"] + assert bridge_kwargs["litellm_params"]["safety_identifier"] == data["safety_identifier"] assert bridge_kwargs["litellm_params"]["guardrails"] == ["request-guardrail", "agent-guardrail"] assert bridge_kwargs["litellm_params"]["extra_headers"] == { "X-Tenant": "tenant-1",