mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
feat(a2a/): ensure a2a guardrails works on response output
This commit is contained in:
parent
b77f2e046b
commit
f30f883744
5 changed files with 54 additions and 25 deletions
|
|
@ -74,6 +74,7 @@ from litellm.llms.vertex_ai.cost_calculator import (
|
|||
from litellm.llms.vertex_ai.cost_calculator import cost_router as google_cost_router
|
||||
from litellm.llms.xai.cost_calculator import cost_per_token as xai_cost_per_token
|
||||
from litellm.responses.utils import ResponseAPILoggingUtils
|
||||
from litellm.types.agents import LiteLLMSendMessageResponse
|
||||
from litellm.types.llms.openai import (
|
||||
HttpxBinaryResponseContent,
|
||||
ImageGenerationRequestQuality,
|
||||
|
|
@ -150,32 +151,33 @@ def _get_additional_costs(
|
|||
) -> Optional[dict]:
|
||||
"""
|
||||
Calculate additional costs beyond standard token costs.
|
||||
|
||||
|
||||
This function delegates to provider-specific config classes to calculate
|
||||
any additional costs like routing fees, infrastructure costs, etc.
|
||||
|
||||
|
||||
Args:
|
||||
model: The model name
|
||||
custom_llm_provider: The provider name (optional)
|
||||
prompt_tokens: Number of prompt tokens
|
||||
completion_tokens: Number of completion tokens
|
||||
|
||||
|
||||
Returns:
|
||||
Optional dictionary with cost names and amounts, or None if no additional costs
|
||||
"""
|
||||
if not custom_llm_provider:
|
||||
return None
|
||||
|
||||
|
||||
try:
|
||||
config_class = None
|
||||
if custom_llm_provider == "azure_ai":
|
||||
from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo
|
||||
|
||||
config_class = AzureFoundryModelInfo.get_azure_ai_config_for_model(model)
|
||||
# Add more providers here as needed
|
||||
# elif custom_llm_provider == "other_provider":
|
||||
# config_class = get_other_provider_config(model)
|
||||
|
||||
if config_class and hasattr(config_class, 'calculate_additional_costs'):
|
||||
|
||||
if config_class and hasattr(config_class, "calculate_additional_costs"):
|
||||
return config_class.calculate_additional_costs(
|
||||
model=model,
|
||||
prompt_tokens=prompt_tokens,
|
||||
|
|
@ -183,7 +185,7 @@ def _get_additional_costs(
|
|||
)
|
||||
except Exception as e:
|
||||
verbose_logger.debug(f"Error calculating additional costs: {e}")
|
||||
|
||||
|
||||
return None
|
||||
|
||||
|
||||
|
|
@ -748,6 +750,8 @@ def _infer_call_type(
|
|||
return "image_generation"
|
||||
elif isinstance(completion_response, TextCompletionResponse):
|
||||
return "text_completion"
|
||||
elif isinstance(completion_response, LiteLLMSendMessageResponse):
|
||||
return "send_message"
|
||||
|
||||
return call_type
|
||||
|
||||
|
|
@ -1037,9 +1041,9 @@ def completion_cost( # noqa: PLR0915
|
|||
or isinstance(completion_response, dict)
|
||||
): # tts returns a custom class
|
||||
if isinstance(completion_response, dict):
|
||||
usage_obj: Optional[
|
||||
Union[dict, Usage]
|
||||
] = completion_response.get("usage", {})
|
||||
usage_obj: Optional[Union[dict, Usage]] = (
|
||||
completion_response.get("usage", {})
|
||||
)
|
||||
else:
|
||||
usage_obj = getattr(completion_response, "usage", {})
|
||||
if isinstance(usage_obj, BaseModel) and not _is_known_usage_objects(
|
||||
|
|
@ -1393,7 +1397,7 @@ def completion_cost( # noqa: PLR0915
|
|||
service_tier=service_tier,
|
||||
response=completion_response,
|
||||
)
|
||||
|
||||
|
||||
# Get additional costs from provider (e.g., routing fees, infrastructure costs)
|
||||
additional_costs = _get_additional_costs(
|
||||
model=model,
|
||||
|
|
@ -1401,7 +1405,7 @@ def completion_cost( # noqa: PLR0915
|
|||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
)
|
||||
|
||||
|
||||
_final_cost = (
|
||||
prompt_tokens_cost_usd_dollar + completion_tokens_cost_usd_dollar
|
||||
)
|
||||
|
|
|
|||
|
|
@ -323,8 +323,18 @@ async def invoke_agent_a2a(
|
|||
metadata=data.get("metadata", {}),
|
||||
proxy_server_request=data.get("proxy_server_request"),
|
||||
)
|
||||
|
||||
response = await proxy_logging_obj.post_call_success_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
data=data,
|
||||
response=response,
|
||||
)
|
||||
return JSONResponse(
|
||||
content=response.model_dump(mode="json", exclude_none=True)
|
||||
content=(
|
||||
response.model_dump(mode="json", exclude_none=True) # type: ignore
|
||||
if hasattr(response, "model_dump")
|
||||
else response
|
||||
)
|
||||
)
|
||||
|
||||
elif method == "message/stream":
|
||||
|
|
|
|||
|
|
@ -67,7 +67,6 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
if call_type == CallTypes.call_mcp_tool.value:
|
||||
event_type = GuardrailEventHooks.pre_mcp_call
|
||||
|
||||
|
||||
if (
|
||||
guardrail_to_apply.should_run_guardrail(data=data, event_type=event_type)
|
||||
is not True
|
||||
|
|
|
|||
|
|
@ -950,6 +950,7 @@ class ProxyLogging:
|
|||
data: dict,
|
||||
user_api_key_dict: Optional[UserAPIKeyAuth],
|
||||
call_type: CallTypesLiteral,
|
||||
event_type: GuardrailEventHooks,
|
||||
) -> Optional[dict]:
|
||||
"""
|
||||
Process a guardrail callback during pre-call hook.
|
||||
|
|
@ -969,8 +970,10 @@ class ProxyLogging:
|
|||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
# Determine the event type based on call type
|
||||
event_type = GuardrailEventHooks.pre_call
|
||||
if call_type == CallTypes.call_mcp_tool.value:
|
||||
if (
|
||||
event_type is GuardrailEventHooks.pre_call
|
||||
and call_type == CallTypes.call_mcp_tool.value
|
||||
):
|
||||
event_type = GuardrailEventHooks.pre_mcp_call
|
||||
|
||||
# Check if the guardrail should run for this request
|
||||
|
|
@ -1332,6 +1335,7 @@ class ProxyLogging:
|
|||
data=data, # type: ignore
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type=call_type,
|
||||
event_type=GuardrailEventHooks.pre_call,
|
||||
)
|
||||
if result is None:
|
||||
continue
|
||||
|
|
@ -1485,11 +1489,11 @@ class ProxyLogging:
|
|||
# Note: user_info is a CallInfo that can represent user/team/org level info. For team budgets,
|
||||
# alert_emails is populated from team_object.metadata.soft_budget_alerting_emails (see auth_checks.py)
|
||||
is_soft_budget_with_alert_emails = (
|
||||
type == "soft_budget"
|
||||
and user_info.alert_emails is not None
|
||||
type == "soft_budget"
|
||||
and user_info.alert_emails is not None
|
||||
and len(user_info.alert_emails) > 0
|
||||
)
|
||||
|
||||
|
||||
if self.alerting is None and not is_soft_budget_with_alert_emails:
|
||||
# do nothing if alerting is not switched on (unless it's a soft_budget alert with team-specific emails)
|
||||
return
|
||||
|
|
@ -1505,10 +1509,9 @@ class ProxyLogging:
|
|||
# 1. "email" is in alerting config, OR
|
||||
# 2. It's a soft_budget alert with team-specific alert_emails (bypasses global alerting config)
|
||||
should_send_email = (
|
||||
(self.alerting is not None and "email" in self.alerting)
|
||||
or is_soft_budget_with_alert_emails
|
||||
)
|
||||
|
||||
self.alerting is not None and "email" in self.alerting
|
||||
) or is_soft_budget_with_alert_emails
|
||||
|
||||
if should_send_email and self.email_logging_instance is not None:
|
||||
await self.email_logging_instance.budget_alerts(
|
||||
type=type,
|
||||
|
|
@ -1872,6 +1875,7 @@ class ProxyLogging:
|
|||
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
|
||||
guardrail_callbacks: List[CustomGuardrail] = []
|
||||
other_callbacks: List[CustomLogger] = []
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -33,6 +33,7 @@ from litellm.types.llms.base import (
|
|||
from litellm.types.mcp import MCPServerCostInfo
|
||||
|
||||
from ..litellm_core_utils.core_helpers import map_finish_reason
|
||||
from .agents import LiteLLMSendMessageResponse
|
||||
from .guardrails import GuardrailEventHooks
|
||||
from .llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse
|
||||
from .llms.base import HiddenParams
|
||||
|
|
@ -777,6 +778,7 @@ API_ROUTE_TO_CALL_TYPES = {
|
|||
"/mcp/call_tool": [CallTypes.call_mcp_tool],
|
||||
# A2A (Agent-to-Agent)
|
||||
"/a2a/{agent_id}": [CallTypes.asend_message, CallTypes.send_message],
|
||||
"/a2a/{agent_id}/message/send": [CallTypes.asend_message, CallTypes.send_message],
|
||||
# Passthrough endpoints
|
||||
"/llm_passthrough": [
|
||||
CallTypes.llm_passthrough_route,
|
||||
|
|
@ -2139,7 +2141,14 @@ class ImageObject(OpenAIImage):
|
|||
revised_prompt: Optional[str] = None
|
||||
provider_specific_fields: Optional[Dict[str, Any]] = None
|
||||
|
||||
def __init__(self, b64_json=None, url=None, revised_prompt=None, provider_specific_fields=None, **kwargs):
|
||||
def __init__(
|
||||
self,
|
||||
b64_json=None,
|
||||
url=None,
|
||||
revised_prompt=None,
|
||||
provider_specific_fields=None,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(b64_json=b64_json, url=url, revised_prompt=revised_prompt) # type: ignore
|
||||
if provider_specific_fields:
|
||||
self.provider_specific_fields = provider_specific_fields
|
||||
|
|
@ -2641,7 +2650,9 @@ class CostBreakdown(TypedDict, total=False):
|
|||
)
|
||||
total_cost: float # Total cost (input + output + tool usage)
|
||||
tool_usage_cost: float # Cost of usage of built-in tools
|
||||
additional_costs: Dict[str, float] # Free-form additional costs (e.g., {"azure_model_router_flat_cost": 0.00014})
|
||||
additional_costs: Dict[
|
||||
str, float
|
||||
] # Free-form additional costs (e.g., {"azure_model_router_flat_cost": 0.00014})
|
||||
original_cost: float # Cost before discount (optional)
|
||||
discount_percent: float # Discount percentage applied (e.g., 0.05 = 5%) (optional)
|
||||
discount_amount: float # Discount amount in USD (optional)
|
||||
|
|
@ -3355,6 +3366,7 @@ LLMResponseTypes = Union[
|
|||
LiteLLMFineTuningJob,
|
||||
AnthropicMessagesResponse,
|
||||
ResponsesAPIResponse,
|
||||
LiteLLMSendMessageResponse,
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue