fix: feedback changes

This commit is contained in:
Harshit Jain 2026-02-24 18:03:05 +05:30
parent cb256f5bc8
commit 42c2f55010
No known key found for this signature in database
GPG key ID: 36C392CD4415B4CF
3 changed files with 115 additions and 64 deletions

View file

@ -133,9 +133,11 @@ class CustomGuardrail(CustomLogger):
Optional[List[int]],
]:
"""Filter to only the latest message for target_role."""
import copy
for index in range(len(messages) - 1, -1, -1):
if messages[index].get("role") == target_role:
return [messages[index]], list(messages), [index]
return [copy.deepcopy(messages[index])], list(messages), [index]
return None, None, None
def merge_filtered_messages(

View file

@ -242,13 +242,22 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
self,
messages: Optional[List[AllMessageValues]],
) -> GuardrailMessageFilterResult:
"""Return payload + merge metadata for the latest user message."""
"""Return payload + merge metadata for the latest user message.
If the proxy has already filtered the messages (e.g. len == 1), this avoids
redundant calculation of original_messages/indices.
"""
if messages is None:
return GuardrailMessageFilterResult(None, None, None)
if self.experimental_use_latest_role_message_only is not True:
return GuardrailMessageFilterResult(messages, None, None)
# If the Proxy already filtered this to a single message, don't re-filter
# This prevents redundant nested merging.
if len(messages) == 1 and messages[0].get("role") == "user":
return GuardrailMessageFilterResult(messages, None, None)
(
filtered_messages,
original_messages,

View file

@ -25,23 +25,31 @@ from typing import (
)
from litellm import _custom_logger_compatible_callbacks_literal
from litellm.constants import (DEFAULT_MODEL_CREATED_AT_TIME,
MAX_TEAM_LIST_LIMIT)
from litellm.proxy._types import (DB_CONNECTION_ERROR_TYPES, CommonProxyErrors,
ProxyErrorTypes, ProxyException,
SpendLogsMetadata, SpendLogsPayload)
from litellm.constants import DEFAULT_MODEL_CREATED_AT_TIME, MAX_TEAM_LIST_LIMIT
from litellm.proxy._types import (
DB_CONNECTION_ERROR_TYPES,
CommonProxyErrors,
ProxyErrorTypes,
ProxyException,
SpendLogsMetadata,
SpendLogsPayload,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import CallTypes, CallTypesLiteral
try:
from litellm_enterprise.enterprise_callbacks.send_emails.base_email import \
BaseEmailLogger
from litellm_enterprise.enterprise_callbacks.send_emails.resend_email import \
ResendEmailLogger
from litellm_enterprise.enterprise_callbacks.send_emails.sendgrid_email import \
SendGridEmailLogger
from litellm_enterprise.enterprise_callbacks.send_emails.smtp_email import \
SMTPEmailLogger
from litellm_enterprise.enterprise_callbacks.send_emails.base_email import (
BaseEmailLogger,
)
from litellm_enterprise.enterprise_callbacks.send_emails.resend_email import (
ResendEmailLogger,
)
from litellm_enterprise.enterprise_callbacks.send_emails.sendgrid_email import (
SendGridEmailLogger,
)
from litellm_enterprise.enterprise_callbacks.send_emails.smtp_email import (
SMTPEmailLogger,
)
except ImportError:
BaseEmailLogger = None # type: ignore
SendGridEmailLogger = None # type: ignore
@ -60,56 +68,70 @@ from fastapi import HTTPException, status
import litellm
import litellm.litellm_core_utils
import litellm.litellm_core_utils.litellm_logging
from litellm import (EmbeddingResponse, ImageResponse, ModelResponse,
ModelResponseStream, Router)
from litellm import (
EmbeddingResponse,
ImageResponse,
ModelResponse,
ModelResponseStream,
Router,
)
from litellm._logging import verbose_proxy_logger
from litellm._service_logger import ServiceLogging, ServiceTypes
from litellm.caching.caching import DualCache, RedisCache
from litellm.caching.dual_cache import LimitedSizeOrderedDict
from litellm.exceptions import RejectedRequestError
from litellm.integrations.custom_guardrail import (CustomGuardrail,
ModifyResponseException)
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
ModifyResponseException,
)
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
from litellm.integrations.SlackAlerting.utils import \
_add_langfuse_trace_id_to_alert
from litellm.integrations.SlackAlerting.utils import _add_langfuse_trace_id_to_alert
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy._types import (AlertType, CallInfo,
LiteLLM_VerificationTokenView, Member,
UserAPIKeyAuth)
from litellm.proxy._types import (
AlertType,
CallInfo,
LiteLLM_VerificationTokenView,
Member,
UserAPIKeyAuth,
)
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.db.create_views import (create_missing_views,
should_create_missing_views)
from litellm.proxy.db.create_views import (
create_missing_views,
should_create_missing_views,
)
from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.db.log_db_metrics import log_db_metrics
from litellm.proxy.db.prisma_client import PrismaWrapper
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import \
UnifiedLLMGuardrails
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
from litellm.proxy.hooks import PROXY_HOOKS, get_proxy_hook
from litellm.proxy.hooks.cache_control_check import _PROXY_CacheControlCheck
from litellm.proxy.hooks.max_budget_limiter import _PROXY_MaxBudgetLimiter
from litellm.proxy.hooks.parallel_request_limiter import \
_PROXY_MaxParallelRequestsHandler
from litellm.proxy.hooks.parallel_request_limiter import (
_PROXY_MaxParallelRequestsHandler,
)
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.policy_engine.pipeline_executor import PipelineExecutor
from litellm.secret_managers.main import str_to_bool
from litellm.types.integrations.slack_alerting import DEFAULT_ALERT_TYPES
from litellm.types.mcp import (MCPDuringCallResponseObject,
MCPPreCallRequestObject,
MCPPreCallResponseObject)
from litellm.types.proxy.policy_engine.pipeline_types import \
PipelineExecutionResult
from litellm.types.mcp import (
MCPDuringCallResponseObject,
MCPPreCallRequestObject,
MCPPreCallResponseObject,
)
from litellm.types.proxy.policy_engine.pipeline_types import PipelineExecutionResult
from litellm.types.utils import LLMResponseTypes, LoggedLiteLLMParams
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
from litellm.litellm_core_utils.litellm_logging import \
Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
Span = Union[_Span, Any]
else:
@ -1081,9 +1103,10 @@ class ProxyLogging:
"""Process prompt template if applicable."""
from litellm.proxy.prompts.prompt_endpoints import (
construct_versioned_prompt_id, get_latest_version_prompt_id)
from litellm.proxy.prompts.prompt_registry import \
IN_MEMORY_PROMPT_REGISTRY
construct_versioned_prompt_id,
get_latest_version_prompt_id,
)
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
from litellm.utils import get_non_default_completion_params
if prompt_version is None:
@ -1133,8 +1156,9 @@ class ProxyLogging:
def _process_guardrail_metadata(self, data: dict) -> None:
"""Process guardrails from metadata and add to applied_guardrails."""
from litellm.proxy.common_utils.callback_utils import \
add_guardrail_to_applied_guardrails_header
from litellm.proxy.common_utils.callback_utils import (
add_guardrail_to_applied_guardrails_header,
)
metadata_standard = data.get("metadata") or {}
metadata_litellm = data.get("litellm_metadata") or {}
@ -1453,19 +1477,34 @@ class ProxyLogging:
else:
user_api_key_auth_dict = user_api_key_dict
# Add task to list for parallel execution
data_for_guardrail = data
if (
hasattr(callback, "experimental_use_latest_role_message_only")
and callback.experimental_use_latest_role_message_only
and isinstance(data.get("messages"), list)
):
import copy
data_for_guardrail = copy.copy(data)
filtered, _, _ = callback.filter_messages_for_latest_role(
data["messages"]
)
if filtered is not None:
data_for_guardrail["messages"] = filtered
if (
"apply_guardrail" in type(callback).__dict__
and user_api_key_dict is not None
):
data["guardrail_to_apply"] = callback
data_for_guardrail["guardrail_to_apply"] = callback
guardrail_task = unified_guardrail.async_moderation_hook(
user_api_key_dict=user_api_key_dict,
data=data,
data=data_for_guardrail,
call_type=call_type,
)
else:
guardrail_task = callback.async_moderation_hook(
data=data,
data=data_for_guardrail,
user_api_key_dict=user_api_key_auth_dict, # type: ignore
call_type=call_type, # type: ignore
)
@ -2031,8 +2070,7 @@ class ProxyLogging:
if isinstance(response, (ModelResponse, ModelResponseStream)):
response_str = litellm.get_response_string(response_obj=response)
elif isinstance(response, dict) and self.is_a2a_streaming_response(response):
from litellm.llms.a2a.common_utils import \
extract_text_from_a2a_response
from litellm.llms.a2a.common_utils import extract_text_from_a2a_response
response_str = extract_text_from_a2a_response(response)
if response_str is not None:
@ -2041,8 +2079,7 @@ class ProxyLogging:
_callback: Optional[CustomLogger] = None
if isinstance(callback, CustomGuardrail):
# Main - V2 Guardrails implementation
from litellm.types.guardrails import \
GuardrailEventHooks
from litellm.types.guardrails import GuardrailEventHooks
## CHECK FOR MODEL-LEVEL GUARDRAILS
modified_data = _check_and_merge_model_level_guardrails(
@ -4192,9 +4229,9 @@ class ProxyUpdateSpend:
:MAX_LOGS_PER_INTERVAL
]
# Remove the logs we're about to process
prisma_client.spend_log_transactions = prisma_client.spend_log_transactions[
len(logs_to_process) :
]
prisma_client.spend_log_transactions = (
prisma_client.spend_log_transactions[len(logs_to_process) :]
)
popped_batch = True
start_time = time.time()
try:
@ -4335,9 +4372,7 @@ async def update_spend_logs_job(
return
async with prisma_client._spend_log_transactions_lock:
logs_to_process = prisma_client.spend_log_transactions[
:MAX_LOGS_PER_INTERVAL
]
logs_to_process = prisma_client.spend_log_transactions[:MAX_LOGS_PER_INTERVAL]
prisma_client.spend_log_transactions = prisma_client.spend_log_transactions[
len(logs_to_process) :
]
@ -4352,8 +4387,10 @@ async def update_spend_logs_job(
# Guardrail/policy usage tracking (same batch, outside spend-logs update)
try:
from litellm.proxy.guardrails.usage_tracking import \
process_spend_logs_guardrail_usage
from litellm.proxy.guardrails.usage_tracking import (
process_spend_logs_guardrail_usage,
)
await process_spend_logs_guardrail_usage(
prisma_client=prisma_client,
logs_to_process=logs_to_process,
@ -4379,8 +4416,10 @@ async def _monitor_spend_logs_queue(
db_writer_client: Optional HTTP handler for external spend logs endpoint
proxy_logging_obj: Proxy logging object
"""
from litellm.constants import (SPEND_LOG_QUEUE_POLL_INTERVAL,
SPEND_LOG_QUEUE_SIZE_THRESHOLD)
from litellm.constants import (
SPEND_LOG_QUEUE_POLL_INTERVAL,
SPEND_LOG_QUEUE_SIZE_THRESHOLD,
)
threshold = SPEND_LOG_QUEUE_SIZE_THRESHOLD
base_interval = SPEND_LOG_QUEUE_POLL_INTERVAL
@ -4901,11 +4940,12 @@ async def get_available_models_for_user(
List of model names available to the user
"""
from litellm.proxy.auth.auth_checks import get_team_object
from litellm.proxy.auth.model_checks import (get_complete_model_list,
get_key_models,
get_team_models)
from litellm.proxy.management_endpoints.team_endpoints import \
validate_membership
from litellm.proxy.auth.model_checks import (
get_complete_model_list,
get_key_models,
get_team_models,
)
from litellm.proxy.management_endpoints.team_endpoints import validate_membership
# Get proxy model list and access groups
if llm_router is None: