mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Gemini managed agents support (#28270)
* Add support for environment variable in interactions api * Add sdk support for gemini create agent * Add agents endpoint support via proxy * Add outputs of each api * Add routing for model and agents param * Remove redundant condition in get_provider_agents_api_config LlmProviders.GEMINI.value is literally the string "gemini", so the second clause of the or was checking the exact same thing as the first. Co-authored-by: Sameer Kankute <Sameerlite@users.noreply.github.com> * fix: forward query-param credentials to list/get/delete/versions Gemini agent endpoints The list_gemini_agents, get_gemini_agent, delete_gemini_agent, and list_gemini_agent_versions endpoints previously constructed a hardcoded data dict with no mechanism to pass provider credentials. Unlike create_gemini_agent (POST, reads litellm_params_template from body), these GET/DELETE endpoints gave no way for multi-tenant callers to supply a per-request api_key or other LiteLLM params. Fix: - Add _merge_query_params_into_data() helper that reads query parameters from the request and merges them into the data dict without overwriting already-set keys (e.g. path params like 'name'). - Support a JSON-encoded litellm_params_template query parameter (matching the POST body pattern) as well as flat key=value pairs (e.g. api_key=AIza...). - Apply the helper in all four affected endpoints. - Add 13 unit tests covering the helper and each endpoint. Co-authored-by: Sameer Kankute <Sameerlite@users.noreply.github.com> * fix: pass model=None for managed agent proxy endpoints to prevent agent name polluting data["model"] Endpoints acreate_agent, aget_agent, adelete_agent, and alist_agent_versions were passing model=<agent_name> to base_process_llm_request. This caused common_processing_pre_call_logic to write the agent name into self.data["model"], which then triggered 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 is passed correctly to the SDK functions (litellm.interactions.agents.*). There is no reason to also set model=<agent_name>; the correct value is model=None for all five managed-agent management routes. Adds tests/test_litellm/proxy/google_endpoints/test_managed_agents_model_param.py to verify all five managed-agent endpoints pass model=None. Co-authored-by: Sameer Kankute <Sameerlite@users.noreply.github.com> * fix: address greptile P1/P2 review comments P1 (router.py): Restore fallback/retry support for acreate_interaction and create_interaction. Both were silently moved to _init_interactions_api_endpoints (direct call, no fallbacks). Moved them back to _ageneric_api_call_with_fallbacks so users with configured fallback models keep retry behaviour. P1 security (agents_endpoints.py): Remove flat query-param credential path (e.g. ?api_key=AIza...) from _merge_query_params_into_data. Credentials in URL query strings appear verbatim in server access logs, CDN edge logs, and browser history. Only the JSON-encoded litellm_params_template query param (matching the POST body pattern) is retained. P2 (interactions/http_handler.py): Extract _BaseHTTPHandler with shared _handle_error, _sync_client, and _async_client helpers. InteractionsHTTPHandler now extends _BaseHTTPHandler. The _async_client reads the provider from litellm_params instead of hardcoding GEMINI. P2 (interactions/agents/http_handler.py): AgentsHTTPHandler now extends InteractionsHTTPHandler (which inherits _BaseHTTPHandler) so all shared HTTP infrastructure is reused rather than duplicated. Removes the hardcoded LlmProviders.GEMINI from the async client path. Co-authored-by: Cursor <cursoragent@cursor.com> * fix: address CI failures from greptile review fixes - black: format interactions/agents/main.py and utils.py - tests: update test_gemini_agents_endpoints.py to match new _merge_query_params_into_data behaviour (flat credential params are rejected; only JSON-encoded litellm_params_template is accepted) - ci: add test_gemini_agents_endpoints.py to endpoints-and-responses shard in test-unit-proxy-db.yml so assert-shard-coverage passes - tests: add _initialize_managed_agents_endpoints and _init_managed_agents_api_endpoints test coverage so router_code_coverage passes; also fix TestRouterCreateInteractionRouting to reflect that acreate_interaction now correctly routes through _ageneric_api_call_with_fallbacks (restoring fallback support) Co-authored-by: Cursor <cursoragent@cursor.com> * fix: remove InteractionsHTTPHandler._handle_error override to fix type errors AgentsHTTPHandler extends InteractionsHTTPHandler and calls self._handle_error(provider_config=agents_api_config) where agents_api_config is BaseAgentsAPIConfig. Python MRO resolved _handle_error to InteractionsHTTPHandler._handle_error which expected BaseInteractionsAPIConfig, causing 10 mypy arg-type errors in interactions/agents/http_handler.py. Removing the redundant override lets both classes inherit _BaseHTTPHandler._handle_error (provider_config: Any) which is structurally correct for both config types. Co-authored-by: Cursor <cursoragent@cursor.com> * fix: agent-only interactions and managed agents provider routing Resolve None custom_llm_provider in agents HTTP client lookup and set custom_llm_provider on GenericLiteLLMParams for all agent CRUD paths. Stop mapping agent names to proxy model routing; route interactions through _init_interactions_api_endpoints with fallbacks only when model is set. Consolidate duplicate router elif branches for interaction APIs. Co-authored-by: Cursor <cursoragent@cursor.com> * Fix greptile review * test(agents): add unit tests for managed agents SDK and HTTP handler Adds coverage for the new `litellm.interactions.agents` surface area: - main.py: sync/async entry points (create/list/get/delete/list_versions), provider config lookup, logging-obj helper, async error wrapping - http_handler.py: every CRUD method (sync + async paths), `_is_async` dispatch branches, and provider error mapping through GeminiAgentsConfig - utils.py: get_provider_agents_api_config for supported / unsupported providers Brings patch coverage on these files from <25% to ~100% so codecov/patch is satisfied. Co-authored-by: Mateo Wang <mateo-berri@users.noreply.github.com> * docs(gemini-agents): fix misleading credential-passing examples in GET/DELETE docstrings (#28293) The four GET/DELETE endpoint docstrings (list_gemini_agents, get_gemini_agent, delete_gemini_agent, list_gemini_agent_versions) documented passing per-request credentials as flat query parameters (e.g. ?api_key=AIza...). However, _merge_query_params_into_data only reads the JSON-encoded litellm_params_template query parameter and intentionally ignores flat params (URL query strings appear verbatim in access logs, browser history, and Referer headers). Callers following the documented curl examples would have their credentials silently dropped and hit auth failures against Gemini. Update the examples to use the supported JSON-encoded litellm_params_template query parameter, matching _merge_query_params_into_data's own docstring. Co-authored-by: Cursor Agent <cursoragent@cursor.com> Co-authored-by: Mateo Wang <mateo-berri@users.noreply.github.com> * refactor(agents): rename provider-agnostic agent response types Move GeminiAgent{ListResponse,DeleteResult,VersionsResponse} to provider-neutral names (AgentListResponse, AgentDeleteResult, AgentVersionsResponse) so the BaseAgentsAPIConfig interface no longer references Gemini-specific type names. * fix(gemini-agents): close veria-flagged credential-escalation gaps Two high-severity findings from the veria-ai PR review are addressed: 1. **api_base override could leak the shared Gemini key** GeminiAgentsConfig.validate_environment falls back to GOOGLE_API_KEY / GEMINI_API_KEY when no api_key is supplied. Combined with caller-controlled api_base on the proxy CRUD endpoints, an authenticated user could redirect the outbound request to an attacker-controlled host and capture the operator's shared Gemini key from the x-goog-api-key header. The config now refuses env-fallback whenever api_base is explicitly overridden. 2. **Managed-agent CRUD exposed to ordinary LLM keys** The new /v1beta/agents routes live in google_routes (i.e. llm_api_routes), so any non-admin LLM key can reach them. Unlike /v1beta/models/...: generateContent these endpoints are NOT model-routed and have no model_list-supplied credentials, so env-fallback would let any LLM key list / create / delete agents inside the operator's Gemini project. Each endpoint now calls _enforce_caller_supplied_provider_key, which requires non-admin callers to supply their own Gemini api_key via litellm_params_template. Proxy admins keep the env-fallback convenience. Tests cover non-admin rejection, admin allow-through, the api_base override guard, and SDK env-fallback when api_base is not overridden. Co-authored-by: Mateo Wang <mateo-berri@users.noreply.github.com> * test(router): restore strict assert_called_once_with on interactions default-provider test --------- Co-authored-by: Cursor Agent <cursoragent@cursor.com> Co-authored-by: Sameer Kankute <Sameerlite@users.noreply.github.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Co-authored-by: Mateo Wang <mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
f73e76afad
commit
47f1de1b02
33 changed files with 4246 additions and 81 deletions
1
.github/workflows/test-unit-proxy-db.yml
vendored
1
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
39
litellm/interactions/agents/__init__.py
Normal file
39
litellm/interactions/agents/__init__.py
Normal file
|
|
@ -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",
|
||||
]
|
||||
478
litellm/interactions/agents/http_handler.py
Normal file
478
litellm/interactions/agents/http_handler.py
Normal file
|
|
@ -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()
|
||||
523
litellm/interactions/agents/main.py
Normal file
523
litellm/interactions/agents/main.py
Normal file
|
|
@ -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,
|
||||
)
|
||||
23
litellm/interactions/agents/utils.py
Normal file
23
litellm/interactions/agents/utils.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ INTERACTIONS_API_OPTIONAL_PARAMS = {
|
|||
"stream",
|
||||
"store",
|
||||
"background",
|
||||
"environment",
|
||||
"response_modalities",
|
||||
"response_format",
|
||||
"response_mime_type",
|
||||
|
|
|
|||
0
litellm/llms/base_llm/agents/__init__.py
Normal file
0
litellm/llms/base_llm/agents/__init__.py
Normal file
165
litellm/llms/base_llm/agents/transformation.py
Normal file
165
litellm/llms/base_llm/agents/transformation.py
Normal file
|
|
@ -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,
|
||||
)
|
||||
0
litellm/llms/gemini/agents/__init__.py
Normal file
0
litellm/llms/gemini/agents/__init__.py
Normal file
299
litellm/llms/gemini/agents/transformation.py
Normal file
299
litellm/llms/gemini/agents/transformation.py
Normal file
|
|
@ -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"),
|
||||
)
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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 = [
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
*,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
445
litellm/proxy/google_endpoints/agents_endpoints.py
Normal file
445
litellm/proxy/google_endpoints/agents_endpoints.py
Normal file
|
|
@ -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"],
|
||||
)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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]]
|
||||
|
|
|
|||
519
tests/proxy_unit_tests/test_gemini_agents_endpoints.py
Normal file
519
tests/proxy_unit_tests/test_gemini_agents_endpoints.py
Normal file
|
|
@ -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"
|
||||
587
tests/test_litellm/interactions/test_agents_http_handler.py
Normal file
587
tests/test_litellm/interactions/test_agents_http_handler.py
Normal file
|
|
@ -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,
|
||||
)
|
||||
354
tests/test_litellm/interactions/test_agents_main_and_utils.py
Normal file
354
tests/test_litellm/interactions/test_agents_main_and_utils.py
Normal file
|
|
@ -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")
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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=<agent_name>`` 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"
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue