mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
feat: support new 'litellm_agent' model provider
This commit is contained in:
parent
32ce793587
commit
a271cf28ad
7 changed files with 202 additions and 13 deletions
|
|
@ -98,6 +98,7 @@ _custom_logger_compatible_callbacks_literal = Literal[
|
|||
"openmeter",
|
||||
"logfire",
|
||||
"literalai",
|
||||
"litellm_agent",
|
||||
"dynamic_rate_limiter",
|
||||
"dynamic_rate_limiter_v3",
|
||||
"langsmith",
|
||||
|
|
|
|||
5
litellm/integrations/litellm_agent/__init__.py
Normal file
5
litellm/integrations/litellm_agent/__init__.py
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
"""LiteLLM Agent integration - model name resolver for litellm_agent/ prefix."""
|
||||
|
||||
from .litellm_agent_model_resolver import LiteLLMAgentModelResolver
|
||||
|
||||
__all__ = ["LiteLLMAgentModelResolver"]
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue