This commit is contained in:
Vrajesh Sulakhe 2026-10-04 10:33:32 -07:00 • committed by GitHub
commit db986d9b49
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 129 additions and 8 deletions

View file

@ -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],

View file

@ -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,9 +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
if "/" not in full_model_name:
full_model_name = f"{custom_llm_provider}/{full_model_name}"
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,
)
optional_params_for_followup: Final = {**optional_params, **patch.optional_params}
if patch.tools is not None:

View file

@ -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,9 +5688,12 @@ class BaseLLMHTTPHandler:
if patch.messages is None:
raise ValueError("Agentic loop plan missing patched messages")
full_model_name = patch.model or model
if "/" not in full_model_name:
full_model_name = f"{custom_llm_provider}/{full_model_name}"
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,
)
optional_params_for_followup: Final = dict(optional_params)
optional_params_for_followup.update(patch.optional_params)

View file

@ -0,0 +1,92 @@
from typing import Final, Literal
from unittest.mock import AsyncMock, patch
import pytest
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.chat_completion_agentic_loop import (
_execute_chat_completion_agentic_plan,
)
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.types.integrations.custom_logger import (
AgenticLoopPlan,
AgenticLoopRequestPatch,
)
from litellm.types.utils import ModelResponse
@pytest.mark.asyncio
@pytest.mark.parametrize("execution_path", ("http", "sdk"))
@pytest.mark.parametrize(
("request_model", "patch_model", "custom_llm_provider", "expected_model"),
(
("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",
),
)
async def test_agentic_followup_preserves_provider_prefix(
execution_path: Literal["http", "sdk"],
request_model: str,
patch_model: str | None,
custom_llm_provider: str,
expected_model: str,
) -> None:
plan: Final = AgenticLoopPlan(
run_agentic_loop=True,
request_patch=AgenticLoopRequestPatch(
model=patch_model,
messages=[{"role": "user", "content": "Continue"}],
),
)
followup: Final = AsyncMock(wraps=litellm.acompletion)
with patch("litellm.acompletion", new=followup):
response: Final = (
await BaseLLMHTTPHandler()._execute_chat_completion_agentic_plan(
plan=plan,
model=request_model,
messages=[],
optional_params={"mock_response": "Follow-up complete"},
kwargs={"custom_llm_provider": custom_llm_provider},
custom_llm_provider=custom_llm_provider,
depth=0,
max_loops=3,
fingerprints=[],
fingerprint="followup",
)
if execution_path == "http"
else await _execute_chat_completion_agentic_plan(
plan=plan,
callback=CustomLogger(),
model=request_model,
optional_params={"mock_response": "Follow-up complete"},
kwargs={"custom_llm_provider": custom_llm_provider},
logging_obj=None,
custom_llm_provider=custom_llm_provider,
depth=0,
max_loops=3,
fingerprints=[],
fingerprint="followup",
)
)
followup.assert_awaited_once()
assert followup.await_args is not None
assert followup.await_args.kwargs["model"] == expected_model
assert isinstance(response, ModelResponse)
assert response.model == expected_model.split("/", 1)[1]