refactor(agents): share inference model resolution with admission policy

This commit is contained in:
Joshua Valluru 2026-09-26 11:43:17 -07:00
parent 91629ca447
commit db87f4cc69
7 changed files with 402 additions and 34 deletions

View 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

View file

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

View file

@ -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:

View file

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

View file

@ -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
@ -12120,13 +12121,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(
@ -12380,13 +12375,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 []

View file

@ -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"
)

View file

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