diff --git a/server-common/src/main/java/org/a2aproject/sdk/server/AgentCardValidator.java b/server-common/src/main/java/org/a2aproject/sdk/server/AgentCardValidator.java index 856f13328..298ff90f9 100644 --- a/server-common/src/main/java/org/a2aproject/sdk/server/AgentCardValidator.java +++ b/server-common/src/main/java/org/a2aproject/sdk/server/AgentCardValidator.java @@ -5,7 +5,7 @@ import java.util.List; import java.util.ServiceLoader; import java.util.Set; -import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.ConcurrentHashMap; import java.util.function.Consumer; import java.util.function.Supplier; import java.util.logging.Logger; @@ -13,11 +13,13 @@ import jakarta.enterprise.inject.Instance; +import org.a2aproject.sdk.server.multitenancy.AgentCardRouter; +import org.a2aproject.sdk.server.multitenancy.TenantNotFoundException; import org.a2aproject.sdk.server.util.CdiUtils; import org.a2aproject.sdk.spec.AgentCard; -import org.jspecify.annotations.Nullable; import org.a2aproject.sdk.spec.AgentInterface; import org.a2aproject.sdk.spec.TransportProtocol; +import org.jspecify.annotations.Nullable; /** * Validates AgentCard transport configuration against available transport endpoints. @@ -35,41 +37,52 @@ public class AgentCardValidator { public static final String SKIP_GRPC_PROPERTY = "org.a2aproject.sdk.transport.grpc.skipValidation"; public static final String SKIP_REST_PROPERTY = "org.a2aproject.sdk.transport.rest.skipValidation"; + /** + * Creates a new thread-safe set for tracking which {@link AgentCard} instances have already + * been validated. Each distinct card is validated at most once; subsequent requests for the + * same card skip validation. + * + * @return a concurrent set suitable for use as the {@code validatedCards} parameter + */ + public static Set newValidatedCardsSet() { + return ConcurrentHashMap.newKeySet(); + } + /** * Resolves an {@link AgentCard} from the given {@link Instance} and validates its transport - * configuration exactly once using the default {@link #validateTransportConfiguration} check. - * The {@code transportValidated} guard is reset on failure so validation can be retried on - * the next call. + * configuration once per distinct card using the default + * {@link #validateTransportConfiguration} check. On failure the card is removed from the + * set so validation can be retried on the next call. * * @param agentCardInstance the CDI instance holding the agent card - * @param transportValidated atomic guard ensuring validation runs once + * @param validatedCards set tracking which cards have already been validated * @return the resolved agent card */ public static AgentCard resolveAndValidateOnce(Instance agentCardInstance, - AtomicBoolean transportValidated) { - return resolveAndValidateOnce(agentCardInstance::get, transportValidated, + Set validatedCards) { + return resolveAndValidateOnce(agentCardInstance::get, validatedCards, AgentCardValidator::validateTransportConfiguration); } /** * Obtains an {@link AgentCard} from the given supplier and applies the provided validator - * exactly once. The {@code transportValidated} guard is reset on failure so validation can - * be retried on the next call. + * once per distinct card. On failure the card is removed from the set so validation can be + * retried on the next call. * * @param agentCardSupplier supplier that produces the agent card - * @param transportValidated atomic guard ensuring validation runs once + * @param validatedCards set tracking which cards have already been validated * @param validator validation logic to apply on first access * @return the resolved agent card */ public static AgentCard resolveAndValidateOnce(Supplier agentCardSupplier, - AtomicBoolean transportValidated, + Set validatedCards, Consumer validator) { AgentCard card = agentCardSupplier.get(); - if (transportValidated.compareAndSet(false, true)) { + if (validatedCards.add(card)) { try { validator.accept(card); } catch (RuntimeException e) { - transportValidated.set(false); + validatedCards.remove(card); throw e; } } @@ -82,22 +95,118 @@ public static AgentCard resolveAndValidateOnce(Supplier agentCardSupp * * @param publicCard the CDI instance for the {@code @PublicAgentCard} * @param extendedCard the CDI instance for the {@code @ExtendedAgentCard}, may be {@code null} - * @param transportValidated atomic guard ensuring validation runs once + * @param validatedCards set tracking which cards have already been validated * @return the resolved agent card * @throws IllegalStateException if neither card is available */ public static AgentCard resolveWithFallback(Instance publicCard, @Nullable Instance extendedCard, - AtomicBoolean transportValidated) { + Set validatedCards) { + return resolveWithFallback(publicCard, extendedCard, null, null, validatedCards); + } + + /** + * Convenience overload that uses the default {@link #validateTransportConfiguration} validator. + * + * @see #resolveWithFallback(Instance, Instance, AgentCardRouter, String, Set, Consumer) + */ + public static AgentCard resolveWithFallback(Instance publicCard, + @Nullable Instance extendedCard, + @Nullable AgentCardRouter agentCardRouter, + @Nullable String tenant, + Set validatedCards) { + return resolveWithFallback(publicCard, extendedCard, agentCardRouter, tenant, validatedCards, + AgentCardValidator::validateTransportConfiguration); + } + + /** + * Resolves an agent card using one of two resolution strategies depending on whether a + * tenant-scoped request is being made, and validates using the supplied validator. + * + * @param publicCard the CDI instance for the {@code @PublicAgentCard} + * @param extendedCard the CDI instance for the {@code @ExtendedAgentCard}, may be {@code null} + * @param agentCardRouter optional router for tenant-specific card resolution + * @param tenant the tenant identifier, may be {@code null} + * @param validatedCards set tracking which cards have already been validated + * @param validator validation logic to apply on first access of each distinct card + * @return the resolved agent card + * @throws TenantNotFoundException if a non-blank tenant is specified, a router is available, + * but neither a public nor extended card is registered for that tenant + * @throws IllegalStateException if no card can be resolved (non-tenant-scoped path) + * @see #resolveTenantScoped(AgentCardRouter, String, Set, Consumer) + * @see #resolveDefaultWithRouterFallback(Instance, Instance, AgentCardRouter, String, Set, Consumer) + */ + public static AgentCard resolveWithFallback(Instance publicCard, + @Nullable Instance extendedCard, + @Nullable AgentCardRouter agentCardRouter, + @Nullable String tenant, + Set validatedCards, + Consumer validator) { + if (tenant != null && !tenant.isBlank() && agentCardRouter != null) { + return resolveTenantScoped(agentCardRouter, tenant, validatedCards, validator); + } + return resolveDefaultWithRouterFallback(publicCard, extendedCard, agentCardRouter, tenant, + validatedCards, validator); + } + + /** + * Tenant-scoped resolution: only the {@link AgentCardRouter} is consulted. CDI default beans + * are not used as fallbacks — doing so would let the request proceed against the + * wrong tenant's card. + *
    + *
  1. Tenant-specific public card via the router
  2. + *
  3. Tenant-specific extended card via the router
  4. + *
  5. {@link TenantNotFoundException} if neither is registered
  6. + *
+ */ + private static AgentCard resolveTenantScoped(AgentCardRouter agentCardRouter, String tenant, + Set validatedCards, Consumer validator) { + AgentCard routerCard = agentCardRouter.resolvePublicCard(tenant); + if (routerCard != null) { + return resolveAndValidateOnce(() -> routerCard, validatedCards, validator); + } + AgentCard routerExtCard = agentCardRouter.resolveExtendedCard(tenant); + if (routerExtCard != null) { + return resolveAndValidateOnce(() -> routerExtCard, validatedCards, validator); + } + throw new TenantNotFoundException(tenant); + } + + /** + * Non-tenant-scoped resolution with optional router fallback: + *
    + *
  1. Unqualified {@code @PublicAgentCard} CDI bean
  2. + *
  3. Router's public card (when no default bean exists)
  4. + *
  5. Unqualified {@code @ExtendedAgentCard} CDI bean
  6. + *
  7. Router's extended card (when no default bean exists)
  8. + *
  9. {@link IllegalStateException} if nothing resolves
  10. + *
+ */ + private static AgentCard resolveDefaultWithRouterFallback(Instance publicCard, + @Nullable Instance extendedCard, + @Nullable AgentCardRouter agentCardRouter, + @Nullable String tenant, + Set validatedCards, + Consumer validator) { AgentCard resolved = CdiUtils.resolveDefault(publicCard); if (resolved != null) { - return resolveAndValidateOnce(() -> resolved, transportValidated, - AgentCardValidator::validateTransportConfiguration); + return resolveAndValidateOnce(() -> resolved, validatedCards, validator); + } + if (agentCardRouter != null) { + AgentCard routerCard = agentCardRouter.resolvePublicCard(tenant); + if (routerCard != null) { + return resolveAndValidateOnce(() -> routerCard, validatedCards, validator); + } } AgentCard extResolved = CdiUtils.resolveDefault(extendedCard); if (extResolved != null) { - return resolveAndValidateOnce(() -> extResolved, transportValidated, - AgentCardValidator::validateTransportConfiguration); + return resolveAndValidateOnce(() -> extResolved, validatedCards, validator); + } + if (agentCardRouter != null) { + AgentCard routerExtCard = agentCardRouter.resolveExtendedCard(tenant); + if (routerExtCard != null) { + return resolveAndValidateOnce(() -> routerExtCard, validatedCards, validator); + } } throw new IllegalStateException(NO_AGENT_CARD_MESSAGE); } diff --git a/server-common/src/main/java/org/a2aproject/sdk/server/requesthandlers/DefaultRequestHandler.java b/server-common/src/main/java/org/a2aproject/sdk/server/requesthandlers/DefaultRequestHandler.java index 4388b0443..6f72722f3 100644 --- a/server-common/src/main/java/org/a2aproject/sdk/server/requesthandlers/DefaultRequestHandler.java +++ b/server-common/src/main/java/org/a2aproject/sdk/server/requesthandlers/DefaultRequestHandler.java @@ -262,10 +262,15 @@ public class DefaultRequestHandler implements RequestHandler { */ int reconciliationTimeoutSeconds; + // Only used inside initConfig() (CDI lifecycle). In static create() paths this field + // remains null and is never accessed, hence the NullAway suppression. + @SuppressWarnings("NullAway") + private @Nullable Instance agentExecutorInstance; + // Fields set by constructor injection cannot be final. We need a noargs constructor for // Jakarta compatibility, and it seems that making fields set by constructor injection // final, is not proxyable in all runtimes - private AgentExecutor agentExecutor; + private @Nullable AgentExecutor agentExecutor; private TaskStore taskStore; private QueueManager queueManager; private PushNotificationConfigStore pushConfigStore; @@ -298,6 +303,7 @@ public class DefaultRequestHandler implements RequestHandler { @SuppressWarnings("NullAway") protected DefaultRequestHandler() { // For CDI proxy creation + this.agentExecutorInstance = null; this.agentExecutor = null; this.taskStore = null; this.queueManager = null; @@ -313,11 +319,33 @@ protected DefaultRequestHandler() { * releases; application code should use {@link #builder()} to configure and create a handler. */ @Inject - public DefaultRequestHandler(AgentExecutor agentExecutor, TaskStore taskStore, + public DefaultRequestHandler(@Any Instance agentExecutorInstance, TaskStore taskStore, QueueManager queueManager, PushNotificationConfigStore pushConfigStore, MainEventBusProcessor mainEventBusProcessor, @Internal Executor executor, @EventConsumerExecutor Executor eventConsumerExecutor) { + this.agentExecutorInstance = agentExecutorInstance; + this.agentExecutor = null; + this.taskStore = taskStore; + this.queueManager = queueManager; + this.pushConfigStore = pushConfigStore; + this.mainEventBusProcessor = mainEventBusProcessor; + this.executor = executor; + this.eventConsumerExecutor = eventConsumerExecutor; + this.requestContextBuilder = () -> new SimpleRequestContextBuilder(taskStore, false, null); + this.mainEventBusProcessor.start(); + } + + /** + * Constructor used by the {@link Builder} and tests. + * The builder always supplies a concrete executor, so no CDI resolution is needed. + */ + DefaultRequestHandler(AgentExecutor agentExecutor, TaskStore taskStore, + QueueManager queueManager, PushNotificationConfigStore pushConfigStore, + MainEventBusProcessor mainEventBusProcessor, + Executor executor, + Executor eventConsumerExecutor) { + this.agentExecutorInstance = null; this.agentExecutor = agentExecutor; this.taskStore = taskStore; this.queueManager = queueManager; @@ -325,10 +353,6 @@ public DefaultRequestHandler(AgentExecutor agentExecutor, TaskStore taskStore, this.mainEventBusProcessor = mainEventBusProcessor; this.executor = executor; this.eventConsumerExecutor = eventConsumerExecutor; - // TODO In Python this is also a constructor parameter defaulting to this SimpleRequestContextBuilder - // implementation if the parameter is null. Skip that for now, since otherwise I get CDI errors, and - // I am unsure about the correct scope. - // Also reworked to make a Supplier since otherwise the builder gets polluted with wrong tasks this.requestContextBuilder = () -> new SimpleRequestContextBuilder(taskStore, false, null); this.mainEventBusProcessor.start(); } @@ -344,6 +368,9 @@ void initConfig() { configProvider.getValue(A2A_BLOCKING_RECONCILIATION_TIMEOUT_SECONDS)); authorizationProvider = CdiUtils.getIfResolvable(authorizationProviderInstance); agentExecutorRouter = CdiUtils.getIfResolvable(agentExecutorRouterInstance); + if (agentExecutor == null && agentExecutorInstance != null) { + agentExecutor = CdiUtils.resolveDefault(agentExecutorInstance); + } pushNotificationsEnabled = Boolean.parseBoolean( configProvider.getValue(A2A_PUSH_NOTIFICATIONS_ENABLED)); @@ -1436,10 +1463,16 @@ private MessageSendSetup initMessageSend(MessageSendParams params, ServerCallCon return new MessageSendSetup(taskManager, task, requestContext); } - private AgentExecutor resolveAgentExecutor(@Nullable String tenant) { + AgentExecutor resolveAgentExecutor(@Nullable String tenant) { if (agentExecutorRouter != null) { return agentExecutorRouter.resolve(tenant); } + if (agentExecutor == null) { + throw new IllegalStateException( + "No AgentExecutor available. Either provide an unqualified AgentExecutor " + + "CDI bean, or add the multitenancy extension (a2a-java-sdk-extras-multitenancy) " + + "which provides an AgentExecutorRouter for @Tenant-qualified beans."); + } return agentExecutor; } diff --git a/server-common/src/test/java/org/a2aproject/sdk/server/AgentCardValidatorTest.java b/server-common/src/test/java/org/a2aproject/sdk/server/AgentCardValidatorTest.java index b15a0af4a..249438b10 100644 --- a/server-common/src/test/java/org/a2aproject/sdk/server/AgentCardValidatorTest.java +++ b/server-common/src/test/java/org/a2aproject/sdk/server/AgentCardValidatorTest.java @@ -9,16 +9,18 @@ import java.util.Collections; import java.util.List; import java.util.Set; -import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicInteger; import java.util.logging.Handler; import java.util.logging.LogRecord; import java.util.logging.Logger; +import org.a2aproject.sdk.server.multitenancy.AgentCardRouter; +import org.a2aproject.sdk.server.multitenancy.TenantNotFoundException; import org.a2aproject.sdk.spec.AgentCapabilities; import org.a2aproject.sdk.spec.AgentCard; import org.a2aproject.sdk.spec.AgentInterface; import org.a2aproject.sdk.spec.TransportProtocol; +import org.jspecify.annotations.Nullable; import org.junit.jupiter.api.Test; public class AgentCardValidatorTest { @@ -235,23 +237,23 @@ void testSkipPropertiesFilterWarnings() { @Test void resolveAndValidateOnceRetriesAfterFailure() { AgentCard card = createTestAgentCardBuilder().build(); - AtomicBoolean guard = new AtomicBoolean(false); + Set validatedCards = AgentCardValidator.newValidatedCardsSet(); AtomicInteger validationCount = new AtomicInteger(0); assertThrows(IllegalStateException.class, () -> AgentCardValidator.resolveAndValidateOnce( - () -> card, guard, c -> { + () -> card, validatedCards, c -> { validationCount.incrementAndGet(); throw new IllegalStateException("transient failure"); })); - assertFalse(guard.get(), "guard should be reset after failure"); + assertFalse(validatedCards.contains(card), "card should be removed after failure"); assertEquals(1, validationCount.get()); AgentCard result = AgentCardValidator.resolveAndValidateOnce( - () -> card, guard, c -> validationCount.incrementAndGet()); + () -> card, validatedCards, c -> validationCount.incrementAndGet()); - assertTrue(guard.get(), "guard should be set after success"); + assertTrue(validatedCards.contains(card), "card should be in set after success"); assertEquals(2, validationCount.get()); assertEquals(card, result); } @@ -259,15 +261,32 @@ void resolveAndValidateOnceRetriesAfterFailure() { @Test void resolveAndValidateOnceSkipsValidationOnSubsequentCalls() { AgentCard card = createTestAgentCardBuilder().build(); - AtomicBoolean guard = new AtomicBoolean(false); + Set validatedCards = AgentCardValidator.newValidatedCardsSet(); AtomicInteger validationCount = new AtomicInteger(0); AgentCardValidator.resolveAndValidateOnce( - () -> card, guard, c -> validationCount.incrementAndGet()); + () -> card, validatedCards, c -> validationCount.incrementAndGet()); AgentCardValidator.resolveAndValidateOnce( - () -> card, guard, c -> validationCount.incrementAndGet()); + () -> card, validatedCards, c -> validationCount.incrementAndGet()); - assertEquals(1, validationCount.get(), "validation should run only once"); + assertEquals(1, validationCount.get(), "validation should run only once for the same card"); + } + + @Test + void resolveAndValidateOnceValidatesEachDistinctCard() { + AgentCard card1 = createTestAgentCardBuilder().name("card-1").build(); + AgentCard card2 = createTestAgentCardBuilder().name("card-2").build(); + Set validatedCards = AgentCardValidator.newValidatedCardsSet(); + AtomicInteger validationCount = new AtomicInteger(0); + + AgentCardValidator.resolveAndValidateOnce( + () -> card1, validatedCards, c -> validationCount.incrementAndGet()); + AgentCardValidator.resolveAndValidateOnce( + () -> card2, validatedCards, c -> validationCount.incrementAndGet()); + + assertEquals(2, validationCount.get(), "validation should run once per distinct card"); + assertTrue(validatedCards.contains(card1)); + assertTrue(validatedCards.contains(card2)); } @Test @@ -276,10 +295,10 @@ void resolveWithFallbackUsesPublicCardWhenPresent() { try { AgentCard publicCard = createTestAgentCardBuilder().name("public").build(); AgentCard extendedCard = createTestAgentCardBuilder().name("extended").build(); - AtomicBoolean guard = new AtomicBoolean(false); + Set validatedCards = AgentCardValidator.newValidatedCardsSet(); AgentCard result = AgentCardValidator.resolveWithFallback( - new FixedInstance<>(publicCard), new FixedInstance<>(extendedCard), guard); + new FixedInstance<>(publicCard), new FixedInstance<>(extendedCard), validatedCards); assertEquals("public", result.name()); } finally { @@ -292,10 +311,10 @@ void resolveWithFallbackFallsBackToExtendedCard() { System.setProperty(AgentCardValidator.SKIP_PROPERTY, "true"); try { AgentCard extendedCard = createTestAgentCardBuilder().name("extended").build(); - AtomicBoolean guard = new AtomicBoolean(false); + Set validatedCards = AgentCardValidator.newValidatedCardsSet(); AgentCard result = AgentCardValidator.resolveWithFallback( - FixedInstance.empty(), new FixedInstance<>(extendedCard), guard); + FixedInstance.empty(), new FixedInstance<>(extendedCard), validatedCards); assertEquals("extended", result.name()); } finally { @@ -305,22 +324,22 @@ void resolveWithFallbackFallsBackToExtendedCard() { @Test void resolveWithFallbackThrowsWhenBothAbsent() { - AtomicBoolean guard = new AtomicBoolean(false); + Set validatedCards = AgentCardValidator.newValidatedCardsSet(); IllegalStateException ex = assertThrows(IllegalStateException.class, () -> AgentCardValidator.resolveWithFallback( - FixedInstance.empty(), FixedInstance.empty(), guard)); + FixedInstance.empty(), FixedInstance.empty(), validatedCards)); assertEquals(AgentCardValidator.NO_AGENT_CARD_MESSAGE, ex.getMessage()); } @Test void resolveWithFallbackThrowsWhenExtendedIsNull() { - AtomicBoolean guard = new AtomicBoolean(false); + Set validatedCards = AgentCardValidator.newValidatedCardsSet(); IllegalStateException ex = assertThrows(IllegalStateException.class, () -> AgentCardValidator.resolveWithFallback( - FixedInstance.empty(), null, guard)); + FixedInstance.empty(), null, validatedCards)); assertEquals(AgentCardValidator.NO_AGENT_CARD_MESSAGE, ex.getMessage()); } @@ -348,6 +367,214 @@ void requireFirstThrowsWhenBothNull() { assertEquals(AgentCardValidator.NO_AGENT_CARD_MESSAGE, ex.getMessage()); } + @Test + void resolveWithFallbackUsesRouterWhenNoDefaultBean() { + System.setProperty(AgentCardValidator.SKIP_PROPERTY, "true"); + try { + AgentCard routerCard = createTestAgentCardBuilder().name("router-card").build(); + AgentCardRouter router = new AgentCardRouter() { + @Override + public @Nullable AgentCard resolveExtendedCard(@Nullable String tenant) { + return null; + } + + @Override + public @Nullable AgentCard resolvePublicCard(@Nullable String tenant) { + return routerCard; + } + }; + Set validatedCards = AgentCardValidator.newValidatedCardsSet(); + + AgentCard result = AgentCardValidator.resolveWithFallback( + FixedInstance.empty(), FixedInstance.empty(), router, "tenant-1", validatedCards); + + assertEquals("router-card", result.name()); + } finally { + System.clearProperty(AgentCardValidator.SKIP_PROPERTY); + } + } + + @Test + void resolveWithFallbackPrefersTenantCardOverDefaultBean() { + System.setProperty(AgentCardValidator.SKIP_PROPERTY, "true"); + try { + AgentCard publicCard = createTestAgentCardBuilder().name("default").build(); + AgentCard routerCard = createTestAgentCardBuilder().name("tenant-card").build(); + AgentCardRouter router = new AgentCardRouter() { + @Override + public @Nullable AgentCard resolveExtendedCard(@Nullable String tenant) { + return null; + } + + @Override + public @Nullable AgentCard resolvePublicCard(@Nullable String tenant) { + return "tenant-1".equals(tenant) ? routerCard : null; + } + }; + Set validatedCards = AgentCardValidator.newValidatedCardsSet(); + + AgentCard result = AgentCardValidator.resolveWithFallback( + new FixedInstance<>(publicCard), FixedInstance.empty(), router, "tenant-1", validatedCards); + + assertEquals("tenant-card", result.name()); + } finally { + System.clearProperty(AgentCardValidator.SKIP_PROPERTY); + } + } + + @Test + void resolveWithFallbackPrefersDefaultBeanOverRouterWithoutTenant() { + System.setProperty(AgentCardValidator.SKIP_PROPERTY, "true"); + try { + AgentCard publicCard = createTestAgentCardBuilder().name("default").build(); + AgentCard routerCard = createTestAgentCardBuilder().name("router-default").build(); + AgentCardRouter router = new AgentCardRouter() { + @Override + public @Nullable AgentCard resolveExtendedCard(@Nullable String tenant) { + return null; + } + + @Override + public @Nullable AgentCard resolvePublicCard(@Nullable String tenant) { + return routerCard; + } + }; + Set validatedCards = AgentCardValidator.newValidatedCardsSet(); + + AgentCard result = AgentCardValidator.resolveWithFallback( + new FixedInstance<>(publicCard), FixedInstance.empty(), router, null, validatedCards); + + assertEquals("default", result.name()); + } finally { + System.clearProperty(AgentCardValidator.SKIP_PROPERTY); + } + } + + @Test + void resolveWithFallbackPrefersTenantPublicCardOverDefaultExtendedCard() { + System.setProperty(AgentCardValidator.SKIP_PROPERTY, "true"); + try { + AgentCard defaultExtended = createTestAgentCardBuilder().name("default-extended").build(); + AgentCard tenantPublic = createTestAgentCardBuilder().name("tenant-public").build(); + AgentCardRouter router = new AgentCardRouter() { + @Override + public @Nullable AgentCard resolveExtendedCard(@Nullable String tenant) { + return null; + } + + @Override + public @Nullable AgentCard resolvePublicCard(@Nullable String tenant) { + return "tenant-1".equals(tenant) ? tenantPublic : null; + } + }; + Set validatedCards = AgentCardValidator.newValidatedCardsSet(); + + AgentCard result = AgentCardValidator.resolveWithFallback( + FixedInstance.empty(), new FixedInstance<>(defaultExtended), router, "tenant-1", validatedCards); + + assertEquals("tenant-public", result.name()); + } finally { + System.clearProperty(AgentCardValidator.SKIP_PROPERTY); + } + } + + @Test + void resolveWithFallbackThrowsForUnknownTenantEvenWhenDefaultExtendedExists() { + AgentCardRouter router = new AgentCardRouter() { + @Override + public @Nullable AgentCard resolveExtendedCard(@Nullable String tenant) { + return null; + } + + @Override + public @Nullable AgentCard resolvePublicCard(@Nullable String tenant) { + return null; + } + }; + Set validatedCards = AgentCardValidator.newValidatedCardsSet(); + AgentCard defaultExtended = createTestAgentCardBuilder().name("default-extended").build(); + + TenantNotFoundException ex = assertThrows(TenantNotFoundException.class, () -> + AgentCardValidator.resolveWithFallback( + FixedInstance.empty(), new FixedInstance<>(defaultExtended), router, "unknown-tenant", + validatedCards)); + + assertEquals("unknown-tenant", ex.getTenant()); + } + + @Test + void resolveWithFallbackThrowsForUnknownTenantEvenWhenDefaultPublicExists() { + AgentCardRouter router = new AgentCardRouter() { + @Override + public @Nullable AgentCard resolveExtendedCard(@Nullable String tenant) { + return null; + } + + @Override + public @Nullable AgentCard resolvePublicCard(@Nullable String tenant) { + return null; + } + }; + Set validatedCards = AgentCardValidator.newValidatedCardsSet(); + AgentCard publicCard = createTestAgentCardBuilder().name("default").build(); + + TenantNotFoundException ex = assertThrows(TenantNotFoundException.class, () -> + AgentCardValidator.resolveWithFallback( + new FixedInstance<>(publicCard), FixedInstance.empty(), router, "unknown-tenant", + validatedCards)); + + assertEquals("unknown-tenant", ex.getTenant()); + } + + @Test + void resolveWithFallbackUsesRouterExtendedCardWhenPublicCardIsNull() { + System.setProperty(AgentCardValidator.SKIP_PROPERTY, "true"); + try { + AgentCard extCard = createTestAgentCardBuilder().name("router-extended").build(); + AgentCardRouter router = new AgentCardRouter() { + @Override + public @Nullable AgentCard resolveExtendedCard(@Nullable String tenant) { + return extCard; + } + + @Override + public @Nullable AgentCard resolvePublicCard(@Nullable String tenant) { + return null; + } + }; + Set validatedCards = AgentCardValidator.newValidatedCardsSet(); + + AgentCard result = AgentCardValidator.resolveWithFallback( + FixedInstance.empty(), FixedInstance.empty(), router, "tenant-1", validatedCards); + + assertEquals("router-extended", result.name()); + } finally { + System.clearProperty(AgentCardValidator.SKIP_PROPERTY); + } + } + + @Test + void resolveWithFallbackThrowsWhenRouterReturnsNull() { + AgentCardRouter router = new AgentCardRouter() { + @Override + public @Nullable AgentCard resolveExtendedCard(@Nullable String tenant) { + return null; + } + + @Override + public @Nullable AgentCard resolvePublicCard(@Nullable String tenant) { + return null; + } + }; + Set validatedCards = AgentCardValidator.newValidatedCardsSet(); + + TenantNotFoundException ex = assertThrows(TenantNotFoundException.class, () -> + AgentCardValidator.resolveWithFallback( + FixedInstance.empty(), FixedInstance.empty(), router, "unknown", validatedCards)); + + assertEquals("unknown", ex.getTenant()); + } + // A simple log handler for testing private static class TestLogHandler extends Handler { private final List logMessages = new java.util.ArrayList<>(); diff --git a/transport/grpc/src/main/java/org/a2aproject/sdk/transport/grpc/handler/GrpcHandler.java b/transport/grpc/src/main/java/org/a2aproject/sdk/transport/grpc/handler/GrpcHandler.java index eed559666..39a23c289 100644 --- a/transport/grpc/src/main/java/org/a2aproject/sdk/transport/grpc/handler/GrpcHandler.java +++ b/transport/grpc/src/main/java/org/a2aproject/sdk/transport/grpc/handler/GrpcHandler.java @@ -25,11 +25,13 @@ import org.a2aproject.sdk.jsonrpc.common.json.JsonUtil; import org.a2aproject.sdk.jsonrpc.common.wrappers.ListTasksResult; import org.a2aproject.sdk.server.AgentCardValidator; +import org.a2aproject.sdk.server.FixedInstance; import org.a2aproject.sdk.server.ServerCallContext; import org.a2aproject.sdk.server.auth.UnauthenticatedUser; import org.a2aproject.sdk.server.auth.User; import org.a2aproject.sdk.server.extensions.A2AExtensions; import org.a2aproject.sdk.server.multitenancy.AgentCardRouter; +import org.a2aproject.sdk.server.multitenancy.TenantNotFoundException; import org.a2aproject.sdk.server.requesthandlers.RequestHandler; import org.a2aproject.sdk.server.version.A2AVersionValidator; import org.a2aproject.sdk.spec.A2AError; @@ -58,7 +60,6 @@ import org.a2aproject.sdk.spec.TaskNotFoundError; import org.a2aproject.sdk.spec.TaskPushNotificationConfig; import org.a2aproject.sdk.spec.TaskQueryParams; -import org.a2aproject.sdk.spec.AgentInterface; import org.a2aproject.sdk.spec.TransportProtocol; import org.a2aproject.sdk.spec.A2AErrorCodes; import org.a2aproject.sdk.spec.UnsupportedOperationError; @@ -161,7 +162,7 @@ public abstract class GrpcHandler extends A2AServiceGrpc.A2AServiceImplBase { // Without this we get intermittent failures private static volatile @Nullable Runnable streamingSubscribedRunnable; - private final AtomicBoolean transportValidated = new AtomicBoolean(false); + private final Set validatedCards = AgentCardValidator.newValidatedCardsSet(); private static final Logger LOGGER = Logger.getLogger(GrpcHandler.class.getName()); @@ -201,12 +202,15 @@ public GrpcHandler() { public void sendMessage(org.a2aproject.sdk.grpc.SendMessageRequest request, StreamObserver responseObserver) { try { - ServerCallContext context = createCallContext(responseObserver); + String tenant = extractTenant(request.getTenant()); + ServerCallContext context = createCallContext(responseObserver, tenant); MessageSendParams params = FromProto.messageSendParams(request); EventKind taskOrMessage = getRequestHandler().onMessageSend(params, context); org.a2aproject.sdk.grpc.SendMessageResponse response = ToProto.taskOrMessage(taskOrMessage); responseObserver.onNext(response); responseObserver.onCompleted(); + } catch (TenantNotFoundException e) { + handleTenantNotFound(responseObserver, e); } catch (A2AError e) { handleError(responseObserver, e); } catch (SecurityException e) { @@ -220,7 +224,8 @@ public void sendMessage(org.a2aproject.sdk.grpc.SendMessageRequest request, public void getTask(org.a2aproject.sdk.grpc.GetTaskRequest request, StreamObserver responseObserver) { try { - ServerCallContext context = createCallContext(responseObserver); + String tenant = extractTenant(request.getTenant()); + ServerCallContext context = createCallContext(responseObserver, tenant); TaskQueryParams params = FromProto.taskQueryParams(request); Task task = getRequestHandler().onGetTask(params, context); if (task != null) { @@ -229,6 +234,8 @@ public void getTask(org.a2aproject.sdk.grpc.GetTaskRequest request, } else { handleError(responseObserver, new TaskNotFoundError()); } + } catch (TenantNotFoundException e) { + handleTenantNotFound(responseObserver, e); } catch (A2AError e) { handleError(responseObserver, e); } catch (SecurityException e) { @@ -242,11 +249,14 @@ public void getTask(org.a2aproject.sdk.grpc.GetTaskRequest request, public void listTasks(org.a2aproject.sdk.grpc.ListTasksRequest request, StreamObserver responseObserver) { try { - ServerCallContext context = createCallContext(responseObserver); + String tenant = extractTenant(request.getTenant()); + ServerCallContext context = createCallContext(responseObserver, tenant); org.a2aproject.sdk.spec.ListTasksParams params = FromProto.listTasksParams(request); ListTasksResult result = getRequestHandler().onListTasks(params, context); responseObserver.onNext(ToProto.listTasksResult(result)); responseObserver.onCompleted(); + } catch (TenantNotFoundException e) { + handleTenantNotFound(responseObserver, e); } catch (A2AError e) { handleError(responseObserver, e); } catch (SecurityException e) { @@ -260,7 +270,8 @@ public void listTasks(org.a2aproject.sdk.grpc.ListTasksRequest request, public void cancelTask(org.a2aproject.sdk.grpc.CancelTaskRequest request, StreamObserver responseObserver) { try { - ServerCallContext context = createCallContext(responseObserver); + String tenant = extractTenant(request.getTenant()); + ServerCallContext context = createCallContext(responseObserver, tenant); CancelTaskParams params = FromProto.cancelTaskParams(request); Task task = getRequestHandler().onCancelTask(params, context); if (task != null) { @@ -269,6 +280,8 @@ public void cancelTask(org.a2aproject.sdk.grpc.CancelTaskRequest request, } else { handleError(responseObserver, new TaskNotFoundError()); } + } catch (TenantNotFoundException e) { + handleTenantNotFound(responseObserver, e); } catch (A2AError e) { handleError(responseObserver, e); } catch (SecurityException e) { @@ -281,17 +294,20 @@ public void cancelTask(org.a2aproject.sdk.grpc.CancelTaskRequest request, @Override public void createTaskPushNotificationConfig(org.a2aproject.sdk.grpc.TaskPushNotificationConfig request, StreamObserver responseObserver) { - if (!resolveAgentCard().capabilities().pushNotifications()) { - handleError(responseObserver, new PushNotificationNotSupportedError()); - return; - } - try { - ServerCallContext context = createCallContext(responseObserver); + String tenant = extractTenant(request.getTenant()); + if (!resolveAgentCard(tenant).capabilities().pushNotifications()) { + handleError(responseObserver, new PushNotificationNotSupportedError()); + return; + } + + ServerCallContext context = createCallContext(responseObserver, tenant); TaskPushNotificationConfig config = FromProto.createTaskPushNotificationConfig(request); TaskPushNotificationConfig responseConfig = getRequestHandler().onCreateTaskPushNotificationConfig(config, context); responseObserver.onNext(ToProto.taskPushNotificationConfig(responseConfig)); responseObserver.onCompleted(); + } catch (TenantNotFoundException e) { + handleTenantNotFound(responseObserver, e); } catch (A2AError e) { handleError(responseObserver, e); } catch (SecurityException e) { @@ -304,17 +320,20 @@ public void createTaskPushNotificationConfig(org.a2aproject.sdk.grpc.TaskPushNot @Override public void getTaskPushNotificationConfig(org.a2aproject.sdk.grpc.GetTaskPushNotificationConfigRequest request, StreamObserver responseObserver) { - if (!resolveAgentCard().capabilities().pushNotifications()) { - handleError(responseObserver, new PushNotificationNotSupportedError()); - return; - } - try { - ServerCallContext context = createCallContext(responseObserver); + String tenant = extractTenant(request.getTenant()); + if (!resolveAgentCard(tenant).capabilities().pushNotifications()) { + handleError(responseObserver, new PushNotificationNotSupportedError()); + return; + } + + ServerCallContext context = createCallContext(responseObserver, tenant); GetTaskPushNotificationConfigParams params = FromProto.getTaskPushNotificationConfigParams(request); TaskPushNotificationConfig config = getRequestHandler().onGetTaskPushNotificationConfig(params, context); responseObserver.onNext(ToProto.taskPushNotificationConfig(config)); responseObserver.onCompleted(); + } catch (TenantNotFoundException e) { + handleTenantNotFound(responseObserver, e); } catch (A2AError e) { handleError(responseObserver, e); } catch (SecurityException e) { @@ -327,18 +346,21 @@ public void getTaskPushNotificationConfig(org.a2aproject.sdk.grpc.GetTaskPushNot @Override public void listTaskPushNotificationConfigs(org.a2aproject.sdk.grpc.ListTaskPushNotificationConfigsRequest request, StreamObserver responseObserver) { - if (!resolveAgentCard().capabilities().pushNotifications()) { - handleError(responseObserver, new PushNotificationNotSupportedError()); - return; - } - try { - ServerCallContext context = createCallContext(responseObserver); + String tenant = extractTenant(request.getTenant()); + if (!resolveAgentCard(tenant).capabilities().pushNotifications()) { + handleError(responseObserver, new PushNotificationNotSupportedError()); + return; + } + + ServerCallContext context = createCallContext(responseObserver, tenant); ListTaskPushNotificationConfigsParams params = FromProto.listTaskPushNotificationConfigsParams(request); ListTaskPushNotificationConfigsResult result = getRequestHandler().onListTaskPushNotificationConfigs(params, context); org.a2aproject.sdk.grpc.ListTaskPushNotificationConfigsResponse response = ToProto.listTaskPushNotificationConfigsResponse(result); responseObserver.onNext(response); responseObserver.onCompleted(); + } catch (TenantNotFoundException e) { + handleTenantNotFound(responseObserver, e); } catch (A2AError e) { handleError(responseObserver, e); } catch (SecurityException e) { @@ -388,18 +410,21 @@ public void listTaskPushNotificationConfigs(org.a2aproject.sdk.grpc.ListTaskPush @Override public void sendStreamingMessage(org.a2aproject.sdk.grpc.SendMessageRequest request, StreamObserver responseObserver) { - if (!resolveAgentCard().capabilities().streaming()) { - handleError(responseObserver, - new UnsupportedOperationError(null, "Streaming is not supported by the agent", null)); - return; - } - try { - ServerCallContext context = createCallContext(responseObserver); + String tenant = extractTenant(request.getTenant()); + if (!resolveAgentCard(tenant).capabilities().streaming()) { + handleError(responseObserver, + new UnsupportedOperationError(null, "Streaming is not supported by the agent", null)); + return; + } + + ServerCallContext context = createCallContext(responseObserver, tenant); installForkedContextWrapper(context); MessageSendParams params = FromProto.messageSendParams(request); Flow.Publisher publisher = getRequestHandler().onMessageSendStream(params, context); convertToStreamResponse(publisher, responseObserver, context); + } catch (TenantNotFoundException e) { + handleTenantNotFound(responseObserver, e); } catch (A2AError e) { handleError(responseObserver, e); } catch (SecurityException e) { @@ -412,18 +437,21 @@ public void sendStreamingMessage(org.a2aproject.sdk.grpc.SendMessageRequest requ @Override public void subscribeToTask(org.a2aproject.sdk.grpc.SubscribeToTaskRequest request, StreamObserver responseObserver) { - if (!resolveAgentCard().capabilities().streaming()) { - handleError(responseObserver, - new UnsupportedOperationError(null, "Streaming is not supported by the agent", null)); - return; - } - try { - ServerCallContext context = createCallContext(responseObserver); + String tenant = extractTenant(request.getTenant()); + if (!resolveAgentCard(tenant).capabilities().streaming()) { + handleError(responseObserver, + new UnsupportedOperationError(null, "Streaming is not supported by the agent", null)); + return; + } + + ServerCallContext context = createCallContext(responseObserver, tenant); installForkedContextWrapper(context); TaskIdParams params = FromProto.taskIdParams(request); Flow.Publisher publisher = getRequestHandler().onSubscribeToTask(params, context); convertToStreamResponse(publisher, responseObserver, context); + } catch (TenantNotFoundException e) { + handleTenantNotFound(responseObserver, e); } catch (A2AError e) { handleError(responseObserver, e); } catch (SecurityException e) { @@ -565,18 +593,16 @@ private void completeStream() { public void getExtendedAgentCard(org.a2aproject.sdk.grpc.GetExtendedAgentCardRequest request, StreamObserver responseObserver) { try { - if (!resolveAgentCard().capabilities().extendedAgentCard()) { + String tenant = extractTenant(request.getTenant()); + Utils.validateTenant(tenant); + if (!resolveAgentCard(tenant).capabilities().extendedAgentCard()) { handleError(responseObserver, new UnsupportedOperationError()); return; } - // Removing this 2A causes protocol version validation and required extensions validation to no longer run for gRPC getExtendedAgentCard requests. - //This violates Section 3.6.2 of the A2A spec which requires version checking on every request. - createCallContext(responseObserver); - String tenant = request.getTenant().isBlank() ? null : request.getTenant(); - Utils.validateTenant(tenant); + createCallContext(responseObserver, tenant); AgentCardRouter router = getAgentCardRouter(); AgentCard extendedAgentCard; - if (router != null) { + if (tenant != null && !tenant.isBlank() && router != null) { extendedAgentCard = router.resolveExtendedCard(tenant); } else { extendedAgentCard = getExtendedAgentCard(); @@ -585,9 +611,10 @@ public void getExtendedAgentCard(org.a2aproject.sdk.grpc.GetExtendedAgentCardReq responseObserver.onNext(ToProto.agentCard(extendedAgentCard)); responseObserver.onCompleted(); } else { - // Extended agent card not configured - return error instead of hanging handleError(responseObserver, new ExtendedAgentCardNotConfiguredError(null, "Extended agent card not configured", null)); } + } catch (TenantNotFoundException e) { + handleTenantNotFound(responseObserver, e); } catch (Throwable t) { handleInternalError(responseObserver, t); } @@ -596,18 +623,20 @@ public void getExtendedAgentCard(org.a2aproject.sdk.grpc.GetExtendedAgentCardReq @Override public void deleteTaskPushNotificationConfig(org.a2aproject.sdk.grpc.DeleteTaskPushNotificationConfigRequest request, StreamObserver responseObserver) { - if (!resolveAgentCard().capabilities().pushNotifications()) { - handleError(responseObserver, new PushNotificationNotSupportedError()); - return; - } - try { - ServerCallContext context = createCallContext(responseObserver); + String tenant = extractTenant(request.getTenant()); + if (!resolveAgentCard(tenant).capabilities().pushNotifications()) { + handleError(responseObserver, new PushNotificationNotSupportedError()); + return; + } + + ServerCallContext context = createCallContext(responseObserver, tenant); DeleteTaskPushNotificationConfigParams params = FromProto.deleteTaskPushNotificationConfigParams(request); getRequestHandler().onDeleteTaskPushNotificationConfig(params, context); - // void response responseObserver.onNext(Empty.getDefaultInstance()); responseObserver.onCompleted(); + } catch (TenantNotFoundException e) { + handleTenantNotFound(responseObserver, e); } catch (A2AError e) { handleError(responseObserver, e); } catch (SecurityException e) { @@ -653,7 +682,7 @@ public void deleteTaskPushNotificationConfig(org.a2aproject.sdk.grpc.DeleteTaskP * @see CallContextFactory * @see org.a2aproject.sdk.transport.grpc.context.GrpcContextKeys */ - private ServerCallContext createCallContext(StreamObserver responseObserver) { + private ServerCallContext createCallContext(StreamObserver responseObserver, @Nullable String tenant) { CallContextFactory factory = getCallContextFactory(); ServerCallContext context; if (factory == null) { @@ -719,8 +748,9 @@ private ServerCallContext createCallContext(StreamObserver responseObserv context = factory.create(responseObserver); // Fall back to basic create() method for now } - A2AVersionValidator.validateProtocolVersion(resolveAgentCard(), context); - A2AExtensions.validateRequiredExtensions(resolveAgentCard(), context); + AgentCard agentCard = resolveAgentCard(tenant); + A2AVersionValidator.validateProtocolVersion(agentCard, context); + A2AExtensions.validateRequiredExtensions(agentCard, context); return context; } @@ -847,10 +877,22 @@ private void handleInternalError(StreamObserver responseObserver, Throwab } - private AgentCard resolveAgentCard() { - return AgentCardValidator.resolveAndValidateOnce( - () -> AgentCardValidator.requireFirst(getAgentCard(), getExtendedAgentCard()), - transportValidated, this::validateTransportConfigurationWithCorrectClassLoader); + private AgentCard resolveAgentCard(@Nullable String tenant) { + AgentCard publicCard = getAgentCard(); + AgentCard extendedCard = getExtendedAgentCard(); + return AgentCardValidator.resolveWithFallback( + publicCard != null ? new FixedInstance<>(publicCard) : FixedInstance.empty(), + extendedCard != null ? new FixedInstance<>(extendedCard) : FixedInstance.empty(), + getAgentCardRouter(), tenant, validatedCards, + this::validateTransportConfigurationWithCorrectClassLoader); + } + + private static @Nullable String extractTenant(String protoTenant) { + return protoTenant.isBlank() ? null : protoTenant; + } + + private void handleTenantNotFound(StreamObserver responseObserver, TenantNotFoundException e) { + responseObserver.onError(Status.NOT_FOUND.withDescription(e.getResponseMessage()).asRuntimeException()); } private void validateTransportConfigurationWithCorrectClassLoader(AgentCard agentCard) { diff --git a/transport/grpc/src/test/java/org/a2aproject/sdk/transport/grpc/handler/GrpcHandlerTest.java b/transport/grpc/src/test/java/org/a2aproject/sdk/transport/grpc/handler/GrpcHandlerTest.java index 160b7ebb7..69f3aceae 100644 --- a/transport/grpc/src/test/java/org/a2aproject/sdk/transport/grpc/handler/GrpcHandlerTest.java +++ b/transport/grpc/src/test/java/org/a2aproject/sdk/transport/grpc/handler/GrpcHandlerTest.java @@ -842,7 +842,17 @@ protected AgentCardRouter getAgentCardRouter() { public void testExtendedAgentCardWithRouterReturnsNull() throws Exception { AgentCard cardWithExtCapability = AgentCard.builder(AbstractA2ARequestHandlerTest.CARD) .capabilities(AgentCapabilities.builder().extendedAgentCard(true).build()).build(); - AgentCardRouter router = tenant -> null; + AgentCardRouter router = new AgentCardRouter() { + @Override + public AgentCard resolveExtendedCard(String tenant) { + return null; + } + + @Override + public AgentCard resolvePublicCard(String tenant) { + return cardWithExtCapability; + } + }; GrpcHandler handler = new TestGrpcHandler(cardWithExtCapability, requestHandler, internalExecutor) { @Override @@ -859,6 +869,50 @@ protected AgentCardRouter getAgentCardRouter() { assertGrpcError(recorder, Status.Code.FAILED_PRECONDITION); } + @Test + public void testSendMessageReturnsNotFoundForUnknownTenant() throws Exception { + AgentCardRouter router = new AgentCardRouter() { + @Override + public AgentCard resolveExtendedCard(String tenant) { + return null; + } + + @Override + public AgentCard resolvePublicCard(String tenant) { + return "known".equals(tenant) ? AbstractA2ARequestHandlerTest.CARD : null; + } + }; + + GrpcHandler handler = new TestGrpcHandler(null, requestHandler, internalExecutor) { + @Override + protected @org.jspecify.annotations.Nullable AgentCard getAgentCard() { + return null; + } + + @Override + protected AgentCard getExtendedAgentCard() { + return null; + } + + @Override + protected AgentCardRouter getAgentCardRouter() { + return router; + } + }; + + SendMessageRequest request = SendMessageRequest.newBuilder() + .setMessage(GRPC_MESSAGE) + .setTenant("unknown-tenant") + .build(); + StreamRecorder recorder = StreamRecorder.create(); + handler.sendMessage(request, recorder); + + Assertions.assertNotNull(recorder.getError()); + Assertions.assertInstanceOf(StatusRuntimeException.class, recorder.getError()); + StatusRuntimeException sre = (StatusRuntimeException) recorder.getError(); + Assertions.assertEquals(Status.Code.NOT_FOUND, sre.getStatus().getCode()); + } + @Test public void testExtendedAgentCardWithoutRouter() throws Exception { AgentCard cardWithExtCapability = AgentCard.builder(AbstractA2ARequestHandlerTest.CARD) diff --git a/transport/jsonrpc/src/main/java/org/a2aproject/sdk/transport/jsonrpc/handler/JSONRPCHandler.java b/transport/jsonrpc/src/main/java/org/a2aproject/sdk/transport/jsonrpc/handler/JSONRPCHandler.java index 2f51a748a..4eebb2fe3 100644 --- a/transport/jsonrpc/src/main/java/org/a2aproject/sdk/transport/jsonrpc/handler/JSONRPCHandler.java +++ b/transport/jsonrpc/src/main/java/org/a2aproject/sdk/transport/jsonrpc/handler/JSONRPCHandler.java @@ -2,10 +2,10 @@ import static org.a2aproject.sdk.server.util.async.AsyncUtils.createTubeConfig; +import java.util.Set; import java.util.concurrent.CompletableFuture; import java.util.concurrent.Executor; import java.util.concurrent.Flow; -import java.util.concurrent.atomic.AtomicBoolean; import java.util.logging.Level; import java.util.logging.Logger; @@ -56,6 +56,7 @@ import org.a2aproject.sdk.spec.EventKind; import org.a2aproject.sdk.spec.ExtendedAgentCardNotConfiguredError; import org.a2aproject.sdk.spec.InternalError; +import org.a2aproject.sdk.spec.InvalidParamsError; import org.a2aproject.sdk.spec.ListTaskPushNotificationConfigsResult; import org.a2aproject.sdk.spec.PushNotificationNotSupportedError; import org.a2aproject.sdk.spec.StreamingEventKind; @@ -149,7 +150,7 @@ public class JSONRPCHandler { private @Nullable Instance extendedAgentCard; private RequestHandler requestHandler; private Executor executor; - private final AtomicBoolean transportValidated = new AtomicBoolean(false); + private final Set validatedCards = AgentCardValidator.newValidatedCardsSet(); private @Nullable AgentCardRouter agentCardRouter; @@ -244,7 +245,8 @@ public JSONRPCHandler(@Nullable AgentCard agentCard, RequestHandler requestHandl */ public SendMessageResponse onMessageSend(SendMessageRequest request, ServerCallContext context) { try { - validateVersionAndExtensions(context); + String tenant = request.getParams() != null ? request.getParams().tenant() : null; + validateVersionAndExtensions(tenant, context); EventKind taskOrMessage = requestHandler.onMessageSend(request.getParams(), context); return new SendMessageResponse(request.getId(), taskOrMessage); } catch (A2AError e) { @@ -292,15 +294,15 @@ public SendMessageResponse onMessageSend(SendMessageRequest request, ServerCallC */ public Flow.Publisher onMessageSendStream( SendStreamingMessageRequest request, ServerCallContext context) { - if (!resolveAgentCard().capabilities().streaming()) { - return ZeroPublisher.fromItems( - new SendStreamingMessageResponse( - request.getId(), - new UnsupportedOperationError(null, "Streaming is not supported by the agent", null))); - } - + String tenant = request.getParams() != null ? request.getParams().tenant() : null; try { - validateVersionAndExtensions(context); + validateVersionAndExtensions(tenant, context); + if (!resolveAgentCard(tenant).capabilities().streaming()) { + return ZeroPublisher.fromItems( + new SendStreamingMessageResponse( + request.getId(), + new UnsupportedOperationError(null, "Streaming is not supported by the agent", null))); + } Flow.Publisher publisher = requestHandler.onMessageSendStream(request.getParams(), context); // We can't use the convertingProcessor convenience method since that propagates any errors as an error handled @@ -340,7 +342,8 @@ public Flow.Publisher onMessageSendStream( */ public CancelTaskResponse onCancelTask(CancelTaskRequest request, ServerCallContext context) { try { - validateVersionAndExtensions(context); + String tenant = request.getParams() != null ? request.getParams().tenant() : null; + validateVersionAndExtensions(tenant, context); Task task = requestHandler.onCancelTask(request.getParams(), context); if (task != null) { return new CancelTaskResponse(request.getId(), task); @@ -389,15 +392,20 @@ public CancelTaskResponse onCancelTask(CancelTaskRequest request, ServerCallCont */ public Flow.Publisher onSubscribeToTask( SubscribeToTaskRequest request, ServerCallContext context) throws A2AError { - if (!resolveAgentCard().capabilities().streaming()) { - return ZeroPublisher.fromItems( - new SendStreamingMessageResponse( - request.getId(), - new UnsupportedOperationError(null, "Streaming is not supported by the agent", null))); + String tenant = request.getParams() != null ? request.getParams().tenant() : null; + try { + validateVersionAndExtensions(tenant, context); + if (!resolveAgentCard(tenant).capabilities().streaming()) { + return ZeroPublisher.fromItems( + new SendStreamingMessageResponse( + request.getId(), + new UnsupportedOperationError(null, "Streaming is not supported by the agent", null))); + } + } catch (A2AError e) { + return ZeroPublisher.fromItems(new SendStreamingMessageResponse(request.getId(), e)); } requestHandler.authorizeTaskAccess(request.getParams().id(), context, TaskOperation.SUBSCRIBE_TO_TASK); try { - validateVersionAndExtensions(context); Flow.Publisher publisher = requestHandler.onSubscribeToTask(request.getParams(), context); // We can't use the convertingProcessor convenience method since that propagates any errors as an error handled @@ -437,12 +445,13 @@ public Flow.Publisher onSubscribeToTask( */ public GetTaskPushNotificationConfigResponse getPushNotificationConfig( GetTaskPushNotificationConfigRequest request, ServerCallContext context) { - if (!resolveAgentCard().capabilities().pushNotifications()) { - return new GetTaskPushNotificationConfigResponse(request.getId(), - new PushNotificationNotSupportedError()); - } + String tenant = request.getParams() != null ? request.getParams().tenant() : null; try { - validateVersionAndExtensions(context); + validateVersionAndExtensions(tenant, context); + if (!resolveAgentCard(tenant).capabilities().pushNotifications()) { + return new GetTaskPushNotificationConfigResponse(request.getId(), + new PushNotificationNotSupportedError()); + } TaskPushNotificationConfig config = requestHandler.onGetTaskPushNotificationConfig(request.getParams(), context); return new GetTaskPushNotificationConfigResponse(request.getId(), config); @@ -480,12 +489,13 @@ public GetTaskPushNotificationConfigResponse getPushNotificationConfig( */ public CreateTaskPushNotificationConfigResponse setPushNotificationConfig( CreateTaskPushNotificationConfigRequest request, ServerCallContext context) { - if (!resolveAgentCard().capabilities().pushNotifications()) { - return new CreateTaskPushNotificationConfigResponse(request.getId(), - new PushNotificationNotSupportedError()); - } + String tenant = request.getParams() != null ? request.getParams().tenant() : null; try { - validateVersionAndExtensions(context); + validateVersionAndExtensions(tenant, context); + if (!resolveAgentCard(tenant).capabilities().pushNotifications()) { + return new CreateTaskPushNotificationConfigResponse(request.getId(), + new PushNotificationNotSupportedError()); + } TaskPushNotificationConfig config = requestHandler.onCreateTaskPushNotificationConfig(request.getParams(), context); return new CreateTaskPushNotificationConfigResponse(request.getId(), config); @@ -522,7 +532,8 @@ public CreateTaskPushNotificationConfigResponse setPushNotificationConfig( */ public GetTaskResponse onGetTask(GetTaskRequest request, ServerCallContext context) { try { - validateVersionAndExtensions(context); + String tenant = request.getParams() != null ? request.getParams().tenant() : null; + validateVersionAndExtensions(tenant, context); Task task = requestHandler.onGetTask(request.getParams(), context); return new GetTaskResponse(request.getId(), task); } catch (A2AError e) { @@ -570,7 +581,8 @@ public GetTaskResponse onGetTask(GetTaskRequest request, ServerCallContext conte */ public ListTasksResponse onListTasks(ListTasksRequest request, ServerCallContext context) { try { - validateVersionAndExtensions(context); + String tenant = request.getParams() != null ? request.getParams().tenant() : null; + validateVersionAndExtensions(tenant, context); ListTasksResult result = requestHandler.onListTasks(request.getParams(), context); return new ListTasksResponse(request.getId(), result); } catch (A2AError e) { @@ -606,12 +618,13 @@ public ListTasksResponse onListTasks(ListTasksRequest request, ServerCallContext */ public ListTaskPushNotificationConfigsResponse listPushNotificationConfigs( ListTaskPushNotificationConfigsRequest request, ServerCallContext context) { - if (!resolveAgentCard().capabilities().pushNotifications()) { - return new ListTaskPushNotificationConfigsResponse(request.getId(), - new PushNotificationNotSupportedError()); - } + String tenant = request.getParams() != null ? request.getParams().tenant() : null; try { - validateVersionAndExtensions(context); + validateVersionAndExtensions(tenant, context); + if (!resolveAgentCard(tenant).capabilities().pushNotifications()) { + return new ListTaskPushNotificationConfigsResponse(request.getId(), + new PushNotificationNotSupportedError()); + } ListTaskPushNotificationConfigsResult result = requestHandler.onListTaskPushNotificationConfigs(request.getParams(), context); return new ListTaskPushNotificationConfigsResponse(request.getId(), result); @@ -649,12 +662,13 @@ public ListTaskPushNotificationConfigsResponse listPushNotificationConfigs( */ public DeleteTaskPushNotificationConfigResponse deletePushNotificationConfig( DeleteTaskPushNotificationConfigRequest request, ServerCallContext context) { - if (!resolveAgentCard().capabilities().pushNotifications()) { - return new DeleteTaskPushNotificationConfigResponse(request.getId(), - new PushNotificationNotSupportedError()); - } + String tenant = request.getParams() != null ? request.getParams().tenant() : null; try { - validateVersionAndExtensions(context); + validateVersionAndExtensions(tenant, context); + if (!resolveAgentCard(tenant).capabilities().pushNotifications()) { + return new DeleteTaskPushNotificationConfigResponse(request.getId(), + new PushNotificationNotSupportedError()); + } requestHandler.onDeleteTaskPushNotificationConfig(request.getParams(), context); return new DeleteTaskPushNotificationConfigResponse(request.getId()); } catch (A2AError e) { @@ -690,13 +704,13 @@ public DeleteTaskPushNotificationConfigResponse deletePushNotificationConfig( // TODO: Add authentication (https://github.com/a2aproject/a2a-java/issues/77) public GetExtendedAgentCardResponse onGetExtendedCardRequest( GetExtendedAgentCardRequest request, ServerCallContext context) { - if (!resolveAgentCard().capabilities().extendedAgentCard()) { - return new GetExtendedAgentCardResponse(request.getId(), new UnsupportedOperationError()); - } + String tenant = request.getParams() != null ? request.getParams().tenant() : null; try { - validateVersionAndExtensions(context); - String tenant = request.getParams() != null ? request.getParams().tenant() : null; - if (agentCardRouter != null) { + validateVersionAndExtensions(tenant, context); + if (!resolveAgentCard(tenant).capabilities().extendedAgentCard()) { + return new GetExtendedAgentCardResponse(request.getId(), new UnsupportedOperationError()); + } + if (tenant != null && !tenant.isBlank() && agentCardRouter != null) { AgentCard card = agentCardRouter.resolveExtendedCard(tenant); if (card == null) { return new GetExtendedAgentCardResponse(request.getId(), @@ -704,13 +718,16 @@ public GetExtendedAgentCardResponse onGetExtendedCardRequest( } return new GetExtendedAgentCardResponse(request.getId(), card); } - if (extendedAgentCard == null || !extendedAgentCard.isResolvable()) { + AgentCard extCard = CdiUtils.resolveDefault(extendedAgentCard); + if (extCard == null) { return new GetExtendedAgentCardResponse(request.getId(), new ExtendedAgentCardNotConfiguredError(null, "Extended Card not configured", null)); } - return new GetExtendedAgentCardResponse(request.getId(), extendedAgentCard.get()); + return new GetExtendedAgentCardResponse(request.getId(), extCard); } catch (A2AError e) { return new GetExtendedAgentCardResponse(request.getId(), e); + } catch(TenantNotFoundException ex) { + return new GetExtendedAgentCardResponse(request.getId(), new InvalidParamsError(ex.getResponseMessage())); } catch (Throwable t) { return new GetExtendedAgentCardResponse(request.getId(), internalError(t)); } @@ -725,10 +742,14 @@ public GetExtendedAgentCardResponse onGetExtendedCardRequest( * @param context the server call context carrying the requested version and extensions * @throws A2AError if the requested version or a required extension is not supported */ - private void validateVersionAndExtensions(ServerCallContext context) throws A2AError { - AgentCard agentCard = resolveAgentCard(); - A2AVersionValidator.validateProtocolVersion(agentCard, context); - A2AExtensions.validateRequiredExtensions(agentCard, context); + private void validateVersionAndExtensions(@Nullable String tenant, ServerCallContext context) throws A2AError { + try { + AgentCard agentCard = resolveAgentCard(tenant); + A2AVersionValidator.validateProtocolVersion(agentCard, context); + A2AExtensions.validateRequiredExtensions(agentCard, context); + } catch (TenantNotFoundException ex) { + throw new InvalidParamsError(ex.getResponseMessage()); + } } /** @@ -773,15 +794,24 @@ private void validateVersionAndExtensions(ServerCallContext context) throws A2AE LOGGER.fine(() -> "No AgentCardRouter configured; serving default public card for tenant '" + tenant + "'"); } AgentCard card = CdiUtils.resolveDefault(agentCardInstance); + if (card == null && agentCardRouter != null) { + card = agentCardRouter.resolvePublicCard(tenant); + } if (card == null) { return null; } - return AgentCardValidator.resolveAndValidateOnce(() -> card, transportValidated, + AgentCard validatedCard = card; + return AgentCardValidator.resolveAndValidateOnce(() -> validatedCard, validatedCards, AgentCardValidator::validateTransportConfiguration); } - private AgentCard resolveAgentCard() { - return AgentCardValidator.resolveWithFallback(agentCardInstance, extendedAgentCard, transportValidated); + private AgentCard resolveAgentCard(@Nullable String tenant) { + try { + return AgentCardValidator.resolveWithFallback(agentCardInstance, extendedAgentCard, + agentCardRouter, tenant, validatedCards); + } catch (TenantNotFoundException ex) { + throw new InvalidParamsError(ex.getResponseMessage()); + } } private Flow.Publisher convertToSendStreamingMessageResponse( diff --git a/transport/jsonrpc/src/test/java/org/a2aproject/sdk/transport/jsonrpc/handler/JSONRPCHandlerTest.java b/transport/jsonrpc/src/test/java/org/a2aproject/sdk/transport/jsonrpc/handler/JSONRPCHandlerTest.java index 12936f776..3e3fed012 100644 --- a/transport/jsonrpc/src/test/java/org/a2aproject/sdk/transport/jsonrpc/handler/JSONRPCHandlerTest.java +++ b/transport/jsonrpc/src/test/java/org/a2aproject/sdk/transport/jsonrpc/handler/JSONRPCHandlerTest.java @@ -1608,7 +1608,17 @@ public void testExtendedAgentCardWithRouterKnownTenant() throws Exception { public void testExtendedAgentCardWithRouterReturnsNull() throws Exception { AgentCard cardWithExtCapability = AgentCard.builder(CARD) .capabilities(AgentCapabilities.builder().extendedAgentCard(true).build()).build(); - AgentCardRouter router = tenant -> null; + AgentCardRouter router = new AgentCardRouter() { + @Override + public AgentCard resolveExtendedCard(String tenant) { + return null; + } + + @Override + public AgentCard resolvePublicCard(String tenant) { + return cardWithExtCapability; + } + }; JSONRPCHandler handler = new JSONRPCHandler(new FixedInstance<>(cardWithExtCapability), null, requestHandler, internalExecutor, new FixedInstance<>(router)); diff --git a/transport/rest/src/main/java/org/a2aproject/sdk/transport/rest/handler/RestHandler.java b/transport/rest/src/main/java/org/a2aproject/sdk/transport/rest/handler/RestHandler.java index 2b5c338c8..019f486da 100644 --- a/transport/rest/src/main/java/org/a2aproject/sdk/transport/rest/handler/RestHandler.java +++ b/transport/rest/src/main/java/org/a2aproject/sdk/transport/rest/handler/RestHandler.java @@ -9,10 +9,10 @@ import java.util.HashMap; import java.util.List; import java.util.Map; +import java.util.Set; import java.util.concurrent.CompletableFuture; import java.util.concurrent.Executor; import java.util.concurrent.Flow; -import java.util.concurrent.atomic.AtomicBoolean; import java.util.logging.Level; import java.util.logging.Logger; import java.util.stream.Collectors; @@ -136,7 +136,7 @@ public class RestHandler { private AgentCardCacheMetadata cacheMetadata; private RequestHandler requestHandler; private Executor executor; - private final AtomicBoolean transportValidated = new AtomicBoolean(false); + private final Set validatedCards = AgentCardValidator.newValidatedCardsSet(); private @Nullable AgentCardRouter agentCardRouter; @@ -238,7 +238,7 @@ public RestHandler(@Nullable AgentCard agentCard, AgentCardCacheMetadata cacheMe */ public HTTPRestResponse sendMessage(ServerCallContext context, String tenant, String body) { try { - validateVersionAndExtensions(context); + validateVersionAndExtensions(tenant, context); org.a2aproject.sdk.grpc.SendMessageRequest.Builder request = org.a2aproject.sdk.grpc.SendMessageRequest.newBuilder(); parseRequestBody(body, request); request.setTenant(tenant); @@ -302,10 +302,10 @@ public HTTPRestResponse sendMessage(ServerCallContext context, String tenant, St */ public HTTPRestResponse sendStreamingMessage(ServerCallContext context, String tenant, String body) { try { - if (!resolveAgentCard().capabilities().streaming()) { + validateVersionAndExtensions(tenant, context); + if (!resolveAgentCard(tenant).capabilities().streaming()) { return createErrorResponse(new UnsupportedOperationError(null, "Streaming is not supported by the agent", null)); } - validateVersionAndExtensions(context); org.a2aproject.sdk.grpc.SendMessageRequest.Builder request = org.a2aproject.sdk.grpc.SendMessageRequest.newBuilder(); parseRequestBody(body, request); request.setTenant(tenant); @@ -350,7 +350,7 @@ public HTTPRestResponse sendStreamingMessage(ServerCallContext context, String t @SuppressWarnings("unchecked") public HTTPRestResponse cancelTask(ServerCallContext context, String tenant, String body, String taskId) { try { - validateVersionAndExtensions(context); + validateVersionAndExtensions(tenant, context); if (taskId == null || taskId.isEmpty()) { throw new InvalidParamsError(); } @@ -379,10 +379,10 @@ public HTTPRestResponse cancelTask(ServerCallContext context, String tenant, Str */ public HTTPRestResponse createTaskPushNotificationConfiguration(ServerCallContext context, String tenant, String body, String taskId) { try { - if (!resolveAgentCard().capabilities().pushNotifications()) { + validateVersionAndExtensions(tenant, context); + if (!resolveAgentCard(tenant).capabilities().pushNotifications()) { throw new PushNotificationNotSupportedError(); } - validateVersionAndExtensions(context); org.a2aproject.sdk.grpc.TaskPushNotificationConfig.Builder builder = org.a2aproject.sdk.grpc.TaskPushNotificationConfig.newBuilder(); parseRequestBody(body, builder); @@ -435,10 +435,10 @@ public HTTPRestResponse createTaskPushNotificationConfiguration(ServerCallContex */ public HTTPRestResponse subscribeToTask(ServerCallContext context, String tenant, String taskId) { try { - if (!resolveAgentCard().capabilities().streaming()) { + validateVersionAndExtensions(tenant, context); + if (!resolveAgentCard(tenant).capabilities().streaming()) { return createErrorResponse(new UnsupportedOperationError(null, "Streaming is not supported by the agent", null)); } - validateVersionAndExtensions(context); TaskIdParams params = TaskIdParams.builder().id(taskId).tenant(tenant).build(); try { requestHandler.authorizeTaskAccess(params.id(), context, TaskOperation.SUBSCRIBE_TO_TASK); @@ -465,7 +465,7 @@ public HTTPRestResponse subscribeToTask(ServerCallContext context, String tenant */ public HTTPRestResponse getTask(ServerCallContext context, String tenant, String taskId, @Nullable Integer historyLength) { try { - validateVersionAndExtensions(context); + validateVersionAndExtensions(tenant, context); TaskQueryParams params = new TaskQueryParams(taskId, historyLength, tenant); Task task = requestHandler.onGetTask(params, context); if (task != null) { @@ -527,7 +527,7 @@ public HTTPRestResponse listTasks(ServerCallContext context, String tenant, @Nullable Integer historyLength, @Nullable String statusTimestampAfter, @Nullable Boolean includeArtifacts) { try { - validateVersionAndExtensions(context); + validateVersionAndExtensions(tenant, context); // Build params ListTasksParams.Builder paramsBuilder = ListTasksParams.builder(); if (contextId != null) { @@ -595,10 +595,10 @@ public HTTPRestResponse listTasks(ServerCallContext context, String tenant, */ public HTTPRestResponse getTaskPushNotificationConfiguration(ServerCallContext context, String tenant, String taskId, String configId) { try { - if (!resolveAgentCard().capabilities().pushNotifications()) { + validateVersionAndExtensions(tenant, context); + if (!resolveAgentCard(tenant).capabilities().pushNotifications()) { throw new PushNotificationNotSupportedError(); } - validateVersionAndExtensions(context); GetTaskPushNotificationConfigParams params = new GetTaskPushNotificationConfigParams(taskId, configId, tenant); TaskPushNotificationConfig config = requestHandler.onGetTaskPushNotificationConfig(params, context); return createSuccessResponse(200, org.a2aproject.sdk.grpc.TaskPushNotificationConfig.newBuilder(ProtoUtils.ToProto.taskPushNotificationConfig(config))); @@ -621,10 +621,10 @@ public HTTPRestResponse getTaskPushNotificationConfiguration(ServerCallContext c */ public HTTPRestResponse listTaskPushNotificationConfigurations(ServerCallContext context, String tenant, String taskId, int pageSize, String pageToken) { try { - if (!resolveAgentCard().capabilities().pushNotifications()) { + validateVersionAndExtensions(tenant, context); + if (!resolveAgentCard(tenant).capabilities().pushNotifications()) { throw new PushNotificationNotSupportedError(); } - validateVersionAndExtensions(context); ListTaskPushNotificationConfigsParams params = new ListTaskPushNotificationConfigsParams(taskId, pageSize, pageToken, tenant); ListTaskPushNotificationConfigsResult result = requestHandler.onListTaskPushNotificationConfigs(params, context); return createSuccessResponse(200, org.a2aproject.sdk.grpc.ListTaskPushNotificationConfigsResponse.newBuilder(ProtoUtils.ToProto.listTaskPushNotificationConfigsResponse(result))); @@ -646,10 +646,10 @@ public HTTPRestResponse listTaskPushNotificationConfigurations(ServerCallContext */ public HTTPRestResponse deleteTaskPushNotificationConfiguration(ServerCallContext context, String tenant, String taskId, String configId) { try { - if (!resolveAgentCard().capabilities().pushNotifications()) { + validateVersionAndExtensions(tenant, context); + if (!resolveAgentCard(tenant).capabilities().pushNotifications()) { throw new PushNotificationNotSupportedError(); } - validateVersionAndExtensions(context); DeleteTaskPushNotificationConfigParams params = new DeleteTaskPushNotificationConfigParams(taskId, configId, tenant); requestHandler.onDeleteTaskPushNotificationConfig(params, context); return new HTTPRestResponse(204, APPLICATION_JSON, ""); @@ -669,10 +669,14 @@ public HTTPRestResponse deleteTaskPushNotificationConfiguration(ServerCallContex * @param context the server call context carrying the requested version and extensions * @throws A2AError if the requested version or a required extension is not supported */ - private void validateVersionAndExtensions(ServerCallContext context) throws A2AError { - AgentCard agentCard = resolveAgentCard(); - A2AVersionValidator.validateProtocolVersion(agentCard, context); - A2AExtensions.validateRequiredExtensions(agentCard, context); + private void validateVersionAndExtensions(@Nullable String tenant, ServerCallContext context) throws A2AError { + try { + AgentCard agentCard = resolveAgentCard(tenant); + A2AVersionValidator.validateProtocolVersion(agentCard, context); + A2AExtensions.validateRequiredExtensions(agentCard, context); + } catch (TenantNotFoundException ex) { + throw new InvalidParamsError(ex.getResponseMessage()); + } } private void parseRequestBody(String body, com.google.protobuf.Message.Builder builder) throws A2AError { @@ -848,12 +852,11 @@ private static int mapErrorToHttpStatus(A2AError error) { public HTTPRestResponse getExtendedAgentCard(ServerCallContext context, @Nullable String tenant) { try { Utils.validateTenant(tenant); - if (!resolveAgentCard().capabilities().extendedAgentCard()) { + validateVersionAndExtensions(tenant, context); + if (!resolveAgentCard(tenant).capabilities().extendedAgentCard()) { throw new UnsupportedOperationError(); } - // Validate version before card lookup so version errors take precedence - validateVersionAndExtensions(context); - if (agentCardRouter != null) { + if (tenant != null && !tenant.isBlank() && agentCardRouter != null) { AgentCard card = agentCardRouter.resolveExtendedCard(tenant); if (card == null) { throw new ExtendedAgentCardNotConfiguredError(null, "Extended Card not configured", null); @@ -907,8 +910,13 @@ public HTTPRestResponse getExtendedAgentCard(ServerCallContext context, @Nullabl * @see AgentCard * @see #getExtendedAgentCard(ServerCallContext, String) */ - private AgentCard resolveAgentCard() { - return AgentCardValidator.resolveWithFallback(agentCardInstance, extendedAgentCard, transportValidated); + private AgentCard resolveAgentCard(@Nullable String tenant) { + try { + return AgentCardValidator.resolveWithFallback(agentCardInstance, extendedAgentCard, + agentCardRouter, tenant, validatedCards); + } catch (TenantNotFoundException ex) { + throw new InvalidParamsError(ex.getResponseMessage()); + } } public HTTPRestResponse getAgentCard() { @@ -942,11 +950,15 @@ public HTTPRestResponse getAgentCard(@Nullable String tenant) { LOGGER.fine(() -> "No AgentCardRouter configured; serving default public card for tenant '" + tenant + "'"); } AgentCard card = CdiUtils.resolveDefault(agentCardInstance); + if (card == null && agentCardRouter != null) { + card = agentCardRouter.resolvePublicCard(tenant); + } if (card == null) { return new HTTPRestResponse(404, "text/plain", "Public agent card not configured"); } + AgentCard validatedCard = card; return new HTTPRestResponse(200, APPLICATION_JSON, - JsonUtil.toJson(AgentCardValidator.resolveAndValidateOnce(() -> card, transportValidated, + JsonUtil.toJson(AgentCardValidator.resolveAndValidateOnce(() -> validatedCard, validatedCards, AgentCardValidator::validateTransportConfiguration)), cacheMetadata.getHttpHeadersMap()); } catch (TenantNotFoundException e) { diff --git a/transport/rest/src/test/java/org/a2aproject/sdk/transport/rest/handler/RestHandlerTest.java b/transport/rest/src/test/java/org/a2aproject/sdk/transport/rest/handler/RestHandlerTest.java index b170e28ee..c33273ee9 100644 --- a/transport/rest/src/test/java/org/a2aproject/sdk/transport/rest/handler/RestHandlerTest.java +++ b/transport/rest/src/test/java/org/a2aproject/sdk/transport/rest/handler/RestHandlerTest.java @@ -1272,7 +1272,17 @@ public void testExtendedAgentCardWithRouterKnownTenant() { public void testExtendedAgentCardWithRouterReturnsNull() { AgentCard cardWithExtCapability = AgentCard.builder(CARD) .capabilities(AgentCapabilities.builder().extendedAgentCard(true).build()).build(); - AgentCardRouter router = tenant -> null; + AgentCardRouter router = new AgentCardRouter() { + @Override + public AgentCard resolveExtendedCard(String tenant) { + return null; + } + + @Override + public AgentCard resolvePublicCard(String tenant) { + return cardWithExtCapability; + } + }; RestHandler handler = new RestHandler(new FixedInstance<>(cardWithExtCapability), null, createCacheMetadata(cardWithExtCapability), requestHandler, internalExecutor,