mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(agentic): keep provider prefix for org-named models on follow-up calls
This commit is contained in:
parent
efa37234ec
commit
dba346d816
4 changed files with 57 additions and 29 deletions
|
|
@ -4,6 +4,23 @@ from types import MappingProxyType
|
|||
from typing import Final
|
||||
|
||||
|
||||
def resolve_agentic_followup_model(
|
||||
*,
|
||||
request_model: str,
|
||||
patch_model: str | None,
|
||||
custom_llm_provider: str,
|
||||
known_providers: Collection[str],
|
||||
) -> str:
|
||||
"""The request model is already provider-stripped, so a leading "openai/" there is an org name.
|
||||
Only a callback-supplied model may carry its own provider prefix and switch providers"""
|
||||
model: Final = patch_model or request_model
|
||||
if not custom_llm_provider or model.startswith(f"{custom_llm_provider}/"):
|
||||
return model
|
||||
if patch_model and "/" in patch_model and patch_model.split("/", 1)[0] in known_providers:
|
||||
return patch_model
|
||||
return f"{custom_llm_provider}/{model}"
|
||||
|
||||
|
||||
def build_agentic_followup_kwargs(
|
||||
*,
|
||||
request_kwargs: Mapping[str, object],
|
||||
|
|
|
|||
|
|
@ -8,7 +8,10 @@ from typing import Final, cast
|
|||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.agentic_followup_kwargs import build_agentic_followup_kwargs
|
||||
from litellm.litellm_core_utils.agentic_followup_kwargs import (
|
||||
build_agentic_followup_kwargs,
|
||||
resolve_agentic_followup_model,
|
||||
)
|
||||
from litellm.litellm_core_utils.agentic_loop_settings import (
|
||||
DEFAULT_MAX_AGENTIC_LOOPS,
|
||||
validated_max_agentic_loops,
|
||||
|
|
@ -170,14 +173,12 @@ async def _execute_chat_completion_agentic_plan(
|
|||
if patch.messages is None:
|
||||
raise ValueError("Agentic loop plan missing patched messages")
|
||||
|
||||
full_model_name = patch.model or model
|
||||
known_providers: Final = getattr(litellm, "provider_list", [])
|
||||
has_provider_prefix: Final = (
|
||||
full_model_name.startswith(f"{custom_llm_provider}/")
|
||||
or ("/" in full_model_name and full_model_name.split("/", 1)[0] in known_providers)
|
||||
full_model_name: Final = resolve_agentic_followup_model(
|
||||
request_model=model,
|
||||
patch_model=patch.model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
known_providers=litellm.provider_list,
|
||||
)
|
||||
if custom_llm_provider and not has_provider_prefix:
|
||||
full_model_name = f"{custom_llm_provider}/{full_model_name}"
|
||||
|
||||
optional_params_for_followup: Final = {**optional_params, **patch.optional_params}
|
||||
if patch.tools is not None:
|
||||
|
|
|
|||
|
|
@ -37,7 +37,10 @@ from litellm._logging import _redact_string, verbose_logger
|
|||
from litellm.anthropic_beta_headers_manager import update_headers_with_filtered_beta
|
||||
from litellm.constants import MAX_FILE_LIST_LIMIT, REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
|
||||
from litellm.files.types import FileContentStreamingResult
|
||||
from litellm.litellm_core_utils.agentic_followup_kwargs import build_agentic_followup_kwargs
|
||||
from litellm.litellm_core_utils.agentic_followup_kwargs import (
|
||||
build_agentic_followup_kwargs,
|
||||
resolve_agentic_followup_model,
|
||||
)
|
||||
from litellm.litellm_core_utils.agentic_loop_settings import (
|
||||
DEFAULT_MAX_AGENTIC_LOOPS,
|
||||
validated_max_agentic_loops,
|
||||
|
|
@ -5685,14 +5688,12 @@ class BaseLLMHTTPHandler:
|
|||
if patch.messages is None:
|
||||
raise ValueError("Agentic loop plan missing patched messages")
|
||||
|
||||
full_model_name = patch.model or model
|
||||
known_providers: Final = getattr(litellm, "provider_list", [])
|
||||
has_provider_prefix: Final = (
|
||||
full_model_name.startswith(f"{custom_llm_provider}/")
|
||||
or ("/" in full_model_name and full_model_name.split("/", 1)[0] in known_providers)
|
||||
full_model_name: Final = resolve_agentic_followup_model(
|
||||
request_model=model,
|
||||
patch_model=patch.model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
known_providers=litellm.provider_list,
|
||||
)
|
||||
if custom_llm_provider and not has_provider_prefix:
|
||||
full_model_name = f"{custom_llm_provider}/{full_model_name}"
|
||||
|
||||
optional_params_for_followup: Final = dict(optional_params)
|
||||
optional_params_for_followup.update(patch.optional_params)
|
||||
|
|
|
|||
|
|
@ -18,29 +18,38 @@ from litellm.types.utils import ModelResponse
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("execution_path", ("http", "sdk"))
|
||||
@pytest.mark.parametrize("use_patch_model", (False, True), ids=("original-model", "patch-model"))
|
||||
@pytest.mark.parametrize(
|
||||
("model_name", "custom_llm_provider", "expected_model"),
|
||||
("request_model", "patch_model", "custom_llm_provider", "expected_model"),
|
||||
(
|
||||
("zai-org/GLM-5.3-Flash", "hosted_vllm", "hosted_vllm/zai-org/GLM-5.3-Flash"),
|
||||
("hosted_vllm/zai-org/GLM-5.3-Flash", "hosted_vllm", "hosted_vllm/zai-org/GLM-5.3-Flash"),
|
||||
("hosted_vllm/zai-org/GLM-5.3-Flash", "", "hosted_vllm/zai-org/GLM-5.3-Flash"),
|
||||
("openai/gpt-4o", "hosted_vllm", "openai/gpt-4o"),
|
||||
("zai-org/GLM-5.3-Flash", None, "hosted_vllm", "hosted_vllm/zai-org/GLM-5.3-Flash"),
|
||||
("hosted_vllm/zai-org/GLM-5.3-Flash", None, "hosted_vllm", "hosted_vllm/zai-org/GLM-5.3-Flash"),
|
||||
("hosted_vllm/zai-org/GLM-5.3-Flash", None, "", "hosted_vllm/zai-org/GLM-5.3-Flash"),
|
||||
("openai/whisper-large-v3", None, "hosted_vllm", "hosted_vllm/openai/whisper-large-v3"),
|
||||
("original-model", "zai-org/GLM-5.3-Flash", "hosted_vllm", "hosted_vllm/zai-org/GLM-5.3-Flash"),
|
||||
("original-model", "hosted_vllm/zai-org/GLM-5.3-Flash", "hosted_vllm", "hosted_vllm/zai-org/GLM-5.3-Flash"),
|
||||
("original-model", "openai/gpt-4o", "hosted_vllm", "openai/gpt-4o"),
|
||||
),
|
||||
ids=(
|
||||
"organization-model",
|
||||
"already-prefixed",
|
||||
"no-provider",
|
||||
"organization-named-after-provider",
|
||||
"patch-organization-model",
|
||||
"patch-already-prefixed",
|
||||
"patch-cross-provider",
|
||||
),
|
||||
ids=("organization-model", "already-prefixed", "no-provider", "cross-provider"),
|
||||
)
|
||||
async def test_agentic_followup_preserves_provider_prefix(
|
||||
execution_path: Literal["http", "sdk"],
|
||||
use_patch_model: bool,
|
||||
model_name: str,
|
||||
request_model: str,
|
||||
patch_model: str | None,
|
||||
custom_llm_provider: str,
|
||||
expected_model: str,
|
||||
) -> None:
|
||||
model: Final = "original-model" if use_patch_model else model_name
|
||||
plan: Final = AgenticLoopPlan(
|
||||
run_agentic_loop=True,
|
||||
request_patch=AgenticLoopRequestPatch(
|
||||
model=model_name if use_patch_model else None,
|
||||
model=patch_model,
|
||||
messages=[{"role": "user", "content": "Continue"}],
|
||||
),
|
||||
)
|
||||
|
|
@ -50,7 +59,7 @@ async def test_agentic_followup_preserves_provider_prefix(
|
|||
response: Final = (
|
||||
await BaseLLMHTTPHandler()._execute_chat_completion_agentic_plan(
|
||||
plan=plan,
|
||||
model=model,
|
||||
model=request_model,
|
||||
messages=[],
|
||||
optional_params={"mock_response": "Follow-up complete"},
|
||||
kwargs={"custom_llm_provider": custom_llm_provider},
|
||||
|
|
@ -64,7 +73,7 @@ async def test_agentic_followup_preserves_provider_prefix(
|
|||
else await _execute_chat_completion_agentic_plan(
|
||||
plan=plan,
|
||||
callback=CustomLogger(),
|
||||
model=model,
|
||||
model=request_model,
|
||||
optional_params={"mock_response": "Follow-up complete"},
|
||||
kwargs={"custom_llm_provider": custom_llm_provider},
|
||||
logging_obj=None,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue