diff --git a/litellm/proxy/public_endpoints/public_endpoints.py b/litellm/proxy/public_endpoints/public_endpoints.py index 7d9da543c75..d12e7a35fbf 100644 --- a/litellm/proxy/public_endpoints/public_endpoints.py +++ b/litellm/proxy/public_endpoints/public_endpoints.py @@ -5,7 +5,7 @@ from importlib.resources import files from typing import Any, Dict, List, Optional import litellm -from fastapi import APIRouter, HTTPException +from fastapi import APIRouter, HTTPException, Request from litellm._logging import verbose_logger from litellm.litellm_core_utils.get_blog_posts import ( @@ -211,7 +211,7 @@ async def public_model_hub(): tags=["[beta] Agents", "public"], response_model=List[AgentCard], ) -async def get_agents(): +async def get_agents(request: Request): import litellm from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry @@ -219,12 +219,16 @@ async def get_agents(): if litellm.public_agent_groups is None: return [] - agent_card_list = [ - agent.agent_card_params + + proxy_base = str(request.base_url).rstrip("/") + return [ + { + **(agent.agent_card_params or {}), + "url": f"{proxy_base}/a2a/{agent.agent_id}", + } for agent in agents if agent.agent_id in litellm.public_agent_groups ] - return agent_card_list @router.get( diff --git a/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py b/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py index f82da59899b..6cff91d2c74 100644 --- a/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py +++ b/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py @@ -463,6 +463,70 @@ def test_public_model_hub_mixed_health_statuses(): app.dependency_overrides.clear() +# --------------------------------------------------------------------------- +# /public/agent_hub +# --------------------------------------------------------------------------- + + +def test_public_agent_hub_rewrites_upstream_url_to_proxy(): + """Public agent hub must not leak the upstream backend URL retained on the + stored card. The ``url`` field has to be overwritten with the proxy + ``/a2a/{agent_id}`` entrypoint, matching the well-known card endpoint, so + an unauthenticated client cannot call the backend directly.""" + from litellm.types.agents import AgentResponse + + upstream_url = "https://upstream.internal.example.com/a2a" + agent = AgentResponse( + agent_id="agent-123", + agent_name="public-agent", + agent_card_params={"name": "public-agent", "url": upstream_url}, + ) + + app = FastAPI() + app.include_router(router) + client = TestClient(app) + + mock_registry = MagicMock() + mock_registry.get_public_agent_list.return_value = [agent] + + with ( + patch("litellm.public_agent_groups", ["agent-123"]), + patch( + "litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", + mock_registry, + ), + ): + response = client.get("/public/agent_hub") + + assert response.status_code == 200, response.text + payload = response.json() + assert len(payload) == 1 + card = payload[0] + assert upstream_url not in card.get("url", "") + assert card["url"].endswith("/a2a/agent-123") + + +def test_public_agent_hub_returns_empty_when_no_public_groups(): + app = FastAPI() + app.include_router(router) + client = TestClient(app) + + mock_registry = MagicMock() + mock_registry.get_public_agent_list.return_value = [] + + with ( + patch("litellm.public_agent_groups", None), + patch( + "litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", + mock_registry, + ), + ): + response = client.get("/public/agent_hub") + + assert response.status_code == 200 + assert response.json() == [] + + # --------------------------------------------------------------------------- # /public/endpoints # ---------------------------------------------------------------------------