mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
refactor(agents): share inference model resolution with admission policy
This commit is contained in:
parent
2d78a3aeec
commit
bb7830a63d
7 changed files with 402 additions and 34 deletions
145
litellm/proxy/agent_endpoints/auth/managed_authorization.py
Normal file
145
litellm/proxy/agent_endpoints/auth/managed_authorization.py
Normal file
|
|
@ -0,0 +1,145 @@
|
|||
from collections.abc import Mapping
|
||||
from itertools import product
|
||||
from typing import Final
|
||||
|
||||
from litellm.proxy._types import LiteLLMRoutes
|
||||
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityFailure, ManagedAgentContext
|
||||
|
||||
_MANAGED_REALTIME_ROUTES: Final = frozenset(("/realtime", "/v1/realtime", "/openai/v1/realtime"))
|
||||
_MANAGED_MODEL_ROUTES: Final = frozenset(
|
||||
f"{prefix}/{operation}"
|
||||
for prefix, operation in product(
|
||||
("", "/v1"),
|
||||
(
|
||||
"chat/completions",
|
||||
"completions",
|
||||
"embeddings",
|
||||
"responses",
|
||||
"messages",
|
||||
"messages/count_tokens",
|
||||
"images/generations",
|
||||
"images/edits",
|
||||
"audio/transcriptions",
|
||||
"audio/speech",
|
||||
"moderations",
|
||||
"rerank",
|
||||
"ocr",
|
||||
),
|
||||
)
|
||||
) | frozenset(
|
||||
(
|
||||
"/openai/v1/responses",
|
||||
"/v2/rerank",
|
||||
"/claude_code_gateway/v1/messages",
|
||||
"/claude_code_gateway/v1/messages/count_tokens",
|
||||
"/cursor/chat/completions",
|
||||
)
|
||||
)
|
||||
_MANAGED_MODEL_PATHS: Final = (
|
||||
"/engines/{model:path}/chat/completions",
|
||||
"/engines/{model:path}/completions",
|
||||
"/engines/{model:path}/embeddings",
|
||||
"/openai/deployments/{model:path}/chat/completions",
|
||||
"/openai/deployments/{model:path}/completions",
|
||||
"/openai/deployments/{model:path}/embeddings",
|
||||
"/openai/deployments/{model:path}/images/generations",
|
||||
"/openai/deployments/{model:path}/images/edits",
|
||||
"/v1beta/models/{model_name:path}:countTokens",
|
||||
"/v1beta/models/{model_name:path}:generateContent",
|
||||
"/v1beta/models/{model_name:path}:streamGenerateContent",
|
||||
"/models/{model_name:path}:countTokens",
|
||||
"/models/{model_name:path}:generateContent",
|
||||
"/models/{model_name:path}:streamGenerateContent",
|
||||
)
|
||||
_MANAGED_MCP_ROUTES: Final = tuple(
|
||||
route for route in LiteLLMRoutes.mcp_inference_routes.value if route not in ("/token", "/introspect")
|
||||
)
|
||||
|
||||
|
||||
def managed_agent_route_allowed(route: str, method: str | None) -> bool:
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
|
||||
if route in ("/agents", "/v1/agents"):
|
||||
return method in (None, "GET", "HEAD")
|
||||
if route in _MANAGED_REALTIME_ROUTES:
|
||||
return method in (None, "GET")
|
||||
if route in _MANAGED_MODEL_ROUTES or RouteChecks.check_route_access(route, _MANAGED_MODEL_PATHS):
|
||||
return method in (None, "POST")
|
||||
return RouteChecks.check_route_access(route, _MANAGED_MCP_ROUTES) or RouteChecks.check_route_access(
|
||||
route, LiteLLMRoutes.agent_inference_routes.value
|
||||
)
|
||||
|
||||
|
||||
def managed_inference_request(
|
||||
route: str,
|
||||
body: Mapping[str, object],
|
||||
settings: Mapping[str, object],
|
||||
cli_model: str | None,
|
||||
path_model: object = None,
|
||||
query_model: object = None,
|
||||
) -> dict[str, object]:
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
|
||||
if route in _MANAGED_REALTIME_ROUTES:
|
||||
model: Final = query_model or body.get("model")
|
||||
if not isinstance(model, str) or not model:
|
||||
raise_identity_failure(
|
||||
AgentIdentityFailure(message="Managed inference requires an explicit or configured model")
|
||||
)
|
||||
return {**body, "model": model}
|
||||
if route not in _MANAGED_MODEL_ROUTES and not RouteChecks.check_route_access(route, _MANAGED_MODEL_PATHS):
|
||||
return dict(body)
|
||||
from litellm.proxy.common_utils.http_parsing_utils import resolve_inference_model
|
||||
|
||||
kind: Final = (
|
||||
"image_generation"
|
||||
if route.endswith("/images/generations")
|
||||
else "image_edit"
|
||||
if route.endswith("/images/edits")
|
||||
else "moderation"
|
||||
if route.endswith(("/moderations", "/audio/transcriptions"))
|
||||
else "speech"
|
||||
if route.endswith("/audio/speech")
|
||||
else "body"
|
||||
if route.endswith(("/rerank", "/messages/count_tokens"))
|
||||
else "path"
|
||||
if route.endswith(":countTokens")
|
||||
else "completion"
|
||||
)
|
||||
endpoint_model: Final = path_model or (
|
||||
query_model if route.endswith(("/completions", "/embeddings", "/images/generations", "/images/edits")) else None
|
||||
)
|
||||
effective: Final = resolve_inference_model(body.get("model"), settings, cli_model, endpoint_model, kind=kind)
|
||||
if not isinstance(effective, str) or not effective:
|
||||
raise_identity_failure(
|
||||
AgentIdentityFailure(message="Managed inference requires an explicit or configured model")
|
||||
)
|
||||
return {**body, "model": effective}
|
||||
|
||||
|
||||
def actor_admission_failure(
|
||||
agent: AgentResponse,
|
||||
context: ManagedAgentContext | None,
|
||||
) -> AgentIdentityFailure | None:
|
||||
if not agent.enabled or agent.identity is None or not agent.identity.active:
|
||||
return AgentIdentityFailure(message="Agent execution is disabled")
|
||||
if context is None:
|
||||
return AgentIdentityFailure(message="This agent requires its bound identity provider token")
|
||||
if context.agent_id != agent.agent_id or context.binding_revision != agent.identity.revision:
|
||||
return AgentIdentityFailure(message="Agent identity changed during authentication; retry")
|
||||
if agent.execution_mode not in (context.mode, "both"):
|
||||
return AgentIdentityFailure(message="Agent is not enabled for this execution mode")
|
||||
if context.mode == "delegated" and not context.user_id:
|
||||
return AgentIdentityFailure(message="A verified human subject is required")
|
||||
return None
|
||||
|
||||
|
||||
def invocation_target(route: str, body: Mapping[str, object]) -> str | None:
|
||||
model: Final = body.get("model")
|
||||
if isinstance(model, str) and model.startswith("a2a/"):
|
||||
return model.removeprefix("a2a/") or None
|
||||
components: Final = tuple(route.strip("/").split("/"))
|
||||
path: Final = components[1:] if components and components[0] == "v1" else components
|
||||
return path[1] if len(path) >= 2 and path[0] == "a2a" else None
|
||||
|
|
@ -93,6 +93,7 @@ from litellm.proxy.common_utils.error_body_call_id import JSON_OBJECT, error_bod
|
|||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
get_client_requested_model,
|
||||
get_tags_from_request_body,
|
||||
resolve_inference_model,
|
||||
)
|
||||
from litellm.proxy.common_utils.openai_error_payload import (
|
||||
LITELLM_CALL_ID_HEADER,
|
||||
|
|
@ -2068,11 +2069,12 @@ class ProxyBaseLLMRequestProcessing:
|
|||
if isinstance(model, str):
|
||||
reject_url_valued_destination("model", model)
|
||||
|
||||
self.data["model"] = (
|
||||
general_settings.get("completion_model", None) # server default
|
||||
or user_model # model name passed via cli args
|
||||
or model # for azure deployments
|
||||
or self.data.get("model", None) # default passed in http request
|
||||
self.data["model"] = resolve_inference_model(
|
||||
self.data.get("model"),
|
||||
general_settings,
|
||||
user_model,
|
||||
model,
|
||||
kind="image_edit" if route_type == "aimage_edit" else "completion",
|
||||
)
|
||||
|
||||
# override with user settings, these are params passed via cli
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ import json
|
|||
import re
|
||||
from collections.abc import Collection, Mapping
|
||||
from types import MappingProxyType, UnionType
|
||||
from typing import Annotated, Any, Final, Union, get_args, get_origin
|
||||
from typing import Annotated, Any, Final, Literal, Union, get_args, get_origin
|
||||
|
||||
import orjson
|
||||
from fastapi import Request, UploadFile, status
|
||||
|
|
@ -25,6 +25,39 @@ _FORM_CONTENT_TYPES: Final[frozenset[str]] = frozenset({"application/x-www-form-
|
|||
_ANNOTATION_QUALIFIERS: Final[frozenset[object]] = frozenset({Annotated, NotRequired, ReadOnly, Required})
|
||||
|
||||
|
||||
def resolve_inference_model(
|
||||
body_model: object,
|
||||
settings: Mapping[str, object],
|
||||
cli_model: str | None,
|
||||
endpoint_model: object = None,
|
||||
*,
|
||||
kind: Literal[
|
||||
"completion", "image_generation", "image_edit", "moderation", "speech", "body", "path"
|
||||
] = "completion",
|
||||
) -> object:
|
||||
match kind:
|
||||
case "image_generation":
|
||||
return cli_model or endpoint_model or settings.get("image_generation_model") or body_model
|
||||
case "image_edit":
|
||||
return (
|
||||
settings.get("completion_model")
|
||||
or cli_model
|
||||
or endpoint_model
|
||||
or settings.get("image_generation_model")
|
||||
or body_model
|
||||
)
|
||||
case "moderation":
|
||||
return cli_model or settings.get("moderation_model") or body_model
|
||||
case "speech":
|
||||
return cli_model or body_model
|
||||
case "body":
|
||||
return body_model
|
||||
case "path":
|
||||
return endpoint_model
|
||||
case "completion":
|
||||
return settings.get("completion_model") or cli_model or endpoint_model or body_model
|
||||
|
||||
|
||||
def _normalize_media_type(content_type: str) -> str:
|
||||
"""Return the bare media type per RFC 7231: strip params, trim, lowercase."""
|
||||
if not content_type:
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ from litellm.proxy.common_request_processing import (
|
|||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
coerce_numeric_form_fields,
|
||||
numeric_form_fields,
|
||||
resolve_inference_model,
|
||||
)
|
||||
from litellm.proxy.common_utils.openai_error_payload import (
|
||||
error_status_code,
|
||||
|
|
@ -118,14 +119,9 @@ async def image_generation(
|
|||
if isinstance(model, str):
|
||||
reject_url_valued_destination("model", model)
|
||||
|
||||
data["model"] = (
|
||||
model
|
||||
or general_settings.get("image_generation_model", None) # server default
|
||||
or user_model # model name passed via cli args
|
||||
or data.get("model", None) # default passed in http request
|
||||
data["model"] = resolve_inference_model(
|
||||
data.get("model"), general_settings, user_model, model, kind="image_generation"
|
||||
)
|
||||
if user_model:
|
||||
data["model"] = user_model
|
||||
|
||||
### MODEL ALIAS MAPPING ###
|
||||
# check if model name in model alias map
|
||||
|
|
@ -324,12 +320,6 @@ async def image_edit_api(
|
|||
if "prompt" not in data:
|
||||
data["prompt"] = None
|
||||
|
||||
data["model"] = (
|
||||
model
|
||||
or general_settings.get("image_generation_model", None) # server default
|
||||
or user_model # model name passed via cli args
|
||||
or data.get("model", None) # default passed in http request
|
||||
)
|
||||
#########################################################
|
||||
# Process request
|
||||
#########################################################
|
||||
|
|
@ -346,7 +336,7 @@ async def image_edit_api(
|
|||
general_settings=general_settings,
|
||||
proxy_config=proxy_config,
|
||||
select_data_generator=select_data_generator,
|
||||
model=None,
|
||||
model=model,
|
||||
user_model=user_model,
|
||||
user_temperature=user_temperature,
|
||||
user_request_timeout=user_request_timeout,
|
||||
|
|
|
|||
|
|
@ -419,6 +419,7 @@ from litellm.proxy.common_utils.http_parsing_utils import (
|
|||
_safe_get_request_headers,
|
||||
check_file_size_under_limit,
|
||||
get_form_data,
|
||||
resolve_inference_model,
|
||||
)
|
||||
from litellm.proxy.common_utils.load_config_utils import get_config_from_bucket
|
||||
from litellm.proxy.common_utils.model_deprecation import collect_model_deprecations
|
||||
|
|
@ -12139,13 +12140,7 @@ async def moderations(
|
|||
proxy_config=proxy_config,
|
||||
)
|
||||
|
||||
data["model"] = (
|
||||
general_settings.get("moderation_model", None) # server default
|
||||
or user_model # model name passed via cli args
|
||||
or data.get("model") # default passed in http request
|
||||
)
|
||||
if user_model:
|
||||
data["model"] = user_model
|
||||
data["model"] = resolve_inference_model(data.get("model"), general_settings, user_model, kind="moderation")
|
||||
|
||||
### CALL HOOKS ### - modify incoming data / reject request before calling the model
|
||||
data = await proxy_logging_obj.pre_call_hook(
|
||||
|
|
@ -12399,13 +12394,7 @@ async def audio_transcriptions(
|
|||
if data.get("user", None) is None and user_api_key_dict.user_id is not None:
|
||||
data["user"] = user_api_key_dict.user_id
|
||||
|
||||
data["model"] = (
|
||||
general_settings.get("moderation_model", None) # server default
|
||||
or user_model # model name passed via cli args
|
||||
or data.get("model", None) # default passed in http request
|
||||
)
|
||||
if user_model:
|
||||
data["model"] = user_model
|
||||
data["model"] = resolve_inference_model(data.get("model"), general_settings, user_model, kind="moderation")
|
||||
|
||||
router_model_names: Final = llm_router.model_names if llm_router is not None else []
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,184 @@
|
|||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import (
|
||||
actor_admission_failure,
|
||||
invocation_target,
|
||||
)
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityBinding, AgentIdentityFailure, ManagedAgentContext
|
||||
|
||||
BINDING: Final = AgentIdentityBinding(
|
||||
agent_id="agent",
|
||||
provider="microsoft_entra",
|
||||
tenant_id="tenant",
|
||||
client_id="client",
|
||||
service_principal_id="principal",
|
||||
issuer="issuer",
|
||||
revision="current",
|
||||
)
|
||||
|
||||
|
||||
def agent(**overrides: object) -> AgentResponse:
|
||||
return AgentResponse.model_validate(
|
||||
{
|
||||
"agent_id": "agent",
|
||||
"agent_name": "Agent",
|
||||
"agent_card_params": {},
|
||||
"identity": BINDING,
|
||||
"identity_managed": True,
|
||||
"execution_mode": "both",
|
||||
**overrides,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"state",
|
||||
[
|
||||
{"enabled": False},
|
||||
{"identity": None},
|
||||
{"identity": BINDING.model_copy(update={"active": False})},
|
||||
{"execution_mode": "delegated"},
|
||||
],
|
||||
)
|
||||
def test_keys_cannot_bypass_lifecycle_or_delegated_only_mode(state: dict[str, object]) -> None:
|
||||
assert isinstance(actor_admission_failure(agent(**state), None), AgentIdentityFailure)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("mode", ["autonomous", "both", "delegated"])
|
||||
def test_keys_cannot_impersonate_an_entra_bound_agent(mode: str) -> None:
|
||||
assert isinstance(actor_admission_failure(agent(execution_mode=mode), None), AgentIdentityFailure)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"context",
|
||||
[
|
||||
ManagedAgentContext(agent_id="agent", binding_revision="previous", mode="autonomous"),
|
||||
ManagedAgentContext(agent_id="another", binding_revision="current", mode="autonomous"),
|
||||
ManagedAgentContext(agent_id="agent", binding_revision="current", mode="delegated"),
|
||||
],
|
||||
)
|
||||
def test_stale_binding_and_unverified_delegation_cannot_pass_admission(context: ManagedAgentContext) -> None:
|
||||
assert isinstance(actor_admission_failure(agent(), context), AgentIdentityFailure)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route,body,expected",
|
||||
[
|
||||
("/a2a/agent", {}, "agent"),
|
||||
("/v1/a2a/agent/", {}, "agent"),
|
||||
("/v1/chat/completions", {"model": "a2a/Readable name"}, "Readable name"),
|
||||
("/v1/chat/completions", {"model": "a2a/"}, None),
|
||||
("/v1/chat/completions", {"model": "ordinary-model"}, None),
|
||||
("/a2a", {}, None),
|
||||
],
|
||||
)
|
||||
def test_invocation_routes_resolve_the_same_target(route: str, body: dict[str, object], expected: str | None) -> None:
|
||||
assert invocation_target(route, body) == expected
|
||||
|
||||
|
||||
def test_execution_mode_must_match_verified_token_mode() -> None:
|
||||
context: Final = ManagedAgentContext(agent_id="agent", binding_revision="current", mode="autonomous")
|
||||
failure: Final = actor_admission_failure(agent(execution_mode="delegated"), context)
|
||||
assert isinstance(failure, AgentIdentityFailure)
|
||||
assert "execution mode" in failure.message
|
||||
|
||||
|
||||
@pytest.mark.parametrize("mode,user", [("autonomous", None), ("delegated", "verified-human")])
|
||||
def test_matching_identity_revision_and_execution_mode_pass_admission(mode: str, user: str | None) -> None:
|
||||
context: Final = ManagedAgentContext.model_validate(
|
||||
{"agent_id": "agent", "binding_revision": "current", "mode": mode, "user_id": user}
|
||||
)
|
||||
assert actor_admission_failure(agent(), context) is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route,method,allowed",
|
||||
[
|
||||
("/v1/agents", "GET", True),
|
||||
("/v1/agents", "POST", False),
|
||||
("/v1/chat/completions", "POST", True),
|
||||
("/v1/chat/completions", "DELETE", False),
|
||||
("/openai/deployments/model/chat/completions", "POST", True),
|
||||
("/engines/openai/model/chat/completions", "POST", True),
|
||||
("/openai/deployments/openai/model/images/generations", "POST", True),
|
||||
("/openai/deployments/openai/model/images/edits", "POST", True),
|
||||
("/v1beta/models/gemini-model:generateContent", "POST", True),
|
||||
("/v1/realtime", "GET", True),
|
||||
("/v1/realtime", "POST", False),
|
||||
("/v1/realtime/client_secrets", "POST", False),
|
||||
("/mcp/tools/call", "POST", True),
|
||||
("/a2a/target/message/send", "POST", True),
|
||||
("/v1/agents/target", "PATCH", False),
|
||||
("/v1/responses/other-response", "GET", False),
|
||||
("/v1/files", "GET", False),
|
||||
("/v1/files", "POST", False),
|
||||
("/openai/v1/files", "GET", False),
|
||||
("/anthropic/v1/files", "GET", False),
|
||||
],
|
||||
)
|
||||
def test_managed_route_scope_excludes_provider_resources(route: str, method: str, allowed: bool) -> None:
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_route_allowed
|
||||
|
||||
assert managed_agent_route_allowed(route, method) is allowed
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route,body,settings,cli_model,path_model,expected",
|
||||
[
|
||||
("/v1/chat/completions", {"model": "body"}, {"completion_model": "default"}, "cli", "path", "default"),
|
||||
("/v1/moderations", {"model": "body"}, {"moderation_model": "default"}, "cli", None, "cli"),
|
||||
("/v1/audio/speech", {"model": "body"}, {"completion_model": "ignored"}, None, None, "body"),
|
||||
("/openai/deployments/path/embeddings", {"model": "body"}, {}, None, "path", "path"),
|
||||
("/v1/messages/count_tokens", {"model": "body"}, {"completion_model": "ignored"}, "cli", None, "body"),
|
||||
("/mcp/tools/call", {}, {"completion_model": "ignored"}, "cli", None, None),
|
||||
("/v1/images/generations", {"model": "image"}, {"completion_model": "text"}, None, None, "image"),
|
||||
("/v1/images/generations", {}, {"image_generation_model": "image"}, None, None, "image"),
|
||||
("/v1/images/edits", {}, {"image_generation_model": "image"}, None, None, "image"),
|
||||
("/v1/rerank", {"model": "reranker"}, {"completion_model": "text"}, "cli", None, "reranker"),
|
||||
("/v1beta/models/path:countTokens", {"model": "body"}, {"completion_model": "text"}, "cli", "path", "path"),
|
||||
],
|
||||
)
|
||||
def test_managed_inference_resolves_dispatch_precedence(route, body, settings, cli_model, path_model, expected):
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request
|
||||
|
||||
assert managed_inference_request(route, body, settings, cli_model, path_model).get("model") == expected
|
||||
|
||||
|
||||
def test_managed_inference_without_any_model_cannot_skip_model_grants():
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request
|
||||
|
||||
with pytest.raises(HTTPException, match="explicit or configured model"):
|
||||
managed_inference_request("/v1/moderations", {}, {}, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("route", ["/v1/chat/completions", "/v1/images/generations", "/v1/images/edits"])
|
||||
def test_managed_inference_query_model_takes_precedence_over_body(route: str):
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request
|
||||
|
||||
assert managed_inference_request(route, {"model": "body"}, {}, None, query_model="query")["model"] == "query"
|
||||
|
||||
|
||||
def test_managed_inference_ignores_unsupported_query_model():
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request
|
||||
|
||||
assert (
|
||||
managed_inference_request("/v1/messages", {"model": "body"}, {}, None, query_model="query")["model"] == "body"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("route", ["/realtime", "/v1/realtime", "/openai/v1/realtime"])
|
||||
def test_managed_realtime_requires_a_model_and_ignores_completion_defaults(route: str) -> None:
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request
|
||||
|
||||
with pytest.raises(HTTPException, match="explicit or configured model"):
|
||||
managed_inference_request(route, {}, {"completion_model": "allowed-default"}, "cli")
|
||||
assert (
|
||||
managed_inference_request(route, {"model": "requested"}, {"completion_model": "allowed-default"}, "cli")[
|
||||
"model"
|
||||
]
|
||||
== "requested"
|
||||
)
|
||||
|
|
@ -1210,3 +1210,28 @@ class TestCoerceNumericFormFields:
|
|||
numeric_fields=self.numeric_fields,
|
||||
)
|
||||
assert result == {"n": 3, "temperature": None, "image": buffer}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"kind,settings,cli,path,body,expected",
|
||||
[
|
||||
("completion", {"completion_model": "default"}, "cli", "path", "body", "default"),
|
||||
("completion", {}, "cli", "path", "body", "cli"),
|
||||
("completion", {}, None, "path", "body", "path"),
|
||||
("completion", {}, None, None, "body", "body"),
|
||||
("image_generation", {"completion_model": "text", "image_generation_model": "image"}, None, None, "body", "image"),
|
||||
("image_generation", {"image_generation_model": "image"}, "cli", "path", "body", "cli"),
|
||||
("image_generation", {"image_generation_model": "image"}, None, "path", "body", "path"),
|
||||
("image_edit", {"completion_model": "text", "image_generation_model": "image"}, None, None, "body", "text"),
|
||||
("image_edit", {"image_generation_model": "image"}, None, "path", "body", "path"),
|
||||
("image_edit", {"image_generation_model": "image"}, None, None, "body", "image"),
|
||||
("moderation", {"moderation_model": "mod"}, "cli", None, "body", "cli"),
|
||||
("speech", {"completion_model": "text"}, None, None, "body", "body"),
|
||||
("body", {"completion_model": "text"}, "cli", None, "body", "body"),
|
||||
("path", {"completion_model": "text"}, "cli", "path", "body", "path"),
|
||||
],
|
||||
)
|
||||
def test_shared_inference_model_selection_preserves_handler_precedence(kind, settings, cli, path, body, expected):
|
||||
from litellm.proxy.common_utils.http_parsing_utils import resolve_inference_model
|
||||
|
||||
assert resolve_inference_model(body, settings, cli, path, kind=kind) == expected
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue