diff --git a/litellm/__init__.py b/litellm/__init__.py index a994db85b11..97f36a9b00f 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -98,6 +98,7 @@ _custom_logger_compatible_callbacks_literal = Literal[ "openmeter", "logfire", "literalai", + "litellm_agent", "dynamic_rate_limiter", "dynamic_rate_limiter_v3", "langsmith", diff --git a/litellm/integrations/litellm_agent/__init__.py b/litellm/integrations/litellm_agent/__init__.py new file mode 100644 index 00000000000..f09434080ed --- /dev/null +++ b/litellm/integrations/litellm_agent/__init__.py @@ -0,0 +1,5 @@ +"""LiteLLM Agent integration - model name resolver for litellm_agent/ prefix.""" + +from .litellm_agent_model_resolver import LiteLLMAgentModelResolver + +__all__ = ["LiteLLMAgentModelResolver"] diff --git a/litellm/integrations/litellm_agent/litellm_agent_model_resolver.py b/litellm/integrations/litellm_agent/litellm_agent_model_resolver.py new file mode 100644 index 00000000000..85d209da5b1 --- /dev/null +++ b/litellm/integrations/litellm_agent/litellm_agent_model_resolver.py @@ -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, + ) diff --git a/litellm/litellm_core_utils/custom_logger_registry.py b/litellm/litellm_core_utils/custom_logger_registry.py index a3c25ab65e9..fc73701ea9d 100644 --- a/litellm/litellm_core_utils/custom_logger_registry.py +++ b/litellm/litellm_core_utils/custom_logger_registry.py @@ -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, diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 76e37010109..b961f7470ae 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -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: diff --git a/litellm/router.py b/litellm/router.py index 69e3e994dcd..a24532762fd 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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)}" ) diff --git a/tests/test_litellm/integrations/litellm_agent/test_litellm_agent_model_resolver.py b/tests/test_litellm/integrations/litellm_agent/test_litellm_agent_model_resolver.py new file mode 100644 index 00000000000..c91da9a4b89 --- /dev/null +++ b/tests/test_litellm/integrations/litellm_agent/test_litellm_agent_model_resolver.py @@ -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