mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
refactor: extract helper methods from pre_call_hook to fix PLR0915 linting error (#17613)
- Extract _process_prompt_template method for prompt template processing - Extract _process_guardrail_metadata method for guardrail metadata processing - Reduces statement count in pre_call_hook from 68 to under 50
This commit is contained in:
parent
a1d114a8ad
commit
3359daf621
1 changed files with 86 additions and 79 deletions
|
|
@ -824,6 +824,85 @@ class ProxyLogging:
|
|||
|
||||
return data
|
||||
|
||||
def _process_prompt_template(
|
||||
self, data: dict, litellm_logging_obj: Any, prompt_id: Any, prompt_version: Any, call_type: CallTypesLiteral
|
||||
) -> None:
|
||||
"""Process prompt template if applicable."""
|
||||
from litellm.utils import get_non_default_completion_params
|
||||
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
|
||||
|
||||
if prompt_version is None:
|
||||
lookup_prompt_id = get_latest_version_prompt_id(
|
||||
prompt_id=prompt_id,
|
||||
all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS,
|
||||
)
|
||||
else:
|
||||
lookup_prompt_id = construct_versioned_prompt_id(
|
||||
prompt_id=prompt_id, version=prompt_version
|
||||
)
|
||||
|
||||
custom_logger = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id(
|
||||
lookup_prompt_id
|
||||
)
|
||||
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(lookup_prompt_id)
|
||||
litellm_prompt_id: Optional[str] = None
|
||||
if prompt_spec is not None:
|
||||
litellm_prompt_id = prompt_spec.litellm_params.prompt_id
|
||||
|
||||
if custom_logger and litellm_prompt_id is not None:
|
||||
(
|
||||
model,
|
||||
messages,
|
||||
optional_params,
|
||||
) = litellm_logging_obj.get_chat_completion_prompt(
|
||||
model=data.get("model", ""),
|
||||
messages=data.get("messages", []),
|
||||
non_default_params=get_non_default_completion_params(kwargs=data),
|
||||
prompt_id=litellm_prompt_id,
|
||||
prompt_management_logger=custom_logger,
|
||||
prompt_variables=data.get("prompt_variables", None),
|
||||
prompt_label=data.get("prompt_label", None),
|
||||
prompt_version=data.get("prompt_version", None),
|
||||
)
|
||||
|
||||
data.update(optional_params)
|
||||
data["model"] = model
|
||||
data["messages"] = messages
|
||||
|
||||
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,
|
||||
)
|
||||
metadata_standard = data.get("metadata") or {}
|
||||
metadata_litellm = data.get("litellm_metadata") or {}
|
||||
|
||||
guardrails_in_metadata = []
|
||||
if isinstance(metadata_standard, dict) and "guardrails" in metadata_standard:
|
||||
guardrails_in_metadata = metadata_standard.get("guardrails", [])
|
||||
elif isinstance(metadata_litellm, dict) and "guardrails" in metadata_litellm:
|
||||
guardrails_in_metadata = metadata_litellm.get("guardrails", [])
|
||||
|
||||
if guardrails_in_metadata and isinstance(guardrails_in_metadata, list):
|
||||
applied_guardrails = []
|
||||
if isinstance(metadata_standard, dict) and "applied_guardrails" in metadata_standard:
|
||||
applied_guardrails = metadata_standard.get("applied_guardrails", [])
|
||||
elif isinstance(metadata_litellm, dict) and "applied_guardrails" in metadata_litellm:
|
||||
applied_guardrails = metadata_litellm.get("applied_guardrails", [])
|
||||
|
||||
if not isinstance(applied_guardrails, list):
|
||||
applied_guardrails = []
|
||||
|
||||
for guardrail_name in guardrails_in_metadata:
|
||||
if isinstance(guardrail_name, str) and guardrail_name not in applied_guardrails:
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=data, guardrail_name=guardrail_name
|
||||
)
|
||||
|
||||
# The actual implementation of the function
|
||||
@overload
|
||||
async def pre_call_hook(
|
||||
|
|
@ -857,8 +936,6 @@ class ProxyLogging:
|
|||
2. /embeddings
|
||||
3. /image/generation
|
||||
"""
|
||||
from litellm.utils import get_non_default_completion_params
|
||||
|
||||
verbose_proxy_logger.debug("Inside Proxy Logging Pre-call hook!")
|
||||
|
||||
self._init_response_taking_too_long_task(data=data)
|
||||
|
|
@ -878,51 +955,13 @@ class ProxyLogging:
|
|||
and prompt_id is not None
|
||||
and (call_type == "completion" or call_type == "acompletion")
|
||||
):
|
||||
from litellm.proxy.prompts.prompt_endpoints import (
|
||||
construct_versioned_prompt_id,
|
||||
get_latest_version_prompt_id,
|
||||
self._process_prompt_template(
|
||||
data=data,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
prompt_id=prompt_id,
|
||||
prompt_version=prompt_version,
|
||||
call_type=call_type,
|
||||
)
|
||||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||||
|
||||
# If no version is specified, find the latest version
|
||||
if prompt_version is None:
|
||||
lookup_prompt_id = get_latest_version_prompt_id(
|
||||
prompt_id=prompt_id,
|
||||
all_prompt_ids=IN_MEMORY_PROMPT_REGISTRY.IN_MEMORY_PROMPTS,
|
||||
)
|
||||
else:
|
||||
# Construct versioned prompt_id if prompt_version is provided
|
||||
lookup_prompt_id = construct_versioned_prompt_id(
|
||||
prompt_id=prompt_id, version=prompt_version
|
||||
)
|
||||
|
||||
custom_logger = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id(
|
||||
lookup_prompt_id
|
||||
)
|
||||
prompt_spec = IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id(lookup_prompt_id)
|
||||
litellm_prompt_id: Optional[str] = None
|
||||
if prompt_spec is not None:
|
||||
litellm_prompt_id = prompt_spec.litellm_params.prompt_id
|
||||
|
||||
if custom_logger and litellm_prompt_id is not None:
|
||||
(
|
||||
model,
|
||||
messages,
|
||||
optional_params,
|
||||
) = litellm_logging_obj.get_chat_completion_prompt(
|
||||
model=data.get("model", ""),
|
||||
messages=data.get("messages", []),
|
||||
non_default_params=get_non_default_completion_params(kwargs=data),
|
||||
prompt_id=litellm_prompt_id,
|
||||
prompt_management_logger=custom_logger,
|
||||
prompt_variables=data.get("prompt_variables", None),
|
||||
prompt_label=data.get("prompt_label", None),
|
||||
prompt_version=data.get("prompt_version", None),
|
||||
)
|
||||
|
||||
data.update(optional_params)
|
||||
data["model"] = model
|
||||
data["messages"] = messages
|
||||
|
||||
try:
|
||||
for callback in litellm.callbacks:
|
||||
|
|
@ -984,40 +1023,8 @@ class ProxyLogging:
|
|||
end_time=end_time,
|
||||
)
|
||||
|
||||
# After processing all callbacks, check if there are guardrails in metadata
|
||||
# that weren't processed (e.g., guardrails from API key metadata that don't have callbacks)
|
||||
# and add them to applied_guardrails
|
||||
if data is not None:
|
||||
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 {}
|
||||
|
||||
# Check both metadata locations for guardrails
|
||||
guardrails_in_metadata = []
|
||||
if isinstance(metadata_standard, dict) and "guardrails" in metadata_standard:
|
||||
guardrails_in_metadata = metadata_standard.get("guardrails", [])
|
||||
elif isinstance(metadata_litellm, dict) and "guardrails" in metadata_litellm:
|
||||
guardrails_in_metadata = metadata_litellm.get("guardrails", [])
|
||||
|
||||
if guardrails_in_metadata and isinstance(guardrails_in_metadata, list):
|
||||
# Get guardrails that are already in applied_guardrails (check both locations)
|
||||
applied_guardrails = []
|
||||
if isinstance(metadata_standard, dict) and "applied_guardrails" in metadata_standard:
|
||||
applied_guardrails = metadata_standard.get("applied_guardrails", [])
|
||||
elif isinstance(metadata_litellm, dict) and "applied_guardrails" in metadata_litellm:
|
||||
applied_guardrails = metadata_litellm.get("applied_guardrails", [])
|
||||
|
||||
if not isinstance(applied_guardrails, list):
|
||||
applied_guardrails = []
|
||||
|
||||
# Add any guardrails from metadata that aren't already in applied_guardrails
|
||||
for guardrail_name in guardrails_in_metadata:
|
||||
if isinstance(guardrail_name, str) and guardrail_name not in applied_guardrails:
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=data, guardrail_name=guardrail_name
|
||||
)
|
||||
self._process_guardrail_metadata(data)
|
||||
|
||||
return data
|
||||
except Exception as e:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue