diff --git a/src/agentex/lib/utils/registration.py b/src/agentex/lib/utils/registration.py index 36b5f9a04..ef25743f8 100644 --- a/src/agentex/lib/utils/registration.py +++ b/src/agentex/lib/utils/registration.py @@ -12,12 +12,13 @@ logger = make_logger(__name__) + def get_auth_principal(env_vars: EnvironmentVariables): if not env_vars.AUTH_PRINCIPAL_B64: return None try: - decoded_str = base64.b64decode(env_vars.AUTH_PRINCIPAL_B64).decode('utf-8') + decoded_str = base64.b64decode(env_vars.AUTH_PRINCIPAL_B64).decode("utf-8") return json.loads(decoded_str) except Exception: return None @@ -53,10 +54,7 @@ async def register_agent(env_vars: EnvironmentVariables, agent_card=None): # Build the agent's own URL full_acp_url = f"{env_vars.ACP_URL.rstrip('/')}:{env_vars.ACP_PORT}" - description = ( - env_vars.AGENT_DESCRIPTION - or f"Generic description for agent: {env_vars.AGENT_NAME}" - ) + description = env_vars.AGENT_DESCRIPTION or f"Generic description for agent: {env_vars.AGENT_NAME}" registration_metadata = build_registration_metadata(env_vars, agent_card) @@ -77,6 +75,7 @@ async def register_agent(env_vars: EnvironmentVariables, agent_card=None): # Make the registration request registration_url = f"{env_vars.AGENTEX_BASE_URL.rstrip('/')}/agents/register" + registration_headers = {"x-agent-api-key": env_vars.AGENT_API_KEY} if env_vars.AGENT_API_KEY else {} # Retry logic with configurable attempts and delay max_retries = 3 base_delay = 5 # seconds @@ -87,46 +86,43 @@ async def register_agent(env_vars: EnvironmentVariables, agent_card=None): try: async with httpx.AsyncClient() as client: response = await client.post( - registration_url, json=registration_data, timeout=30.0 + registration_url, json=registration_data, headers=registration_headers, timeout=30.0 ) if response.status_code == 200: agent = response.json() agent_id, agent_name = agent["id"], agent["name"] - agent_api_key = agent["agent_api_key"] + returned_key = agent.get("agent_api_key") + if returned_key is not None and not isinstance(returned_key, str): + raise ValueError("Unexpected API key type in registration response") + agent_api_key = returned_key or env_vars.AGENT_API_KEY os.environ["AGENT_ID"] = agent_id os.environ["AGENT_NAME"] = agent_name - os.environ["AGENT_API_KEY"] = agent_api_key env_vars.AGENT_ID = agent_id env_vars.AGENT_NAME = agent_name - env_vars.AGENT_API_KEY = agent_api_key + if agent_api_key: + os.environ["AGENT_API_KEY"] = agent_api_key + env_vars.AGENT_API_KEY = agent_api_key global refreshed_environment_variables refreshed_environment_variables = env_vars logger.info( - f"Successfully registered agent '{env_vars.AGENT_NAME}' with Agentex server with acp_url: {full_acp_url}. Registration data: {registration_data}" + f"Successfully registered agent '{env_vars.AGENT_NAME}' with Agentex server with acp_url: {full_acp_url}." ) return # Success, exit the retry loop else: - error_msg = f"Failed to register agent. Status: {response.status_code}, Response: {response.text}" + error_msg = f"Failed to register agent. Status: {response.status_code}" logger.error(error_msg) - last_exception = Exception( - f"Failed to startup agent: {response.text}" - ) + last_exception = RuntimeError(error_msg) except Exception as e: - logger.error( - f"Exception during agent registration attempt {attempt + 1}: {e}" - ) - last_exception = e + logger.error(f"Exception during agent registration attempt {attempt + 1}: {type(e).__name__}") + # Transport errors and response bodies may contain request credentials. + last_exception = RuntimeError(f"Failed to register agent ({type(e).__name__})") attempt += 1 if attempt < max_retries: delay = (attempt) * base_delay # 5, 10, 15 seconds - logger.info( - f"Retrying in {delay} seconds... (attempt {attempt}/{max_retries})" - ) + logger.info(f"Retrying in {delay} seconds... (attempt {attempt}/{max_retries})") await asyncio.sleep(delay) # If we get here, all retries failed - raise last_exception or Exception( - f"Failed to register agent after {max_retries} attempts" - ) + raise last_exception or Exception(f"Failed to register agent after {max_retries} attempts") diff --git a/tests/lib/test_agentex_worker.py b/tests/lib/test_agentex_worker.py index b0bf47a63..8b2fe50c5 100644 --- a/tests/lib/test_agentex_worker.py +++ b/tests/lib/test_agentex_worker.py @@ -135,6 +135,7 @@ def _env_vars_mock(): env.ACP_PORT = 8000 env.AGENT_DESCRIPTION = "test description" env.AGENT_NAME = "test-agent" + env.AGENT_API_KEY = None env.ACP_TYPE = "agentic" env.AUTH_PRINCIPAL_B64 = None env.AGENTEX_DEPLOYMENT_ID = None @@ -145,7 +146,7 @@ def _env_vars_mock(): return env @staticmethod - def _httpx_client_mock(captured_payloads): + def _httpx_client_mock(captured_payloads, captured_headers): response = MagicMock() response.status_code = 200 response.json.return_value = { @@ -154,8 +155,9 @@ def _httpx_client_mock(captured_payloads): "agent_api_key": "api-key", } - async def post(url, json=None, timeout=None): # noqa: ARG001 + async def post(url, json=None, headers=None, timeout=None): # noqa: ARG001 captured_payloads.append(json) + captured_headers.append(headers) return response client = MagicMock() @@ -227,7 +229,8 @@ async def test_supplied_card_forwarded_exactly_once_by_run_lifecycle(self): mock_register.assert_awaited_once_with(env, agent_card=card) - async def test_worker_and_fastacp_paths_serialize_the_same_card_shape(self): + @pytest.mark.parametrize("configured_key", [None, "configured-agent-key"]) + async def test_worker_and_fastacp_paths_serialize_the_same_card_shape(self, configured_key): """The worker path and the FastACP/BaseACPServer lifespan path hand the same card to register_agent, so the registration payload's registration_metadata.agent_card is identical.""" @@ -238,6 +241,7 @@ async def test_worker_and_fastacp_paths_serialize_the_same_card_shape(self): card = AgentCard(metadata={"permits_capable": True, "region": "us"}) worker_payloads = [] + worker_headers = [] worker = AgentexWorker( task_queue="test-queue", health_check_port=8080, agent_card=card ) @@ -248,12 +252,14 @@ async def test_worker_and_fastacp_paths_serialize_the_same_card_shape(self): "agentex.lib.core.temporal.workers.worker.EnvironmentVariables" ) as mock_env_cls, patch( "agentex.lib.utils.registration.httpx.AsyncClient", - new=self._httpx_client_mock(worker_payloads), + new=self._httpx_client_mock(worker_payloads, worker_headers), ): mock_env_cls.refresh.return_value = self._env_vars_mock() + mock_env_cls.refresh.return_value.AGENT_API_KEY = configured_key await worker._register_agent() acp_payloads = [] + acp_headers = [] server = BaseACPServer.create() server._agent_card = card lifespan = server.get_lifespan_function() @@ -267,9 +273,10 @@ async def test_worker_and_fastacp_paths_serialize_the_same_card_shape(self): new=AsyncMock(), ), patch( "agentex.lib.utils.registration.httpx.AsyncClient", - new=self._httpx_client_mock(acp_payloads), + new=self._httpx_client_mock(acp_payloads, acp_headers), ): mock_env_cls.refresh.return_value = self._env_vars_mock() + mock_env_cls.refresh.return_value.AGENT_API_KEY = configured_key async with lifespan(MagicMock()): pass @@ -278,6 +285,8 @@ async def test_worker_and_fastacp_paths_serialize_the_same_card_shape(self): worker_card = worker_payloads[0]["registration_metadata"]["agent_card"] acp_card = acp_payloads[0]["registration_metadata"]["agent_card"] assert worker_card == acp_card == card.model_dump() + expected_headers = {"x-agent-api-key": configured_key} if configured_key else {} + assert worker_headers == acp_headers == [expected_headers] class TestGetTemporalClientMetricsConfig: diff --git a/tests/lib/utils/test_registration.py b/tests/lib/utils/test_registration.py index 65960d757..2a3f54bef 100644 --- a/tests/lib/utils/test_registration.py +++ b/tests/lib/utils/test_registration.py @@ -1,10 +1,16 @@ -"""Registration metadata: what an agent reports about itself at startup.""" +"""Startup registration metadata and credential compatibility.""" from __future__ import annotations +import os +import json +from unittest.mock import AsyncMock, call + +import httpx import pytest -from agentex.lib.utils.registration import build_registration_metadata +from agentex.lib.utils import registration +from agentex.lib.utils.registration import register_agent, build_registration_metadata from agentex.lib.environment_variables import EnvironmentVariables SHA = "b362b171a9c4e1f09d8e7a6b5c4d3e2f1a0b9c8d" @@ -47,3 +53,154 @@ def model_dump(self): "deployment_id": "dep-1", "agent_card": {"name": "sample"}, } + + +@pytest.fixture +def registration_env(monkeypatch): + for name in ("AGENT_ID", "AGENT_NAME", "AGENT_API_KEY"): + monkeypatch.setenv(name, "") + monkeypatch.delenv(name, raising=False) + monkeypatch.setattr(registration, "refreshed_environment_variables", None, raising=False) + return _env(AGENTEX_BASE_URL="https://agentex.example.test/") + + +@pytest.fixture +def retry_sleep(monkeypatch): + sleep = AsyncMock() + monkeypatch.setattr(registration.asyncio, "sleep", sleep) + return sleep + + +def _success(**overrides): + return httpx.Response(200, json={"id": "agent-1", "name": "sample-agent", **overrides}) + + +@pytest.mark.parametrize("configured_key", [None, "", "configured-agent-key"]) +@pytest.mark.parametrize("response_key", [{}, {"agent_api_key": None}, {"agent_api_key": ""}]) +async def test_registration_preserves_configured_key_when_response_has_no_key( + registration_env, respx_mock, configured_key, response_key +): + registration_env.AGENT_API_KEY = configured_key + if configured_key is not None: + os.environ["AGENT_API_KEY"] = configured_key + route = respx_mock.post("https://agentex.example.test/agents/register").mock(return_value=_success(**response_key)) + + await register_agent(registration_env) + + request = route.calls.last.request + assert route.call_count == 1 + assert request.headers.get("x-agent-api-key") == (configured_key or None) + assert "authorization" not in request.headers + assert "agent_api_key" not in json.loads(request.content) + assert registration_env.AGENT_API_KEY == configured_key + assert os.environ.get("AGENT_API_KEY") == configured_key + assert registration_env.AGENT_ID == os.environ["AGENT_ID"] == "agent-1" + assert registration_env.AGENT_NAME == os.environ["AGENT_NAME"] == "sample-agent" + + +@pytest.mark.parametrize("configured_key", [None, "older-agent-key"]) +async def test_registration_accepts_returned_key(registration_env, respx_mock, configured_key): + registration_env.AGENT_API_KEY = configured_key + route = respx_mock.post("https://agentex.example.test/agents/register").mock( + return_value=_success(agent_api_key="returned-agent-key") + ) + + await register_agent(registration_env) + + assert route.calls.last.request.headers.get("x-agent-api-key") == configured_key + assert registration_env.AGENT_API_KEY == os.environ["AGENT_API_KEY"] == "returned-agent-key" + + +@pytest.mark.parametrize("status", [401, 403]) +async def test_registration_retries_preserve_configured_headers( + registration_env, respx_mock, retry_sleep, caplog, status +): + key = "configured-agent-key" + registration_env.AGENT_API_KEY = key + os.environ["AGENT_API_KEY"] = key + route = respx_mock.post("https://agentex.example.test/agents/register").mock( + return_value=httpx.Response(status, text=f"Rejected credential: {key}") + ) + + with pytest.raises(RuntimeError, match=f"Status: {status}") as exc: + await register_agent(registration_env) + + assert route.call_count == 3 + assert all(item.request.headers["x-agent-api-key"] == key for item in route.calls) + assert registration_env.AGENT_API_KEY == os.environ["AGENT_API_KEY"] == key + assert registration_env.AGENT_ID is None + assert "AGENT_ID" not in os.environ + assert retry_sleep.await_args_list == [call(5), call(10)] + assert key not in caplog.text + assert key not in str(exc.value) + + +async def test_transient_failure_retries_with_configured_key(registration_env, respx_mock, retry_sleep): + registration_env.AGENT_API_KEY = "configured-agent-key" + route = respx_mock.post("https://agentex.example.test/agents/register").mock( + side_effect=[httpx.Response(503), _success(agent_api_key="returned-agent-key")] + ) + + await register_agent(registration_env) + + assert route.call_count == 2 + assert all(item.request.headers["x-agent-api-key"] == "configured-agent-key" for item in route.calls) + assert registration_env.AGENT_API_KEY == os.environ["AGENT_API_KEY"] == "returned-agent-key" + retry_sleep.assert_awaited_once_with(5) + + +async def test_transport_error_does_not_expose_credential(registration_env, respx_mock, retry_sleep, caplog): + key = "configured-agent-key" + registration_env.AGENT_API_KEY = key + route = respx_mock.post("https://agentex.example.test/agents/register").mock( + side_effect=httpx.ReadTimeout(f"Timeout sending {key}") + ) + + with pytest.raises(RuntimeError, match="ReadTimeout") as exc: + await register_agent(registration_env) + + assert route.call_count == 3 + assert all(item.request.headers["x-agent-api-key"] == key for item in route.calls) + assert key not in caplog.text + assert key not in str(exc.value) + assert exc.value.__context__ is None + + +async def test_success_does_not_log_registration_payload(registration_env, respx_mock, caplog): + registration_env.AGENT_API_KEY = "configured-agent-key" + registration_env.AGENT_DESCRIPTION = "private-description" + respx_mock.post("https://agentex.example.test/agents/register").mock( + return_value=_success(agent_api_key="returned-agent-key") + ) + + await register_agent(registration_env, agent_card={"name": "private-card"}) + + assert "Successfully registered agent" in caplog.text + for value in ("configured-agent-key", "returned-agent-key", "private-description", "private-card"): + assert value not in caplog.text + + +@pytest.mark.parametrize("returned_key", [123, {"value": "not-a-string"}]) +async def test_invalid_returned_key_does_not_overwrite_configuration( + registration_env, respx_mock, retry_sleep, returned_key +): + registration_env.AGENT_API_KEY = "configured-agent-key" + os.environ["AGENT_API_KEY"] = "configured-agent-key" + respx_mock.post("https://agentex.example.test/agents/register").mock( + return_value=_success(agent_api_key=returned_key) + ) + + with pytest.raises(RuntimeError, match="ValueError"): + await register_agent(registration_env) + + assert registration_env.AGENT_API_KEY == os.environ["AGENT_API_KEY"] == "configured-agent-key" + assert registration_env.AGENT_ID is None + assert "AGENT_ID" not in os.environ + + +async def test_missing_base_url_skips_registration(registration_env, respx_mock): + registration_env.AGENTEX_BASE_URL = None + + await register_agent(registration_env) + + assert not respx_mock.calls