mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix: allow using credentials with amoderation (#10723)
This commit is contained in:
parent
7210b713dc
commit
214a427038
2 changed files with 40 additions and 0 deletions
|
|
@ -3136,6 +3136,13 @@ class Router:
|
|||
request_kwargs=kwargs,
|
||||
)
|
||||
kwargs["model"] = deployment["litellm_params"]["model"]
|
||||
data = deployment["litellm_params"].copy()
|
||||
self._update_kwargs_with_deployment(
|
||||
deployment=deployment,
|
||||
kwargs=kwargs,
|
||||
)
|
||||
kwargs.update(data)
|
||||
|
||||
return await original_function(**kwargs)
|
||||
|
||||
def factory_function(
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import copy
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
|
@ -148,6 +149,7 @@ async def test_router_acreate_file():
|
|||
# assert that the mock_acreate_file was called twice
|
||||
assert mock_acreate_file.call_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_acreate_file_with_jsonl():
|
||||
"""
|
||||
|
|
@ -263,3 +265,34 @@ async def test_router_async_get_healthy_deployments():
|
|||
assert result[0]["model_name"] == "gpt-3.5-turbo"
|
||||
assert result[0]["litellm_params"]["model"] == "gpt-3.5-turbo"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.amoderation")
|
||||
async def test_router_amoderation_with_credential_name(mock_amoderation):
|
||||
"""
|
||||
Test that router.amoderation passes litellm_credential_name to the underlying litellm.amoderation call
|
||||
"""
|
||||
mock_amoderation.return_value = AsyncMock()
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "text-moderation-stable",
|
||||
"litellm_params": {
|
||||
"model": "text-moderation-stable",
|
||||
"litellm_credential_name": "my-custom-auth",
|
||||
},
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
await router.amoderation(input="I love everyone!", model="text-moderation-stable")
|
||||
|
||||
mock_amoderation.assert_called_once()
|
||||
call_kwargs = mock_amoderation.call_args[1] # Get the kwargs of the call
|
||||
print(
|
||||
"call kwargs for router.amoderation=",
|
||||
json.dumps(call_kwargs, indent=4, default=str),
|
||||
)
|
||||
assert call_kwargs["litellm_credential_name"] == "my-custom-auth"
|
||||
assert call_kwargs["model"] == "text-moderation-stable"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue