diff --git a/agentex/src/domain/services/agent_acp_service.py b/agentex/src/domain/services/agent_acp_service.py index 3a1bfddc..acb8b1f7 100644 --- a/agentex/src/domain/services/agent_acp_service.py +++ b/agentex/src/domain/services/agent_acp_service.py @@ -278,6 +278,10 @@ async def get_headers( agent: AgentEntity, request_headers: dict[str, str] | None = None, ) -> dict[str, str]: + # Fall back to inbound request headers when callers don't pass them, so + # allowlisted client x-* headers are forwarded to the agent. + if request_headers is None: + request_headers = dict(self._request.headers) filtered_request_headers = filter_request_headers(request_headers) delegation_headers = self.get_delegation_headers(agent) auth_headers = await self.get_agent_auth_headers(agent) diff --git a/agentex/tests/unit/services/test_agent_acp_service.py b/agentex/tests/unit/services/test_agent_acp_service.py index d4dce6d6..85bd96e7 100644 --- a/agentex/tests/unit/services/test_agent_acp_service.py +++ b/agentex/tests/unit/services/test_agent_acp_service.py @@ -417,6 +417,61 @@ async def test_get_headers_server_request_id_wins_over_passthrough( assert headers["x-request-id"] != "client-request-id" assert len(headers["x-request-id"]) > 0 + async def test_get_headers_falls_back_to_inbound_request_headers( + self, + agent_acp_service, + mock_request, + sample_agent, + ): + """When request_headers is omitted, safe inbound x-* headers are forwarded.""" + mock_request.state.principal_context = None + mock_request.state.agent_identity = None + mock_request.headers = { + "x-trace-id": "trace-789", + "x-user-id": "user-123", + "x-api-key": "must-not-forward", + "authorization": "Bearer must-not-forward", + "user-agent": "must-not-forward", + } + + with patch.object( + agent_acp_service, + "get_agent_auth_headers", + new=AsyncMock(return_value={}), + ): + headers = await agent_acp_service.get_headers(sample_agent) + + assert headers["x-trace-id"] == "trace-789" + assert headers["x-user-id"] == "user-123" + assert "x-api-key" not in headers + assert "authorization" not in headers + assert "user-agent" not in headers + + async def test_get_headers_empty_request_headers_forwards_none( + self, + agent_acp_service, + mock_request, + sample_agent, + ): + """Passing an explicit empty dict forwards no client headers (no fallback).""" + mock_request.state.principal_context = None + mock_request.state.agent_identity = None + mock_request.headers = {"x-trace-id": "should-not-be-used"} + + with patch.object( + agent_acp_service, + "get_agent_auth_headers", + new=AsyncMock(return_value={}), + ): + headers = await agent_acp_service.get_headers( + sample_agent, + request_headers={}, + ) + + assert "x-trace-id" not in headers + # Only the server-generated x-request-id should remain. + assert set(headers.keys()) == {"x-request-id"} + async def test_send_message_success_data( self, agent_acp_service,