mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #14107 from uc4w6c/feat/add_guardrails_anthropic
feat: Add guardrail to the Anthropic API endpoint
This commit is contained in:
commit
62d4623d1c
4 changed files with 138 additions and 25 deletions
|
|
@ -119,11 +119,8 @@ class CustomGuardrail(CustomLogger):
|
|||
"""
|
||||
if "guardrails" in data:
|
||||
return data["guardrails"]
|
||||
metadata = data.get("metadata") or {}
|
||||
requested_guardrails = metadata.get("guardrails") or []
|
||||
if requested_guardrails:
|
||||
return requested_guardrails
|
||||
return requested_guardrails
|
||||
metadata = data.get("litellm_metadata") or data.get("metadata", {})
|
||||
return metadata.get("guardrails") or []
|
||||
|
||||
def _guardrail_is_in_requested_guardrails(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -90,6 +90,17 @@ async def anthropic_response( # noqa: PLR0915
|
|||
user_api_key_dict=user_api_key_dict, data=data, call_type="text_completion"
|
||||
)
|
||||
|
||||
tasks = []
|
||||
tasks.append(
|
||||
proxy_logging_obj.during_call_hook(
|
||||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type=ProxyBaseLLMRequestProcessing._get_pre_call_type(
|
||||
route_type="anthropic_messages" # type: ignore
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
### ROUTE THE REQUESTs ###
|
||||
router_model_names = llm_router.model_names if llm_router is not None else []
|
||||
|
||||
|
|
@ -97,23 +108,21 @@ async def anthropic_response( # noqa: PLR0915
|
|||
if (
|
||||
llm_router is not None and data["model"] in router_model_names
|
||||
): # model in router model list
|
||||
llm_response = asyncio.create_task(llm_router.aanthropic_messages(**data))
|
||||
llm_coro = llm_router.aanthropic_messages(**data)
|
||||
elif (
|
||||
llm_router is not None
|
||||
and llm_router.model_group_alias is not None
|
||||
and data["model"] in llm_router.model_group_alias
|
||||
): # model set in model_group_alias
|
||||
llm_response = asyncio.create_task(llm_router.aanthropic_messages(**data))
|
||||
llm_coro = llm_router.aanthropic_messages(**data)
|
||||
elif (
|
||||
llm_router is not None and data["model"] in llm_router.deployment_names
|
||||
): # model in router deployments, calling a specific deployment on the router
|
||||
llm_response = asyncio.create_task(
|
||||
llm_router.aanthropic_messages(**data, specific_deployment=True)
|
||||
)
|
||||
llm_coro = llm_router.aanthropic_messages(**data, specific_deployment=True)
|
||||
elif (
|
||||
llm_router is not None and data["model"] in llm_router.get_model_ids()
|
||||
): # model in router model list
|
||||
llm_response = asyncio.create_task(llm_router.aanthropic_messages(**data))
|
||||
llm_coro = llm_router.aanthropic_messages(**data)
|
||||
elif (
|
||||
llm_router is not None
|
||||
and data["model"] not in router_model_names
|
||||
|
|
@ -122,9 +131,9 @@ async def anthropic_response( # noqa: PLR0915
|
|||
or len(llm_router.pattern_router.patterns) > 0
|
||||
)
|
||||
): # model in router deployments, calling a specific deployment on the router
|
||||
llm_response = asyncio.create_task(llm_router.aanthropic_messages(**data))
|
||||
llm_coro = llm_router.aanthropic_messages(**data)
|
||||
elif user_model is not None: # `litellm --model <your-model-name>`
|
||||
llm_response = asyncio.create_task(litellm.anthropic_messages(**data))
|
||||
llm_coro = litellm.anthropic_messages(**data)
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
|
|
@ -134,8 +143,16 @@ async def anthropic_response( # noqa: PLR0915
|
|||
},
|
||||
)
|
||||
|
||||
# Await the llm_response task
|
||||
response = await llm_response
|
||||
tasks.append(llm_coro)
|
||||
|
||||
# wait for call to end
|
||||
llm_responses = asyncio.gather(
|
||||
*tasks
|
||||
) # run the moderation check in parallel to the actual llm api call
|
||||
|
||||
responses = await llm_responses
|
||||
|
||||
response = responses[1]
|
||||
|
||||
hidden_params = getattr(response, "_hidden_params", {}) or {}
|
||||
model_id = hidden_params.get("model_id", None) or ""
|
||||
|
|
@ -183,6 +200,11 @@ async def anthropic_response( # noqa: PLR0915
|
|||
headers=dict(fastapi_response.headers),
|
||||
)
|
||||
|
||||
### CALL HOOKS ### - modify outgoing data
|
||||
response = await proxy_logging_obj.post_call_success_hook(
|
||||
data=data, user_api_key_dict=user_api_key_dict, response=response # type: ignore
|
||||
)
|
||||
|
||||
verbose_proxy_logger.info("\nResponse from Litellm:\n{}".format(response))
|
||||
return response
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -109,7 +109,6 @@ async def create_streaming_response(
|
|||
final_status_code = default_status_code
|
||||
|
||||
try:
|
||||
|
||||
# Handle coroutine that returns a generator
|
||||
if asyncio.iscoroutine(generator):
|
||||
generator = await generator
|
||||
|
|
@ -118,7 +117,6 @@ async def create_streaming_response(
|
|||
first_chunk_value = await generator.__anext__()
|
||||
|
||||
if first_chunk_value is not None:
|
||||
|
||||
try:
|
||||
error_code_from_chunk = await _parse_event_data_for_error(
|
||||
first_chunk_value
|
||||
|
|
@ -132,7 +130,6 @@ async def create_streaming_response(
|
|||
verbose_proxy_logger.debug(f"Error parsing first chunk value: {e}")
|
||||
|
||||
except StopAsyncIteration:
|
||||
|
||||
# Generator was empty. Default status
|
||||
async def empty_gen() -> AsyncGenerator[str, None]:
|
||||
if False:
|
||||
|
|
@ -145,7 +142,6 @@ async def create_streaming_response(
|
|||
status_code=default_status_code,
|
||||
)
|
||||
except Exception as e:
|
||||
|
||||
# Unexpected error consuming first chunk.
|
||||
verbose_proxy_logger.exception(
|
||||
f"Error consuming first chunk from generator: {e}"
|
||||
|
|
@ -168,7 +164,6 @@ async def create_streaming_response(
|
|||
with tracer.trace(DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE):
|
||||
yield first_chunk_value
|
||||
async for chunk in generator:
|
||||
|
||||
with tracer.trace(DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE):
|
||||
yield chunk
|
||||
|
||||
|
|
@ -462,7 +457,6 @@ class ProxyBaseLLMRequestProcessing:
|
|||
) or self._is_streaming_response(
|
||||
response
|
||||
): # use generate_responses to stream responses
|
||||
|
||||
custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_id=logging_obj.litellm_call_id,
|
||||
|
|
@ -480,7 +474,6 @@ class ProxyBaseLLMRequestProcessing:
|
|||
if route_type == "allm_passthrough_route":
|
||||
# Check if response is an async generator
|
||||
if self._is_streaming_response(response):
|
||||
|
||||
if asyncio.iscoroutine(response):
|
||||
generator = await response
|
||||
else:
|
||||
|
|
@ -501,7 +494,6 @@ class ProxyBaseLLMRequestProcessing:
|
|||
headers=custom_headers,
|
||||
)
|
||||
else:
|
||||
|
||||
selected_data_generator = select_data_generator(
|
||||
response=response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
@ -740,7 +732,11 @@ class ProxyBaseLLMRequestProcessing:
|
|||
verbose_proxy_logger.debug("inside generator")
|
||||
try:
|
||||
str_so_far = ""
|
||||
async for chunk in response:
|
||||
async for chunk in proxy_logging_obj.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=response,
|
||||
request_data=request_data,
|
||||
):
|
||||
verbose_proxy_logger.debug(
|
||||
"async_data_generator: received streaming chunk - {}".format(chunk)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from unittest.mock import AsyncMock, MagicMock
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -82,3 +82,101 @@ class TestCustomGuardrailDeploymentHook:
|
|||
# Verify messages were updated in result
|
||||
assert result["messages"] == mock_result["messages"]
|
||||
assert result["messages"] != original_messages
|
||||
|
||||
|
||||
class TestCustomGuardrailShouldRunGuardrail:
|
||||
|
||||
def test_should_run_guardrail_with_litellm_metadata(self):
|
||||
"""Test that should_run_guardrail works with litellm_metadata pattern"""
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
custom_guardrail = CustomGuardrail(
|
||||
guardrail_name="test_guardrail",
|
||||
default_on=False,
|
||||
event_hook=GuardrailEventHooks.pre_call
|
||||
)
|
||||
|
||||
# Test with guardrails in litellm_metadata
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"litellm_metadata": {
|
||||
"guardrails": ["test_guardrail"]
|
||||
}
|
||||
}
|
||||
|
||||
result = custom_guardrail.should_run_guardrail(
|
||||
data=data, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
|
||||
assert result is True
|
||||
|
||||
def test_should_run_guardrail_with_metadata(self):
|
||||
"""Test that should_run_guardrail works with metadata pattern"""
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
custom_guardrail = CustomGuardrail(
|
||||
guardrail_name="test_guardrail",
|
||||
default_on=False,
|
||||
event_hook=GuardrailEventHooks.pre_call
|
||||
)
|
||||
|
||||
# Test with guardrails in metadata
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"metadata": {
|
||||
"guardrails": ["test_guardrail"]
|
||||
}
|
||||
}
|
||||
|
||||
result = custom_guardrail.should_run_guardrail(
|
||||
data=data, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
|
||||
assert result is True
|
||||
|
||||
def test_should_run_guardrail_with_root_level_guardrails(self):
|
||||
"""Test that should_run_guardrail works with root level guardrails"""
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
custom_guardrail = CustomGuardrail(
|
||||
guardrail_name="test_guardrail",
|
||||
default_on=False,
|
||||
event_hook=GuardrailEventHooks.pre_call
|
||||
)
|
||||
|
||||
# Test with guardrails at root level
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"guardrails": ["test_guardrail"]
|
||||
}
|
||||
|
||||
result = custom_guardrail.should_run_guardrail(
|
||||
data=data, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
|
||||
assert result is True
|
||||
|
||||
|
||||
def test_should_run_guardrail_no_matching_guardrail(self):
|
||||
"""Test that should_run_guardrail returns False when guardrail name doesn't match"""
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
custom_guardrail = CustomGuardrail(
|
||||
guardrail_name="test_guardrail",
|
||||
default_on=False,
|
||||
event_hook=GuardrailEventHooks.pre_call
|
||||
)
|
||||
|
||||
# Test with different guardrail name
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"litellm_metadata": {
|
||||
"guardrails": ["different_guardrail"]
|
||||
}
|
||||
}
|
||||
|
||||
result = custom_guardrail.should_run_guardrail(
|
||||
data=data, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
|
||||
assert result is False
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue