[Bug Fix] Fix model_group tracked for /v1/messages and /moderations (#11933)

* fixes _get_router_metadata_variable_name

* fixes _update_kwargs_before_fallbacks

* test_anthropic_messages_litellm_router_non_streaming_with_logging

* test_moderations_api_logging

* fix _pass_through_moderation_endpoint_factory
This commit is contained in:
Ishaan Jaff 2025-06-20 14:51:50 -07:00 • committed by GitHub
parent 0d486120bc
commit 931b2e4875
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 53 additions and 9 deletions

View file

@ -2454,6 +2454,14 @@ class Router:
The response from the handler function
"""
handler_name = original_function.__name__
function_name = "_ageneric_api_call_with_fallbacks"
self._update_kwargs_before_fallbacks(
model=model,
kwargs=kwargs,
metadata_variable_name = _get_router_metadata_variable_name(
function_name=function_name
)
)
try:
verbose_router_logger.debug(
f"Inside _ageneric_api_call() - handler: {handler_name}, model: {model}; kwargs: {kwargs}"
@ -2467,7 +2475,7 @@ class Router:
)
self._update_kwargs_with_deployment(
deployment=deployment, kwargs=kwargs, function_name="generic_api_call"
deployment=deployment, kwargs=kwargs, function_name=function_name
)
data = deployment["litellm_params"].copy()
@ -3180,6 +3188,11 @@ class Router:
original_function: Callable,
**kwargs,
):
# update kwargs with model_group
self._update_kwargs_before_fallbacks(
model=kwargs.get("model", ""),
kwargs=kwargs,
)
if kwargs.get("model") and self.get_model_list(model_name=kwargs["model"]):
deployment = await self.async_get_available_deployment(
model=kwargs["model"],

View file

@ -79,7 +79,7 @@ def _get_router_metadata_variable_name(function_name: Optional[str]) -> str:
For ALL other endpoints we call this "metadata
"""
ROUTER_METHODS_USING_LITELLM_METADATA = set(
["batch", "generic_api_call", "_acreate_batch", "file"]
["batch", "generic_api_call", "_acreate_batch", "file", "_ageneric_api_call_with_fallbacks"]
)
if function_name and any(
method in function_name for method in ROUTER_METHODS_USING_LITELLM_METADATA

View file

@ -17,6 +17,7 @@ sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
import litellm
from litellm.router import Router
import asyncio
from typing import Optional
from litellm.types.utils import StandardLoggingPayload, Usage, ModelInfoBase
@ -41,7 +42,8 @@ class TestCustomLogger(CustomLogger):
@pytest.mark.asyncio
@pytest.mark.parametrize("model", [
None,
"omni-moderation-latest"
"omni-moderation-latest",
"router-internal-moderation-model"
])
async def test_moderations_api_logging(model):
@ -51,12 +53,31 @@ async def test_moderations_api_logging(model):
custom_logger = TestCustomLogger()
litellm.logging_callback_manager.add_litellm_callback(custom_logger)
input_content = "Hello, how are you?"
response = await litellm.amoderation(
input=input_content,
model=model,
MODEL_GROUP = "internal-moderation-model"
router = Router(
model_list=[
{
"model_name": MODEL_GROUP,
"litellm_params": {
"model": "openai/omni-moderation-latest",
},
}
]
)
input_content = "Hello, how are you?"
if model == "router-internal-moderation-model":
response = await router.amoderation(
input=input_content,
model=MODEL_GROUP,
)
else:
response = await litellm.amoderation(
input=input_content,
model=model,
)
print("response", json.dumps(response, indent=4, default=str))
await asyncio.sleep(2)
@ -76,3 +97,8 @@ async def test_moderations_api_logging(model):
# assert the logged response == response user received client side
assert dict(standard_logging_payload["response"]) == response.model_dump()
# if router used, validate model_group is logged as expected
if model == "router-internal-moderation-model":
assert standard_logging_payload["model_group"] == MODEL_GROUP

View file

@ -222,10 +222,11 @@ async def test_anthropic_messages_litellm_router_non_streaming_with_logging():
test_custom_logger = TestCustomLogger()
litellm.callbacks = [test_custom_logger]
litellm._turn_on_debug()
MODEL_GROUP = "claude-special-alias"
router = Router(
model_list=[
{
"model_name": "claude-special-alias",
"model_name": MODEL_GROUP,
"litellm_params": {
"model": "claude-3-haiku-20240307",
"api_key": os.getenv("ANTHROPIC_API_KEY"),
@ -240,7 +241,7 @@ async def test_anthropic_messages_litellm_router_non_streaming_with_logging():
# Call the handler
response = await router.aanthropic_messages(
messages=messages,
model="claude-special-alias",
model=MODEL_GROUP,
max_tokens=100,
)
@ -252,6 +253,7 @@ async def test_anthropic_messages_litellm_router_non_streaming_with_logging():
await asyncio.sleep(1)
assert test_custom_logger.logged_standard_logging_payload is not None, "Logging payload should not be None"
print("tracked standard logging payload", json.dumps(test_custom_logger.logged_standard_logging_payload, indent=4, default=str))
assert test_custom_logger.logged_standard_logging_payload["messages"] == messages
assert test_custom_logger.logged_standard_logging_payload["response"] is not None
assert (
@ -270,6 +272,9 @@ async def test_anthropic_messages_litellm_router_non_streaming_with_logging():
== response["usage"]["output_tokens"]
)
# assert model_group
assert test_custom_logger.logged_standard_logging_payload["model_group"] == MODEL_GROUP
@pytest.mark.asyncio
async def test_anthropic_messages_litellm_router_streaming_with_logging():