diff --git a/sdk/core/azure-core/CHANGELOG.md b/sdk/core/azure-core/CHANGELOG.md index b6533eee3f70..a3faaedbe612 100644 --- a/sdk/core/azure-core/CHANGELOG.md +++ b/sdk/core/azure-core/CHANGELOG.md @@ -8,6 +8,8 @@ ### Bugs Fixed +- Fixed synchronous streaming of non-replayable `BinaryData` response bodies. + ### Other Changes ## 1.59.0 (2026-08-12) diff --git a/sdk/core/azure-core/src/main/java/com/azure/core/http/policy/HttpLoggingPolicy.java b/sdk/core/azure-core/src/main/java/com/azure/core/http/policy/HttpLoggingPolicy.java index 20e2c2f7c9d9..6a4e3f172c51 100644 --- a/sdk/core/azure-core/src/main/java/com/azure/core/http/policy/HttpLoggingPolicy.java +++ b/sdk/core/azure-core/src/main/java/com/azure/core/http/policy/HttpLoggingPolicy.java @@ -21,6 +21,7 @@ import com.azure.core.implementation.util.ByteArrayContent; import com.azure.core.implementation.util.ByteBufferContent; import com.azure.core.implementation.util.HttpHeadersAccessHelper; +import com.azure.core.implementation.util.HttpUtils; import com.azure.core.implementation.util.InputStreamContent; import com.azure.core.implementation.util.SerializableContent; import com.azure.core.implementation.util.StringContent; @@ -506,8 +507,8 @@ private Long getAndLogContentLength(HttpHeaders headers, LoggingEventBuilder log /* * Determines if the request or response body should be logged. * - *

The request or response body is logged if the Content-Type is not "application/octet-stream" and the body - * isn't empty and is less than 16KB in size.

+ *

The request or response body is logged if the Content-Type is neither "application/octet-stream" nor + * "text/event-stream" and the body isn't empty and is less than 16KB in size.

* * @param contentTypeHeader Content-Type header value. * @@ -518,6 +519,7 @@ private Long getAndLogContentLength(HttpHeaders headers, LoggingEventBuilder log private static boolean shouldBodyBeLogged(String contentTypeHeader, Long contentLength) { return contentLength != null && !ContentType.APPLICATION_OCTET_STREAM.equalsIgnoreCase(contentTypeHeader) + && !HttpUtils.isTextEventStreamContentType(contentTypeHeader) && contentLength != 0 && contentLength < MAX_BODY_LOG_SIZE; } diff --git a/sdk/core/azure-core/src/main/java/com/azure/core/implementation/FluxInputStream.java b/sdk/core/azure-core/src/main/java/com/azure/core/implementation/FluxInputStream.java index aa962955a1b2..a5b06be1f9f6 100644 --- a/sdk/core/azure-core/src/main/java/com/azure/core/implementation/FluxInputStream.java +++ b/sdk/core/azure-core/src/main/java/com/azure/core/implementation/FluxInputStream.java @@ -29,7 +29,7 @@ public class FluxInputStream extends InputStream { private final Flux data; // Subscription to request more data from as needed - private Subscription subscription; + private volatile Subscription subscription; private ByteArrayInputStream buffer; @@ -141,6 +141,12 @@ public int read(byte[] b, int off, int len) throws IOException { } } + /** + * Closes the stream and cancels its Flux subscription. If the stream has not been read, closing subscribes and + * immediately cancels without requesting data so publisher cleanup can run. + * + * @throws IOException if the stream cannot be closed. + */ @Override public void close() throws IOException { closed = true; @@ -151,6 +157,9 @@ public void close() throws IOException { // Unblock any thread waiting in blockForData(). lock.lock(); try { + if (!subscribed) { + subscribeToData(); + } waitingForData = false; dataAvailable.signal(); } finally { @@ -227,13 +236,13 @@ private void subscribeToData() { this::signalOnCompleteOrError, // Subscription consumer subscription -> { + this.subscription = subscription; + this.subscribed = true; if (this.closed) { subscription.cancel(); return; } - this.subscription = subscription; - this.subscribed = true; - this.subscription.request(1); + subscription.request(1); }); } diff --git a/sdk/core/azure-core/src/main/java/com/azure/core/implementation/http/rest/AsyncRestProxy.java b/sdk/core/azure-core/src/main/java/com/azure/core/implementation/http/rest/AsyncRestProxy.java index 2548bec2a64f..fd89f5d98fc3 100644 --- a/sdk/core/azure-core/src/main/java/com/azure/core/implementation/http/rest/AsyncRestProxy.java +++ b/sdk/core/azure-core/src/main/java/com/azure/core/implementation/http/rest/AsyncRestProxy.java @@ -13,6 +13,7 @@ import com.azure.core.http.rest.StreamResponse; import com.azure.core.implementation.TypeUtil; import com.azure.core.implementation.serializer.HttpResponseDecoder; +import com.azure.core.implementation.util.HttpUtils; import com.azure.core.util.Base64Url; import com.azure.core.util.BinaryData; import com.azure.core.util.Context; @@ -46,8 +47,6 @@ */ public class AsyncRestProxy extends RestProxyBase { - private static final String TEXT_EVENT_STREAM = "text/event-stream"; - /** * Create a RestProxy. * @@ -143,23 +142,28 @@ private Mono ensureExpectedStatus( private Mono handleRestResponseReturnType(final HttpResponseDecoder.HttpDecodedResponse response, final SwaggerMethodParser methodParser, final Type entityType) { + final boolean isTextEventStream = HttpUtils.isTextEventStreamContentType( + response.getSourceResponse().getHeaders().getValue(HttpHeaderName.CONTENT_TYPE)); + final ResponseBodyOwner responseBodyOwner + = isTextEventStream ? new ResponseBodyOwner(response.getSourceResponse()) : null; if (methodParser.isStreamResponse()) { return Mono.fromSupplier(() -> new StreamResponse(response.getSourceResponse())); } else if (TypeUtil.isTypeOrSubTypeOf(entityType, Response.class)) { final Type bodyType = TypeUtil.getRestResponseBodyType(entityType); if (TypeUtil.isTypeOrSubTypeOf(bodyType, Void.class)) { - return response.getSourceResponse() - .getBody() - .ignoreElements() + Flux responseBody + = responseBodyOwner == null ? response.getSourceResponse().getBody() : responseBodyOwner.getBody(); + return responseBody.ignoreElements() .then(Mono.fromCallable(() -> createResponse(response, entityType, null))); } else { - return handleBodyReturnType(response.getSourceResponse(), decodeBytes(response), methodParser, bodyType) - .map(bodyAsObject -> createResponse(response, entityType, bodyAsObject)) - .switchIfEmpty(Mono.fromCallable(() -> createResponse(response, entityType, null))); + return handleBodyReturnType(response.getSourceResponse(), decodeBytes(response), methodParser, bodyType, + responseBodyOwner).map(bodyAsObject -> createResponse(response, entityType, bodyAsObject)) + .switchIfEmpty(Mono.fromCallable(() -> createResponse(response, entityType, null))); } } else { // For now, we're just throwing if the Maybe didn't emit a value. - return handleBodyReturnType(response.getSourceResponse(), decodeBytes(response), methodParser, entityType); + return handleBodyReturnType(response.getSourceResponse(), decodeBytes(response), methodParser, entityType, + responseBodyOwner); } } @@ -177,10 +181,17 @@ private static Function> decodeBytes(HttpResponseDecoder.Ht } static Mono handleBodyReturnType(HttpResponse sourceResponse, Function> getDecodedBody, - SwaggerMethodParser methodParser, Type entityType) { + SwaggerMethodParser methodParser, Type entityType, ResponseBodyOwner responseBodyOwner) { final int responseStatusCode = sourceResponse.getStatusCode(); final HttpMethod httpMethod = methodParser.getHttpMethod(); final Type returnValueWireType = methodParser.getReturnValueWireType(); + if (responseBodyOwner == null + && HttpUtils + .isTextEventStreamContentType(sourceResponse.getHeaders().getValue(HttpHeaderName.CONTENT_TYPE))) { + responseBodyOwner = new ResponseBodyOwner(sourceResponse); + } + final Flux responseBody + = responseBodyOwner == null ? sourceResponse.getBody() : responseBodyOwner.getBody(); final Mono asyncResult; if (httpMethod == HttpMethod.HEAD @@ -199,20 +210,19 @@ static Mono handleBodyReturnType(HttpResponse sourceResponse, Function> - asyncResult = Mono.just(sourceResponse.getBody()); + asyncResult = Mono.just(responseBody); } else if (TypeUtil.isTypeOrSubTypeOf(entityType, BinaryData.class)) { - String contentType = sourceResponse.getHeaders().getValue(HttpHeaderName.CONTENT_TYPE); // Mono // The raw response is directly used to create an instance of BinaryData which then provides // different methods to read the response. The reading of the response is delayed until BinaryData // is read and depending on which format the content is converted into, the response is not necessarily // fully copied into memory resulting in lesser overall memory usage. - if (contentType != null && contentType.startsWith(TEXT_EVENT_STREAM)) { - // if the response content type is a stream, create a BinaryData instance with bufferContent set to - // false. - asyncResult = BinaryData.fromFlux(sourceResponse.getBody(), null, false); + if (responseBodyOwner != null) { + // If the response content type identifies a stream, create a BinaryData instance with bufferContent + // set to false. + asyncResult = BinaryData.fromFlux(responseBody, null, false); } else { - asyncResult = BinaryData.fromFlux(sourceResponse.getBody()); + asyncResult = BinaryData.fromFlux(responseBody); } } else if (TypeUtil.isTypeOrSubTypeOf(entityType, InputStream.class)) { // Corresponds to the Open API 2.0 type "file" which is mapped to an InputStream. @@ -224,6 +234,11 @@ static Mono handleBodyReturnType(HttpResponse sourceResponse, Function handleBodyReturnType(HttpResponse sourceResponse, Function> getDecodedBody, + SwaggerMethodParser methodParser, Type entityType) { + return handleBodyReturnType(sourceResponse, getDecodedBody, methodParser, entityType, null); + } + /** * Handle the provided asynchronous HTTP response and return the deserialized value. * @@ -238,7 +253,6 @@ private Object handleRestReturnType(Mono errorOptionsSet) { final Mono asyncExpectedResponse = endSpanWhenDone( ensureExpectedStatus(asyncHttpDecodedResponse, methodParser, options, errorOptionsSet), context); - final Object result; if (TypeUtil.isTypeOrSubTypeOf(returnType, Mono.class)) { final Type monoTypeParam = TypeUtil.getTypeArgument(returnType); diff --git a/sdk/core/azure-core/src/main/java/com/azure/core/implementation/http/rest/RestProxyBase.java b/sdk/core/azure-core/src/main/java/com/azure/core/implementation/http/rest/RestProxyBase.java index a33abec128a8..b1b370f8757f 100644 --- a/sdk/core/azure-core/src/main/java/com/azure/core/implementation/http/rest/RestProxyBase.java +++ b/sdk/core/azure-core/src/main/java/com/azure/core/implementation/http/rest/RestProxyBase.java @@ -35,13 +35,17 @@ import com.azure.core.util.serializer.SerializerAdapter; import com.azure.core.util.tracing.Tracer; import reactor.core.Exceptions; +import reactor.core.publisher.Flux; +import java.io.Closeable; import java.io.IOException; import java.lang.reflect.Method; import java.lang.reflect.Type; import java.net.URL; +import java.nio.ByteBuffer; import java.nio.charset.StandardCharsets; import java.util.EnumSet; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.function.Consumer; import static com.azure.core.util.FluxUtil.monoError; @@ -213,6 +217,30 @@ public Response createResponse(HttpResponseDecoder.HttpDecodedResponse response, return RESPONSE_CONSTRUCTORS_CACHE.invoke(constructorReflectiveInvoker, response, bodyAsObject); } + static final class ResponseBodyOwner implements Closeable { + private final AtomicBoolean closed = new AtomicBoolean(); + private final HttpResponse response; + + ResponseBodyOwner(HttpResponse response) { + this.response = response; + } + + Flux getBody() { + return getBody(response.getBody()); + } + + Flux getBody(Flux responseBody) { + return Flux.using(() -> this, ignored -> responseBody, ResponseBodyOwner::close); + } + + @Override + public void close() { + if (closed.compareAndSet(false, true)) { + response.close(); + } + } + } + /** * Starts the tracing span for the current service call, additionally set metadata attributes on the span by passing * additional context information. diff --git a/sdk/core/azure-core/src/main/java/com/azure/core/implementation/http/rest/SyncRestProxy.java b/sdk/core/azure-core/src/main/java/com/azure/core/implementation/http/rest/SyncRestProxy.java index e2b3ca0c95c8..e1f8890317a1 100644 --- a/sdk/core/azure-core/src/main/java/com/azure/core/implementation/http/rest/SyncRestProxy.java +++ b/sdk/core/azure-core/src/main/java/com/azure/core/implementation/http/rest/SyncRestProxy.java @@ -13,6 +13,8 @@ import com.azure.core.implementation.ImplUtils; import com.azure.core.implementation.TypeUtil; import com.azure.core.implementation.serializer.HttpResponseDecoder; +import com.azure.core.implementation.util.BinaryDataHelper; +import com.azure.core.implementation.util.FluxByteBufferContent; import com.azure.core.util.Base64Url; import com.azure.core.util.BinaryData; import com.azure.core.util.Context; @@ -145,7 +147,9 @@ private Object handleRestResponseReturnType(HttpResponseDecoder.HttpDecodedRespo response.getSourceResponse().close(); return createResponse(response, entityType, null); } else { - Object bodyAsObject = handleBodyReturnType(response, methodParser, bodyType); + Object bodyAsObject = TypeUtil.isTypeOrSubTypeOf(bodyType, BinaryData.class) + ? getOwnedResponseBody(response.getSourceResponse()) + : handleBodyReturnType(response, methodParser, bodyType); Response httpResponse = createResponse(response, entityType, bodyAsObject); if (httpResponse == null) { return createResponse(response, entityType, null); @@ -195,6 +199,17 @@ private Object handleBodyReturnType(HttpResponseDecoder.HttpDecodedResponse resp return result; } + private static BinaryData getOwnedResponseBody(HttpResponse response) { + BinaryData responseBody = response.getBodyAsBinaryData(); + if (responseBody == null || responseBody.isReplayable()) { + return responseBody; + } + + ResponseBodyOwner responseBodyOwner = new ResponseBodyOwner(response); + return BinaryDataHelper.createBinaryData(new FluxByteBufferContent( + responseBodyOwner.getBody(responseBody.toFluxByteBuffer()), responseBody.getLength(), false)); + } + /** * Handle the provided asynchronous HTTP response and return the deserialized value. * diff --git a/sdk/core/azure-core/src/main/java/com/azure/core/implementation/util/FluxByteBufferContent.java b/sdk/core/azure-core/src/main/java/com/azure/core/implementation/util/FluxByteBufferContent.java index 1ccc4414be8e..ca0ae3639a7e 100644 --- a/sdk/core/azure-core/src/main/java/com/azure/core/implementation/util/FluxByteBufferContent.java +++ b/sdk/core/azure-core/src/main/java/com/azure/core/implementation/util/FluxByteBufferContent.java @@ -3,6 +3,7 @@ package com.azure.core.implementation.util; +import com.azure.core.implementation.FluxInputStream; import com.azure.core.util.FluxUtil; import com.azure.core.util.logging.ClientLogger; import com.azure.core.util.serializer.ObjectSerializer; @@ -100,9 +101,24 @@ public T toObject(TypeReference typeReference, ObjectSerializer serialize return serializer.deserializeFromBytes(toBytes(), typeReference); } + /** + * Returns an in-memory stream for cached or replayable content. Uncached non-replayable content is streamed + * incrementally without buffering the full Flux. + * + * @return A stream over this content. + */ @Override public InputStream toStream() { - return new ByteArrayInputStream(toBytes()); + byte[] cachedBytes = BYTES_UPDATER.get(this); + if (cachedBytes != null) { + return new ByteArrayInputStream(cachedBytes); + } + + if (isReplayable) { + return new ByteArrayInputStream(toBytes()); + } + + return new FluxInputStream(content); } @Override diff --git a/sdk/core/azure-core/src/main/java/com/azure/core/implementation/util/HttpUtils.java b/sdk/core/azure-core/src/main/java/com/azure/core/implementation/util/HttpUtils.java index 0fd02b8f9f3a..4c14fc3e57c1 100644 --- a/sdk/core/azure-core/src/main/java/com/azure/core/implementation/util/HttpUtils.java +++ b/sdk/core/azure-core/src/main/java/com/azure/core/implementation/util/HttpUtils.java @@ -6,6 +6,9 @@ import com.azure.core.util.logging.ClientLogger; import java.time.Duration; +import java.util.ArrayList; +import java.util.Collections; +import java.util.List; import static com.azure.core.util.Configuration.PROPERTY_AZURE_REQUEST_CONNECT_TIMEOUT; import static com.azure.core.util.Configuration.PROPERTY_AZURE_REQUEST_READ_TIMEOUT; @@ -17,6 +20,7 @@ * Utilities shared with HttpClient implementations. */ public final class HttpUtils { + private static final String TEXT_EVENT_STREAM = "text/event-stream"; private static final ClientLogger LOGGER = new ClientLogger(HttpUtils.class); private static final Duration MINIMUM_TIMEOUT = Duration.ofMillis(1); @@ -60,6 +64,56 @@ public final class HttpUtils { */ public static final String AZURE_EAGERLY_CONVERT_HEADERS = "azure-eagerly-convert-headers"; + /** + * Determines whether a Content-Type header identifies exactly one {@code text/event-stream} representation. + * Charset parameters don't affect this determination as event streams are always decoded as UTF-8. + * + * @param headerValue The Content-Type header value. + * @return Whether the header identifies a {@code text/event-stream} representation. + */ + public static boolean isTextEventStreamContentType(String headerValue) { + if (headerValue == null) { + return false; + } + + List mediaTypeAndParameters = splitHeaderValue(headerValue, ';'); + if (mediaTypeAndParameters.size() == 0 + || !TEXT_EVENT_STREAM.equalsIgnoreCase(mediaTypeAndParameters.get(0).trim()) + || splitHeaderValue(headerValue, ',').size() != 1) { + return false; + } + + return true; + } + + private static List splitHeaderValue(String value, char delimiter) { + List segments = new ArrayList<>(); + int start = 0; + boolean quoted = false; + boolean escaped = false; + + for (int i = 0; i < value.length(); i++) { + char character = value.charAt(i); + if (escaped) { + escaped = false; + } else if (quoted && character == '\\') { + escaped = true; + } else if (character == '"') { + quoted = !quoted; + } else if (!quoted && character == delimiter) { + segments.add(value.substring(start, i)); + start = i + 1; + } + } + + if (quoted) { + return Collections.emptyList(); + } + + segments.add(value.substring(start)); + return segments; + } + /** * Gets the default connect timeout. * diff --git a/sdk/core/azure-core/src/test/java/com/azure/core/http/policy/HttpLoggingPolicyTests.java b/sdk/core/azure-core/src/test/java/com/azure/core/http/policy/HttpLoggingPolicyTests.java index e21f13423dcb..cdad461025c5 100644 --- a/sdk/core/azure-core/src/test/java/com/azure/core/http/policy/HttpLoggingPolicyTests.java +++ b/sdk/core/azure-core/src/test/java/com/azure/core/http/policy/HttpLoggingPolicyTests.java @@ -28,6 +28,7 @@ import com.fasterxml.jackson.databind.ObjectMapper; import org.junit.jupiter.api.parallel.Execution; import org.junit.jupiter.api.parallel.ExecutionMode; +import org.junit.jupiter.api.Test; import org.junit.jupiter.params.ParameterizedTest; import org.junit.jupiter.params.provider.Arguments; import org.junit.jupiter.params.provider.MethodSource; @@ -58,6 +59,7 @@ import static com.azure.core.CoreTestUtils.createUrl; import static com.azure.core.http.HttpHeaderName.X_MS_REQUEST_ID; import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; import static org.junit.jupiter.api.Assertions.assertNotNull; import static org.junit.jupiter.api.Assertions.assertTrue; import static org.junit.jupiter.api.Assertions.fail; @@ -192,6 +194,32 @@ public void validateLoggingDoesNotConsumeRequestSync(BinaryData requestBody, byt expectedRequest.assertEqual(messages.get(0), logOptions, LogLevel.INFORMATIONAL); } + @Test + public void textEventStreamRequestBodiesAreNotLogged() { + byte[] data = "request event data".getBytes(StandardCharsets.UTF_8); + AtomicInteger requestCount = new AtomicInteger(); + HttpPipeline pipeline = new HttpPipelineBuilder() + .policies(new HttpLoggingPolicy(new HttpLogOptions().setLogLevel(HttpLogDetailLevel.BODY))) + .httpClient(request -> FluxUtil.collectBytesInByteBufferStream(request.getBody()).doOnSuccess(bytes -> { + assertArraysEqual(data, bytes); + requestCount.incrementAndGet(); + }).then(Mono.empty())) + .build(); + + HttpRequest asyncRequest = new HttpRequest(HttpMethod.POST, "https://test.com/async") + .setHeader(HttpHeaderName.CONTENT_TYPE, "Text/Event-Stream; charset=utf-8") + .setBody(BinaryData.fromBytes(data)); + StepVerifier.create(pipeline.send(asyncRequest)).verifyComplete(); + + HttpRequest syncRequest = new HttpRequest(HttpMethod.POST, "https://test.com/sync") + .setHeader(HttpHeaderName.CONTENT_TYPE, "Text/Event-Stream; charset=utf-8") + .setBody(BinaryData.fromBytes(data)); + pipeline.sendSync(syncRequest, Context.NONE); + + assertEquals(2, requestCount.get()); + assertFalse(convertOutputStreamToString(logCaptureStream).contains(new String(data, StandardCharsets.UTF_8))); + } + /** * Tests that logging the response body doesn't consume the stream before it is returned from the service call. */ @@ -244,6 +272,43 @@ public void validateLoggingDoesNotConsumeResponseSync(BinaryData responseBody, b assertTrue(logString.contains(new String(data, StandardCharsets.UTF_8))); } + @ParameterizedTest(name = "[{index}] {displayName}") + @MethodSource("responseLoggingSupplier") + public void responseLoggingUsesActualContentType(String contentType, int expectedBufferCount, + boolean expectBodyLogged) { + byte[] data = "streaming response".getBytes(StandardCharsets.UTF_8); + AtomicInteger bufferCount = new AtomicInteger(); + HttpRequest request = new HttpRequest(HttpMethod.GET, "https://test.com/responseLoggingUsesActualContentType"); + HttpHeaders responseHeaders = new HttpHeaders().set(HttpHeaderName.CONTENT_TYPE, contentType) + .set(HttpHeaderName.CONTENT_LENGTH, Integer.toString(data.length)); + + HttpPipeline pipeline = new HttpPipelineBuilder() + .policies(new HttpLoggingPolicy(new HttpLogOptions().setLogLevel(HttpLogDetailLevel.BODY))) + .httpClient(ignored -> Mono.just(new BufferTrackingHttpResponse(ignored, responseHeaders, + Flux.just(ByteBuffer.wrap(data)), bufferCount))) + .build(); + + Context context = getCallerMethodContext("streamingResponsesAreNotBuffered", LogLevel.INFORMATIONAL); + + try (HttpResponse response = pipeline.send(request, context).block()) { + assertNotNull(response); + assertArraysEqual(data, response.getBodyAsBinaryData().toBytes()); + } + + try (HttpResponse response = pipeline.sendSync(request, context)) { + assertArraysEqual(data, response.getBodyAsBinaryData().toBytes()); + } + + assertEquals(expectedBufferCount, bufferCount.get()); + assertEquals(expectBodyLogged, + convertOutputStreamToString(logCaptureStream).contains(new String(data, StandardCharsets.UTF_8))); + } + + private static Stream responseLoggingSupplier() { + return Stream.of(Arguments.of("Text/Event-Stream; charset=utf-8", 0, false), + Arguments.of(ContentType.APPLICATION_JSON, 2, true)); + } + private static Stream validateLoggingDoesNotConsumeSupplierSync() { byte[] data = "this is a test".getBytes(StandardCharsets.UTF_8); @@ -351,6 +416,22 @@ public Mono getBodyAsString(Charset charset) { } } + private static final class BufferTrackingHttpResponse extends MockHttpResponse { + private final AtomicInteger bufferCount; + + private BufferTrackingHttpResponse(HttpRequest request, HttpHeaders headers, Flux body, + AtomicInteger bufferCount) { + super(request, headers, body); + this.bufferCount = bufferCount; + } + + @Override + public HttpResponse buffer() { + bufferCount.incrementAndGet(); + return super.buffer(); + } + } + @ParameterizedTest(name = "[{index}] {displayName}") @MethodSource("logOptionsSupplier") public void loggingIncludesRetryCount(HttpLogOptions logOptions) { diff --git a/sdk/core/azure-core/src/test/java/com/azure/core/implementation/FluxInputStreamTests.java b/sdk/core/azure-core/src/test/java/com/azure/core/implementation/FluxInputStreamTests.java index 9408b4b83217..7ef7330252bd 100644 --- a/sdk/core/azure-core/src/test/java/com/azure/core/implementation/FluxInputStreamTests.java +++ b/sdk/core/azure-core/src/test/java/com/azure/core/implementation/FluxInputStreamTests.java @@ -9,6 +9,7 @@ import org.junit.jupiter.params.ParameterizedTest; import org.junit.jupiter.params.provider.MethodSource; import org.junit.jupiter.params.provider.ValueSource; +import org.reactivestreams.Subscription; import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; @@ -19,10 +20,18 @@ import java.nio.charset.Charset; import java.util.ArrayList; import java.util.List; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicLong; +import java.util.concurrent.atomic.AtomicReference; import java.util.stream.Stream; import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNull; import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; public class FluxInputStreamTests { private static final int KB = 1024; @@ -93,6 +102,85 @@ public void fluxInputStreamWithEmptyByteBuffers() throws IOException { } } + @Test + public void closeBeforeFirstReadSubscribesAndCancelsWithoutDemand() throws IOException { + AtomicInteger subscribeCalls = new AtomicInteger(); + AtomicInteger cancelCalls = new AtomicInteger(); + AtomicInteger disposeCalls = new AtomicInteger(); + AtomicLong requested = new AtomicLong(); + Flux data = Flux.using(Object::new, + ignored -> Flux.never() + .doOnSubscribe(subscription -> subscribeCalls.incrementAndGet()) + .doOnRequest(requested::addAndGet) + .doOnCancel(cancelCalls::incrementAndGet), + ignored -> disposeCalls.incrementAndGet()); + + FluxInputStream stream = new FluxInputStream(data); + stream.close(); + stream.close(); + + assertEquals(1, subscribeCalls.get()); + assertEquals(0, requested.get()); + assertEquals(1, cancelCalls.get()); + assertEquals(1, disposeCalls.get()); + } + + @Test + public void closeDuringFirstSubscriptionCancelsPublishedSubscription() throws Exception { + CountDownLatch subscribeEntered = new CountDownLatch(1); + AtomicInteger cancelCalls = new AtomicInteger(); + AtomicInteger disposeCalls = new AtomicInteger(); + AtomicLong requested = new AtomicLong(); + AtomicReference closeThreadReference = new AtomicReference<>(); + Flux data = Flux.using(Object::new, ignored -> Flux.from(subscriber -> { + subscribeEntered.countDown(); + awaitThreadWaiting(closeThreadReference); + subscriber.onSubscribe(new Subscription() { + @Override + public void request(long count) { + requested.addAndGet(count); + } + + @Override + public void cancel() { + cancelCalls.incrementAndGet(); + } + }); + }), ignored -> disposeCalls.incrementAndGet()); + FluxInputStream stream = new FluxInputStream(data); + AtomicReference readError = new AtomicReference<>(); + AtomicReference closeError = new AtomicReference<>(); + Thread readThread = new Thread(() -> { + try { + stream.read(); + } catch (Throwable throwable) { + readError.set(throwable); + } + }); + Thread closeThread = new Thread(() -> { + try { + stream.close(); + } catch (Throwable throwable) { + closeError.set(throwable); + } + }); + closeThreadReference.set(closeThread); + + readThread.start(); + assertTrue(subscribeEntered.await(5, TimeUnit.SECONDS)); + closeThread.start(); + readThread.join(TimeUnit.SECONDS.toMillis(5)); + closeThread.join(TimeUnit.SECONDS.toMillis(5)); + + assertFalse(readThread.isAlive()); + assertFalse(closeThread.isAlive()); + assertEquals(1, cancelCalls.get()); + assertEquals(0, requested.get()); + assertEquals(1, disposeCalls.get()); + assertNull(closeError.get()); + assertTrue(readError.get() instanceof IllegalStateException); + } + @ParameterizedTest @MethodSource("fluxInputStreamErrorSupplier") public void fluxInputStreamError(RuntimeException exception) { @@ -145,4 +233,16 @@ public Mono getBodyAsString(Charset charset) { new HttpResponseException("Mock exception", httpResponse, null), new UncheckedIOException(new IOException("Mock IO Exception."))); } + + private static void awaitThreadWaiting(AtomicReference threadReference) { + long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(5); + while (System.nanoTime() < deadline) { + Thread thread = threadReference.get(); + if (thread != null && thread.getState() == Thread.State.WAITING) { + return; + } + Thread.yield(); + } + throw new AssertionError("Close thread didn't wait for the stream lock."); + } } diff --git a/sdk/core/azure-core/src/test/java/com/azure/core/implementation/http/rest/AsyncRestProxyTests.java b/sdk/core/azure-core/src/test/java/com/azure/core/implementation/http/rest/AsyncRestProxyTests.java index af0ddd076d72..ed366c1be3c2 100644 --- a/sdk/core/azure-core/src/test/java/com/azure/core/implementation/http/rest/AsyncRestProxyTests.java +++ b/sdk/core/azure-core/src/test/java/com/azure/core/implementation/http/rest/AsyncRestProxyTests.java @@ -3,22 +3,30 @@ package com.azure.core.implementation.http.rest; +import com.azure.core.annotation.ExpectedResponses; import com.azure.core.annotation.Get; import com.azure.core.annotation.Head; import com.azure.core.annotation.Host; import com.azure.core.annotation.ServiceInterface; import com.azure.core.http.ContentType; +import com.azure.core.http.HttpClient; import com.azure.core.http.HttpHeaderName; import com.azure.core.http.HttpHeaders; +import com.azure.core.http.HttpPipelineBuilder; import com.azure.core.http.HttpResponse; import com.azure.core.http.MockHttpResponse; +import com.azure.core.http.rest.RequestOptions; +import com.azure.core.http.rest.Response; +import com.azure.core.http.rest.RestProxy; import com.azure.core.util.BinaryData; +import com.azure.core.util.Context; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.junit.jupiter.params.ParameterizedTest; import org.junit.jupiter.params.provider.Arguments; import org.junit.jupiter.params.provider.MethodSource; import reactor.core.publisher.Flux; +import reactor.core.publisher.Mono; import reactor.test.StepVerifier; import java.io.IOException; @@ -26,6 +34,7 @@ import java.lang.reflect.Type; import java.nio.ByteBuffer; import java.nio.charset.StandardCharsets; +import java.util.concurrent.atomic.AtomicInteger; import java.util.stream.Stream; import static com.azure.core.CoreTestUtils.assertArraysEqual; @@ -54,6 +63,10 @@ private interface MockService { @Get("getStreamResponse") Flux getStreamResponse(); + + @Get("getStreamingResponse") + @ExpectedResponses({ 200 }) + Mono> getStreamingResponse(RequestOptions options, Context context); } @BeforeEach @@ -193,4 +206,65 @@ public static Stream getResponseHeaderAndReplayability() { Arguments.of(new HttpHeaders().set(HttpHeaderName.CONTENT_TYPE, ContentType.APPLICATION_JSON), true), Arguments.of(new HttpHeaders().set(HttpHeaderName.CONTENT_TYPE, "application/xml"), true)); } + + @ParameterizedTest + @MethodSource("streamingResponseOwnershipSupplier") + public void streamingResponseUsesResponseContentType(String accept, String contentType, boolean replayable, + int expectedCloseCount) { + byte[] expectedBytes = "hello".getBytes(StandardCharsets.UTF_8); + AtomicInteger responseCloseCount = new AtomicInteger(); + HttpClient client = request -> Mono.just(new MockHttpResponse(request, 200, + new HttpHeaders().set(HttpHeaderName.CONTENT_TYPE, contentType), expectedBytes) { + @Override + public void close() { + responseCloseCount.incrementAndGet(); + super.close(); + } + }); + MockService service = RestProxy.create(MockService.class, new HttpPipelineBuilder().httpClient(client).build()); + RequestOptions options = new RequestOptions(); + if (accept != null) { + options.setHeader(HttpHeaderName.ACCEPT, accept); + } + + Response response = service.getStreamingResponse(options, Context.NONE).block(); + + assertEquals(replayable, response.getValue().isReplayable()); + assertArraysEqual(expectedBytes, response.getValue().toBytes()); + assertEquals(expectedCloseCount, responseCloseCount.get()); + } + + @Test + public void cancellingStreamingBinaryDataClosesResponse() { + byte[] expectedBytes = "hello".getBytes(StandardCharsets.UTF_8); + AtomicInteger responseCloseCount = new AtomicInteger(); + HttpClient client = request -> Mono.just(new MockHttpResponse(request, 200, + new HttpHeaders().set(HttpHeaderName.CONTENT_TYPE, "text/event-stream"), expectedBytes) { + @Override + public Flux getBody() { + return Flux.concat(Flux.just(ByteBuffer.wrap(expectedBytes)), Flux.never()); + } + + @Override + public void close() { + responseCloseCount.incrementAndGet(); + super.close(); + } + }); + MockService service = RestProxy.create(MockService.class, new HttpPipelineBuilder().httpClient(client).build()); + + Flux body = service.getStreamingResponse(new RequestOptions(), Context.NONE) + .flatMapMany(response -> response.getValue().toFluxByteBuffer()); + + StepVerifier.create(body) + .assertNext(buffer -> assertArraysEqual(expectedBytes, buffer.array())) + .thenCancel() + .verify(); + assertEquals(1, responseCloseCount.get()); + } + + private static Stream streamingResponseOwnershipSupplier() { + return Stream.of(Arguments.of("text/event-stream", ContentType.APPLICATION_JSON, true, 0), + Arguments.of(null, "text/event-stream; charset=utf-8", false, 1)); + } } diff --git a/sdk/core/azure-core/src/test/java/com/azure/core/implementation/http/rest/SyncRestProxyTests.java b/sdk/core/azure-core/src/test/java/com/azure/core/implementation/http/rest/SyncRestProxyTests.java index 682ec7dc5f7f..98099311266d 100644 --- a/sdk/core/azure-core/src/test/java/com/azure/core/implementation/http/rest/SyncRestProxyTests.java +++ b/sdk/core/azure-core/src/test/java/com/azure/core/implementation/http/rest/SyncRestProxyTests.java @@ -23,24 +23,37 @@ import com.azure.core.http.rest.Response; import com.azure.core.http.rest.RestProxy; import com.azure.core.http.rest.StreamResponse; +import com.azure.core.implementation.util.BinaryDataHelper; +import com.azure.core.implementation.util.FluxByteBufferContent; import com.azure.core.util.BinaryData; import com.azure.core.util.Context; import org.junit.jupiter.api.Test; import org.junit.jupiter.params.ParameterizedTest; import org.junit.jupiter.params.provider.Arguments; import org.junit.jupiter.params.provider.MethodSource; +import reactor.core.Disposable; +import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; +import reactor.test.StepVerifier; import java.io.ByteArrayInputStream; import java.io.IOException; import java.io.InputStream; +import java.nio.ByteBuffer; +import java.time.Duration; import java.util.Collections; import java.util.HashMap; import java.util.Map; import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicInteger; import java.util.stream.Stream; import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertArrayEquals; +import static org.junit.jupiter.api.Assertions.assertInstanceOf; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTimeoutPreemptively; import static org.junit.jupiter.api.Assertions.assertTrue; /** @@ -67,6 +80,10 @@ Response testMethod(@BodyParam("application/octet-stream") BinaryData data @Put("my/url/path") @ExpectedResponses({ 200 }) Response testInputStreamResponse(Context context); + + @Get("my/url/path") + @ExpectedResponses({ 200 }) + Response testBinaryDataResponse(Context context); } @Test @@ -193,6 +210,129 @@ public void testInputStream() throws IOException { assertEquals("hello", new String(bytes)); } + @Test + public void binaryDataResponseClosesOnCompletion() { + AtomicInteger responseCloseCount = new AtomicInteger(); + TestInterface testInterface + = createBinaryDataService(Flux.just(ByteBuffer.wrap("hello".getBytes())), responseCloseCount); + + BinaryData responseBody = testInterface.testBinaryDataResponse(Context.NONE).getValue(); + + assertFalse(responseBody.isReplayable()); + assertEquals("hello", responseBody.toString()); + assertEquals(1, responseCloseCount.get()); + assertEquals("hello", responseBody.toString()); + assertEquals(1, responseCloseCount.get()); + } + + @Test + public void binaryDataResponseClosesOnCancellation() { + AtomicInteger responseCloseCount = new AtomicInteger(); + Flux responseBody = Flux.concat(Flux.just(ByteBuffer.wrap("hello".getBytes())), Flux.never()); + TestInterface testInterface = createBinaryDataService(responseBody, responseCloseCount); + + Disposable subscription + = testInterface.testBinaryDataResponse(Context.NONE).getValue().toFluxByteBuffer().subscribe(); + subscription.dispose(); + + assertEquals(1, responseCloseCount.get()); + } + + @Test + public void binaryDataResponseToStreamReadsWithoutWaitingForCompletion() throws IOException { + byte[] firstChunk = new byte[8192]; + byte[] secondChunk = new byte[8192]; + firstChunk[0] = 1; + secondChunk[0] = 2; + AtomicInteger responseCloseCount = new AtomicInteger(); + Flux responseBody + = Flux.concat(Flux.just(ByteBuffer.wrap(firstChunk), ByteBuffer.wrap(secondChunk)), Flux.never()); + TestInterface testInterface = createBinaryDataService(responseBody, responseCloseCount); + + BinaryData binaryData = testInterface.testBinaryDataResponse(Context.NONE).getValue(); + assertFalse(binaryData.isReplayable()); + + assertTimeoutPreemptively(Duration.ofSeconds(5), () -> { + try (InputStream stream = binaryData.toStream()) { + byte[] actualFirstChunk = new byte[firstChunk.length]; + byte[] actualSecondChunk = new byte[secondChunk.length]; + assertEquals(firstChunk.length, stream.read(actualFirstChunk)); + assertEquals(secondChunk.length, stream.read(actualSecondChunk)); + assertArrayEquals(firstChunk, actualFirstChunk); + assertArrayEquals(secondChunk, actualSecondChunk); + assertEquals(0, responseCloseCount.get()); + } + }); + + assertEquals(1, responseCloseCount.get()); + } + + @Test + public void binaryDataResponseToStreamClosesOnError() throws IOException { + AtomicInteger responseCloseCount = new AtomicInteger(); + TestInterface testInterface + = createBinaryDataService(Flux.error(new IllegalStateException("Body read failed.")), responseCloseCount); + + try (InputStream stream = testInterface.testBinaryDataResponse(Context.NONE).getValue().toStream()) { + assertThrows(IOException.class, stream::read); + assertEquals(1, responseCloseCount.get()); + } + + assertEquals(1, responseCloseCount.get()); + } + + @Test + public void binaryDataResponseToStreamClosesBeforeFirstRead() throws IOException { + AtomicInteger responseCloseCount = new AtomicInteger(); + TestInterface testInterface = createBinaryDataService(Flux.never(), responseCloseCount); + + BinaryData responseBody = testInterface.testBinaryDataResponse(Context.NONE).getValue(); + assertInstanceOf(FluxByteBufferContent.class, BinaryDataHelper.getContent(responseBody)); + responseBody.toStream().close(); + + assertEquals(1, responseCloseCount.get()); + } + + @Test + public void binaryDataResponseClosesOnError() { + AtomicInteger responseCloseCount = new AtomicInteger(); + TestInterface testInterface + = createBinaryDataService(Flux.error(new IllegalStateException("Body read failed.")), responseCloseCount); + + StepVerifier.create(testInterface.testBinaryDataResponse(Context.NONE).getValue().toFluxByteBuffer()) + .expectErrorMessage("Body read failed.") + .verify(); + + assertEquals(1, responseCloseCount.get()); + } + + private static TestInterface createBinaryDataService(Flux responseBody, + AtomicInteger responseCloseCount) { + HttpClient client = new HttpClient() { + @Override + public Mono send(HttpRequest request) { + return Mono.error(new IllegalStateException("Async Send API was Invoked.")); + } + + @Override + public HttpResponse sendSync(HttpRequest request, Context context) { + return new MockHttpResponse(request, 200) { + @Override + public BinaryData getBodyAsBinaryData() { + return BinaryData.fromFlux(responseBody, null, false).block(); + } + + @Override + public void close() { + responseCloseCount.incrementAndGet(); + super.close(); + } + }; + } + }; + return RestProxy.create(TestInterface.class, new HttpPipelineBuilder().httpClient(client).build()); + } + private static Stream mergeRequestOptionsContextSupplier() { Map twoValuesMap = new HashMap<>(); twoValuesMap.put("key", "value"); diff --git a/sdk/core/azure-core/src/test/java/com/azure/core/implementation/util/HttpUtilsTests.java b/sdk/core/azure-core/src/test/java/com/azure/core/implementation/util/HttpUtilsTests.java new file mode 100644 index 000000000000..430a3772e8fe --- /dev/null +++ b/sdk/core/azure-core/src/test/java/com/azure/core/implementation/util/HttpUtilsTests.java @@ -0,0 +1,23 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +package com.azure.core.implementation.util; + +import org.junit.jupiter.api.Test; + +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertTrue; + +public class HttpUtilsTests { + @Test + public void textEventStreamContentTypeRequiresSingleMediaType() { + assertTrue(HttpUtils.isTextEventStreamContentType("Text/Event-Stream; charset=utf-8")); + assertTrue(HttpUtils.isTextEventStreamContentType("text/event-stream; charset=\"UTF-8\"")); + assertTrue(HttpUtils.isTextEventStreamContentType("text/event-stream; charset=iso-8859-1")); + assertTrue(HttpUtils.isTextEventStreamContentType("text/event-stream; charset=utf-16")); + assertTrue(HttpUtils.isTextEventStreamContentType("text/event-stream; charset=not-a-charset; charset=utf-16")); + assertFalse(HttpUtils.isTextEventStreamContentType("application/json, text/event-stream")); + assertTrue(HttpUtils.isTextEventStreamContentType("text/event-stream; note=\"x,y;z\"")); + assertFalse(HttpUtils.isTextEventStreamContentType("text/event-stream; note=\"unterminated")); + } +} diff --git a/sdk/core/azure-core/src/test/java/com/azure/core/util/BinaryDataTest.java b/sdk/core/azure-core/src/test/java/com/azure/core/util/BinaryDataTest.java index 187a09f9cd2d..3b03704e0aaf 100644 --- a/sdk/core/azure-core/src/test/java/com/azure/core/util/BinaryDataTest.java +++ b/sdk/core/azure-core/src/test/java/com/azure/core/util/BinaryDataTest.java @@ -82,6 +82,7 @@ import static org.junit.jupiter.api.Assertions.assertNotSame; import static org.junit.jupiter.api.Assertions.assertSame; import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTimeoutPreemptively; import static org.junit.jupiter.api.Assertions.assertTrue; /** @@ -567,6 +568,24 @@ public void fluxContent() { .verifyComplete(); } + @Test + public void nonReplayableFluxToStreamReadsWithoutWaitingForCompletion() { + byte[] expected = "hello".getBytes(StandardCharsets.UTF_8); + AtomicBoolean cancelled = new AtomicBoolean(); + Flux content + = Flux.concat(Flux.just(ByteBuffer.wrap(expected)), Flux.never()).doOnCancel(() -> cancelled.set(true)); + BinaryData data = BinaryData.fromFlux(content, null, false).block(); + + assertTimeoutPreemptively(Duration.ofSeconds(5), () -> { + try (InputStream stream = data.toStream()) { + byte[] actual = new byte[expected.length]; + assertEquals(expected.length, stream.read(actual)); + assertTrue(Arrays.equals(expected, actual)); + } + }); + assertTrue(cancelled.get()); + } + @Test public void testFromFile() throws Exception { Path file = Files.createTempFile("binaryDataFromFile" + UUID.randomUUID(), ".txt");