Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
import java.util.List;
import java.util.Locale;
import java.util.Map;
import java.util.Set;

import org.a2aproject.sdk.client.transport.spi.interceptors.ClientCallContext;
import org.a2aproject.sdk.client.transport.spi.interceptors.ClientCallInterceptor;
Expand All @@ -28,6 +29,22 @@ public class AuthInterceptor extends ClientCallInterceptor {
public static final String AUTHORIZATION = "Authorization";
private static final String BEARER = "Bearer ";
private static final String BASIC = "Basic ";

/**
* Allowlist of header names that are safe for API key injection.
* This prevents credential leakage through malicious header names that could
* be forwarded to third-party origins during redirects or other scenarios.
* Only standard authentication-related headers are permitted.
* Header names are stored in lowercase for case-insensitive comparison.
*/
private static final Set<String> SAFE_API_KEY_HEADER_NAMES = Set.of(
"authorization",
"x-api-key",
"api-key",
"x-auth-token",
"x-authentication"
);

private final CredentialService credentialService;

public AuthInterceptor(final CredentialService credentialService) {
Expand Down Expand Up @@ -66,8 +83,13 @@ public PayloadAndHeaders intercept(String methodName, @Nullable Object payload,
updatedHeaders.put(AUTHORIZATION, getBearerValue(credential));
return new PayloadAndHeaders(payload, updatedHeaders);
} else if (securityScheme instanceof APIKeySecurityScheme apiKeySecurityScheme) {
updatedHeaders.put(apiKeySecurityScheme.name(), credential);
return new PayloadAndHeaders(payload, updatedHeaders);
// Only inject API key if it's intended for header transport and the header name is safe
if (apiKeySecurityScheme.location() == APIKeySecurityScheme.Location.HEADER
&& isSafeHeaderName(apiKeySecurityScheme.name())) {
updatedHeaders.put(apiKeySecurityScheme.name(), credential);
return new PayloadAndHeaders(payload, updatedHeaders);
}
// Skip credential injection for unsafe header names or non-header locations
}
}
}
Expand All @@ -82,4 +104,17 @@ private static String getBearerValue(String credential) {
private static String getBasicValue(String credential) {
return BASIC + credential;
}

/**
* Validates that a header name is safe for API key injection.
* This prevents credential leakage by rejecting header names that could
* be exploited to forward credentials to unintended destinations.
* Header name comparison is case-insensitive per RFC 7230.
*
* @param headerName the header name to validate
* @return true if the header name is in the safe allowlist, false otherwise
*/
private static boolean isSafeHeaderName(String headerName) {
return SAFE_API_KEY_HEADER_NAMES.contains(headerName.toLowerCase(Locale.ROOT));
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,50 @@ private static class AuthTestCase {

@Test
public void testAPIKeySecurityScheme() {
AuthTestCase authTestCase = new AuthTestCase(
"http://agent.com/rpc",
"session-id",
APIKeySecurityScheme.TYPE,
"secret-api-key",
new APIKeySecurityScheme(APIKeySecurityScheme.Location.HEADER, "X-API-Key", "API Key authentication"),
"X-API-Key",
"secret-api-key"
);
testSecurityScheme(authTestCase);
}

@Test
public void testAPIKeySecurityScheme_SafeHeaderName_Authorization() {
AuthTestCase authTestCase = new AuthTestCase(
"http://agent.com/rpc",
"session-id",
APIKeySecurityScheme.TYPE,
"secret-api-key",
new APIKeySecurityScheme(APIKeySecurityScheme.Location.HEADER, "Authorization", "API Key authentication"),
"Authorization",
"secret-api-key"
);
testSecurityScheme(authTestCase);
}

@Test
public void testAPIKeySecurityScheme_SafeHeaderName_XAuthToken() {
AuthTestCase authTestCase = new AuthTestCase(
"http://agent.com/rpc",
"session-id",
APIKeySecurityScheme.TYPE,
"secret-api-key",
new APIKeySecurityScheme(APIKeySecurityScheme.Location.HEADER, "X-Auth-Token", "API Key authentication"),
"X-Auth-Token",
"secret-api-key"
);
testSecurityScheme(authTestCase);
}


@Test
public void testAPIKeySecurityScheme_CaseInsensitiveHeaderName() {
// Test that lowercase header names are accepted (case-insensitive comparison)
AuthTestCase authTestCase = new AuthTestCase(
"http://agent.com/rpc",
"session-id",
Expand All @@ -87,6 +131,124 @@ public void testAPIKeySecurityScheme() {
testSecurityScheme(authTestCase);
}

@Test
public void testAPIKeySecurityScheme_CaseInsensitiveHeaderName_Authorization() {
// Test that lowercase "authorization" is accepted
AuthTestCase authTestCase = new AuthTestCase(
"http://agent.com/rpc",
"session-id",
APIKeySecurityScheme.TYPE,
"secret-api-key",
new APIKeySecurityScheme(APIKeySecurityScheme.Location.HEADER, "authorization", "API Key authentication"),
"authorization",
"secret-api-key"
);
testSecurityScheme(authTestCase);
}


@Test
public void testAPIKeySecurityScheme_UnsafeHeaderName_Rejected() {
String sessionId = "session-id";
String schemeName = APIKeySecurityScheme.TYPE;
String credential = "secret-api-key";

credentialStore.setCredential(sessionId, schemeName, credential);

// Use an unsafe header name that's not in the allowlist
SecurityScheme securityScheme = new APIKeySecurityScheme(
APIKeySecurityScheme.Location.HEADER,
"X-Malicious-Redirect-Header",
"Unsafe header"
);
AgentCard agentCard = createAgentCard(schemeName, securityScheme);

Map<String, Object> requestPayload = Map.of("test", "payload");
Map<String, String> headers = Map.of();
ClientCallContext context = new ClientCallContext(Map.of("sessionId", sessionId), Map.of());

PayloadAndHeaders result = authInterceptor.intercept(
"SendMessage",
requestPayload,
headers,
agentCard,
context
);

assertEquals(requestPayload, result.getPayload());
// Credential should NOT be injected for unsafe header name
assertNull(result.getHeaders().get("X-Malicious-Redirect-Header"));
assertEquals(0, result.getHeaders().size());
}

@Test
public void testAPIKeySecurityScheme_QueryLocation_NotInjectedInHeader() {
String sessionId = "session-id";
String schemeName = APIKeySecurityScheme.TYPE;
String credential = "secret-api-key";

credentialStore.setCredential(sessionId, schemeName, credential);

// API key in query parameter location should not be injected as header
SecurityScheme securityScheme = new APIKeySecurityScheme(
APIKeySecurityScheme.Location.QUERY,
"api_key",
"Query parameter API key"
);
AgentCard agentCard = createAgentCard(schemeName, securityScheme);

Map<String, Object> requestPayload = Map.of("test", "payload");
Map<String, String> headers = Map.of();
ClientCallContext context = new ClientCallContext(Map.of("sessionId", sessionId), Map.of());

PayloadAndHeaders result = authInterceptor.intercept(
"SendMessage",
requestPayload,
headers,
agentCard,
context
);

assertEquals(requestPayload, result.getPayload());
// Credential should NOT be injected as header for query location
assertNull(result.getHeaders().get("api_key"));
assertEquals(0, result.getHeaders().size());
}

@Test
public void testAPIKeySecurityScheme_CookieLocation_NotInjectedInHeader() {
String sessionId = "session-id";
String schemeName = APIKeySecurityScheme.TYPE;
String credential = "secret-api-key";

credentialStore.setCredential(sessionId, schemeName, credential);

// API key in cookie location should not be injected as header
SecurityScheme securityScheme = new APIKeySecurityScheme(
APIKeySecurityScheme.Location.COOKIE,
"session_token",
"Cookie-based API key"
);
AgentCard agentCard = createAgentCard(schemeName, securityScheme);

Map<String, Object> requestPayload = Map.of("test", "payload");
Map<String, String> headers = Map.of();
ClientCallContext context = new ClientCallContext(Map.of("sessionId", sessionId), Map.of());

PayloadAndHeaders result = authInterceptor.intercept(
"SendMessage",
requestPayload,
headers,
agentCard,
context
);

assertEquals(requestPayload, result.getPayload());
// Credential should NOT be injected as header for cookie location
assertNull(result.getHeaders().get("session_token"));
assertEquals(0, result.getHeaders().size());
}

@Test
public void testOAuth2SecurityScheme() {
AuthTestCase authTestCase = new AuthTestCase(
Expand Down Expand Up @@ -238,9 +400,9 @@ void testAvailableSecuritySchemeNotInAgentCardSecuritySchemes() {
String schemeName = "missing";
String sessionId = "session-id";
String credential = "dummy-token";

credentialStore.setCredential(sessionId, schemeName, credential);

// Create agent card with security requirement but no scheme definition
AgentCard agentCard = AgentCard.builder()
.name("missing")
Expand All @@ -254,7 +416,7 @@ void testAvailableSecuritySchemeNotInAgentCardSecuritySchemes() {
.securityRequirements(List.of(SecurityRequirement.builder().scheme(schemeName, List.of()).build()))
.securitySchemes(Map.of()) // no security schemes
.build();

Map<String, Object> requestPayload = Map.of("foo", "bar");
Map<String, String> headers = Map.of("fizz", "buzz");
ClientCallContext context = new ClientCallContext(Map.of("sessionId", sessionId), Map.of());
Expand All @@ -276,7 +438,7 @@ void testNoCredentialAvailable() {
String schemeName = "apikey";
SecurityScheme securityScheme = new APIKeySecurityScheme(APIKeySecurityScheme.Location.HEADER, "X-API-Key", "API Key authentication");
AgentCard agentCard = createAgentCard(schemeName, securityScheme);

Map<String, Object> requestPayload = Map.of("test", "payload");
Map<String, String> headers = Map.of();
ClientCallContext context = new ClientCallContext(Map.of("sessionId", "session-id"), Map.of());
Expand Down Expand Up @@ -307,7 +469,7 @@ void testNoAgentCardSecuritySpecified() {
.skills(List.of())
.securityRequirements(null) // no security info
.build();

Map<String, Object> requestPayload = Map.of("test", "payload");
Map<String, String> headers = Map.of();
ClientCallContext context = new ClientCallContext(Map.of("sessionId", "session-id"), Map.of());
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
import java.util.List;
import java.util.Locale;
import java.util.Map;
import java.util.Set;

import org.a2aproject.sdk.compat03.client.transport.spi.interceptors.ClientCallContext_v0_3;
import org.a2aproject.sdk.compat03.client.transport.spi.interceptors.ClientCallInterceptor_v0_3;
Expand All @@ -27,6 +28,22 @@ public class AuthInterceptor_v0_3 extends ClientCallInterceptor_v0_3 {
public static final String AUTHORIZATION = "Authorization";
private static final String BEARER = "Bearer ";
private static final String BASIC = "Basic ";

/**
* Allowlist of header names that are safe for API key injection.
* This prevents credential leakage through malicious header names that could
* be forwarded to third-party origins during redirects or other scenarios.
* Only standard authentication-related headers are permitted.
* Header names are stored in lowercase for case-insensitive comparison.
*/
private static final Set<String> SAFE_API_KEY_HEADER_NAMES = Set.of(
"authorization",
"x-api-key",
"api-key",
"x-auth-token",
"x-authentication"
);

private final CredentialService_v0_3 credentialService;

public AuthInterceptor_v0_3(final CredentialService_v0_3 credentialService) {
Expand Down Expand Up @@ -62,8 +79,13 @@ public PayloadAndHeaders_v0_3 intercept(String methodName, @Nullable Object payl
updatedHeaders.put(AUTHORIZATION, getBearerValue(credential));
return new PayloadAndHeaders_v0_3(payload, updatedHeaders);
} else if (securityScheme instanceof APIKeySecurityScheme_v0_3 apiKeySecurityScheme) {
updatedHeaders.put(apiKeySecurityScheme.name(), credential);
return new PayloadAndHeaders_v0_3(payload, updatedHeaders);
// Only inject API key if it's intended for header transport and the header name is safe
if ("header".equals(apiKeySecurityScheme.in())
&& isSafeHeaderName(apiKeySecurityScheme.name())) {
updatedHeaders.put(apiKeySecurityScheme.name(), credential);
return new PayloadAndHeaders_v0_3(payload, updatedHeaders);
}
// Skip credential injection for unsafe header names or non-header locations
}
}
}
Expand All @@ -78,4 +100,17 @@ private static String getBearerValue(String credential) {
private static String getBasicValue(String credential) {
return BASIC + credential;
}

/**
* Validates that a header name is safe for API key injection.
* This prevents credential leakage by rejecting header names that could
* be exploited to forward credentials to unintended destinations.
* Header name comparison is case-insensitive per RFC 7230.
*
* @param headerName the header name to validate
* @return true if the header name is in the safe allowlist, false otherwise
*/
private static boolean isSafeHeaderName(String headerName) {
return SAFE_API_KEY_HEADER_NAMES.contains(headerName.toLowerCase(Locale.ROOT));
}
}
Loading
Loading