feat: support new 'litellm_agent' model provider

This commit is contained in:
Krrish Dholakia 2026-02-20 21:09:01 -08:00
parent 32ce793587
commit a271cf28ad
7 changed files with 202 additions and 13 deletions

View file

@ -98,6 +98,7 @@ _custom_logger_compatible_callbacks_literal = Literal[
"openmeter",
"logfire",
"literalai",
"litellm_agent",
"dynamic_rate_limiter",
"dynamic_rate_limiter_v3",
"langsmith",

View file

@ -0,0 +1,5 @@
"""LiteLLM Agent integration - model name resolver for litellm_agent/ prefix."""
from .litellm_agent_model_resolver import LiteLLMAgentModelResolver
__all__ = ["LiteLLMAgentModelResolver"]

View file

@ -0,0 +1,79 @@
"""
Hook for LiteLLM that strips the litellm_agent/ prefix from model names.
When model is litellm_agent/gpt-3.5-turbo, this hook replaces it with gpt-3.5-turbo
before the completion call, similar to langfuse/model resolution.
"""
from typing import Dict, List, Optional, Tuple
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.llms.openai import AllMessageValues
from litellm.types.prompts.init_prompts import PromptSpec
from litellm.types.utils import StandardCallbackDynamicParams
LITELLM_AGENT_PREFIX = "litellm_agent/"
class LiteLLMAgentModelResolver(CustomLogger):
"""
CustomLogger that strips litellm_agent/ prefix from model names.
Enables model configs like litellm_agent/gpt-3.5-turbo to resolve to gpt-3.5-turbo.
"""
def get_chat_completion_prompt(
self,
model: str,
messages: List[AllMessageValues],
non_default_params: dict,
prompt_id: Optional[str],
prompt_variables: Optional[dict],
dynamic_callback_params: StandardCallbackDynamicParams,
prompt_spec: Optional[PromptSpec] = None,
prompt_label: Optional[str] = None,
prompt_version: Optional[int] = None,
ignore_prompt_manager_model: Optional[bool] = False,
ignore_prompt_manager_optional_params: Optional[bool] = False,
) -> Tuple[str, List[AllMessageValues], dict]:
"""
Strip litellm_agent/ prefix from model name.
Returns:
(resolved_model, messages, non_default_params)
"""
if ignore_prompt_manager_model:
return model, messages, non_default_params
resolved_model = model.replace(LITELLM_AGENT_PREFIX, "", 1)
return resolved_model, messages, non_default_params
async def async_get_chat_completion_prompt(
self,
model: str,
messages: List[AllMessageValues],
non_default_params: dict,
prompt_id: Optional[str],
prompt_variables: Optional[dict],
dynamic_callback_params: StandardCallbackDynamicParams,
litellm_logging_obj: object,
prompt_spec: Optional[PromptSpec] = None,
tools: Optional[List[Dict]] = None,
prompt_label: Optional[str] = None,
prompt_version: Optional[int] = None,
ignore_prompt_manager_model: Optional[bool] = False,
ignore_prompt_manager_optional_params: Optional[bool] = False,
) -> Tuple[str, List[AllMessageValues], dict]:
"""Async delegate to get_chat_completion_prompt."""
return self.get_chat_completion_prompt(
model=model,
messages=messages,
non_default_params=non_default_params,
prompt_id=prompt_id,
prompt_variables=prompt_variables,
dynamic_callback_params=dynamic_callback_params,
prompt_spec=prompt_spec,
prompt_label=prompt_label,
prompt_version=prompt_version,
ignore_prompt_manager_model=ignore_prompt_manager_model,
ignore_prompt_manager_optional_params=ignore_prompt_manager_optional_params,
)

View file

@ -18,11 +18,11 @@ from litellm.integrations.azure_storage.azure_storage import AzureBlobStorageLog
from litellm.integrations.bitbucket import BitBucketPromptManager
from litellm.integrations.braintrust_logging import BraintrustLogger
from litellm.integrations.cloudzero.cloudzero import CloudZeroLogger
from litellm.integrations.focus.focus_logger import FocusLogger
from litellm.integrations.datadog.datadog import DataDogLogger
from litellm.integrations.datadog.datadog_llm_obs import DataDogLLMObsLogger
from litellm.integrations.deepeval import DeepEvalLogger
from litellm.integrations.dotprompt import DotpromptManager
from litellm.integrations.focus.focus_logger import FocusLogger
from litellm.integrations.galileo import GalileoObserve
from litellm.integrations.gcs_bucket.gcs_bucket import GCSBucketLogger
from litellm.integrations.gcs_pubsub.pub_sub import GcsPubSubLogger
@ -33,6 +33,7 @@ from litellm.integrations.langfuse.langfuse_prompt_management import (
LangfusePromptManagement,
)
from litellm.integrations.langsmith import LangsmithLogger
from litellm.integrations.litellm_agent import LiteLLMAgentModelResolver
from litellm.integrations.literal_ai import LiteralAILogger
from litellm.integrations.mlflow import MlflowLogger
from litellm.integrations.openmeter import OpenMeterLogger
@ -61,6 +62,7 @@ class CustomLoggerRegistry:
"galileo": GalileoObserve,
"langsmith": LangsmithLogger,
"literalai": LiteralAILogger,
"litellm_agent": LiteLLMAgentModelResolver,
"prometheus": PrometheusLogger,
"datadog": DataDogLogger,
"datadog_llm_observability": DataDogLLMObsLogger,

View file

@ -147,6 +147,7 @@ from ..integrations.langfuse.langfuse import LangFuseLogger
from ..integrations.langfuse.langfuse_handler import LangFuseHandler
from ..integrations.langfuse.langfuse_prompt_management import LangfusePromptManagement
from ..integrations.langsmith import LangsmithLogger
from ..integrations.litellm_agent import LiteLLMAgentModelResolver
from ..integrations.literal_ai import LiteralAILogger
from ..integrations.logfire_logger import LogfireLevel, LogfireLogger
from ..integrations.lunary import LunaryLogger
@ -587,6 +588,11 @@ class Logging(LiteLLMLoggingBaseClass):
if prompt_id:
return True
# Check if model uses litellm_agent prefix (model replacement without prompt_id)
model = non_default_params.get("model", "")
if isinstance(model, str) and model.startswith("litellm_agent/"):
return True
if self._should_run_prompt_management_hooks_without_prompt_id(
non_default_params=non_default_params,
tools=tools,
@ -3579,6 +3585,14 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
_literalai_logger = LiteralAILogger()
_in_memory_loggers.append(_literalai_logger)
return _literalai_logger # type: ignore
elif logging_integration == "litellm_agent":
for callback in _in_memory_loggers:
if isinstance(callback, LiteLLMAgentModelResolver):
return callback # type: ignore
_litellm_agent_resolver = LiteLLMAgentModelResolver()
_in_memory_loggers.append(_litellm_agent_resolver)
return _litellm_agent_resolver # type: ignore
elif logging_integration == "prometheus":
PrometheusLogger = _get_cached_prometheus_logger()
@ -4133,6 +4147,10 @@ def get_custom_logger_compatible_class( # noqa: PLR0915
for callback in _in_memory_loggers:
if isinstance(callback, LiteralAILogger):
return callback
elif logging_integration == "litellm_agent":
for callback in _in_memory_loggers:
if isinstance(callback, LiteLLMAgentModelResolver):
return callback
elif logging_integration == "prometheus":
PrometheusLogger = _get_cached_prometheus_logger()
for callback in _in_memory_loggers:

View file

@ -162,11 +162,7 @@ from litellm.types.utils import (
)
from litellm.types.utils import ModelInfo
from litellm.types.utils import ModelInfo as ModelMapInfo
from litellm.types.utils import (
ModelResponseStream,
StandardLoggingPayload,
Usage,
)
from litellm.types.utils import ModelResponseStream, StandardLoggingPayload, Usage
from litellm.utils import (
CustomStreamWrapper,
EmbeddingResponse,
@ -1956,12 +1952,11 @@ class Router:
When both have tools, concatenate them (deployment tools first, then request tools).
tool_choice: use request value if provided, else deployment's.
"""
dep_params = deployment.get("litellm_params", {}) or {}
dep_params = (
dep_params.model_dump(exclude_none=True)
if hasattr(dep_params, "model_dump")
else dep_params
)
dep_params_raw = deployment.get("litellm_params", {}) or {}
if isinstance(dep_params_raw, dict):
dep_params = dep_params_raw
else:
dep_params = dep_params_raw.model_dump(exclude_none=True)
dep_tools = dep_params.get("tools") or []
req_tools = kwargs.get("tools") or []
if dep_tools or req_tools:
@ -2532,6 +2527,12 @@ class Router:
litellm_model = data.get("model", None)
# litellm_agent/ prefix only strips the model name, no prompt_id needed
is_litellm_agent_model = (
isinstance(litellm_model, str)
and litellm_model.startswith("litellm_agent/")
)
prompt_id = kwargs.get("prompt_id") or prompt_management_deployment[
"litellm_params"
].get("prompt_id", None)
@ -2544,7 +2545,9 @@ class Router:
"litellm_params"
].get("prompt_label", None)
if prompt_id is None or not isinstance(prompt_id, str):
if not is_litellm_agent_model and (
prompt_id is None or not isinstance(prompt_id, str)
):
raise ValueError(
f"Prompt ID is not set or not a string. Got={prompt_id}, type={type(prompt_id)}"
)

View file

@ -0,0 +1,81 @@
"""Unit tests for LiteLLMAgentModelResolver - litellm_agent/ prefix model resolution."""
from unittest.mock import MagicMock
import pytest
from litellm.integrations.litellm_agent import LiteLLMAgentModelResolver
class TestLiteLLMAgentModelResolver:
def test_get_chat_completion_prompt_strips_prefix(self):
"""Verify get_chat_completion_prompt strips litellm_agent/ prefix from model."""
resolver = LiteLLMAgentModelResolver()
messages = [{"role": "user", "content": "Hello"}]
resolved_model, out_messages, out_params = resolver.get_chat_completion_prompt(
model="litellm_agent/gpt-3.5-turbo",
messages=messages,
non_default_params={},
prompt_id=None,
prompt_variables=None,
dynamic_callback_params={},
)
assert resolved_model == "gpt-3.5-turbo"
assert out_messages == messages
assert out_params == {}
def test_get_chat_completion_prompt_preserves_rest_of_model(self):
"""Verify model name after prefix is preserved (e.g. openai/gpt-3.5-turbo)."""
resolver = LiteLLMAgentModelResolver()
messages = [{"role": "user", "content": "Test"}]
resolved_model, _, _ = resolver.get_chat_completion_prompt(
model="litellm_agent/openai/gpt-3.5-turbo",
messages=messages,
non_default_params={},
prompt_id=None,
prompt_variables=None,
dynamic_callback_params={},
)
assert resolved_model == "openai/gpt-3.5-turbo"
def test_get_chat_completion_prompt_respects_ignore_prompt_manager_model(self):
"""Verify model is unchanged when ignore_prompt_manager_model is True."""
resolver = LiteLLMAgentModelResolver()
messages = [{"role": "user", "content": "Hello"}]
resolved_model, _, _ = resolver.get_chat_completion_prompt(
model="litellm_agent/gpt-3.5-turbo",
messages=messages,
non_default_params={},
prompt_id=None,
prompt_variables=None,
dynamic_callback_params={},
ignore_prompt_manager_model=True,
)
assert resolved_model == "litellm_agent/gpt-3.5-turbo"
@pytest.mark.asyncio
async def test_async_get_chat_completion_prompt_strips_prefix(self):
"""Verify async_get_chat_completion_prompt strips prefix."""
resolver = LiteLLMAgentModelResolver()
messages = [{"role": "user", "content": "Hello"}]
resolved_model, out_messages, _ = (
await resolver.async_get_chat_completion_prompt(
model="litellm_agent/gpt-3.5-turbo",
messages=messages,
non_default_params={},
prompt_id=None,
prompt_variables=None,
dynamic_callback_params={},
litellm_logging_obj=MagicMock(),
)
)
assert resolved_model == "gpt-3.5-turbo"
assert out_messages == messages