diff --git a/.github/workflows/test-unit-proxy-db.yml b/.github/workflows/test-unit-proxy-db.yml index 8ee3a1ed0cd..49a36aa23f0 100644 --- a/.github/workflows/test-unit-proxy-db.yml +++ b/.github/workflows/test-unit-proxy-db.yml @@ -215,6 +215,7 @@ jobs: tests/proxy_unit_tests/test_models_fallback_endpoint.py tests/proxy_unit_tests/test_google_endpoint_routing.py tests/proxy_unit_tests/test_google_gemini_proxy_request.py + tests/proxy_unit_tests/test_gemini_agents_endpoints.py tests/proxy_unit_tests/test_get_favicon.py tests/proxy_unit_tests/test_get_image.py tests/proxy_unit_tests/test_ui_path_detection.py diff --git a/litellm/__init__.py b/litellm/__init__.py index c868ae55b4f..1e8f8613fba 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1288,6 +1288,18 @@ from .responses.main import * # Interactions API is available as litellm.interactions module # Usage: litellm.interactions.create(), litellm.interactions.get(), etc. from . import interactions +from .interactions.agents.main import ( + acreate as acreate_agent, + create as create_agent, + alist as alist_agents, + list as list_agents, + aget as aget_agent, + get as get_agent, + adelete as adelete_agent, + delete as delete_agent, + alist_versions as alist_agent_versions, + list_versions as list_agent_versions, +) from .skills.main import ( create_skill, acreate_skill, diff --git a/litellm/interactions/__init__.py b/litellm/interactions/__init__.py index e1125b649a6..ed01462cba6 100644 --- a/litellm/interactions/__init__.py +++ b/litellm/interactions/__init__.py @@ -5,31 +5,40 @@ This module provides SDK methods for Google's Interactions API. Usage: import litellm - + # Create an interaction with a model response = litellm.interactions.create( model="gemini-2.5-flash", input="Hello, how are you?" ) - + # Create an interaction with an agent response = litellm.interactions.create( agent="deep-research-pro-preview-12-2025", input="Research the current state of cancer research" ) - + # Async version response = await litellm.interactions.acreate(...) - + # Get an interaction response = litellm.interactions.get(interaction_id="...") - + # Delete an interaction result = litellm.interactions.delete(interaction_id="...") - + # Cancel an interaction result = litellm.interactions.cancel(interaction_id="...") + # Create a managed agent on the provider side + result = litellm.interactions.agents.create( + name="waverunner", + custom_llm_provider="gemini", + api_key="...", + base_agent="gemini-2.5-flash", + instructions="You are a helpful assistant.", + ) + Methods: - create(): Sync create interaction - acreate(): Async create interaction @@ -39,8 +48,12 @@ Methods: - adelete(): Async delete interaction - cancel(): Sync cancel interaction - acancel(): Async cancel interaction + +Sub-modules: +- agents: Provider-side agent creation (litellm.interactions.agents.create) """ +from litellm.interactions import agents from litellm.interactions.main import ( acancel, acreate, @@ -65,4 +78,6 @@ __all__ = [ # Cancel "cancel", "acancel", + # Sub-modules + "agents", ] diff --git a/litellm/interactions/agents/__init__.py b/litellm/interactions/agents/__init__.py new file mode 100644 index 00000000000..711a54fdcbb --- /dev/null +++ b/litellm/interactions/agents/__init__.py @@ -0,0 +1,39 @@ +""" +litellm.interactions.agents + +Full CRUD SDK for provider-side managed agents (e.g. Gemini v1beta/agents). + + litellm.interactions.agents.create(name=..., ...) + litellm.interactions.agents.list(api_key=...) + litellm.interactions.agents.get(name=..., ...) + litellm.interactions.agents.delete(name=..., ...) + litellm.interactions.agents.list_versions(name=..., ...) + +Async counterparts: acreate, alist, aget, adelete, alist_versions +""" + +from litellm.interactions.agents.main import ( + acreate, + adelete, + aget, + alist, + alist_versions, + create, + delete, + get, + list, + list_versions, +) + +__all__ = [ + "create", + "acreate", + "list", + "alist", + "get", + "aget", + "delete", + "adelete", + "list_versions", + "alist_versions", +] diff --git a/litellm/interactions/agents/http_handler.py b/litellm/interactions/agents/http_handler.py new file mode 100644 index 00000000000..d45ca6f4346 --- /dev/null +++ b/litellm/interactions/agents/http_handler.py @@ -0,0 +1,478 @@ +""" +HTTP handler for the Agents API. + +Extends InteractionsHTTPHandler so that the shared HTTP infrastructure +(_handle_error, _sync_client, _async_client) is reused rather than +duplicated. BaseAgentsAPIConfig stays as pure transform code. +""" + +from typing import Any, Coroutine, Dict, Optional, Union + +import httpx + +from litellm.constants import request_timeout +from litellm.interactions.http_handler import InteractionsHTTPHandler +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.base_llm.agents.transformation import BaseAgentsAPIConfig +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.types.agents import ( + AgentCreateResponse, + AgentDeleteResult, + AgentListResponse, + AgentVersionsResponse, +) +from litellm.types.router import GenericLiteLLMParams + + +class AgentsHTTPHandler(InteractionsHTTPHandler): + """HTTP handler for Agents API CRUD requests.""" + + # ------------------------------------------------------------------ # + # CREATE # + # ------------------------------------------------------------------ # + + def create_agent( + self, + agents_api_config: BaseAgentsAPIConfig, + name: str, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + extra_headers: Optional[Dict[str, Any]] = None, + extra_body: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[HTTPHandler] = None, + _is_async: bool = False, + ) -> Union[AgentCreateResponse, Coroutine[Any, Any, AgentCreateResponse]]: + if _is_async: + return self.async_create_agent( + agents_api_config=agents_api_config, + name=name, + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers=extra_headers, + extra_body=extra_body, + timeout=timeout, + ) + + sync_httpx_client = self._sync_client(litellm_params, client) + headers = agents_api_config.validate_environment( + headers=extra_headers or {}, litellm_params=dict(litellm_params) + ) + url = agents_api_config.get_complete_url( + api_base=litellm_params.get("api_base"), + litellm_params=dict(litellm_params), + ) + data = agents_api_config.transform_create_request( + name=name, litellm_params=dict(litellm_params) + ) + if extra_body: + data.update(extra_body) + + logging_obj.pre_call( + input=name, + api_key="", + additional_args={ + "complete_input_dict": data, + "api_base": url, + "headers": headers, + }, + ) + try: + response = sync_httpx_client.post( + url=url, headers=headers, json=data, timeout=timeout or request_timeout + ) + except Exception as e: + raise self._handle_error(e=e, provider_config=agents_api_config) + + logging_obj.post_call( + original_response=response.text, + additional_args={"complete_input_dict": data}, + ) + return agents_api_config.transform_create_response( + raw_response=response, name=name + ) + + async def async_create_agent( + self, + agents_api_config: BaseAgentsAPIConfig, + name: str, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + extra_headers: Optional[Dict[str, Any]] = None, + extra_body: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[AsyncHTTPHandler] = None, + ) -> AgentCreateResponse: + async_httpx_client = self._async_client(litellm_params, client) + headers = agents_api_config.validate_environment( + headers=extra_headers or {}, litellm_params=dict(litellm_params) + ) + url = agents_api_config.get_complete_url( + api_base=litellm_params.get("api_base"), + litellm_params=dict(litellm_params), + ) + data = agents_api_config.transform_create_request( + name=name, litellm_params=dict(litellm_params) + ) + if extra_body: + data.update(extra_body) + + logging_obj.pre_call( + input=name, + api_key="", + additional_args={ + "complete_input_dict": data, + "api_base": url, + "headers": headers, + }, + ) + try: + response = await async_httpx_client.post( + url=url, headers=headers, json=data, timeout=timeout or request_timeout + ) + except Exception as e: + raise self._handle_error(e=e, provider_config=agents_api_config) + + logging_obj.post_call( + original_response=response.text, + additional_args={"complete_input_dict": data}, + ) + return agents_api_config.transform_create_response( + raw_response=response, name=name + ) + + # ------------------------------------------------------------------ # + # LIST # + # ------------------------------------------------------------------ # + + def list_agents( + self, + agents_api_config: BaseAgentsAPIConfig, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[HTTPHandler] = None, + _is_async: bool = False, + ) -> Union[AgentListResponse, Coroutine[Any, Any, AgentListResponse]]: + if _is_async: + return self.async_list_agents( + agents_api_config=agents_api_config, + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers=extra_headers, + timeout=timeout, + ) + + sync_httpx_client = self._sync_client(litellm_params, client) + headers = agents_api_config.validate_environment( + headers=extra_headers or {}, litellm_params=dict(litellm_params) + ) + url, params = agents_api_config.transform_list_request( + api_base=litellm_params.get("api_base"), + litellm_params=dict(litellm_params), + ) + logging_obj.pre_call( + input="list_agents", + api_key="", + additional_args={"api_base": url, "headers": headers}, + ) + try: + response = sync_httpx_client.get(url=url, headers=headers, params=params) + except Exception as e: + raise self._handle_error(e=e, provider_config=agents_api_config) + + logging_obj.post_call(original_response=response.text, additional_args={}) + return agents_api_config.transform_list_response(raw_response=response) + + async def async_list_agents( + self, + agents_api_config: BaseAgentsAPIConfig, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[AsyncHTTPHandler] = None, + ) -> AgentListResponse: + async_httpx_client = self._async_client(litellm_params, client) + headers = agents_api_config.validate_environment( + headers=extra_headers or {}, litellm_params=dict(litellm_params) + ) + url, params = agents_api_config.transform_list_request( + api_base=litellm_params.get("api_base"), + litellm_params=dict(litellm_params), + ) + logging_obj.pre_call( + input="list_agents", + api_key="", + additional_args={"api_base": url, "headers": headers}, + ) + try: + response = await async_httpx_client.get( + url=url, headers=headers, params=params + ) + except Exception as e: + raise self._handle_error(e=e, provider_config=agents_api_config) + + logging_obj.post_call(original_response=response.text, additional_args={}) + return agents_api_config.transform_list_response(raw_response=response) + + # ------------------------------------------------------------------ # + # GET # + # ------------------------------------------------------------------ # + + def get_agent( + self, + agents_api_config: BaseAgentsAPIConfig, + name: str, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[HTTPHandler] = None, + _is_async: bool = False, + ) -> Union[AgentCreateResponse, Coroutine[Any, Any, AgentCreateResponse]]: + if _is_async: + return self.async_get_agent( + agents_api_config=agents_api_config, + name=name, + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers=extra_headers, + timeout=timeout, + ) + + sync_httpx_client = self._sync_client(litellm_params, client) + headers = agents_api_config.validate_environment( + headers=extra_headers or {}, litellm_params=dict(litellm_params) + ) + url, params = agents_api_config.transform_get_request( + name=name, + api_base=litellm_params.get("api_base"), + litellm_params=dict(litellm_params), + ) + logging_obj.pre_call( + input=name, + api_key="", + additional_args={"api_base": url, "headers": headers}, + ) + try: + response = sync_httpx_client.get(url=url, headers=headers, params=params) + except Exception as e: + raise self._handle_error(e=e, provider_config=agents_api_config) + + logging_obj.post_call(original_response=response.text, additional_args={}) + return agents_api_config.transform_get_response( + raw_response=response, name=name + ) + + async def async_get_agent( + self, + agents_api_config: BaseAgentsAPIConfig, + name: str, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[AsyncHTTPHandler] = None, + ) -> AgentCreateResponse: + async_httpx_client = self._async_client(litellm_params, client) + headers = agents_api_config.validate_environment( + headers=extra_headers or {}, litellm_params=dict(litellm_params) + ) + url, params = agents_api_config.transform_get_request( + name=name, + api_base=litellm_params.get("api_base"), + litellm_params=dict(litellm_params), + ) + logging_obj.pre_call( + input=name, + api_key="", + additional_args={"api_base": url, "headers": headers}, + ) + try: + response = await async_httpx_client.get( + url=url, headers=headers, params=params + ) + except Exception as e: + raise self._handle_error(e=e, provider_config=agents_api_config) + + logging_obj.post_call(original_response=response.text, additional_args={}) + return agents_api_config.transform_get_response( + raw_response=response, name=name + ) + + # ------------------------------------------------------------------ # + # DELETE # + # ------------------------------------------------------------------ # + + def delete_agent( + self, + agents_api_config: BaseAgentsAPIConfig, + name: str, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[HTTPHandler] = None, + _is_async: bool = False, + ) -> Union[AgentDeleteResult, Coroutine[Any, Any, AgentDeleteResult]]: + if _is_async: + return self.async_delete_agent( + agents_api_config=agents_api_config, + name=name, + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers=extra_headers, + timeout=timeout, + ) + + sync_httpx_client = self._sync_client(litellm_params, client) + headers = agents_api_config.validate_environment( + headers=extra_headers or {}, litellm_params=dict(litellm_params) + ) + url = agents_api_config.transform_delete_request( + name=name, + api_base=litellm_params.get("api_base"), + litellm_params=dict(litellm_params), + ) + logging_obj.pre_call( + input=name, + api_key="", + additional_args={"api_base": url, "headers": headers}, + ) + try: + response = sync_httpx_client.delete( + url=url, headers=headers, timeout=timeout or request_timeout + ) + except Exception as e: + raise self._handle_error(e=e, provider_config=agents_api_config) + + logging_obj.post_call(original_response=response.text, additional_args={}) + return agents_api_config.transform_delete_response( + raw_response=response, name=name + ) + + async def async_delete_agent( + self, + agents_api_config: BaseAgentsAPIConfig, + name: str, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[AsyncHTTPHandler] = None, + ) -> AgentDeleteResult: + async_httpx_client = self._async_client(litellm_params, client) + headers = agents_api_config.validate_environment( + headers=extra_headers or {}, litellm_params=dict(litellm_params) + ) + url = agents_api_config.transform_delete_request( + name=name, + api_base=litellm_params.get("api_base"), + litellm_params=dict(litellm_params), + ) + logging_obj.pre_call( + input=name, + api_key="", + additional_args={"api_base": url, "headers": headers}, + ) + try: + response = await async_httpx_client.delete( + url=url, headers=headers, timeout=timeout or request_timeout + ) + except Exception as e: + raise self._handle_error(e=e, provider_config=agents_api_config) + + logging_obj.post_call(original_response=response.text, additional_args={}) + return agents_api_config.transform_delete_response( + raw_response=response, name=name + ) + + # ------------------------------------------------------------------ # + # LIST VERSIONS # + # ------------------------------------------------------------------ # + + def list_agent_versions( + self, + agents_api_config: BaseAgentsAPIConfig, + name: str, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[HTTPHandler] = None, + _is_async: bool = False, + ) -> Union[AgentVersionsResponse, Coroutine[Any, Any, AgentVersionsResponse]]: + if _is_async: + return self.async_list_agent_versions( + agents_api_config=agents_api_config, + name=name, + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers=extra_headers, + timeout=timeout, + ) + + sync_httpx_client = self._sync_client(litellm_params, client) + headers = agents_api_config.validate_environment( + headers=extra_headers or {}, litellm_params=dict(litellm_params) + ) + url, params = agents_api_config.transform_list_versions_request( + name=name, + api_base=litellm_params.get("api_base"), + litellm_params=dict(litellm_params), + ) + logging_obj.pre_call( + input=name, + api_key="", + additional_args={"api_base": url, "headers": headers}, + ) + try: + response = sync_httpx_client.get(url=url, headers=headers, params=params) + except Exception as e: + raise self._handle_error(e=e, provider_config=agents_api_config) + + logging_obj.post_call(original_response=response.text, additional_args={}) + return agents_api_config.transform_list_versions_response( + raw_response=response, name=name + ) + + async def async_list_agent_versions( + self, + agents_api_config: BaseAgentsAPIConfig, + name: str, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[AsyncHTTPHandler] = None, + ) -> AgentVersionsResponse: + async_httpx_client = self._async_client(litellm_params, client) + headers = agents_api_config.validate_environment( + headers=extra_headers or {}, litellm_params=dict(litellm_params) + ) + url, params = agents_api_config.transform_list_versions_request( + name=name, + api_base=litellm_params.get("api_base"), + litellm_params=dict(litellm_params), + ) + logging_obj.pre_call( + input=name, + api_key="", + additional_args={"api_base": url, "headers": headers}, + ) + try: + response = await async_httpx_client.get( + url=url, headers=headers, params=params + ) + except Exception as e: + raise self._handle_error(e=e, provider_config=agents_api_config) + + logging_obj.post_call(original_response=response.text, additional_args={}) + return agents_api_config.transform_list_versions_response( + raw_response=response, name=name + ) + + +agents_http_handler = AgentsHTTPHandler() diff --git a/litellm/interactions/agents/main.py b/litellm/interactions/agents/main.py new file mode 100644 index 00000000000..7375fd6273f --- /dev/null +++ b/litellm/interactions/agents/main.py @@ -0,0 +1,523 @@ +""" +LiteLLM Agents API - Main Module + +Usage: + import litellm + + # Create + response = litellm.interactions.agents.create( + name="waverunner", + custom_llm_provider="gemini", + api_key="...", + base_agent="gemini-2.5-flash", + instructions="You are a helpful assistant.", + ) + + # List + response = litellm.interactions.agents.list(api_key="...", custom_llm_provider="gemini") + + # Get + response = litellm.interactions.agents.get(name="waverunner", api_key="...") + + # Delete + result = litellm.interactions.agents.delete(name="waverunner", api_key="...") + + # List versions + result = litellm.interactions.agents.list_versions(name="waverunner", api_key="...") + + # Async versions: acreate, alist, aget, adelete, alist_versions +""" + +import asyncio +import contextvars +from functools import partial +from typing import Any, Coroutine, Dict, Optional, Union + +import httpx + +import litellm +from litellm.interactions.agents.http_handler import agents_http_handler +from litellm.interactions.agents.utils import get_provider_agents_api_config +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.types.agents import ( + AgentCreateResponse, + AgentDeleteResult, + AgentListResponse, + AgentVersionsResponse, +) +from litellm.types.interactions import InteractionEnvironment +from litellm.types.router import GenericLiteLLMParams +from litellm.utils import client + + +# ------------------------------------------------------------------ # +# Shared helpers # +# ------------------------------------------------------------------ # + + +def _get_agents_api_config(custom_llm_provider: str): + config = get_provider_agents_api_config(custom_llm_provider) + if config is None: + raise litellm.BadRequestError( + message=( + f"Provider '{custom_llm_provider}' does not have a native " + "agents API. Use the proxy POST /v1/agents endpoint to store " + "agents locally." + ), + model="", + llm_provider=custom_llm_provider, + ) + return config + + +def _make_logging_obj( + kwargs: Dict[str, Any], + model: str, + custom_llm_provider: str, + call_type: str, + optional_params: Dict[str, Any], +) -> LiteLLMLoggingObj: + litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore + litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_logging_obj.update_from_kwargs( + kwargs=kwargs, + model=model, + optional_params=optional_params, + litellm_params={"litellm_call_id": litellm_call_id}, + custom_llm_provider=custom_llm_provider, + ) + return litellm_logging_obj + + +# ================================================================== # +# CREATE # +# ================================================================== # + + +@client +async def acreate( + name: str, + base_agent: Optional[str] = None, + instructions: Optional[str] = None, + base_environment: Optional[InteractionEnvironment] = None, + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + extra_body: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + **kwargs, +) -> AgentCreateResponse: + """Async: Create a managed agent on the provider side.""" + local_vars = locals() + try: + loop = asyncio.get_event_loop() + kwargs["acreate_agent"] = True + func = partial( + create, + name=name, + base_agent=base_agent, + instructions=instructions, + base_environment=base_environment, + custom_llm_provider=custom_llm_provider or "gemini", + extra_headers=extra_headers, + extra_body=extra_body, + timeout=timeout, + **kwargs, + ) + ctx = contextvars.copy_context() + init_response = await loop.run_in_executor(None, partial(ctx.run, func)) + if asyncio.iscoroutine(init_response): + return await init_response + return init_response + except Exception as e: + raise litellm.exception_type( + model=name, + custom_llm_provider=custom_llm_provider or "gemini", + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +@client +def create( + name: str, + base_agent: Optional[str] = None, + instructions: Optional[str] = None, + base_environment: Optional[InteractionEnvironment] = None, + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + extra_body: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + **kwargs, +) -> Union[AgentCreateResponse, Coroutine[Any, Any, AgentCreateResponse]]: + """ + Sync: Create a managed agent on the provider side. + + Args: + name: Name for the agent (required). + base_agent: Base agent to derive from (e.g. "waverunner"). + instructions: System instructions for the agent. + base_environment: Environment to fork from — an env_id string or a + dict like ``{"type": "remote", "sources": [...]}``. + custom_llm_provider: Provider to use, e.g. "gemini". + extra_headers: Additional HTTP headers. + extra_body: Additional request body fields. + timeout: Request timeout. + **kwargs: Forwarded to GenericLiteLLMParams (api_key, api_base, etc.). + """ + local_vars = locals() + custom_llm_provider = ( + custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini" + ) + try: + _is_async = kwargs.pop("acreate_agent", False) is True + if base_agent is not None: + kwargs["base_agent"] = base_agent + if instructions is not None: + kwargs["instructions"] = instructions + if base_environment is not None: + kwargs["base_environment"] = base_environment + kwargs.setdefault("custom_llm_provider", custom_llm_provider) + litellm_params = GenericLiteLLMParams(**kwargs) + logging_obj = _make_logging_obj( + kwargs, name, custom_llm_provider, "create_agent", {} + ) + config = _get_agents_api_config(custom_llm_provider) + return agents_http_handler.create_agent( + agents_api_config=config, + name=name, + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers=extra_headers, + extra_body=extra_body, + timeout=timeout, + _is_async=_is_async, + ) + except Exception as e: + raise litellm.exception_type( + model=name, + custom_llm_provider=custom_llm_provider, + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +# ================================================================== # +# LIST # +# ================================================================== # + + +@client +async def alist( + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + **kwargs, +) -> AgentListResponse: + """Async: List all agents on the provider side.""" + local_vars = locals() + try: + loop = asyncio.get_event_loop() + kwargs["alist_agents"] = True + func = partial( + list, + custom_llm_provider=custom_llm_provider or "gemini", + extra_headers=extra_headers, + timeout=timeout, + **kwargs, + ) + ctx = contextvars.copy_context() + init_response = await loop.run_in_executor(None, partial(ctx.run, func)) + if asyncio.iscoroutine(init_response): + return await init_response + return init_response + except Exception as e: + raise litellm.exception_type( + model="", + custom_llm_provider=custom_llm_provider or "gemini", + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +@client +def list( + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + **kwargs, +) -> Union[AgentListResponse, Coroutine[Any, Any, AgentListResponse]]: + """Sync: List all agents on the provider side.""" + local_vars = locals() + custom_llm_provider = ( + custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini" + ) + try: + _is_async = kwargs.pop("alist_agents", False) is True + kwargs.setdefault("custom_llm_provider", custom_llm_provider) + litellm_params = GenericLiteLLMParams(**kwargs) + logging_obj = _make_logging_obj( + kwargs, "", custom_llm_provider, "list_agents", {} + ) + config = _get_agents_api_config(custom_llm_provider) + return agents_http_handler.list_agents( + agents_api_config=config, + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers=extra_headers, + timeout=timeout, + _is_async=_is_async, + ) + except Exception as e: + raise litellm.exception_type( + model="", + custom_llm_provider=custom_llm_provider, + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +# ================================================================== # +# GET # +# ================================================================== # + + +@client +async def aget( + name: str, + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + **kwargs, +) -> AgentCreateResponse: + """Async: Get a specific agent by name.""" + local_vars = locals() + try: + loop = asyncio.get_event_loop() + kwargs["aget_agent"] = True + func = partial( + get, + name=name, + custom_llm_provider=custom_llm_provider or "gemini", + extra_headers=extra_headers, + timeout=timeout, + **kwargs, + ) + ctx = contextvars.copy_context() + init_response = await loop.run_in_executor(None, partial(ctx.run, func)) + if asyncio.iscoroutine(init_response): + return await init_response + return init_response + except Exception as e: + raise litellm.exception_type( + model=name, + custom_llm_provider=custom_llm_provider or "gemini", + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +@client +def get( + name: str, + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + **kwargs, +) -> Union[AgentCreateResponse, Coroutine[Any, Any, AgentCreateResponse]]: + """Sync: Get a specific agent by name.""" + local_vars = locals() + custom_llm_provider = ( + custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini" + ) + try: + _is_async = kwargs.pop("aget_agent", False) is True + kwargs.setdefault("custom_llm_provider", custom_llm_provider) + litellm_params = GenericLiteLLMParams(**kwargs) + logging_obj = _make_logging_obj( + kwargs, name, custom_llm_provider, "get_agent", {"name": name} + ) + config = _get_agents_api_config(custom_llm_provider) + return agents_http_handler.get_agent( + agents_api_config=config, + name=name, + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers=extra_headers, + timeout=timeout, + _is_async=_is_async, + ) + except Exception as e: + raise litellm.exception_type( + model=name, + custom_llm_provider=custom_llm_provider, + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +# ================================================================== # +# DELETE # +# ================================================================== # + + +@client +async def adelete( + name: str, + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + **kwargs, +) -> AgentDeleteResult: + """Async: Delete a specific agent by name.""" + local_vars = locals() + try: + loop = asyncio.get_event_loop() + kwargs["adelete_agent"] = True + func = partial( + delete, + name=name, + custom_llm_provider=custom_llm_provider or "gemini", + extra_headers=extra_headers, + timeout=timeout, + **kwargs, + ) + ctx = contextvars.copy_context() + init_response = await loop.run_in_executor(None, partial(ctx.run, func)) + if asyncio.iscoroutine(init_response): + return await init_response + return init_response + except Exception as e: + raise litellm.exception_type( + model=name, + custom_llm_provider=custom_llm_provider or "gemini", + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +@client +def delete( + name: str, + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + **kwargs, +) -> Union[AgentDeleteResult, Coroutine[Any, Any, AgentDeleteResult]]: + """Sync: Delete a specific agent by name.""" + local_vars = locals() + custom_llm_provider = ( + custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini" + ) + try: + _is_async = kwargs.pop("adelete_agent", False) is True + kwargs.setdefault("custom_llm_provider", custom_llm_provider) + litellm_params = GenericLiteLLMParams(**kwargs) + logging_obj = _make_logging_obj( + kwargs, name, custom_llm_provider, "delete_agent", {"name": name} + ) + config = _get_agents_api_config(custom_llm_provider) + return agents_http_handler.delete_agent( + agents_api_config=config, + name=name, + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers=extra_headers, + timeout=timeout, + _is_async=_is_async, + ) + except Exception as e: + raise litellm.exception_type( + model=name, + custom_llm_provider=custom_llm_provider, + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +# ================================================================== # +# LIST VERSIONS # +# ================================================================== # + + +@client +async def alist_versions( + name: str, + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + **kwargs, +) -> AgentVersionsResponse: + """Async: List versions of a specific agent.""" + local_vars = locals() + try: + loop = asyncio.get_event_loop() + kwargs["alist_agent_versions"] = True + func = partial( + list_versions, + name=name, + custom_llm_provider=custom_llm_provider or "gemini", + extra_headers=extra_headers, + timeout=timeout, + **kwargs, + ) + ctx = contextvars.copy_context() + init_response = await loop.run_in_executor(None, partial(ctx.run, func)) + if asyncio.iscoroutine(init_response): + return await init_response + return init_response + except Exception as e: + raise litellm.exception_type( + model=name, + custom_llm_provider=custom_llm_provider or "gemini", + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +@client +def list_versions( + name: str, + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + **kwargs, +) -> Union[AgentVersionsResponse, Coroutine[Any, Any, AgentVersionsResponse]]: + """Sync: List versions of a specific agent.""" + local_vars = locals() + custom_llm_provider = ( + custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini" + ) + try: + _is_async = kwargs.pop("alist_agent_versions", False) is True + kwargs.setdefault("custom_llm_provider", custom_llm_provider) + litellm_params = GenericLiteLLMParams(**kwargs) + logging_obj = _make_logging_obj( + kwargs, name, custom_llm_provider, "list_agent_versions", {"name": name} + ) + config = _get_agents_api_config(custom_llm_provider) + return agents_http_handler.list_agent_versions( + agents_api_config=config, + name=name, + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers=extra_headers, + timeout=timeout, + _is_async=_is_async, + ) + except Exception as e: + raise litellm.exception_type( + model=name, + custom_llm_provider=custom_llm_provider, + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) diff --git a/litellm/interactions/agents/utils.py b/litellm/interactions/agents/utils.py new file mode 100644 index 00000000000..d16a9597f53 --- /dev/null +++ b/litellm/interactions/agents/utils.py @@ -0,0 +1,23 @@ +""" +Utility functions for the Agents API SDK. +""" + +from typing import Optional + +from litellm.llms.base_llm.agents.transformation import BaseAgentsAPIConfig + + +def get_provider_agents_api_config( + custom_llm_provider: Optional[str], +) -> Optional[BaseAgentsAPIConfig]: + """ + Return a provider-specific BaseAgentsAPIConfig if the provider has a + native agent-creation API, or None otherwise. + """ + from litellm.types.utils import LlmProviders + + if custom_llm_provider == LlmProviders.GEMINI.value: + from litellm.llms.gemini.agents.transformation import GeminiAgentsConfig + + return GeminiAgentsConfig() + return None diff --git a/litellm/interactions/http_handler.py b/litellm/interactions/http_handler.py index 7fead07043f..695da2be89a 100644 --- a/litellm/interactions/http_handler.py +++ b/litellm/interactions/http_handler.py @@ -41,27 +41,55 @@ from litellm.types.interactions import ( from litellm.types.router import GenericLiteLLMParams -class InteractionsHTTPHandler: +class _BaseHTTPHandler: + """ + Shared HTTP infrastructure for LiteLLM handler classes. + + Provides common client resolution and error-mapping helpers so that + handler subclasses (InteractionsHTTPHandler, AgentsHTTPHandler, …) do + not duplicate this boilerplate. + """ + + def _handle_error(self, e: Exception, provider_config: Any) -> Exception: + if isinstance(e, httpx.HTTPStatusError): + return provider_config.get_error_class( + error_message=e.response.text, + status_code=e.response.status_code, + headers=dict(e.response.headers), + ) + return e + + def _sync_client( + self, + litellm_params: GenericLiteLLMParams, + client: Optional[HTTPHandler], + ) -> HTTPHandler: + return client or _get_httpx_client( + params={"ssl_verify": litellm_params.get("ssl_verify", None)} + ) + + def _async_client( + self, + litellm_params: GenericLiteLLMParams, + client: Optional[AsyncHTTPHandler], + ) -> AsyncHTTPHandler: + # GenericLiteLLMParams.get uses getattr; an unset field is None, not the default. + custom_llm_provider = litellm_params.get("custom_llm_provider") or "gemini" + return client or get_async_httpx_client( + llm_provider=litellm.LlmProviders(custom_llm_provider), + params={"ssl_verify": litellm_params.get("ssl_verify", None)}, + ) + + +class InteractionsHTTPHandler(_BaseHTTPHandler): """ HTTP handler for Interactions API requests. """ - def _handle_error( - self, - e: Exception, - provider_config: BaseInteractionsAPIConfig, - ) -> Exception: - """Handle errors from HTTP requests.""" - if isinstance(e, httpx.HTTPStatusError): - error_message = e.response.text - status_code = e.response.status_code - headers = dict(e.response.headers) - return provider_config.get_error_class( - error_message=error_message, - status_code=status_code, - headers=headers, - ) - return e + # _handle_error is inherited from _BaseHTTPHandler (accepts Any provider_config). + # AgentsHTTPHandler also extends this class and passes BaseAgentsAPIConfig, which + # is structurally compatible but a different type — keeping the override here with + # BaseInteractionsAPIConfig would cause type errors in the subclass. # ========================================================= # CREATE INTERACTION diff --git a/litellm/interactions/main.py b/litellm/interactions/main.py index ab429ef6db5..c6eca410fa7 100644 --- a/litellm/interactions/main.py +++ b/litellm/interactions/main.py @@ -48,6 +48,7 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.types.interactions import ( CancelInteractionResult, DeleteInteractionResult, + InteractionEnvironment, InteractionInput, InteractionsAPIResponse, InteractionsAPIStreamingResponse, @@ -80,6 +81,8 @@ async def acreate( store: Optional[bool] = None, # Background execution background: Optional[bool] = None, + # Agent execution environment ("remote", env id, or remote config object) + environment: Optional[InteractionEnvironment] = None, # Response format response_modalities: Optional[List[str]] = None, response_format: Optional[Dict[str, Any]] = None, @@ -109,6 +112,10 @@ async def acreate( stream: Whether to stream the response store: Whether to store the response for later retrieval background: Whether to run in background + environment: Agent execution environment — ``"remote"``, an existing env id + string, or a config object such as + ``{"type": "remote", "sources": [...]}`` / + ``{"type": "remote", "network": {...}}`` response_modalities: Requested response modalities (TEXT, IMAGE, AUDIO) response_format: JSON schema for response format response_mime_type: MIME type of the response @@ -144,6 +151,7 @@ async def acreate( stream=stream, store=store, background=background, + environment=environment, response_modalities=response_modalities, response_format=response_format, response_mime_type=response_mime_type, @@ -194,6 +202,8 @@ def create( store: Optional[bool] = None, # Background execution background: Optional[bool] = None, + # Agent execution environment ("remote", env id, or remote config object) + environment: Optional[InteractionEnvironment] = None, # Response format response_modalities: Optional[List[str]] = None, response_format: Optional[Dict[str, Any]] = None, @@ -231,6 +241,10 @@ def create( stream: Whether to stream the response store: Whether to store the response for later retrieval background: Whether to run in background + environment: Agent execution environment — ``"remote"``, an existing env id + string, or a config object such as + ``{"type": "remote", "sources": [...]}`` / + ``{"type": "remote", "network": {...}}`` response_modalities: Requested response modalities (TEXT, IMAGE, AUDIO) response_format: JSON schema for response format response_mime_type: MIME type of the response @@ -252,7 +266,14 @@ def create( litellm_params = GenericLiteLLMParams(**kwargs) - if model: + # Routing logic: + # - agent provided (no model, or model accidentally set to agent name) → gemini + # - model provided → resolve provider via get_llm_provider (normal routing) + if agent and model == agent: + model = None + if agent and not model: + custom_llm_provider = custom_llm_provider or "gemini" + elif model: model, custom_llm_provider, _, _ = litellm.get_llm_provider( model=model, custom_llm_provider=custom_llm_provider, diff --git a/litellm/interactions/utils.py b/litellm/interactions/utils.py index 3a18ddf52fe..84437f4d3d8 100644 --- a/litellm/interactions/utils.py +++ b/litellm/interactions/utils.py @@ -15,6 +15,7 @@ INTERACTIONS_API_OPTIONAL_PARAMS = { "stream", "store", "background", + "environment", "response_modalities", "response_format", "response_mime_type", diff --git a/litellm/llms/base_llm/agents/__init__.py b/litellm/llms/base_llm/agents/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/base_llm/agents/transformation.py b/litellm/llms/base_llm/agents/transformation.py new file mode 100644 index 00000000000..508e54cb7ab --- /dev/null +++ b/litellm/llms/base_llm/agents/transformation.py @@ -0,0 +1,165 @@ +""" +Base transformation class for provider-side Agents API. + +Providers that have a native agents CRUD API (e.g. Gemini v1beta/agents) +subclass BaseAgentsAPIConfig and implement the abstract methods. + +The HTTP calls are handled by AgentsHTTPHandler — this class is pure +transform logic (same separation as BaseInteractionsAPIConfig / +InteractionsHTTPHandler). +""" + +from abc import ABC, abstractmethod +from typing import Any, Dict, Optional, Tuple, Union + +import httpx + +from litellm.types.agents import ( + AgentCreateResponse, + AgentDeleteResult, + AgentListResponse, + AgentVersionsResponse, +) + + +class BaseAgentsAPIConfig(ABC): + """ + Minimal interface for providers that expose a native agents CRUD API. + """ + + # ------------------------------------------------------------------ # + # CREATE # + # ------------------------------------------------------------------ # + + @abstractmethod + def get_complete_url( + self, + api_base: Optional[str], + litellm_params: Dict[str, Any], + ) -> str: + """Return the full URL for POST /agents (create).""" + + @abstractmethod + def validate_environment( + self, + headers: Dict[str, str], + litellm_params: Dict[str, Any], + ) -> Dict[str, str]: + """Validate credentials and return auth headers.""" + + @abstractmethod + def transform_create_request( + self, + name: str, + litellm_params: Dict[str, Any], + ) -> Dict[str, Any]: + """Map name + litellm_params to the provider's create-agent body.""" + + @abstractmethod + def transform_create_response( + self, + raw_response: httpx.Response, + name: str, + ) -> AgentCreateResponse: + """Parse create response. Raise on non-2xx.""" + + # ------------------------------------------------------------------ # + # LIST # + # ------------------------------------------------------------------ # + + @abstractmethod + def transform_list_request( + self, + api_base: Optional[str], + litellm_params: Dict[str, Any], + ) -> Tuple[str, Dict[str, Any]]: + """Return (url, query_params) for GET /agents.""" + + @abstractmethod + def transform_list_response( + self, + raw_response: httpx.Response, + ) -> AgentListResponse: + """Parse list-agents response. Raise on non-2xx.""" + + # ------------------------------------------------------------------ # + # GET # + # ------------------------------------------------------------------ # + + @abstractmethod + def transform_get_request( + self, + name: str, + api_base: Optional[str], + litellm_params: Dict[str, Any], + ) -> Tuple[str, Dict[str, Any]]: + """Return (url, query_params) for GET /agents/{name}.""" + + @abstractmethod + def transform_get_response( + self, + raw_response: httpx.Response, + name: str, + ) -> AgentCreateResponse: + """Parse get-agent response. Raise on non-2xx.""" + + # ------------------------------------------------------------------ # + # DELETE # + # ------------------------------------------------------------------ # + + @abstractmethod + def transform_delete_request( + self, + name: str, + api_base: Optional[str], + litellm_params: Dict[str, Any], + ) -> str: + """Return the URL for DELETE /agents/{name}.""" + + @abstractmethod + def transform_delete_response( + self, + raw_response: httpx.Response, + name: str, + ) -> AgentDeleteResult: + """Parse delete-agent response. Raise on non-2xx.""" + + # ------------------------------------------------------------------ # + # LIST VERSIONS # + # ------------------------------------------------------------------ # + + @abstractmethod + def transform_list_versions_request( + self, + name: str, + api_base: Optional[str], + litellm_params: Dict[str, Any], + ) -> Tuple[str, Dict[str, Any]]: + """Return (url, query_params) for GET /agents/{name}/versions.""" + + @abstractmethod + def transform_list_versions_response( + self, + raw_response: httpx.Response, + name: str, + ) -> AgentVersionsResponse: + """Parse list-versions response. Raise on non-2xx.""" + + # ------------------------------------------------------------------ # + # ERROR HANDLING # + # ------------------------------------------------------------------ # + + def get_error_class( + self, + error_message: str, + status_code: int, + headers: Union[dict, httpx.Headers], + ) -> Exception: + """Map HTTP error status codes to provider-specific exceptions.""" + from litellm.llms.base_llm.chat.transformation import BaseLLMException + + return BaseLLMException( + status_code=status_code, + message=error_message, + headers=headers, + ) diff --git a/litellm/llms/gemini/agents/__init__.py b/litellm/llms/gemini/agents/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/gemini/agents/transformation.py b/litellm/llms/gemini/agents/transformation.py new file mode 100644 index 00000000000..150918c4737 --- /dev/null +++ b/litellm/llms/gemini/agents/transformation.py @@ -0,0 +1,299 @@ +""" +Google AI Studio Agents API configuration. + +Proxies the Gemini v1beta Agents API: + POST /v1beta/agents create + GET /v1beta/agents list + GET /v1beta/agents/{name} get + DELETE /v1beta/agents/{name} delete + GET /v1beta/agents/{name}/versions list versions +""" + +from typing import Any, Dict, Optional, Tuple, Union + +import httpx + +from litellm._logging import verbose_logger +from litellm.llms.base_llm.agents.transformation import BaseAgentsAPIConfig +from litellm.llms.gemini.common_utils import GeminiError, GeminiModelInfo +from litellm.types.agents import ( + AgentCreateResponse, + AgentDeleteResult, + AgentListResponse, + AgentVersionsResponse, +) + + +# Keys inside litellm_params that should be forwarded to the Gemini +# create-agent body verbatim. +_GEMINI_AGENT_BODY_KEYS = ("base_agent", "instructions", "base_environment") + +# LiteLLM-internal keys that must never be forwarded to Gemini. +_LITELLM_INTERNAL_KEYS = frozenset( + { + "custom_llm_provider", + "api_key", + "api_base", + "make_public", + "cost_per_query", + "input_cost_per_token", + "output_cost_per_token", + "require_trace_id_on_calls_to_agent", + "require_trace_id_on_calls_by_agent", + "max_iterations", + "max_budget_per_session", + "guardrails", + "is_public", + "agent_name", + "agent_id", + "agent_card_params", + "provider_agent_response", + } +) + + +class GeminiAgentsConfig(BaseAgentsAPIConfig): + """ + Configuration for the Google AI Studio (Gemini) native Agents API. + + Authentication uses x-goog-api-key, resolved from (in order): + 1. litellm_params["api_key"] + 2. GOOGLE_API_KEY env var + 3. GEMINI_API_KEY env var + """ + + @property + def api_version(self) -> str: + return "v1beta" + + def _base_url(self, api_base: Optional[str]) -> str: + return f"{GeminiModelInfo.get_api_base(api_base)}/{self.api_version}" + + # ------------------------------------------------------------------ # + # Shared helpers # + # ------------------------------------------------------------------ # + + def get_error_class( + self, + error_message: str, + status_code: int, + headers: Union[dict, httpx.Headers], + ) -> Exception: + return GeminiError( + message=error_message, + status_code=status_code, + headers=dict(headers), + ) + + def get_complete_url( + self, + api_base: Optional[str], + litellm_params: Dict[str, Any], + ) -> str: + return f"{self._base_url(api_base)}/agents" + + def validate_environment( + self, + headers: Dict[str, str], + litellm_params: Dict[str, Any], + ) -> Dict[str, str]: + headers = dict(headers) + headers["Content-Type"] = "application/json" + explicit_api_key = litellm_params.get("api_key") + # SECURITY: when the caller overrides ``api_base``, refuse to fall back + # to the process-wide GOOGLE_API_KEY / GEMINI_API_KEY env vars. Otherwise + # an authenticated proxy user could set ``api_base`` to an attacker- + # controlled host and have the proxy ship its shared Gemini key in the + # ``x-goog-api-key`` header. + if litellm_params.get("api_base") and not explicit_api_key: + raise ValueError( + "When overriding api_base for Gemini agents, you must also " + "supply an explicit api_key. Falling back to GOOGLE_API_KEY / " + "GEMINI_API_KEY env vars with a custom api_base is refused " + "to prevent leaking the shared provider key to arbitrary hosts." + ) + api_key = GeminiModelInfo.get_api_key(explicit_api_key) + if not api_key: + raise ValueError( + "Google API key is required. " + "Set GOOGLE_API_KEY or GEMINI_API_KEY, or pass api_key." + ) + headers["x-goog-api-key"] = api_key + return headers + + def _raise_for_status(self, raw_response: httpx.Response) -> None: + if not (200 <= raw_response.status_code < 300): + raise GeminiError( + message=raw_response.text, + status_code=raw_response.status_code, + headers=dict(raw_response.headers), + ) + + # ------------------------------------------------------------------ # + # CREATE # + # ------------------------------------------------------------------ # + + def transform_create_request( + self, + name: str, + litellm_params: Dict[str, Any], + ) -> Dict[str, Any]: + body: Dict[str, Any] = {"name": name} + for key in _GEMINI_AGENT_BODY_KEYS: + value = litellm_params.get(key) + if value is not None: + body[key] = value + verbose_logger.debug("GeminiAgentsConfig create body: %s", body) + return body + + def transform_create_response( + self, + raw_response: httpx.Response, + name: str, + ) -> AgentCreateResponse: + """ + Gemini returns: + {"id": "my-agent", "base_agent": "waverunner", + "system_instruction": "...", "base_environment": {...}} + """ + self._raise_for_status(raw_response) + try: + data: Dict[str, Any] = raw_response.json() + except Exception: + verbose_logger.warning( + "GeminiAgentsConfig: non-JSON create response (status=%d).", + raw_response.status_code, + ) + data = {"id": name} + # Gemini uses "id" as the identifier; normalise to both fields. + data.setdefault("id", name) + data.setdefault("name", data["id"]) + verbose_logger.debug("GeminiAgentsConfig create response: %s", data) + return AgentCreateResponse(**data) + + # ------------------------------------------------------------------ # + # LIST # + # ------------------------------------------------------------------ # + + def transform_list_request( + self, + api_base: Optional[str], + litellm_params: Dict[str, Any], + ) -> Tuple[str, Dict[str, Any]]: + url = f"{self._base_url(api_base)}/agents" + params: Dict[str, Any] = {} + if litellm_params.get("page_size"): + params["pageSize"] = litellm_params["page_size"] + if litellm_params.get("page_token"): + params["pageToken"] = litellm_params["page_token"] + return url, params + + def transform_list_response( + self, + raw_response: httpx.Response, + ) -> AgentListResponse: + self._raise_for_status(raw_response) + try: + data = raw_response.json() + except Exception: + data = {} + verbose_logger.debug("GeminiAgentsConfig list response: %s", data) + return AgentListResponse( + agents=data.get("agents", []), + next_page_token=data.get("nextPageToken"), + ) + + # ------------------------------------------------------------------ # + # GET # + # ------------------------------------------------------------------ # + + def transform_get_request( + self, + name: str, + api_base: Optional[str], + litellm_params: Dict[str, Any], + ) -> Tuple[str, Dict[str, Any]]: + url = f"{self._base_url(api_base)}/agents/{name}" + return url, {} + + def transform_get_response( + self, + raw_response: httpx.Response, + name: str, + ) -> AgentCreateResponse: + """Same shape as create response — Gemini returns "id" as identifier.""" + self._raise_for_status(raw_response) + try: + data = raw_response.json() + except Exception: + data = {"id": name} + data.setdefault("id", name) + data.setdefault("name", data["id"]) + verbose_logger.debug("GeminiAgentsConfig get response: %s", data) + return AgentCreateResponse(**data) + + # ------------------------------------------------------------------ # + # DELETE # + # ------------------------------------------------------------------ # + + def transform_delete_request( + self, + name: str, + api_base: Optional[str], + litellm_params: Dict[str, Any], + ) -> str: + return f"{self._base_url(api_base)}/agents/{name}" + + def transform_delete_response( + self, + raw_response: httpx.Response, + name: str, + ) -> AgentDeleteResult: + """Gemini returns an empty body ``{}`` with HTTP 200 on success.""" + self._raise_for_status(raw_response) + verbose_logger.debug( + "GeminiAgentsConfig delete (status=%d) agent '%s'", + raw_response.status_code, + name, + ) + return AgentDeleteResult(name=name, deleted=True) + + # ------------------------------------------------------------------ # + # LIST VERSIONS # + # ------------------------------------------------------------------ # + + def transform_list_versions_request( + self, + name: str, + api_base: Optional[str], + litellm_params: Dict[str, Any], + ) -> Tuple[str, Dict[str, Any]]: + url = f"{self._base_url(api_base)}/agents/{name}/versions" + params: Dict[str, Any] = {} + if litellm_params.get("page_size"): + params["pageSize"] = litellm_params["page_size"] + if litellm_params.get("page_token"): + params["pageToken"] = litellm_params["page_token"] + return url, params + + def transform_list_versions_response( + self, + raw_response: httpx.Response, + name: str, + ) -> AgentVersionsResponse: + """ + Gemini returns: + {"agentVersions": [{"agent": "waverunner", "name": "agents/.../versions/uuid", ...}]} + """ + self._raise_for_status(raw_response) + try: + data = raw_response.json() + except Exception: + data = {} + verbose_logger.debug( + "GeminiAgentsConfig list_versions response for '%s': %s", name, data + ) + return AgentVersionsResponse( + agent_versions=data.get("agentVersions", []), + next_page_token=data.get("nextPageToken"), + ) diff --git a/litellm/llms/gemini/interactions/transformation.py b/litellm/llms/gemini/interactions/transformation.py index 593cbf7c2cf..73435c8db6a 100644 --- a/litellm/llms/gemini/interactions/transformation.py +++ b/litellm/llms/gemini/interactions/transformation.py @@ -64,6 +64,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig): "stream", "store", "background", + "environment", "response_modalities", "response_format", "response_mime_type", @@ -142,6 +143,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig): "stream", "store", "background", + "environment", "response_modalities", "response_format", "response_mime_type", diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index 9f034575222..a70c5b3a920 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -84,6 +84,11 @@ LAZY_FEATURES: Tuple[LazyFeature, ...] = ( module_path="litellm.proxy.agent_endpoints.endpoints", path_prefixes=("/v1/agents", "/agents", "/agent/"), ), + LazyFeature( + name="gemini_agents", + module_path="litellm.proxy.google_endpoints.agents_endpoints", + path_prefixes=("/v1beta/agents",), + ), LazyFeature( name="a2a", module_path="litellm.proxy.agent_endpoints.a2a_endpoints", diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index c989f5dff13..9337aa7c8ea 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -481,6 +481,10 @@ class LiteLLMRoutes(enum.Enum): "/v1beta/interactions/{interaction_id}", "/interactions/{interaction_id}/cancel", "/v1beta/interactions/{interaction_id}/cancel", + # Google Managed Agents API + "/v1beta/agents", + "/v1beta/agents/{name}", + "/v1beta/agents/{name}/versions", ] apply_guardrail_routes = [ diff --git a/litellm/proxy/agent_endpoints/utils.py b/litellm/proxy/agent_endpoints/utils.py index 2b968de54be..393f5934fd9 100644 --- a/litellm/proxy/agent_endpoints/utils.py +++ b/litellm/proxy/agent_endpoints/utils.py @@ -2,6 +2,12 @@ from typing import Dict, Mapping, Optional +# Re-export from the canonical SDK location so the proxy and SDK always +# share the same provider-config lookup logic. +from litellm.interactions.agents.utils import ( # noqa: F401 + get_provider_agents_api_config, +) + def merge_agent_headers( *, diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 71413668e23..0c7df1d7216 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -807,6 +807,11 @@ class ProxyBaseLLMRequestProcessing: "aget_interaction", "adelete_interaction", "acancel_interaction", + "acreate_agent", + "alist_agents", + "aget_agent", + "adelete_agent", + "alist_agent_versions", "asend_message", "call_mcp_tool", "acreate_eval", @@ -1074,6 +1079,11 @@ class ProxyBaseLLMRequestProcessing: "aget_interaction", "adelete_interaction", "acancel_interaction", + "acreate_agent", + "alist_agents", + "aget_agent", + "adelete_agent", + "alist_agent_versions", "asend_message", "call_mcp_tool", "acreate_eval", diff --git a/litellm/proxy/google_endpoints/agents_endpoints.py b/litellm/proxy/google_endpoints/agents_endpoints.py new file mode 100644 index 00000000000..779284023a0 --- /dev/null +++ b/litellm/proxy/google_endpoints/agents_endpoints.py @@ -0,0 +1,445 @@ +""" +Google AI Studio Managed Agents API Proxy Endpoints. + +Exposes Gemini's /v1beta/agents surface through the LiteLLM proxy so that +user curl commands transfer 1-to-1 by swapping the host + auth header. + +Routes: + POST /v1beta/agents -> acreate_agent + GET /v1beta/agents -> alist_agents + GET /v1beta/agents/{name} -> aget_agent + DELETE /v1beta/agents/{name} -> adelete_agent + GET /v1beta/agents/{name}/versions -> alist_agent_versions + +These are distinct from the A2A agent registry at /v1/agents. +""" + +import json + +from fastapi import APIRouter, Depends, HTTPException, Request, Response, status +from fastapi.responses import ORJSONResponse + +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing +from litellm.proxy.common_utils.http_parsing_utils import ( + _read_request_body, + _safe_get_request_query_params, +) + +router = APIRouter(tags=["gemini managed agents"]) + + +def _is_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> bool: + return ( + user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN + or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value + ) + + +def _enforce_caller_supplied_provider_key( + data: dict, + user_api_key_dict: UserAPIKeyAuth, +) -> None: + """ + SECURITY: refuse to use the proxy's shared GOOGLE_API_KEY / GEMINI_API_KEY + env fallback for non-admin callers on Gemini managed-agent CRUD endpoints. + + These endpoints are part of ``llm_api_routes`` so any authenticated LLM key + can reach them, but unlike ``/v1beta/models/...:generateContent`` they are + *not* routed through ``model_list`` — the only credential source is either + the per-request ``litellm_params_template`` or the env var fallback. Without + this guard, any ordinary proxy user could list, create, or delete managed + agents inside the operator's Gemini project using the operator's key. + + Proxy admins (master key) keep the env-fallback convenience for ops use. + """ + if _is_proxy_admin(user_api_key_dict): + return + if data.get("api_key"): + return + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=( + "Gemini managed-agent endpoints require a caller-supplied " + "Gemini api_key (via 'litellm_params_template'). Falling back to " + "the proxy's GOOGLE_API_KEY / GEMINI_API_KEY env vars is only " + "permitted for proxy admins." + ), + ) + + +def _merge_query_params_into_data(data: dict, request: Request) -> dict: + """ + For GET/DELETE endpoints that cannot carry a JSON body, read a + JSON-encoded ``litellm_params_template`` query parameter and merge its + contents into *data*, without overwriting keys that are already present + (e.g. path params like ``name`` or the fixed ``custom_llm_provider``). + + This mirrors the ``litellm_params_template`` handling in + ``create_gemini_agent`` and is the supported way for multi-tenant + callers to supply per-request credentials on non-POST endpoints: + + .. code-block:: bash + + curl "http://localhost:4000/v1beta/agents?litellm_params_template=%7B%22api_key%22%3A%22AIza...%22%7D" \\ + -H "Authorization: Bearer sk-..." + + Credentials MUST NOT be passed as plain flat query parameters (e.g. + ``?api_key=AIza...``) because URL query strings appear verbatim in + web-server access logs, CDN edge logs, browser history, and Referer + headers. Use the ``litellm_params_template`` JSON body field on POST + requests, or the JSON-encoded query parameter above for GET/DELETE. + """ + query_params = _safe_get_request_query_params(request) + if not query_params: + return data + + raw_template = query_params.get("litellm_params_template") + if raw_template: + try: + template = ( + json.loads(raw_template) + if isinstance(raw_template, str) + else raw_template + ) + except (json.JSONDecodeError, ValueError): + template = {} + if isinstance(template, dict): + for key, value in template.items(): + data.setdefault(key, value) + + return data + + +def _proxy_server_imports(): + from litellm.proxy.proxy_server import ( # noqa: PLC0415 + general_settings, + llm_router, + proxy_config, + proxy_logging_obj, + select_data_generator, + user_api_base, + user_max_tokens, + user_model, + user_request_timeout, + user_temperature, + version, + ) + + return dict( + general_settings=general_settings, + llm_router=llm_router, + proxy_config=proxy_config, + proxy_logging_obj=proxy_logging_obj, + select_data_generator=select_data_generator, + user_api_base=user_api_base, + user_max_tokens=user_max_tokens, + user_model=user_model, + user_request_timeout=user_request_timeout, + user_temperature=user_temperature, + version=version, + ) + + +@router.post( + "/v1beta/agents", + dependencies=[Depends(user_api_key_auth)], + response_class=ORJSONResponse, +) +async def create_gemini_agent( + request: Request, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Create a named custom agent on the Gemini side. + + Example: + ```bash + curl -X POST "http://localhost:4000/v1beta/agents" \\ + -H "Authorization: Bearer sk-..." \\ + -H "Content-Type: application/json" \\ + -d '{ + "name": "my-custom-slides-agent", + "base_agent": "waverunner", + "instructions": "You are a helpful assistant that creates slides.", + "base_environment": { + "type": "remote", + "sources": [ + {"type": "gcs", "source": "gs://eap-templates/slides-skill", + "target": "/.agents/skills/slides-skill"} + ] + } + }' + ``` + """ + srv = _proxy_server_imports() + data = await _read_request_body(request=request) + # Merge litellm_params_template (e.g. custom_llm_provider, api_key) into the request + litellm_params_template = data.pop("litellm_params_template", None) or {} + if isinstance(litellm_params_template, dict): + for key, value in litellm_params_template.items(): + if key not in data: + data[key] = value + data.setdefault("custom_llm_provider", "gemini") + _enforce_caller_supplied_provider_key(data, user_api_key_dict) + + processor = ProxyBaseLLMRequestProcessing(data=data) + try: + return await processor.base_process_llm_request( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + route_type="acreate_agent", + proxy_logging_obj=srv["proxy_logging_obj"], + llm_router=srv["llm_router"], + general_settings=srv["general_settings"], + proxy_config=srv["proxy_config"], + select_data_generator=srv["select_data_generator"], + model=None, + user_model=srv["user_model"], + user_temperature=srv["user_temperature"], + user_request_timeout=srv["user_request_timeout"], + user_max_tokens=srv["user_max_tokens"], + user_api_base=srv["user_api_base"], + version=srv["version"], + ) + except Exception as e: + raise await processor._handle_llm_api_exception( + e=e, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=srv["proxy_logging_obj"], + version=srv["version"], + ) + + +@router.get( + "/v1beta/agents", + dependencies=[Depends(user_api_key_auth)], + response_class=ORJSONResponse, +) +async def list_gemini_agents( + request: Request, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + List all custom agents on the Gemini side. + + Pass per-request Gemini credentials via the JSON-encoded + ``litellm_params_template`` query parameter. Flat query parameters + (e.g. ``?api_key=AIza...``) are intentionally ignored — see + ``_merge_query_params_into_data`` for the rationale. + + ```bash + curl "http://localhost:4000/v1beta/agents?litellm_params_template=%7B%22api_key%22%3A%22AIza...%22%7D" \\ + -H "Authorization: Bearer sk-..." + ``` + """ + srv = _proxy_server_imports() + data: dict = {"custom_llm_provider": "gemini"} + _merge_query_params_into_data(data, request) + _enforce_caller_supplied_provider_key(data, user_api_key_dict) + + processor = ProxyBaseLLMRequestProcessing(data=data) + try: + return await processor.base_process_llm_request( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + route_type="alist_agents", + proxy_logging_obj=srv["proxy_logging_obj"], + llm_router=srv["llm_router"], + general_settings=srv["general_settings"], + proxy_config=srv["proxy_config"], + select_data_generator=srv["select_data_generator"], + model=None, + user_model=srv["user_model"], + user_temperature=srv["user_temperature"], + user_request_timeout=srv["user_request_timeout"], + user_max_tokens=srv["user_max_tokens"], + user_api_base=srv["user_api_base"], + version=srv["version"], + ) + except Exception as e: + raise await processor._handle_llm_api_exception( + e=e, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=srv["proxy_logging_obj"], + version=srv["version"], + ) + + +@router.get( + "/v1beta/agents/{name}", + dependencies=[Depends(user_api_key_auth)], + response_class=ORJSONResponse, +) +async def get_gemini_agent( + request: Request, + name: str, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Get a specific custom agent by name. + + Pass per-request Gemini credentials via the JSON-encoded + ``litellm_params_template`` query parameter. Flat query parameters + (e.g. ``?api_key=AIza...``) are intentionally ignored — see + ``_merge_query_params_into_data`` for the rationale. + + ```bash + curl "http://localhost:4000/v1beta/agents/my-custom-slides-agent?litellm_params_template=%7B%22api_key%22%3A%22AIza...%22%7D" \\ + -H "Authorization: Bearer sk-..." + ``` + """ + srv = _proxy_server_imports() + data = {"name": name, "custom_llm_provider": "gemini"} + _merge_query_params_into_data(data, request) + _enforce_caller_supplied_provider_key(data, user_api_key_dict) + + processor = ProxyBaseLLMRequestProcessing(data=data) + try: + return await processor.base_process_llm_request( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + route_type="aget_agent", + proxy_logging_obj=srv["proxy_logging_obj"], + llm_router=srv["llm_router"], + general_settings=srv["general_settings"], + proxy_config=srv["proxy_config"], + select_data_generator=srv["select_data_generator"], + model=None, + user_model=srv["user_model"], + user_temperature=srv["user_temperature"], + user_request_timeout=srv["user_request_timeout"], + user_max_tokens=srv["user_max_tokens"], + user_api_base=srv["user_api_base"], + version=srv["version"], + ) + except Exception as e: + raise await processor._handle_llm_api_exception( + e=e, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=srv["proxy_logging_obj"], + version=srv["version"], + ) + + +@router.delete( + "/v1beta/agents/{name}", + dependencies=[Depends(user_api_key_auth)], + response_class=ORJSONResponse, +) +async def delete_gemini_agent( + request: Request, + name: str, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Delete a custom agent by name. + + Pass per-request Gemini credentials via the JSON-encoded + ``litellm_params_template`` query parameter. Flat query parameters + (e.g. ``?api_key=AIza...``) are intentionally ignored — see + ``_merge_query_params_into_data`` for the rationale. + + ```bash + curl -X DELETE "http://localhost:4000/v1beta/agents/my-custom-slides-agent?litellm_params_template=%7B%22api_key%22%3A%22AIza...%22%7D" \\ + -H "Authorization: Bearer sk-..." + ``` + """ + srv = _proxy_server_imports() + data = {"name": name, "custom_llm_provider": "gemini"} + _merge_query_params_into_data(data, request) + _enforce_caller_supplied_provider_key(data, user_api_key_dict) + + processor = ProxyBaseLLMRequestProcessing(data=data) + try: + return await processor.base_process_llm_request( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + route_type="adelete_agent", + proxy_logging_obj=srv["proxy_logging_obj"], + llm_router=srv["llm_router"], + general_settings=srv["general_settings"], + proxy_config=srv["proxy_config"], + select_data_generator=srv["select_data_generator"], + model=None, + user_model=srv["user_model"], + user_temperature=srv["user_temperature"], + user_request_timeout=srv["user_request_timeout"], + user_max_tokens=srv["user_max_tokens"], + user_api_base=srv["user_api_base"], + version=srv["version"], + ) + except Exception as e: + raise await processor._handle_llm_api_exception( + e=e, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=srv["proxy_logging_obj"], + version=srv["version"], + ) + + +@router.get( + "/v1beta/agents/{name}/versions", + dependencies=[Depends(user_api_key_auth)], + response_class=ORJSONResponse, +) +async def list_gemini_agent_versions( + request: Request, + name: str, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + List versions of a custom agent. + + Pass per-request Gemini credentials via the JSON-encoded + ``litellm_params_template`` query parameter. Flat query parameters + (e.g. ``?api_key=AIza...``) are intentionally ignored — see + ``_merge_query_params_into_data`` for the rationale. + + ```bash + curl "http://localhost:4000/v1beta/agents/my-custom-slides-agent/versions?litellm_params_template=%7B%22api_key%22%3A%22AIza...%22%7D" \\ + -H "Authorization: Bearer sk-..." + ``` + """ + srv = _proxy_server_imports() + data = {"name": name, "custom_llm_provider": "gemini"} + _merge_query_params_into_data(data, request) + _enforce_caller_supplied_provider_key(data, user_api_key_dict) + + processor = ProxyBaseLLMRequestProcessing(data=data) + try: + return await processor.base_process_llm_request( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + route_type="alist_agent_versions", + proxy_logging_obj=srv["proxy_logging_obj"], + llm_router=srv["llm_router"], + general_settings=srv["general_settings"], + proxy_config=srv["proxy_config"], + select_data_generator=srv["select_data_generator"], + model=None, + user_model=srv["user_model"], + user_temperature=srv["user_temperature"], + user_request_timeout=srv["user_request_timeout"], + user_max_tokens=srv["user_max_tokens"], + user_api_base=srv["user_api_base"], + version=srv["version"], + ) + except Exception as e: + raise await processor._handle_llm_api_exception( + e=e, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=srv["proxy_logging_obj"], + version=srv["version"], + ) diff --git a/litellm/proxy/google_endpoints/endpoints.py b/litellm/proxy/google_endpoints/endpoints.py index 967ac9f0ac4..1f503247bf4 100644 --- a/litellm/proxy/google_endpoints/endpoints.py +++ b/litellm/proxy/google_endpoints/endpoints.py @@ -285,7 +285,7 @@ async def create_interaction( general_settings=general_settings, proxy_config=proxy_config, select_data_generator=select_data_generator, - model=data.get("model") or data.get("agent"), + model=data.get("model"), user_model=user_model, user_temperature=user_temperature, user_request_timeout=user_request_timeout, diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 06fd35448f1..8f6f7084a0c 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -94,6 +94,12 @@ ROUTE_ENDPOINT_MAPPING = { "aget_interaction": "/interactions/{interaction_id}", "adelete_interaction": "/interactions/{interaction_id}", "acancel_interaction": "/interactions/{interaction_id}/cancel", + # Google Managed Agents API routes + "acreate_agent": "/v1beta/agents", + "alist_agents": "/v1beta/agents", + "aget_agent": "/v1beta/agents/{name}", + "adelete_agent": "/v1beta/agents/{name}", + "alist_agent_versions": "/v1beta/agents/{name}/versions", # OpenAI Evals API routes "acreate_eval": "/evals", "alist_evals": "/evals", @@ -311,6 +317,11 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin "aget_interaction", "adelete_interaction", "acancel_interaction", + "acreate_agent", + "alist_agents", + "aget_agent", + "adelete_agent", + "alist_agent_versions", "asend_message", "call_mcp_tool", "acancel_batch", @@ -468,6 +479,15 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin "acancel_interaction", ]: return getattr(llm_router, f"{route_type}")(**data) + # Managed Agents API: these don't need model routing + if route_type in [ + "acreate_agent", + "alist_agents", + "aget_agent", + "adelete_agent", + "alist_agent_versions", + ]: + return getattr(llm_router, f"{route_type}")(**data) if route_type in [ "avideo_list", "avideo_status", diff --git a/litellm/router.py b/litellm/router.py index a51d9d2ad26..69823b4bd0e 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -1579,6 +1579,44 @@ class Router: cancel_interaction, call_type="cancel_interaction" ) + def _initialize_managed_agents_endpoints(self): + """Initialize Google Managed Agents API endpoints (v1beta/agents).""" + from litellm.interactions.agents import acreate as acreate_agent + from litellm.interactions.agents import adelete as adelete_agent + from litellm.interactions.agents import aget as aget_agent + from litellm.interactions.agents import alist as alist_agents + from litellm.interactions.agents import alist_versions as alist_agent_versions + from litellm.interactions.agents import create as create_agent + from litellm.interactions.agents import delete as delete_agent + from litellm.interactions.agents import get as get_agent + from litellm.interactions.agents import list as list_agents + from litellm.interactions.agents import list_versions as list_agent_versions + + self.acreate_agent = self.factory_function( + acreate_agent, call_type="acreate_agent" + ) + self.create_agent = self.factory_function( + create_agent, call_type="create_agent" + ) + self.alist_agents = self.factory_function( + alist_agents, call_type="alist_agents" + ) + self.list_agents = self.factory_function(list_agents, call_type="list_agents") + self.aget_agent = self.factory_function(aget_agent, call_type="aget_agent") + self.get_agent = self.factory_function(get_agent, call_type="get_agent") + self.adelete_agent = self.factory_function( + adelete_agent, call_type="adelete_agent" + ) + self.delete_agent = self.factory_function( + delete_agent, call_type="delete_agent" + ) + self.alist_agent_versions = self.factory_function( + alist_agent_versions, call_type="alist_agent_versions" + ) + self.list_agent_versions = self.factory_function( + list_agent_versions, call_type="list_agent_versions" + ) + def _initialize_specialized_endpoints(self): """Helper to initialize specialized router endpoints (vector store, OCR, search, video, container, skills, interactions).""" self._initialize_vector_store_endpoints() @@ -1591,6 +1629,7 @@ class Router: self._initialize_container_endpoints() self._initialize_skills_endpoints() self._initialize_interactions_endpoints() + self._initialize_managed_agents_endpoints() def initialize_router_endpoints(self): self._initialize_core_endpoints() @@ -5361,6 +5400,16 @@ class Router: "delete_interaction", "acancel_interaction", "cancel_interaction", + "acreate_agent", + "create_agent", + "alist_agents", + "list_agents", + "aget_agent", + "get_agent", + "adelete_agent", + "delete_agent", + "alist_agent_versions", + "list_agent_versions", ] = "assistants", ): """ @@ -5445,6 +5494,27 @@ class Router: return vector_store_file_sync_wrapper + if call_type in ( + "create_agent", + "list_agents", + "get_agent", + "delete_agent", + "list_agent_versions", + ): + + def managed_agents_sync_wrapper( + custom_llm_provider: Optional[str] = None, + client: Optional[Any] = None, + **kwargs, + ): + if custom_llm_provider and "custom_llm_provider" not in kwargs: + kwargs["custom_llm_provider"] = custom_llm_provider + if "custom_llm_provider" not in kwargs: + kwargs["custom_llm_provider"] = "gemini" + return original_function(**kwargs) + + return managed_agents_sync_wrapper + # Handle asynchronous call types async def async_wrapper( custom_llm_provider: Optional[str] = None, @@ -5508,8 +5578,6 @@ class Router: "alist_skills", "aget_skill", "adelete_skill", - "acreate_interaction", - "create_interaction", ): return await self._ageneric_api_call_with_fallbacks( original_function=original_function, @@ -5569,6 +5637,8 @@ class Router: **kwargs, ) elif call_type in ( + "acreate_interaction", + "create_interaction", "aget_interaction", "adelete_interaction", "acancel_interaction", @@ -5578,6 +5648,18 @@ class Router: custom_llm_provider=custom_llm_provider, **kwargs, ) + elif call_type in ( + "acreate_agent", + "alist_agents", + "aget_agent", + "adelete_agent", + "alist_agent_versions", + ): + return await self._init_managed_agents_api_endpoints( + original_function=original_function, + custom_llm_provider=custom_llm_provider, + **kwargs, + ) return async_wrapper @@ -5682,6 +5764,34 @@ class Router: if custom_llm_provider and "custom_llm_provider" not in kwargs: kwargs["custom_llm_provider"] = custom_llm_provider # Default to gemini for interactions API + if "custom_llm_provider" not in kwargs: + kwargs["custom_llm_provider"] = "gemini" + # If the proxy accidentally passed agent name as model, clear it + if kwargs.get("agent") and kwargs.get("model") == kwargs.get("agent"): + kwargs["model"] = None + # Model-based interactions use deployment routing + fallbacks; agent-only calls + # must not enter model-group lookup (agent name is not a LiteLLM deployment). + if kwargs.get("model"): + return await self._ageneric_api_call_with_fallbacks( + original_function=original_function, + **kwargs, + ) + return await original_function(**kwargs) + + async def _init_managed_agents_api_endpoints( + self, + original_function: Callable, + custom_llm_provider: Optional[str] = None, + **kwargs, + ): + """ + Initialize the Managed Agents API endpoints on the router (v1beta/agents). + + CRUD operations for Gemini managed agents don't need model-based routing, + so we call the original function directly with the custom_llm_provider. + """ + if custom_llm_provider and "custom_llm_provider" not in kwargs: + kwargs["custom_llm_provider"] = custom_llm_provider if "custom_llm_provider" not in kwargs: kwargs["custom_llm_provider"] = "gemini" return await original_function(**kwargs) diff --git a/litellm/types/agents.py b/litellm/types/agents.py index efb2e73bfb5..8556b6bac93 100644 --- a/litellm/types/agents.py +++ b/litellm/types/agents.py @@ -228,6 +228,66 @@ class ListAgentsResponse(BaseModel): agents: List[AgentResponse] +class AgentCreateResponse(LiteLLMPydanticObjectBase): + """ + Response from a provider-side agent creation or get call (e.g. Gemini v1beta/agents). + + Gemini returns ``"id"`` as the agent identifier; we surface both ``id`` + (Gemini's value) and ``name`` (the user-supplied name, equal to ``id`` for + Gemini) so callers can use either. All extra fields returned by the + provider (e.g. ``base_agent``, ``system_instruction``, ``base_environment``) + are preserved via extra="allow". + """ + + id: Optional[str] = None + name: Optional[str] = None + model_config = {"extra": "allow"} + + _hidden_params: dict = PrivateAttr(default_factory=dict) + + +class AgentDeleteResult(LiteLLMPydanticObjectBase): + """Result of a provider-side agent deletion (e.g. Gemini DELETE /v1beta/agents/{name}). + + Gemini returns an empty body ``{}`` on success; we synthesise ``name`` and + ``deleted`` so callers always get a consistent response object. + """ + + name: str + deleted: bool = True + model_config = {"extra": "allow"} + + _hidden_params: dict = PrivateAttr(default_factory=dict) + + +class AgentListResponse(LiteLLMPydanticObjectBase): + """Response from listing agents on the provider side (e.g. Gemini GET /v1beta/agents). + + Gemini returns ``{"agents": [{"id": "..."}, ...]}``; each item is kept as + a plain dict so no fields are silently dropped. + """ + + agents: List[Dict[str, Any]] = [] + next_page_token: Optional[str] = None + model_config = {"extra": "allow"} + + _hidden_params: dict = PrivateAttr(default_factory=dict) + + +class AgentVersionsResponse(LiteLLMPydanticObjectBase): + """Response from listing versions of an agent (e.g. Gemini GET /v1beta/agents/{name}/versions). + + Gemini returns ``{"agentVersions": [...]}``; each version has a ``name`` + field of the form ``agents/{agent_id}/versions/{uuid}``. + """ + + agent_versions: List[Dict[str, Any]] = [] + next_page_token: Optional[str] = None + model_config = {"extra": "allow"} + + _hidden_params: dict = PrivateAttr(default_factory=dict) + + class AgentMakePublicResponse(BaseModel): message: str public_agent_groups: List[str] diff --git a/litellm/types/interactions/__init__.py b/litellm/types/interactions/__init__.py index a3acdc4cb1f..0f934fa0152 100644 --- a/litellm/types/interactions/__init__.py +++ b/litellm/types/interactions/__init__.py @@ -37,6 +37,7 @@ from litellm.types.interactions.generated import ( ImageContent, Interaction, InteractionEvent, + InteractionEnvironment, InteractionInput, InteractionsAPIOptionalRequestParams, InteractionsAPIResponse, @@ -115,6 +116,7 @@ __all__ = [ "ResponseModality", "Annotation", # LiteLLM types + "InteractionEnvironment", "InteractionInput", "InteractionsAPIResponse", "InteractionsAPIStreamingResponse", diff --git a/litellm/types/interactions/generated.py b/litellm/types/interactions/generated.py index ed626b0b7c8..2ce6331b448 100644 --- a/litellm/types/interactions/generated.py +++ b/litellm/types/interactions/generated.py @@ -1257,3 +1257,6 @@ class CancelInteractionResult(BaseLiteLLMOpenAIResponseObject): InteractionTool = Tool InteractionToolChoiceConfig = ToolChoiceConfig InteractionsAPIOptionalRequestParams = Dict[str, Any] + +# Agent interaction execution environment +InteractionEnvironment = Union[str, Dict[str, Any]] diff --git a/tests/proxy_unit_tests/test_gemini_agents_endpoints.py b/tests/proxy_unit_tests/test_gemini_agents_endpoints.py new file mode 100644 index 00000000000..bdac9348f71 --- /dev/null +++ b/tests/proxy_unit_tests/test_gemini_agents_endpoints.py @@ -0,0 +1,519 @@ +""" +Unit tests for litellm/proxy/google_endpoints/agents_endpoints.py + +Focus: verify that list_gemini_agents, get_gemini_agent, delete_gemini_agent, +and list_gemini_agent_versions correctly forward per-request credentials +(api_key, api_base, …) supplied via the JSON-encoded litellm_params_template +query parameter. Flat credential query params (e.g. ?api_key=…) are no +longer accepted — they would appear in server logs. +""" + +import json +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi import Request +from fastapi.datastructures import Headers, QueryParams + +sys.path.insert(0, os.path.abspath("../..")) + +from litellm.proxy.google_endpoints.agents_endpoints import ( + _merge_query_params_into_data, +) + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _make_request(query_string: str = "") -> MagicMock: + """Build a minimal mock Request whose query_params match *query_string*.""" + req = MagicMock(spec=Request) + req.query_params = QueryParams(query_string) + req.headers = Headers({}) + return req + + +# --------------------------------------------------------------------------- +# _merge_query_params_into_data – unit tests for the helper +# --------------------------------------------------------------------------- + + +class TestMergeQueryParamsIntoData: + def test_no_query_params_leaves_data_unchanged(self): + data = {"custom_llm_provider": "gemini"} + request = _make_request("") + result = _merge_query_params_into_data(data, request) + assert result == {"custom_llm_provider": "gemini"} + + def test_flat_api_key_is_ignored(self): + """Flat credential params must NOT be merged (they leak into server logs).""" + data = {"custom_llm_provider": "gemini"} + request = _make_request("api_key=AIzaSyTest123") + _merge_query_params_into_data(data, request) + assert "api_key" not in data + assert data["custom_llm_provider"] == "gemini" + + def test_flat_params_are_silently_dropped(self): + """Flat params (including name injection attempts) are ignored entirely.""" + data = {"name": "my-agent", "custom_llm_provider": "gemini"} + request = _make_request("name=INJECTED&api_key=AIzaSyTest") + _merge_query_params_into_data(data, request) + assert data["name"] == "my-agent" + assert "api_key" not in data + + def test_litellm_params_template_json_is_expanded(self): + template = json.dumps( + {"api_key": "AIzaFromTemplate", "api_base": "https://example.com"} + ) + from urllib.parse import quote + + request = _make_request(f"litellm_params_template={quote(template)}") + data = {"custom_llm_provider": "gemini"} + _merge_query_params_into_data(data, request) + assert data["api_key"] == "AIzaFromTemplate" + assert data["api_base"] == "https://example.com" + # The raw template key itself must NOT appear in data + assert "litellm_params_template" not in data + + def test_litellm_params_template_does_not_overwrite_existing(self): + template = json.dumps( + {"api_key": "FromTemplate", "custom_llm_provider": "openai"} + ) + from urllib.parse import quote + + request = _make_request(f"litellm_params_template={quote(template)}") + data = {"custom_llm_provider": "gemini"} + _merge_query_params_into_data(data, request) + # custom_llm_provider was already set; template must not override it + assert data["custom_llm_provider"] == "gemini" + assert data["api_key"] == "FromTemplate" + + def test_invalid_litellm_params_template_json_is_ignored(self): + request = _make_request("litellm_params_template=NOT_VALID_JSON") + data = {"custom_llm_provider": "gemini"} + _merge_query_params_into_data(data, request) + # Bad JSON is silently skipped; other data stays intact + assert data == {"custom_llm_provider": "gemini"} + + def test_template_only_no_flat_params_merged(self): + """Only litellm_params_template is expanded; unknown flat params are dropped.""" + template = json.dumps({"api_key": "FromTemplate"}) + from urllib.parse import quote + + qs = f"litellm_params_template={quote(template)}&vertex_project=my-project" + request = _make_request(qs) + data = {"custom_llm_provider": "gemini"} + _merge_query_params_into_data(data, request) + assert data["api_key"] == "FromTemplate" + # flat vertex_project is ignored since it wasn't in litellm_params_template + assert "vertex_project" not in data + assert "litellm_params_template" not in data + + +# --------------------------------------------------------------------------- +# Endpoint-level smoke tests: data dict is populated before the processor call +# --------------------------------------------------------------------------- + + +@pytest.fixture +def mock_srv(): + """Patch _proxy_server_imports to return lightweight fakes.""" + srv = { + "general_settings": {}, + "llm_router": MagicMock(), + "proxy_config": MagicMock(), + "proxy_logging_obj": MagicMock(), + "select_data_generator": MagicMock(), + "user_api_base": None, + "user_max_tokens": None, + "user_model": None, + "user_request_timeout": None, + "user_temperature": None, + "version": "0.0.0", + } + with patch( + "litellm.proxy.google_endpoints.agents_endpoints._proxy_server_imports", + return_value=srv, + ): + yield srv + + +@pytest.fixture +def user_api_key_dict(): + from litellm.proxy._types import UserAPIKeyAuth + + return UserAPIKeyAuth(api_key="test-key") + + +def _make_endpoint_request(query_string: str = "") -> MagicMock: + req = MagicMock(spec=Request) + req.query_params = QueryParams(query_string) + req.headers = Headers({}) + req.scope = {} + + async def _body(): + return b"" + + req.body = _body + return req + + +@pytest.mark.asyncio +async def test_list_gemini_agents_passes_api_key_to_processor( + mock_srv, user_api_key_dict +): + from urllib.parse import quote + + from litellm.proxy.google_endpoints.agents_endpoints import list_gemini_agents + + template = json.dumps({"api_key": "AIzaListTest"}) + + with patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" + ) as MockProcessor: + instance = MockProcessor.return_value + instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) + + request = _make_endpoint_request(f"litellm_params_template={quote(template)}") + await list_gemini_agents( + request=request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + init_data = MockProcessor.call_args[1]["data"] + assert init_data.get("api_key") == "AIzaListTest" + assert init_data.get("custom_llm_provider") == "gemini" + + +@pytest.mark.asyncio +async def test_get_gemini_agent_passes_api_key_to_processor( + mock_srv, user_api_key_dict +): + from urllib.parse import quote + + from litellm.proxy.google_endpoints.agents_endpoints import get_gemini_agent + + template = json.dumps({"api_key": "AIzaGetTest"}) + + with patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" + ) as MockProcessor: + instance = MockProcessor.return_value + instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) + + request = _make_endpoint_request(f"litellm_params_template={quote(template)}") + await get_gemini_agent( + request=request, + name="my-agent", + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + init_data = MockProcessor.call_args[1]["data"] + assert init_data.get("api_key") == "AIzaGetTest" + assert init_data.get("name") == "my-agent" + assert init_data.get("custom_llm_provider") == "gemini" + + +@pytest.mark.asyncio +async def test_delete_gemini_agent_passes_api_key_to_processor( + mock_srv, user_api_key_dict +): + from urllib.parse import quote + + from litellm.proxy.google_endpoints.agents_endpoints import delete_gemini_agent + + template = json.dumps({"api_key": "AIzaDeleteTest"}) + + with patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" + ) as MockProcessor: + instance = MockProcessor.return_value + instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) + + request = _make_endpoint_request(f"litellm_params_template={quote(template)}") + await delete_gemini_agent( + request=request, + name="my-agent", + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + init_data = MockProcessor.call_args[1]["data"] + assert init_data.get("api_key") == "AIzaDeleteTest" + assert init_data.get("name") == "my-agent" + assert init_data.get("custom_llm_provider") == "gemini" + + +@pytest.mark.asyncio +async def test_list_gemini_agent_versions_passes_api_key_to_processor( + mock_srv, user_api_key_dict +): + from urllib.parse import quote + + from litellm.proxy.google_endpoints.agents_endpoints import ( + list_gemini_agent_versions, + ) + + template = json.dumps({"api_key": "AIzaVersionsTest"}) + + with patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" + ) as MockProcessor: + instance = MockProcessor.return_value + instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) + + request = _make_endpoint_request(f"litellm_params_template={quote(template)}") + await list_gemini_agent_versions( + request=request, + name="my-agent", + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + init_data = MockProcessor.call_args[1]["data"] + assert init_data.get("api_key") == "AIzaVersionsTest" + assert init_data.get("name") == "my-agent" + assert init_data.get("custom_llm_provider") == "gemini" + + +@pytest.mark.asyncio +async def test_get_gemini_agent_name_not_overwritten_by_query_param( + mock_srv, user_api_key_dict +): + """Path-param ``name`` must not be replaced by an attacker-controlled query param.""" + from urllib.parse import quote + + from litellm.proxy.google_endpoints.agents_endpoints import get_gemini_agent + + with patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" + ) as MockProcessor: + instance = MockProcessor.return_value + instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) + + # Even if a caller tries to inject "name" via flat query param, it is + # ignored (flat params are not merged). The path-param name wins. + # ``api_key`` is supplied via the JSON template (required for non-admin + # callers — see test_*_non_admin_without_api_key_is_rejected below). + template = json.dumps({"api_key": "AIzaTest"}) + request = _make_endpoint_request( + f"name=INJECTED&litellm_params_template={quote(template)}" + ) + await get_gemini_agent( + request=request, + name="real-agent", + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + init_data = MockProcessor.call_args[1]["data"] + assert init_data["name"] == "real-agent" + + +@pytest.mark.asyncio +async def test_list_agents_template_via_query_param(mock_srv, user_api_key_dict): + """litellm_params_template in query string is expanded.""" + from litellm.proxy.google_endpoints.agents_endpoints import list_gemini_agents + from urllib.parse import quote + + template = json.dumps({"api_key": "TemplateKey", "vertex_project": "proj-x"}) + + with patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" + ) as MockProcessor: + instance = MockProcessor.return_value + instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) + + request = _make_endpoint_request(f"litellm_params_template={quote(template)}") + await list_gemini_agents( + request=request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + init_data = MockProcessor.call_args[1]["data"] + assert init_data["api_key"] == "TemplateKey" + assert init_data["vertex_project"] == "proj-x" + assert "litellm_params_template" not in init_data + + +# --------------------------------------------------------------------------- +# Security guards (veria-flagged findings) +# --------------------------------------------------------------------------- + + +@pytest.fixture +def proxy_admin_user_api_key_dict(): + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + + return UserAPIKeyAuth( + api_key="sk-admin", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + + +@pytest.mark.asyncio +async def test_list_agents_non_admin_without_api_key_is_rejected( + mock_srv, user_api_key_dict +): + """Non-admin callers must supply an explicit api_key — the proxy must not + silently fall back to the operator's shared GOOGLE_API_KEY/GEMINI_API_KEY. + """ + from fastapi import HTTPException + + from litellm.proxy.google_endpoints.agents_endpoints import list_gemini_agents + + with patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" + ) as MockProcessor: + instance = MockProcessor.return_value + instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) + + request = _make_endpoint_request("") + with pytest.raises(HTTPException) as excinfo: + await list_gemini_agents( + request=request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + assert excinfo.value.status_code == 401 + # Processor must never be invoked + instance.base_process_llm_request.assert_not_called() + + +@pytest.mark.asyncio +async def test_delete_agent_non_admin_without_api_key_is_rejected( + mock_srv, user_api_key_dict +): + from fastapi import HTTPException + + from litellm.proxy.google_endpoints.agents_endpoints import delete_gemini_agent + + with patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" + ) as MockProcessor: + instance = MockProcessor.return_value + instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) + + request = _make_endpoint_request("") + with pytest.raises(HTTPException) as excinfo: + await delete_gemini_agent( + request=request, + name="my-agent", + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + assert excinfo.value.status_code == 401 + instance.base_process_llm_request.assert_not_called() + + +@pytest.mark.asyncio +async def test_create_agent_non_admin_without_api_key_is_rejected( + mock_srv, user_api_key_dict +): + from fastapi import HTTPException + + from litellm.proxy.google_endpoints.agents_endpoints import create_gemini_agent + + with ( + patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" + ) as MockProcessor, + patch( + "litellm.proxy.google_endpoints.agents_endpoints._read_request_body", + new=AsyncMock(return_value={"name": "agent-1", "base_agent": "waverunner"}), + ), + ): + instance = MockProcessor.return_value + instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) + + request = _make_endpoint_request("") + with pytest.raises(HTTPException) as excinfo: + await create_gemini_agent( + request=request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + assert excinfo.value.status_code == 401 + instance.base_process_llm_request.assert_not_called() + + +@pytest.mark.asyncio +async def test_list_agents_proxy_admin_may_use_env_fallback( + mock_srv, proxy_admin_user_api_key_dict +): + """Proxy admins (master key) keep the env-fallback convenience.""" + from litellm.proxy.google_endpoints.agents_endpoints import list_gemini_agents + + with patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" + ) as MockProcessor: + instance = MockProcessor.return_value + instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) + + request = _make_endpoint_request("") + await list_gemini_agents( + request=request, + fastapi_response=MagicMock(), + user_api_key_dict=proxy_admin_user_api_key_dict, + ) + + init_data = MockProcessor.call_args[1]["data"] + assert "api_key" not in init_data + instance.base_process_llm_request.assert_awaited_once() + + +def test_validate_environment_rejects_api_base_override_without_explicit_key( + monkeypatch, +): + """SECURITY: caller-supplied api_base must be paired with an explicit + api_key — otherwise the proxy's shared GOOGLE_API_KEY leaks to the + attacker-controlled host via the x-goog-api-key header. + """ + from litellm.llms.gemini.agents.transformation import GeminiAgentsConfig + + # Even if env-fallback is available, api_base override must require api_key. + monkeypatch.setenv("GOOGLE_API_KEY", "AIzaSharedSecret") + + cfg = GeminiAgentsConfig() + with pytest.raises(ValueError, match="api_base"): + cfg.validate_environment( + headers={}, + litellm_params={"api_base": "https://attacker.example"}, + ) + + +def test_validate_environment_allows_api_base_with_explicit_key(monkeypatch): + """api_base override is OK when paired with an explicit api_key.""" + from litellm.llms.gemini.agents.transformation import GeminiAgentsConfig + + monkeypatch.delenv("GOOGLE_API_KEY", raising=False) + monkeypatch.delenv("GEMINI_API_KEY", raising=False) + + cfg = GeminiAgentsConfig() + headers = cfg.validate_environment( + headers={}, + litellm_params={ + "api_base": "https://my-gemini-proxy.example", + "api_key": "AIzaCallerOwned", + }, + ) + assert headers["x-goog-api-key"] == "AIzaCallerOwned" + + +def test_validate_environment_env_fallback_when_no_api_base_override(monkeypatch): + """Without api_base override, env fallback continues to work for SDK use.""" + from litellm.llms.gemini.agents.transformation import GeminiAgentsConfig + + monkeypatch.setenv("GOOGLE_API_KEY", "AIzaFromEnv") + monkeypatch.delenv("GEMINI_API_KEY", raising=False) + + cfg = GeminiAgentsConfig() + headers = cfg.validate_environment(headers={}, litellm_params={}) + assert headers["x-goog-api-key"] == "AIzaFromEnv" diff --git a/tests/test_litellm/interactions/test_agents_http_handler.py b/tests/test_litellm/interactions/test_agents_http_handler.py new file mode 100644 index 00000000000..6947503e0bb --- /dev/null +++ b/tests/test_litellm/interactions/test_agents_http_handler.py @@ -0,0 +1,587 @@ +""" +Unit tests for litellm/interactions/agents/http_handler.py + +These tests exercise both the sync and async branches of every CRUD method +on AgentsHTTPHandler using stub httpx clients, plus the _is_async dispatch +branches, error mapping, and pre/post logging hooks. + +No real HTTP traffic is made. +""" + +import os +import sys +from unittest.mock import AsyncMock, MagicMock + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../..")) + +from litellm.interactions.agents.http_handler import ( + AgentsHTTPHandler, + agents_http_handler, +) +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.llms.gemini.agents.transformation import GeminiAgentsConfig +from litellm.llms.gemini.common_utils import GeminiError +from litellm.types.agents import ( + AgentCreateResponse, + AgentDeleteResult, + AgentListResponse, + AgentVersionsResponse, +) +from litellm.types.router import GenericLiteLLMParams + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _make_response(status_code: int = 200, json_data=None, text: str = "") -> MagicMock: + """Build a stub httpx-like response.""" + response = MagicMock() + response.status_code = status_code + response.text = text or (str(json_data) if json_data is not None else "") + response.headers = {} + if json_data is not None: + response.json.return_value = json_data + else: + response.json.return_value = {} + return response + + +def _make_sync_client() -> MagicMock: + client = MagicMock(spec=HTTPHandler) + return client + + +def _make_async_client() -> MagicMock: + client = MagicMock(spec=AsyncHTTPHandler) + client.post = AsyncMock() + client.get = AsyncMock() + client.delete = AsyncMock() + return client + + +def _make_logging_obj() -> MagicMock: + return MagicMock() + + +@pytest.fixture +def handler() -> AgentsHTTPHandler: + return AgentsHTTPHandler() + + +@pytest.fixture +def config() -> GeminiAgentsConfig: + return GeminiAgentsConfig() + + +@pytest.fixture +def litellm_params() -> GenericLiteLLMParams: + return GenericLiteLLMParams(api_key="AIza-test") + + +# --------------------------------------------------------------------------- +# Module-level singleton sanity check +# --------------------------------------------------------------------------- + + +def test_module_singleton_is_agents_http_handler_instance(): + assert isinstance(agents_http_handler, AgentsHTTPHandler) + + +# --------------------------------------------------------------------------- +# CREATE +# --------------------------------------------------------------------------- + + +class TestCreateAgent: + def test_sync_returns_parsed_create_response(self, handler, config, litellm_params): + client = _make_sync_client() + client.post.return_value = _make_response( + 200, json_data={"id": "agent-x", "base_agent": "gemini-2.5-flash"} + ) + logging_obj = _make_logging_obj() + + result = handler.create_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers={"X-Test": "1"}, + extra_body={"foo": "bar"}, + client=client, + ) + + assert isinstance(result, AgentCreateResponse) + assert result.id == "agent-x" + client.post.assert_called_once() + kwargs = client.post.call_args.kwargs + assert kwargs["url"].endswith("/v1beta/agents") + assert kwargs["json"]["name"] == "agent-x" + assert kwargs["json"]["foo"] == "bar" + assert kwargs["headers"]["X-Test"] == "1" + logging_obj.pre_call.assert_called_once() + logging_obj.post_call.assert_called_once() + + def test_sync_dispatches_to_async_when_is_async( + self, handler, config, litellm_params + ): + client = _make_sync_client() + + result = handler.create_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + _is_async=True, + ) + + import asyncio + + assert asyncio.iscoroutine(result) + result.close() + + def test_sync_maps_http_error_via_config(self, handler, config, litellm_params): + client = _make_sync_client() + bad = _make_response(404, text="not found") + client.post.side_effect = httpx.HTTPStatusError( + "boom", request=MagicMock(), response=bad + ) + + with pytest.raises(GeminiError): + handler.create_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + @pytest.mark.asyncio + async def test_async_returns_parsed_create_response( + self, handler, config, litellm_params + ): + client = _make_async_client() + client.post.return_value = _make_response( + 200, json_data={"id": "agent-y", "base_agent": "gemini-2.5-flash"} + ) + + result = await handler.async_create_agent( + agents_api_config=config, + name="agent-y", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + extra_body={"baz": "qux"}, + client=client, + ) + + assert isinstance(result, AgentCreateResponse) + assert result.id == "agent-y" + client.post.assert_awaited_once() + + @pytest.mark.asyncio + async def test_async_maps_http_error_via_config( + self, handler, config, litellm_params + ): + client = _make_async_client() + bad = _make_response(500, text="server error") + client.post.side_effect = httpx.HTTPStatusError( + "boom", request=MagicMock(), response=bad + ) + + with pytest.raises(GeminiError): + await handler.async_create_agent( + agents_api_config=config, + name="agent-y", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + +# --------------------------------------------------------------------------- +# LIST +# --------------------------------------------------------------------------- + + +class TestListAgents: + def test_sync_returns_list_response(self, handler, config, litellm_params): + client = _make_sync_client() + client.get.return_value = _make_response( + 200, + json_data={ + "agents": [{"id": "a-1"}, {"id": "a-2"}], + "nextPageToken": "tok", + }, + ) + + result = handler.list_agents( + agents_api_config=config, + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + assert isinstance(result, AgentListResponse) + assert len(result.agents) == 2 + assert result.next_page_token == "tok" + client.get.assert_called_once() + + def test_sync_dispatches_to_async_when_is_async( + self, handler, config, litellm_params + ): + client = _make_sync_client() + + result = handler.list_agents( + agents_api_config=config, + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + _is_async=True, + ) + + import asyncio + + assert asyncio.iscoroutine(result) + result.close() + + def test_sync_maps_http_error_via_config(self, handler, config, litellm_params): + client = _make_sync_client() + bad = _make_response(403, text="forbidden") + client.get.side_effect = httpx.HTTPStatusError( + "boom", request=MagicMock(), response=bad + ) + + with pytest.raises(GeminiError): + handler.list_agents( + agents_api_config=config, + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + @pytest.mark.asyncio + async def test_async_returns_list_response(self, handler, config, litellm_params): + client = _make_async_client() + client.get.return_value = _make_response( + 200, json_data={"agents": [{"id": "a-1"}]} + ) + + result = await handler.async_list_agents( + agents_api_config=config, + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + assert isinstance(result, AgentListResponse) + assert len(result.agents) == 1 + client.get.assert_awaited_once() + + @pytest.mark.asyncio + async def test_async_maps_http_error_via_config( + self, handler, config, litellm_params + ): + client = _make_async_client() + bad = _make_response(429, text="rate limited") + client.get.side_effect = httpx.HTTPStatusError( + "boom", request=MagicMock(), response=bad + ) + + with pytest.raises(GeminiError): + await handler.async_list_agents( + agents_api_config=config, + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + +# --------------------------------------------------------------------------- +# GET +# --------------------------------------------------------------------------- + + +class TestGetAgent: + def test_sync_returns_get_response(self, handler, config, litellm_params): + client = _make_sync_client() + client.get.return_value = _make_response(200, json_data={"id": "agent-x"}) + + result = handler.get_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + assert isinstance(result, AgentCreateResponse) + assert result.id == "agent-x" + kwargs = client.get.call_args.kwargs + assert kwargs["url"].endswith("/v1beta/agents/agent-x") + + def test_sync_dispatches_to_async_when_is_async( + self, handler, config, litellm_params + ): + result = handler.get_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=_make_sync_client(), + _is_async=True, + ) + import asyncio + + assert asyncio.iscoroutine(result) + result.close() + + def test_sync_maps_http_error_via_config(self, handler, config, litellm_params): + client = _make_sync_client() + bad = _make_response(404, text="not found") + client.get.side_effect = httpx.HTTPStatusError( + "boom", request=MagicMock(), response=bad + ) + + with pytest.raises(GeminiError): + handler.get_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + @pytest.mark.asyncio + async def test_async_returns_get_response(self, handler, config, litellm_params): + client = _make_async_client() + client.get.return_value = _make_response(200, json_data={"id": "agent-y"}) + + result = await handler.async_get_agent( + agents_api_config=config, + name="agent-y", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + assert isinstance(result, AgentCreateResponse) + assert result.id == "agent-y" + + @pytest.mark.asyncio + async def test_async_maps_http_error_via_config( + self, handler, config, litellm_params + ): + client = _make_async_client() + bad = _make_response(404, text="not found") + client.get.side_effect = httpx.HTTPStatusError( + "boom", request=MagicMock(), response=bad + ) + + with pytest.raises(GeminiError): + await handler.async_get_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + +# --------------------------------------------------------------------------- +# DELETE +# --------------------------------------------------------------------------- + + +class TestDeleteAgent: + def test_sync_returns_delete_result(self, handler, config, litellm_params): + client = _make_sync_client() + client.delete.return_value = _make_response(200, json_data={}) + + result = handler.delete_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + assert isinstance(result, AgentDeleteResult) + assert result.name == "agent-x" + assert result.deleted is True + kwargs = client.delete.call_args.kwargs + assert kwargs["url"].endswith("/v1beta/agents/agent-x") + + def test_sync_dispatches_to_async_when_is_async( + self, handler, config, litellm_params + ): + result = handler.delete_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=_make_sync_client(), + _is_async=True, + ) + import asyncio + + assert asyncio.iscoroutine(result) + result.close() + + def test_sync_maps_http_error_via_config(self, handler, config, litellm_params): + client = _make_sync_client() + bad = _make_response(403, text="forbidden") + client.delete.side_effect = httpx.HTTPStatusError( + "boom", request=MagicMock(), response=bad + ) + + with pytest.raises(GeminiError): + handler.delete_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + @pytest.mark.asyncio + async def test_async_returns_delete_result(self, handler, config, litellm_params): + client = _make_async_client() + client.delete.return_value = _make_response(200, json_data={}) + + result = await handler.async_delete_agent( + agents_api_config=config, + name="agent-y", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + assert isinstance(result, AgentDeleteResult) + assert result.name == "agent-y" + assert result.deleted is True + + @pytest.mark.asyncio + async def test_async_maps_http_error_via_config( + self, handler, config, litellm_params + ): + client = _make_async_client() + bad = _make_response(500, text="server error") + client.delete.side_effect = httpx.HTTPStatusError( + "boom", request=MagicMock(), response=bad + ) + + with pytest.raises(GeminiError): + await handler.async_delete_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + +# --------------------------------------------------------------------------- +# LIST VERSIONS +# --------------------------------------------------------------------------- + + +class TestListAgentVersions: + def test_sync_returns_versions_response(self, handler, config, litellm_params): + client = _make_sync_client() + client.get.return_value = _make_response( + 200, + json_data={ + "agentVersions": [ + {"agent": "agent-x", "name": "agents/agent-x/versions/v1"} + ], + "nextPageToken": "tok", + }, + ) + + result = handler.list_agent_versions( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + assert isinstance(result, AgentVersionsResponse) + assert len(result.agent_versions) == 1 + assert result.next_page_token == "tok" + kwargs = client.get.call_args.kwargs + assert kwargs["url"].endswith("/v1beta/agents/agent-x/versions") + + def test_sync_dispatches_to_async_when_is_async( + self, handler, config, litellm_params + ): + result = handler.list_agent_versions( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=_make_sync_client(), + _is_async=True, + ) + import asyncio + + assert asyncio.iscoroutine(result) + result.close() + + def test_sync_maps_http_error_via_config(self, handler, config, litellm_params): + client = _make_sync_client() + bad = _make_response(404, text="not found") + client.get.side_effect = httpx.HTTPStatusError( + "boom", request=MagicMock(), response=bad + ) + + with pytest.raises(GeminiError): + handler.list_agent_versions( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + @pytest.mark.asyncio + async def test_async_returns_versions_response( + self, handler, config, litellm_params + ): + client = _make_async_client() + client.get.return_value = _make_response(200, json_data={"agentVersions": []}) + + result = await handler.async_list_agent_versions( + agents_api_config=config, + name="agent-y", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + assert isinstance(result, AgentVersionsResponse) + assert result.agent_versions == [] + + @pytest.mark.asyncio + async def test_async_maps_http_error_via_config( + self, handler, config, litellm_params + ): + client = _make_async_client() + bad = _make_response(500, text="server error") + client.get.side_effect = httpx.HTTPStatusError( + "boom", request=MagicMock(), response=bad + ) + + with pytest.raises(GeminiError): + await handler.async_list_agent_versions( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) diff --git a/tests/test_litellm/interactions/test_agents_main_and_utils.py b/tests/test_litellm/interactions/test_agents_main_and_utils.py new file mode 100644 index 00000000000..7c0183d20c6 --- /dev/null +++ b/tests/test_litellm/interactions/test_agents_main_and_utils.py @@ -0,0 +1,354 @@ +""" +Unit tests for litellm/interactions/agents/utils.py and main.py +focused on the managed agents SDK surface added in the +"Gemini managed agents support" PR. + +The tests mock the underlying HTTP handler so they cover the public +sync + async create/list/get/delete/list_versions entry points and the +small helper utilities without touching the network. +""" + +import asyncio +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../..")) + +import litellm +from litellm.interactions.agents import ( + acreate, + adelete, + aget, + alist, + alist_versions, + create, + delete, + get, + list as list_agents, + list_versions, +) +from litellm.interactions.agents.main import ( + _get_agents_api_config, + _make_logging_obj, +) +from litellm.interactions.agents.utils import get_provider_agents_api_config +from litellm.llms.base_llm.agents.transformation import BaseAgentsAPIConfig +from litellm.llms.gemini.agents.transformation import GeminiAgentsConfig + + +_HANDLER_PATH = "litellm.interactions.agents.main.agents_http_handler" + + +# --------------------------------------------------------------------------- +# utils.get_provider_agents_api_config +# --------------------------------------------------------------------------- + + +class TestGetProviderAgentsApiConfig: + def test_returns_gemini_config_for_gemini(self): + cfg = get_provider_agents_api_config("gemini") + assert isinstance(cfg, GeminiAgentsConfig) + assert isinstance(cfg, BaseAgentsAPIConfig) + + @pytest.mark.parametrize( + "provider", ["openai", "anthropic", "bedrock", "vertex_ai", "unknown"] + ) + def test_returns_none_for_non_gemini(self, provider): + assert get_provider_agents_api_config(provider) is None + + def test_returns_none_for_none(self): + assert get_provider_agents_api_config(None) is None + + +# --------------------------------------------------------------------------- +# main._get_agents_api_config +# --------------------------------------------------------------------------- + + +class TestGetAgentsApiConfig: + def test_returns_config_for_gemini(self): + cfg = _get_agents_api_config("gemini") + assert isinstance(cfg, GeminiAgentsConfig) + + def test_raises_bad_request_for_unsupported_provider(self): + with pytest.raises(litellm.BadRequestError) as excinfo: + _get_agents_api_config("openai") + assert "does not have a native" in str(excinfo.value) + + +# --------------------------------------------------------------------------- +# main._make_logging_obj +# --------------------------------------------------------------------------- + + +class TestMakeLoggingObj: + def test_calls_update_from_kwargs_and_returns_same_obj(self): + logging_obj = MagicMock() + kwargs = {"litellm_logging_obj": logging_obj, "litellm_call_id": "abc-123"} + + returned = _make_logging_obj( + kwargs=kwargs, + model="my-agent", + custom_llm_provider="gemini", + call_type="create_agent", + optional_params={"foo": "bar"}, + ) + + assert returned is logging_obj + logging_obj.update_from_kwargs.assert_called_once() + kwargs_call = logging_obj.update_from_kwargs.call_args.kwargs + assert kwargs_call["model"] == "my-agent" + assert kwargs_call["optional_params"] == {"foo": "bar"} + assert kwargs_call["custom_llm_provider"] == "gemini" + assert kwargs_call["litellm_params"]["litellm_call_id"] == "abc-123" + + +# --------------------------------------------------------------------------- +# Sync entry points: create / list / get / delete / list_versions +# --------------------------------------------------------------------------- + + +def _stub_handler(return_value): + """Build a stub AgentsHTTPHandler whose CRUD methods return *return_value*.""" + handler = MagicMock() + handler.create_agent.return_value = return_value + handler.list_agents.return_value = return_value + handler.get_agent.return_value = return_value + handler.delete_agent.return_value = return_value + handler.list_agent_versions.return_value = return_value + return handler + + +class TestSyncEntryPoints: + def test_create_passes_args_to_handler(self): + sentinel = MagicMock(name="create_response") + with patch(_HANDLER_PATH, _stub_handler(sentinel)) as handler: + response = create( + name="waverunner", + base_agent="gemini-2.5-flash", + instructions="be helpful", + base_environment={"type": "remote"}, + custom_llm_provider="gemini", + api_key="AIza-test", + extra_headers={"X-Test": "1"}, + extra_body={"foo": "bar"}, + ) + + assert response is sentinel + handler.create_agent.assert_called_once() + kw = handler.create_agent.call_args.kwargs + assert kw["name"] == "waverunner" + assert kw["_is_async"] is False + assert kw["extra_headers"] == {"X-Test": "1"} + assert kw["extra_body"] == {"foo": "bar"} + assert isinstance(kw["agents_api_config"], GeminiAgentsConfig) + + def test_create_defaults_custom_llm_provider_to_gemini(self): + sentinel = MagicMock(name="create_response") + with patch(_HANDLER_PATH, _stub_handler(sentinel)) as handler: + create(name="agent-x", api_key="AIza") + assert handler.create_agent.call_args.kwargs["_is_async"] is False + cfg = handler.create_agent.call_args.kwargs["agents_api_config"] + assert isinstance(cfg, GeminiAgentsConfig) + + def test_create_raises_for_unsupported_provider(self): + with pytest.raises(litellm.exceptions.BadRequestError): + create(name="agent-x", custom_llm_provider="openai", api_key="sk-x") + + def test_list_passes_args_to_handler(self): + sentinel = MagicMock(name="list_response") + with patch(_HANDLER_PATH, _stub_handler(sentinel)) as handler: + response = list_agents(custom_llm_provider="gemini", api_key="AIza") + assert response is sentinel + handler.list_agents.assert_called_once() + assert handler.list_agents.call_args.kwargs["_is_async"] is False + + def test_get_passes_args_to_handler(self): + sentinel = MagicMock(name="get_response") + with patch(_HANDLER_PATH, _stub_handler(sentinel)) as handler: + response = get(name="waverunner", api_key="AIza") + assert response is sentinel + kw = handler.get_agent.call_args.kwargs + assert kw["name"] == "waverunner" + assert kw["_is_async"] is False + + def test_delete_passes_args_to_handler(self): + sentinel = MagicMock(name="delete_response") + with patch(_HANDLER_PATH, _stub_handler(sentinel)) as handler: + response = delete(name="waverunner", api_key="AIza") + assert response is sentinel + kw = handler.delete_agent.call_args.kwargs + assert kw["name"] == "waverunner" + assert kw["_is_async"] is False + + def test_list_versions_passes_args_to_handler(self): + sentinel = MagicMock(name="versions_response") + with patch(_HANDLER_PATH, _stub_handler(sentinel)) as handler: + response = list_versions(name="waverunner", api_key="AIza") + assert response is sentinel + kw = handler.list_agent_versions.call_args.kwargs + assert kw["name"] == "waverunner" + assert kw["_is_async"] is False + + +# --------------------------------------------------------------------------- +# Async entry points +# --------------------------------------------------------------------------- + + +class TestAsyncEntryPoints: + """Async entry points delegate to their sync counterparts via run_in_executor.""" + + @pytest.mark.asyncio + async def test_acreate_dispatches_with_async_flag(self): + sentinel = MagicMock(name="acreate_response") + + def fake_create_agent(**kwargs): + assert kwargs["_is_async"] is True + assert kwargs["name"] == "waverunner" + return sentinel + + handler = MagicMock() + handler.create_agent.side_effect = fake_create_agent + + with patch(_HANDLER_PATH, handler): + response = await acreate( + name="waverunner", + base_agent="gemini-2.5-flash", + api_key="AIza", + ) + assert response is sentinel + + @pytest.mark.asyncio + async def test_acreate_awaits_coroutine_result(self): + async def _coro(): + return "async-value" + + handler = MagicMock() + handler.create_agent.return_value = _coro() + + with patch(_HANDLER_PATH, handler): + response = await acreate(name="waverunner", api_key="AIza") + + assert response == "async-value" + + @pytest.mark.asyncio + async def test_alist_dispatches_with_async_flag(self): + sentinel = MagicMock(name="alist_response") + + def fake_list_agents(**kwargs): + assert kwargs["_is_async"] is True + return sentinel + + handler = MagicMock() + handler.list_agents.side_effect = fake_list_agents + + with patch(_HANDLER_PATH, handler): + response = await alist(api_key="AIza") + assert response is sentinel + + @pytest.mark.asyncio + async def test_aget_dispatches_with_async_flag(self): + sentinel = MagicMock(name="aget_response") + + def fake_get_agent(**kwargs): + assert kwargs["_is_async"] is True + assert kwargs["name"] == "waverunner" + return sentinel + + handler = MagicMock() + handler.get_agent.side_effect = fake_get_agent + + with patch(_HANDLER_PATH, handler): + response = await aget(name="waverunner", api_key="AIza") + assert response is sentinel + + @pytest.mark.asyncio + async def test_adelete_dispatches_with_async_flag(self): + sentinel = MagicMock(name="adelete_response") + + def fake_delete_agent(**kwargs): + assert kwargs["_is_async"] is True + assert kwargs["name"] == "waverunner" + return sentinel + + handler = MagicMock() + handler.delete_agent.side_effect = fake_delete_agent + + with patch(_HANDLER_PATH, handler): + response = await adelete(name="waverunner", api_key="AIza") + assert response is sentinel + + @pytest.mark.asyncio + async def test_alist_versions_dispatches_with_async_flag(self): + sentinel = MagicMock(name="alist_versions_response") + + def fake_versions(**kwargs): + assert kwargs["_is_async"] is True + assert kwargs["name"] == "waverunner" + return sentinel + + handler = MagicMock() + handler.list_agent_versions.side_effect = fake_versions + + with patch(_HANDLER_PATH, handler): + response = await alist_versions(name="waverunner", api_key="AIza") + assert response is sentinel + + +# --------------------------------------------------------------------------- +# Async error wrapping: exception_type must be invoked +# --------------------------------------------------------------------------- + + +class TestAsyncErrorWrapping: + """If the underlying handler raises, async entry points re-raise via + litellm.exception_type so users get a normalised provider error.""" + + @pytest.mark.asyncio + async def test_acreate_wraps_exception(self): + handler = MagicMock() + handler.create_agent.side_effect = RuntimeError("kaboom") + + with patch(_HANDLER_PATH, handler): + with pytest.raises(Exception): + await acreate(name="waverunner", api_key="AIza") + + @pytest.mark.asyncio + async def test_aget_wraps_exception(self): + handler = MagicMock() + handler.get_agent.side_effect = RuntimeError("kaboom") + + with patch(_HANDLER_PATH, handler): + with pytest.raises(Exception): + await aget(name="waverunner", api_key="AIza") + + @pytest.mark.asyncio + async def test_alist_wraps_exception(self): + handler = MagicMock() + handler.list_agents.side_effect = RuntimeError("kaboom") + + with patch(_HANDLER_PATH, handler): + with pytest.raises(Exception): + await alist(api_key="AIza") + + @pytest.mark.asyncio + async def test_adelete_wraps_exception(self): + handler = MagicMock() + handler.delete_agent.side_effect = RuntimeError("kaboom") + + with patch(_HANDLER_PATH, handler): + with pytest.raises(Exception): + await adelete(name="waverunner", api_key="AIza") + + @pytest.mark.asyncio + async def test_alist_versions_wraps_exception(self): + handler = MagicMock() + handler.list_agent_versions.side_effect = RuntimeError("kaboom") + + with patch(_HANDLER_PATH, handler): + with pytest.raises(Exception): + await alist_versions(name="waverunner", api_key="AIza") diff --git a/tests/test_litellm/interactions/test_gemini_interactions_transformation.py b/tests/test_litellm/interactions/test_gemini_interactions_transformation.py index 0ef97e25e44..2e596a72158 100644 --- a/tests/test_litellm/interactions/test_gemini_interactions_transformation.py +++ b/tests/test_litellm/interactions/test_gemini_interactions_transformation.py @@ -122,6 +122,56 @@ class TestGetCompleteUrl: ) +class TestTransformRequest: + def test_passes_environment_to_request_body(self, config): + request_body = config.transform_request( + model=None, + agent="my-custom-slides-agent", + input=[{"type": "text", "text": "Create a 5-slide presentation about AI trends."}], + optional_params={ + "environment": "remote", + "stream": False, + }, + litellm_params=GenericLiteLLMParams(api_key="test-api-key"), + headers={}, + ) + + assert request_body["agent"] == "my-custom-slides-agent" + assert request_body["environment"] == "remote" + assert request_body["stream"] is False + assert request_body["input"] == [ + {"type": "text", "text": "Create a 5-slide presentation about AI trends."} + ] + + def test_passes_environment_object_to_request_body(self, config): + environment_config = { + "type": "remote", + "sources": [{"type": "gcs", "uri": "gs://bucket/skills.zip"}], + "network": {"egress": "allow_all"}, + } + request_body = config.transform_request( + model=None, + agent="waverunner", + input="What is 2 + 2?", + optional_params={"environment": environment_config}, + litellm_params=GenericLiteLLMParams(api_key="test-api-key"), + headers={}, + ) + + assert request_body["environment"] == environment_config + + def test_passes_existing_environment_id_to_request_body(self, config): + env_id = "env-abc123" + request_body = config.transform_request( + model=None, + agent="my-custom-slides-agent", + input="Continue the presentation.", + optional_params={"environment": env_id}, + litellm_params=GenericLiteLLMParams(api_key="test-api-key"), + headers={}, + ) + + assert request_body["environment"] == env_id class TestStreamingIterator: def _make_iterator(self) -> LiteLLMResponsesInteractionsStreamingIterator: return LiteLLMResponsesInteractionsStreamingIterator( diff --git a/tests/test_litellm/proxy/google_endpoints/test_interactions_agent_param.py b/tests/test_litellm/proxy/google_endpoints/test_interactions_agent_param.py index 1063f59afb6..f3cec320532 100644 --- a/tests/test_litellm/proxy/google_endpoints/test_interactions_agent_param.py +++ b/tests/test_litellm/proxy/google_endpoints/test_interactions_agent_param.py @@ -1,75 +1,103 @@ """ -Test for interactions endpoint agent parameter handling. +Tests for managed-agent interaction routing. -Tests that the /v1beta/interactions endpoint correctly extracts -the `agent` parameter as a fallback when `model` is not provided. +Custom Gemini agents are identified by ``agent`` (name/id), not ``model``. +The proxy must not pass the agent name as ``model`` or LiteLLM may route to +openai/* wildcards instead of Gemini interactions. """ +from unittest.mock import MagicMock, patch + import pytest class TestInteractionsAgentParameter: - """Test agent parameter handling in interactions endpoint.""" + """Proxy endpoint must keep agent and model separate.""" - def test_agent_parameter_fallback_logic(self): - """ - Test the core logic: model or agent extraction. - - This tests the fix in endpoints.py line ~267: - model=data.get("model") or data.get("agent") - """ - # Case 1: Only agent provided (Deep Research use case) + def test_create_interaction_uses_model_only_from_body(self): + """POST /v1beta/interactions: model kwarg is only the request's model field.""" data = { - "agent": "deep-research-pro-preview-12-2025", - "input": "Research quantum computing", - "background": True, + "agent": "mqy-custom-slides-agent", + "input": "hello", } - model = data.get("model") or data.get("agent") - assert model == "deep-research-pro-preview-12-2025" + # Fixed behavior: do NOT fall back agent → model + model_for_routing = data.get("model") + assert model_for_routing is None + assert data.get("agent") == "mqy-custom-slides-agent" - # Case 2: Only model provided (normal use case) + def test_model_field_still_used_when_present(self): data = { "model": "gemini-2.5-flash", - "input": "Hello world", + "input": "hello", } - model = data.get("model") or data.get("agent") - assert model == "gemini-2.5-flash" + model_for_routing = data.get("model") + assert model_for_routing == "gemini-2.5-flash" - # Case 3: Both provided (model takes precedence) - data = { - "model": "gemini-2.5-flash", - "agent": "deep-research-pro-preview-12-2025", - "input": "Test", - } - model = data.get("model") or data.get("agent") - assert model == "gemini-2.5-flash" - # Case 4: Neither provided - data = { - "input": "Test", - } - model = data.get("model") or data.get("agent") - assert model is None +class TestInteractionsAgentOnlyProviderRouting: + """SDK: agent-only create must not call get_llm_provider on the agent name.""" - def test_route_type_in_skip_model_routing_list(self): - """ - Test that acreate_interaction is in the list of routes - that skip model-based routing. + @patch("litellm.interactions.main.interactions_http_handler") + @patch("litellm.interactions.main.get_provider_interactions_api_config") + @patch("litellm.get_llm_provider") + def test_agent_only_skips_get_llm_provider( + self, + mock_get_llm_provider, + mock_get_config, + mock_handler, + ): + from litellm.interactions.main import create + from litellm.types.interactions import InteractionsAPIResponse - This tests the fix in route_llm_request.py. - """ - # The list of routes that skip model routing for interactions - skip_model_routing_routes = [ - "acreate_interaction", - "aget_interaction", - "adelete_interaction", - "acancel_interaction", - ] + mock_get_config.return_value = MagicMock() + mock_handler.create_interaction.return_value = InteractionsAPIResponse( + id="int-1", + status="completed", + object="interaction", + ) - # acreate_interaction should be in the list (this is the fix) - assert "acreate_interaction" in skip_model_routing_routes + logging_obj = MagicMock() + create( + agent="mqy-custom-slides-agent", + input="test", + custom_llm_provider="gemini", + litellm_logging_obj=logging_obj, + ) - # All interaction routes should be covered - assert "aget_interaction" in skip_model_routing_routes - assert "adelete_interaction" in skip_model_routing_routes - assert "acancel_interaction" in skip_model_routing_routes + mock_get_llm_provider.assert_not_called() + call_kwargs = mock_handler.create_interaction.call_args.kwargs + assert call_kwargs["agent"] == "mqy-custom-slides-agent" + assert call_kwargs["model"] is None + assert call_kwargs["custom_llm_provider"] == "gemini" + + @patch("litellm.interactions.main.interactions_http_handler") + @patch("litellm.interactions.main.get_provider_interactions_api_config") + @patch("litellm.get_llm_provider") + def test_proxy_mistake_model_equals_agent_is_corrected( + self, + mock_get_llm_provider, + mock_get_config, + mock_handler, + ): + """If model was wrongly set to the agent name, clear it before the HTTP call.""" + from litellm.interactions.main import create + from litellm.types.interactions import InteractionsAPIResponse + + mock_get_config.return_value = MagicMock() + mock_handler.create_interaction.return_value = InteractionsAPIResponse( + id="int-1", + status="completed", + object="interaction", + ) + + logging_obj = MagicMock() + create( + model="mqy-custom-slides-agent", + agent="mqy-custom-slides-agent", + input="test", + custom_llm_provider="gemini", + litellm_logging_obj=logging_obj, + ) + + mock_get_llm_provider.assert_not_called() + assert mock_handler.create_interaction.call_args.kwargs["model"] is None diff --git a/tests/test_litellm/proxy/google_endpoints/test_managed_agents_model_param.py b/tests/test_litellm/proxy/google_endpoints/test_managed_agents_model_param.py new file mode 100644 index 00000000000..5485d0f2929 --- /dev/null +++ b/tests/test_litellm/proxy/google_endpoints/test_managed_agents_model_param.py @@ -0,0 +1,199 @@ +""" +Tests verifying that managed-agent proxy endpoints never pass the agent name +as the ``model`` parameter to ``base_process_llm_request``. + +Passing ``model=`` would cause ``common_processing_pre_call_logic`` +to write the agent name into ``self.data["model"]``, which triggers spurious +model-alias mapping, rate-limiting lookups, and logging tied to a +non-existent model deployment. The agent name is already carried in +``data["name"]`` and must not pollute the ``model`` slot. +""" + +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + + +def _build_agents_client(): + """Build a TestClient whose auth dependency is overridden to a PROXY_ADMIN + user. Using ``dependency_overrides`` is the only reliable way to bypass the + real ``user_api_key_auth`` for FastAPI route tests — patching the module- + level name does not affect the function reference captured by ``Depends``. + The PROXY_ADMIN role also bypasses the caller-supplied-api_key guard so + these tests can focus on the ``model=None`` invariant. + """ + from fastapi import FastAPI + from fastapi.testclient import TestClient + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + from litellm.proxy.google_endpoints.agents_endpoints import router as agents_router + + app = FastAPI() + app.include_router(agents_router) + + async def _fake_user_api_key_auth(): + return UserAPIKeyAuth( + api_key="sk-test", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + + app.dependency_overrides[user_api_key_auth] = _fake_user_api_key_auth + return TestClient(app) + + +def _patch_proxy_server_imports(client=None): + """Return a context-manager that stubs _proxy_server_imports so tests + don't need a running proxy.""" + mock_srv = { + "general_settings": {}, + "llm_router": MagicMock(), + "proxy_config": MagicMock(), + "proxy_logging_obj": MagicMock(), + "select_data_generator": None, + "user_api_base": None, + "user_max_tokens": None, + "user_model": None, + "user_request_timeout": None, + "user_temperature": None, + "version": "0.0.0", + } + return patch( + "litellm.proxy.google_endpoints.agents_endpoints._proxy_server_imports", + return_value=mock_srv, + ) + + +def _patch_base_process(return_value=None): + if return_value is None: + return_value = {"name": "agents/my-agent", "displayName": "My Agent"} + return patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request", + new_callable=AsyncMock, + return_value=return_value, + ) + + +def _patch_auth(): + """Deprecated no-op kept for call-site compatibility. + + ``_build_agents_client`` now installs a FastAPI ``dependency_overrides`` + entry that injects a PROXY_ADMIN ``UserAPIKeyAuth``, so individual tests + no longer need to patch the module-level ``user_api_key_auth`` name. + """ + return patch("os.getpid") + + +class TestManagedAgentsModelParam: + """Endpoints must pass model=None, not the agent name, to base_process_llm_request.""" + + def test_create_agent_passes_model_none(self): + """POST /v1beta/agents: model kwarg must be None, not the name field.""" + try: + client = _build_agents_client() + except ImportError as exc: + pytest.skip(f"Skipping: missing dependency {exc}") + + with ( + _patch_proxy_server_imports(), + _patch_base_process() as mock_process, + _patch_auth(), + ): + client.post( + "/v1beta/agents", + json={ + "name": "my-custom-slides-agent", + "base_agent": "waverunner", + "instructions": "Be helpful.", + }, + ) + + mock_process.assert_called_once() + kwargs = mock_process.call_args.kwargs + assert kwargs["model"] is None, ( + f"create_gemini_agent must not pass model={kwargs['model']!r}; " + "the agent name must stay in data['name'], not pollute data['model']" + ) + assert kwargs["route_type"] == "acreate_agent" + + def test_get_agent_passes_model_none(self): + """GET /v1beta/agents/{name}: model kwarg must be None.""" + try: + client = _build_agents_client() + except ImportError as exc: + pytest.skip(f"Skipping: missing dependency {exc}") + + with ( + _patch_proxy_server_imports(), + _patch_base_process() as mock_process, + _patch_auth(), + ): + client.get("/v1beta/agents/my-custom-slides-agent") + + mock_process.assert_called_once() + kwargs = mock_process.call_args.kwargs + assert ( + kwargs["model"] is None + ), f"get_gemini_agent must not pass model={kwargs['model']!r}" + assert kwargs["route_type"] == "aget_agent" + + def test_delete_agent_passes_model_none(self): + """DELETE /v1beta/agents/{name}: model kwarg must be None.""" + try: + client = _build_agents_client() + except ImportError as exc: + pytest.skip(f"Skipping: missing dependency {exc}") + + with ( + _patch_proxy_server_imports(), + _patch_base_process() as mock_process, + _patch_auth(), + ): + client.delete("/v1beta/agents/my-custom-slides-agent") + + mock_process.assert_called_once() + kwargs = mock_process.call_args.kwargs + assert ( + kwargs["model"] is None + ), f"delete_gemini_agent must not pass model={kwargs['model']!r}" + assert kwargs["route_type"] == "adelete_agent" + + def test_list_agent_versions_passes_model_none(self): + """GET /v1beta/agents/{name}/versions: model kwarg must be None.""" + try: + client = _build_agents_client() + except ImportError as exc: + pytest.skip(f"Skipping: missing dependency {exc}") + + with ( + _patch_proxy_server_imports(), + _patch_base_process() as mock_process, + _patch_auth(), + ): + client.get("/v1beta/agents/my-custom-slides-agent/versions") + + mock_process.assert_called_once() + kwargs = mock_process.call_args.kwargs + assert ( + kwargs["model"] is None + ), f"list_gemini_agent_versions must not pass model={kwargs['model']!r}" + assert kwargs["route_type"] == "alist_agent_versions" + + def test_list_agents_already_passes_model_none(self): + """GET /v1beta/agents: existing list endpoint already passes model=None — keep it so.""" + try: + client = _build_agents_client() + except ImportError as exc: + pytest.skip(f"Skipping: missing dependency {exc}") + + with ( + _patch_proxy_server_imports(), + _patch_base_process(return_value={"agents": []}) as mock_process, + _patch_auth(), + ): + client.get("/v1beta/agents") + + mock_process.assert_called_once() + kwargs = mock_process.call_args.kwargs + assert kwargs["model"] is None + assert kwargs["route_type"] == "alist_agents" diff --git a/tests/test_litellm/router_utils/test_router_interactions_endpoints.py b/tests/test_litellm/router_utils/test_router_interactions_endpoints.py index c5468d73810..91bea170458 100644 --- a/tests/test_litellm/router_utils/test_router_interactions_endpoints.py +++ b/tests/test_litellm/router_utils/test_router_interactions_endpoints.py @@ -140,3 +140,159 @@ class TestInitInteractionsApiEndpoints: custom_llm_provider="vertex_ai", ) assert result == {"result": "success"} + + @pytest.mark.asyncio + async def test_init_interactions_api_endpoints_clears_model_when_equals_agent( + self, + ): + """Managed agent interactions must not pass agent name as model to the SDK.""" + router = Router(model_list=[]) + + mock_function = AsyncMock(return_value={"result": "success"}) + + await router._init_interactions_api_endpoints( + original_function=mock_function, + agent="mqy-custom-slides-agent", + model="mqy-custom-slides-agent", + input="hello", + ) + + mock_function.assert_called_once_with( + custom_llm_provider="gemini", + agent="mqy-custom-slides-agent", + model=None, + input="hello", + ) + + +class TestRouterCreateInteractionRouting: + """acreate_interaction routing: agent-only vs model + fallbacks.""" + + @pytest.mark.asyncio + async def test_acreate_interaction_agent_only_uses_init_interactions(self): + """Agent-only create must not use model-group fallback lookup.""" + router = Router( + model_list=[ + { + "model_name": "openai/*", + "litellm_params": {"model": "gpt-4"}, + } + ] + ) + + with ( + patch.object( + router, + "_init_interactions_api_endpoints", + new_callable=AsyncMock, + return_value={"id": "int-1"}, + ) as mock_init, + patch.object( + router, + "_ageneric_api_call_with_fallbacks", + new_callable=AsyncMock, + ) as mock_generic, + ): + result = await router.acreate_interaction( + agent="mqy-custom-slides-agent", + input="hello", + custom_llm_provider="gemini", + ) + + mock_init.assert_called_once() + mock_generic.assert_not_called() + assert result == {"id": "int-1"} + + @pytest.mark.asyncio + async def test_init_interactions_model_uses_generic_fallbacks(self): + """Model-based create uses _ageneric_api_call_with_fallbacks inside _init_interactions.""" + router = Router(model_list=[]) + + with patch.object( + router, + "_ageneric_api_call_with_fallbacks", + new_callable=AsyncMock, + return_value={"id": "int-1"}, + ) as mock_generic: + result = await router._init_interactions_api_endpoints( + original_function=AsyncMock(), + model="gemini-2.5-flash", + input="hello", + custom_llm_provider="gemini", + ) + + mock_generic.assert_called_once() + assert result == {"id": "int-1"} + + +class TestInitializeManagedAgentsEndpoints: + """Tests for _initialize_managed_agents_endpoints.""" + + def test_initialize_managed_agents_endpoints_creates_methods(self): + router = Router( + model_list=[ + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4"}, + } + ] + ) + + for method_name in ( + "acreate_agent", + "alist_agents", + "aget_agent", + "adelete_agent", + "alist_agent_versions", + ): + assert hasattr(router, method_name), f"missing {method_name}" + assert callable(getattr(router, method_name)), f"{method_name} not callable" + + def test_initialize_managed_agents_endpoints_can_be_called_directly(self): + router = Router(model_list=[]) + router._initialize_managed_agents_endpoints() + assert callable(router.acreate_agent) + assert callable(router.alist_agents) + + +class TestInitManagedAgentsApiEndpoints: + """Tests for _init_managed_agents_api_endpoints.""" + + @pytest.mark.asyncio + async def test_init_managed_agents_api_endpoints_defaults_to_gemini(self): + router = Router(model_list=[]) + mock_fn = AsyncMock(return_value={"agents": []}) + + await router._init_managed_agents_api_endpoints( + original_function=mock_fn, + ) + + call_kwargs = mock_fn.call_args.kwargs + assert call_kwargs["custom_llm_provider"] == "gemini" + + @pytest.mark.asyncio + async def test_init_managed_agents_api_endpoints_passes_custom_provider(self): + router = Router(model_list=[]) + mock_fn = AsyncMock(return_value={"agents": []}) + + await router._init_managed_agents_api_endpoints( + original_function=mock_fn, + custom_llm_provider="vertex_ai", + ) + + call_kwargs = mock_fn.call_args.kwargs + assert call_kwargs["custom_llm_provider"] == "vertex_ai" + + @pytest.mark.asyncio + async def test_init_managed_agents_api_endpoints_does_not_override_existing_provider( + self, + ): + router = Router(model_list=[]) + mock_fn = AsyncMock(return_value={"agents": []}) + + await router._init_managed_agents_api_endpoints( + original_function=mock_fn, + custom_llm_provider="vertex_ai", + ) + + mock_fn.assert_called_once_with(custom_llm_provider="vertex_ai")