mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
refactor(guardrails): delegate get_user_prompt to get_last_user_message
PurviewGuardrailBase duplicated AzureGuardrailBase (and OpenAIGuardrailBase) user-prompt extraction. The same logic already lived in common_utils.get_last_user_message; wire guardrail bases to that helper, fix the helper docstring, and drop its redundant self-import of convert_content_list_to_str. Co-authored-by: Sameer Kankute <Sameerlite@users.noreply.github.com>
This commit is contained in:
parent
40e4b50a86
commit
f09e335375
4 changed files with 15 additions and 89 deletions
|
|
@ -1196,12 +1196,8 @@ def get_last_user_message(messages: List[AllMessageValues]) -> Optional[str]:
|
|||
{"role": "assistant", "content": "I'm good, thank you!"},
|
||||
{"role": "user", "content": "What is the weather in Tokyo?"},
|
||||
]
|
||||
get_user_prompt(messages) -> "What is the weather in Tokyo?"
|
||||
get_last_user_message(messages) -> "What is the weather in Tokyo?"
|
||||
"""
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
convert_content_list_to_str,
|
||||
)
|
||||
|
||||
if not messages:
|
||||
return None
|
||||
|
||||
|
|
|
|||
|
|
@ -2,6 +2,9 @@ import re
|
|||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
get_last_user_message,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
|
|
@ -134,32 +137,4 @@ class AzureGuardrailBase:
|
|||
]
|
||||
get_user_prompt(messages) -> "What is the weather in Tokyo?"
|
||||
"""
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
convert_content_list_to_str,
|
||||
)
|
||||
|
||||
if not messages:
|
||||
return None
|
||||
|
||||
# Iterate from the end to find the last consecutive block of user messages
|
||||
user_messages = []
|
||||
for message in reversed(messages):
|
||||
if message.get("role") == "user":
|
||||
user_messages.append(message)
|
||||
else:
|
||||
# Stop when we hit a non-user message
|
||||
break
|
||||
|
||||
if not user_messages:
|
||||
return None
|
||||
|
||||
# Reverse to get the messages in chronological order
|
||||
user_messages.reverse()
|
||||
|
||||
user_prompt = ""
|
||||
for message in user_messages:
|
||||
text_content = convert_content_list_to_str(message)
|
||||
user_prompt += text_content + "\n"
|
||||
|
||||
result = user_prompt.strip()
|
||||
return result if result else None
|
||||
return get_last_user_message(messages)
|
||||
|
|
|
|||
|
|
@ -4,6 +4,9 @@ from collections import OrderedDict
|
|||
from typing import TYPE_CHECKING, Any, Dict, List, MutableMapping, Optional, Tuple
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
get_last_user_message,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
|
|
@ -279,33 +282,9 @@ class PurviewGuardrailBase:
|
|||
return False
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# User prompt extraction (same pattern as AzureGuardrailBase)
|
||||
# User prompt extraction
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def get_user_prompt(self, messages: List["AllMessageValues"]) -> Optional[str]:
|
||||
"""Get the last consecutive block of user messages as a single string."""
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
convert_content_list_to_str,
|
||||
)
|
||||
|
||||
if not messages:
|
||||
return None
|
||||
|
||||
user_messages = []
|
||||
for message in reversed(messages):
|
||||
if message.get("role") == "user":
|
||||
user_messages.append(message)
|
||||
else:
|
||||
break
|
||||
|
||||
if not user_messages:
|
||||
return None
|
||||
|
||||
user_messages.reverse()
|
||||
user_prompt = ""
|
||||
for message in user_messages:
|
||||
text_content = convert_content_list_to_str(message)
|
||||
user_prompt += text_content + "\n"
|
||||
|
||||
result = user_prompt.strip()
|
||||
return result if result else None
|
||||
return get_last_user_message(messages)
|
||||
|
|
|
|||
|
|
@ -1,5 +1,9 @@
|
|||
from typing import TYPE_CHECKING, List, Optional
|
||||
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
get_last_user_message,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
|
|
@ -21,32 +25,4 @@ class OpenAIGuardrailBase:
|
|||
]
|
||||
get_user_prompt(messages) -> "What is the weather in Tokyo?"
|
||||
"""
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
convert_content_list_to_str,
|
||||
)
|
||||
|
||||
if not messages:
|
||||
return None
|
||||
|
||||
# Iterate from the end to find the last consecutive block of user messages
|
||||
user_messages = []
|
||||
for message in reversed(messages):
|
||||
if message.get("role") == "user":
|
||||
user_messages.append(message)
|
||||
else:
|
||||
# Stop when we hit a non-user message
|
||||
break
|
||||
|
||||
if not user_messages:
|
||||
return None
|
||||
|
||||
# Reverse to get the messages in chronological order
|
||||
user_messages.reverse()
|
||||
|
||||
user_prompt = ""
|
||||
for message in user_messages:
|
||||
text_content = convert_content_list_to_str(message)
|
||||
user_prompt += text_content + "\n"
|
||||
|
||||
result = user_prompt.strip()
|
||||
return result if result else None
|
||||
return get_last_user_message(messages)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue