diff --git a/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/ResponseSubscribers.java b/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/ResponseSubscribers.java index b19904de6..95e6e4868 100644 --- a/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/ResponseSubscribers.java +++ b/mcp-core/src/main/java/io/modelcontextprotocol/client/transport/ResponseSubscribers.java @@ -186,7 +186,7 @@ static class SseLineSubscriber extends BaseSubscriber { * The response information from the HTTP response. Send with each event to * provide context. */ - private ResponseInfo responseInfo; + private final ResponseInfo responseInfo; /** * The maximum number of bytes that may accumulate for a single SSE event. A peer @@ -348,16 +348,15 @@ public AggregateSubscriber(ResponseInfo responseInfo, FluxSink si @Override protected void hookOnSubscribe(Subscription subscription) { + // Register disposal callback to cancel subscription when Flux is disposed + sink.onDispose(subscription::cancel); sink.onRequest(n -> { if (!hasRequestedDemand) { + hasRequestedDemand = true; subscription.request(Long.MAX_VALUE); } - hasRequestedDemand = true; }); - - // Register disposal callback to cancel subscription when Flux is disposed - sink.onDispose(subscription::cancel); } @Override @@ -410,17 +409,14 @@ public BodilessResponseLineSubscriber(ResponseInfo responseInfo, FluxSink { if (!hasRequestedDemand) { + hasRequestedDemand = true; subscription.request(Long.MAX_VALUE); } - hasRequestedDemand = true; - }); - - // Register disposal callback to cancel subscription when Flux is disposed - sink.onDispose(() -> { - subscription.cancel(); }); } diff --git a/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/ResponseSubscribersTests.java b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/ResponseSubscribersTests.java new file mode 100644 index 000000000..8649267c6 --- /dev/null +++ b/mcp-core/src/test/java/io/modelcontextprotocol/client/transport/ResponseSubscribersTests.java @@ -0,0 +1,50 @@ +/* + * Copyright 2024 - 2024 the original author or authors. + */ + +package io.modelcontextprotocol.client.transport; + +import java.net.http.HttpResponse.ResponseInfo; + +import org.junit.jupiter.api.Test; + +import reactor.core.publisher.Flux; +import reactor.test.StepVerifier; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.Mockito.mock; + +class ResponseSubscribersTests { + + @Test + void aggregateSubscriberEmitsResponseWhenRequestCompletesSynchronously() { + ResponseInfo responseInfo = mock(ResponseInfo.class); + + Flux response = Flux.create(sink -> { + var subscriber = new ResponseSubscribers.AggregateSubscriber(responseInfo, sink, Integer.MAX_VALUE); + Flux.just("payload").subscribe(subscriber); + }); + + StepVerifier.create(response).assertNext(event -> { + var aggregate = (ResponseSubscribers.AggregateResponseEvent) event; + assertThat(aggregate.responseInfo()).isSameAs(responseInfo); + assertThat(aggregate.data()).isEqualTo("payload\n"); + }).verifyComplete(); + } + + @Test + void bodilessSubscriberEmitsResponseWhenRequestCompletesSynchronously() { + ResponseInfo responseInfo = mock(ResponseInfo.class); + + Flux response = Flux.create(sink -> { + var subscriber = new ResponseSubscribers.BodilessResponseLineSubscriber(responseInfo, sink); + Flux.empty().subscribe(subscriber); + }); + + StepVerifier.create(response).assertNext(event -> { + var dummy = (ResponseSubscribers.DummyEvent) event; + assertThat(dummy.responseInfo()).isSameAs(responseInfo); + }).verifyComplete(); + } + +}