From b711e0f58f633420e4e9a56e681126fd31c28776 Mon Sep 17 00:00:00 2001 From: Naman Gururani Date: Wed, 23 Sep 2026 23:08:12 +0530 Subject: [PATCH] feat(client): support custom Executor for async operations Add DoclingApiBuilder.asyncExecutor(Executor) so that async operations (task submission, status polling including the delayed re-polls, and result retrieval) can run on a caller-provided executor. - The builder method is a default method throwing UnsupportedOperationException, so existing builder implementations keep compiling. - The executor stays nullable down to AsyncOperations. When it is not set, the 1-arg supplyAsync and 2-arg delayedExecutor are used, so the behaviour is unchanged. - toBuilder() carries the executor over, and the client never shuts it down. - ConvertOperations, ChunkOperations and AsyncOperations gain constructor overloads taking the executor. Closes #664 Signed-off-by: Naman Gururani --- .../ai/docling/serve/api/DoclingServeApi.java | 35 ++- .../serve/api/DoclingServeApiTests.java | 67 ++++- .../serve/client/DoclingServeClient.java | 68 +++-- .../client/operations/AsyncOperations.java | 48 +++- .../client/operations/ChunkOperations.java | 31 ++- .../client/operations/ConvertOperations.java | 41 ++- ...tDoclingServeClientAsyncExecutorTests.java | 247 ++++++++++++++++++ ...ServeJackson2ClientAsyncExecutorTests.java | 27 ++ ...ServeJackson3ClientAsyncExecutorTests.java | 27 ++ docs/src/doc/docs/whats-new.md | 1 + 10 files changed, 535 insertions(+), 57 deletions(-) create mode 100644 docling-serve/docling-serve-client/src/test/java/ai/docling/serve/client/AbstractDoclingServeClientAsyncExecutorTests.java create mode 100644 docling-serve/docling-serve-client/src/test/java/ai/docling/serve/client/DoclingServeJackson2ClientAsyncExecutorTests.java create mode 100644 docling-serve/docling-serve-client/src/test/java/ai/docling/serve/client/DoclingServeJackson3ClientAsyncExecutorTests.java diff --git a/docling-serve/docling-serve-api/src/main/java/ai/docling/serve/api/DoclingServeApi.java b/docling-serve/docling-serve-api/src/main/java/ai/docling/serve/api/DoclingServeApi.java index 24e31d16..491f6fac 100644 --- a/docling-serve/docling-serve-api/src/main/java/ai/docling/serve/api/DoclingServeApi.java +++ b/docling-serve/docling-serve-api/src/main/java/ai/docling/serve/api/DoclingServeApi.java @@ -4,6 +4,7 @@ import java.net.URI; import java.time.Duration; +import java.util.concurrent.Executor; import java.util.stream.Collectors; import org.jspecify.annotations.Nullable; @@ -15,8 +16,7 @@ /** * Docling Serve API interface. */ -public interface DoclingServeApi - extends DoclingServeHealthApi, DoclingServeConvertApi, DoclingServeChunkApi, DoclingServeClearApi, DoclingServeTaskApi { +public interface DoclingServeApi extends DoclingServeHealthApi, DoclingServeConvertApi, DoclingServeChunkApi, DoclingServeClearApi, DoclingServeTaskApi { /** * Creates and returns a builder instance capable of constructing implementations of {@link DoclingServeApi}. @@ -35,12 +35,14 @@ static > B builder( if (factories.isEmpty()) { // No factory found - throw new IllegalStateException("No instance of %s found to build a %s instance. You are probably missing a library on your classpath.".formatted(DoclingServeApiBuilderFactory.class.getName(), DoclingApiBuilder.class.getName())); + throw new IllegalStateException("No instance of %s found to build a %s instance. You are probably missing a library on your classpath." + .formatted(DoclingServeApiBuilderFactory.class.getName(), DoclingApiBuilder.class.getName())); } if (factories.size() > 1) { // Multiple factories found - throw new IllegalStateException("Multiple instances of %s found to build a %s instance: [%s]".formatted(DoclingServeApiBuilderFactory.class.getName(), DoclingApiBuilder.class.getName(), factories.stream().map(f -> f.getClass().getName()).collect(Collectors.joining(", ")))); + throw new IllegalStateException("Multiple instances of %s found to build a %s instance: [%s]".formatted(DoclingServeApiBuilderFactory.class.getName(), DoclingApiBuilder.class + .getName(), factories.stream().map(f -> f.getClass().getName()).collect(Collectors.joining(", ")))); } // Only 1 factory (what we want) @@ -202,6 +204,31 @@ default B prettyPrint() { */ B asyncTimeout(Duration asyncTimeout); + /** + * Sets the {@link Executor} used to run async operations. + * + *

This configures where the work of the async methods (such as + * {@link DoclingServeApi#convertSourceAsync(ConvertDocumentRequest)}) is executed: submitting + * the task, polling for its status and retrieving its result. If not set, async operations run + * on the default async executor of {@link java.util.concurrent.CompletableFuture}. + * + *

The lifecycle of the executor is owned by the caller: the client never shuts it down. + * Avoid direct executors such as {@code Runnable::run}: the blocking HTTP requests would then run + * on the calling thread, making the async methods partially blocking, and on the shared scheduler + * thread of {@link java.util.concurrent.CompletableFuture#delayedExecutor(long, java.util.concurrent.TimeUnit, Executor)}. + * + *

The default implementation throws {@link UnsupportedOperationException}, so that existing + * builder implementations keep compiling; builders supporting a custom executor override it. + * + * @param asyncExecutor the executor to use for async operations (must not be null) + * @return this builder instance for method chaining + * @throws IllegalArgumentException if asyncExecutor is null + * @throws UnsupportedOperationException if this builder does not support a custom executor + */ + default B asyncExecutor(Executor asyncExecutor) { + throw new UnsupportedOperationException("A custom async executor is not supported by " + getClass().getName()); + } + /** * Builds and returns an instance of the specified type, representing the completed configuration * of the builder. The returned instance is typically an implementation of the Docling API. diff --git a/docling-serve/docling-serve-api/src/test/java/ai/docling/serve/api/DoclingServeApiTests.java b/docling-serve/docling-serve-api/src/test/java/ai/docling/serve/api/DoclingServeApiTests.java index e8f66af2..5cbee939 100644 --- a/docling-serve/docling-serve-api/src/test/java/ai/docling/serve/api/DoclingServeApiTests.java +++ b/docling-serve/docling-serve-api/src/test/java/ai/docling/serve/api/DoclingServeApiTests.java @@ -2,6 +2,10 @@ import static org.assertj.core.api.Assertions.assertThatExceptionOfType; +import java.net.URI; +import java.time.Duration; + +import org.jspecify.annotations.Nullable; import org.junit.jupiter.api.Test; import ai.docling.serve.api.DoclingServeApi.DoclingApiBuilder; @@ -12,6 +16,67 @@ class DoclingServeApiTests { void noFactoryFound() { assertThatExceptionOfType(IllegalStateException.class) .isThrownBy(() -> DoclingServeApi.builder()) - .withMessage("No instance of %s found to build a %s instance. You are probably missing a library on your classpath.", DoclingServeApiBuilderFactory.class.getName(), DoclingApiBuilder.class.getName()); + .withMessage("No instance of %s found to build a %s instance. You are probably missing a library on your classpath.", DoclingServeApiBuilderFactory.class + .getName(), DoclingApiBuilder.class.getName()); + } + + @Test + void asyncExecutorIsUnsupportedByDefault() { + assertThatExceptionOfType(UnsupportedOperationException.class) + .isThrownBy(() -> new MinimalBuilder().asyncExecutor(Runnable::run)); + } + + // A builder implementing only the abstract methods of the interface, e.g. one provided through the SPI. + // This class compiling is what guarantees that new DoclingApiBuilder methods don't break such builders. + private static final class MinimalBuilder implements DoclingApiBuilder { + @Override + public MinimalBuilder baseUrl(URI baseUrl) { + return this; + } + + @Override + public MinimalBuilder apiKey(@Nullable String apiKey) { + return this; + } + + @Override + public MinimalBuilder logRequests(boolean logRequests) { + return this; + } + + @Override + public MinimalBuilder logResponses(boolean logResponses) { + return this; + } + + @Override + public MinimalBuilder prettyPrint(boolean prettyPrint) { + return this; + } + + @Override + public MinimalBuilder connectTimeout(Duration connectTimeout) { + return this; + } + + @Override + public MinimalBuilder readTimeout(Duration readTimeout) { + return this; + } + + @Override + public MinimalBuilder asyncPollInterval(Duration asyncPollInterval) { + return this; + } + + @Override + public MinimalBuilder asyncTimeout(Duration asyncTimeout) { + return this; + } + + @Override + public DoclingServeApi build() { + throw new UnsupportedOperationException(); + } } } diff --git a/docling-serve/docling-serve-client/src/main/java/ai/docling/serve/client/DoclingServeClient.java b/docling-serve/docling-serve-client/src/main/java/ai/docling/serve/client/DoclingServeClient.java index acbde249..c8097709 100644 --- a/docling-serve/docling-serve-client/src/main/java/ai/docling/serve/client/DoclingServeClient.java +++ b/docling-serve/docling-serve-client/src/main/java/ai/docling/serve/client/DoclingServeClient.java @@ -20,6 +20,7 @@ import java.util.Objects; import java.util.Optional; import java.util.concurrent.CompletionStage; +import java.util.concurrent.Executor; import java.util.concurrent.Flow.Subscriber; import java.util.stream.Collectors; @@ -90,6 +91,7 @@ public abstract class DoclingServeClient extends HttpOperations implements Docli private final Duration readTimeout; private final Duration asyncPollInterval; private final Duration asyncTimeout; + private final @Nullable Executor asyncExecutor; private final HealthOperations healthOps; private final ConvertOperations convertOps; @@ -129,12 +131,13 @@ protected DoclingServeClient(DoclingServeClientBuilder builder) { this.apiKey = builder.apiKey; this.asyncPollInterval = builder.asyncPollInterval; this.asyncTimeout = builder.asyncTimeout; + this.asyncExecutor = builder.asyncExecutor; // Initialize operations handlers this.healthOps = new HealthOperations(this); this.taskOps = new TaskOperations(this); - this.convertOps = new ConvertOperations(this, this.taskOps, this.asyncPollInterval, this.asyncTimeout); - this.chunkOps = new ChunkOperations(this, this.taskOps, this.asyncPollInterval, this.asyncTimeout); + this.convertOps = new ConvertOperations(this, this.taskOps, this.asyncPollInterval, this.asyncTimeout, this.asyncExecutor); + this.chunkOps = new ChunkOperations(this, this.taskOps, this.asyncPollInterval, this.asyncTimeout, this.asyncExecutor); this.clearOps = new ClearOperations(this); } @@ -176,7 +179,7 @@ protected void logRequest(HttpRequest request) { .stream() .map(this::maskSensitiveHeaderValues) .forEach(entry -> stringBuilder.append(" %s: %s\n".formatted(entry.getKey(), String.join(", ", entry.getValue()))) - ); + ); LOG.info(stringBuilder.toString()); } @@ -188,8 +191,7 @@ private boolean isSensitiveHeader(String headerName) { private Map.Entry> maskSensitiveHeaderValues(Map.Entry> entry) { return Map.entry( - entry.getKey(), - entry.getValue().stream() + entry.getKey(), entry.getValue().stream() .map(value -> isSensitiveHeader(entry.getKey()) ? "*".repeat(value.length()) : value) .toList() ); @@ -201,8 +203,7 @@ protected void logResponse(HttpResponse response, Optional respo stringBuilder.append("\n← RESPONSE: %s\n".formatted(response.statusCode())); stringBuilder.append(" HEADERS:\n"); - response.headers().map().forEach((key, values) -> - stringBuilder.append(" %s: %s\n".formatted(key, String.join(", ", values))) + response.headers().map().forEach((key, values) -> stringBuilder.append(" %s: %s\n".formatted(key, String.join(", ", values))) ); responseBody @@ -221,9 +222,10 @@ protected T execute(HttpRequest request, Class expectedValueType) { try { HttpResponse response = null; - if(StreamResponse.class.equals(expectedValueType)) { + if (StreamResponse.class.equals(expectedValueType)) { response = this.httpClient.send(request, BodyHandlers.ofInputStream()); - } else { + } + else { response = this.httpClient.send(request, BodyHandlers.ofString()); } return getResponse(request, response, expectedValueType); @@ -281,9 +283,9 @@ protected HttpRequest.Builder createRequestBuilder(RequestContext r .header("Accept", "application/json") .timeout(this.readTimeout); - if (Utils.isNotNullOrBlank(this.apiKey)) { - requestBuilder.header(API_KEY_HEADER_NAME, this.apiKey); - } + if (Utils.isNotNullOrBlank(this.apiKey)) { + requestBuilder.header(API_KEY_HEADER_NAME, this.apiKey); + } return requestBuilder; } @@ -309,9 +311,10 @@ protected T getResponse(HttpRequest request, HttpResponse response, Class if (StreamResponse.class.equals(expectedReturnType)) { // typical 4XX & 5XX responses are usually accompanied by JSON response bodies // hence, reading the stream here. - try (InputStream is = (InputStream) body){ + try (InputStream is = (InputStream) body) { body = new String(is.readAllBytes(), StandardCharsets.UTF_8); - } catch (IOException e) { + } + catch (IOException e) { throw new DoclingServeClientException(e); } } @@ -319,16 +322,16 @@ protected T getResponse(HttpRequest request, HttpResponse response, Class if (statusCode == 422) { var validationError = readValue(body.toString(), ValidationError.class); var errorText = validationError.getErrorDetails() - .stream() - .map(ValidationErrorDetail::getMessage) - .filter(Objects::nonNull) - .collect(Collectors.joining("\n")); + .stream() + .map(ValidationErrorDetail::getMessage) + .filter(Objects::nonNull) + .collect(Collectors.joining("\n")); throw new ValidationException( - validationError, - "An error occurred while making %s request to %s:\n%s".formatted(request.method(), request.uri(), errorText) + validationError, "An error occurred while making %s request to %s:\n%s".formatted(request.method(), request.uri(), errorText) ); - } else { + } + else { throw new DoclingServeClientException("An error occurred: %s".formatted(body.toString()), statusCode, body.toString()); } } @@ -337,9 +340,10 @@ protected T getResponse(HttpRequest request, HttpResponse response, Class return (T) StreamResponse .builder() .headers(headerName -> response.headers().firstValue(headerName)) - .body((InputStream)body) + .body((InputStream) body) .build(); - } else { + } + else { return readValue(body.toString(), expectedReturnType); } } @@ -466,6 +470,7 @@ public abstract static class DoclingServeClientBuilderIf not set, async operations run on the default async executor of + * {@link java.util.concurrent.CompletableFuture}. The executor is never shut down by the client. + * + * @param asyncExecutor the executor to use for async operations (must not be null) + * @return this builder instance for method chaining + * @throws IllegalArgumentException if asyncExecutor is null + */ + @Override + public B asyncExecutor(Executor asyncExecutor) { + this.asyncExecutor = ensureNotNull(asyncExecutor, "asyncExecutor"); + return (B) this; + } } } diff --git a/docling-serve/docling-serve-client/src/main/java/ai/docling/serve/client/operations/AsyncOperations.java b/docling-serve/docling-serve-client/src/main/java/ai/docling/serve/client/operations/AsyncOperations.java index df07f39d..c7ca0f97 100644 --- a/docling-serve/docling-serve-client/src/main/java/ai/docling/serve/client/operations/AsyncOperations.java +++ b/docling-serve/docling-serve-client/src/main/java/ai/docling/serve/client/operations/AsyncOperations.java @@ -4,8 +4,11 @@ import java.util.Optional; import java.util.concurrent.CompletableFuture; import java.util.concurrent.CompletionStage; +import java.util.concurrent.Executor; import java.util.concurrent.TimeUnit; +import java.util.function.Supplier; +import org.jspecify.annotations.Nullable; import org.slf4j.Logger; import org.slf4j.LoggerFactory; @@ -27,12 +30,18 @@ public abstract class AsyncOperations { private final DoclingServeTaskApi taskApi; private final Duration asyncPollInterval; private final Duration asyncTimeout; + private final @Nullable Executor asyncExecutor; protected AsyncOperations(HttpOperations httpOperations, DoclingServeTaskApi taskApi, Duration asyncPollInterval, Duration asyncTimeout) { + this(httpOperations, taskApi, asyncPollInterval, asyncTimeout, null); + } + + protected AsyncOperations(HttpOperations httpOperations, DoclingServeTaskApi taskApi, Duration asyncPollInterval, Duration asyncTimeout, @Nullable Executor asyncExecutor) { this.httpOperations = httpOperations; this.taskApi = taskApi; this.asyncPollInterval = asyncPollInterval; this.asyncTimeout = asyncTimeout; + this.asyncExecutor = asyncExecutor; } /** @@ -42,7 +51,7 @@ protected AsyncOperations(HttpOperations httpOperations, DoclingServeTaskApi tas * It uses the information provided in the {@code TaskResultRequest} * to obtain the result of the task execution. * - * @param the type of the result object returned + * @param the type of the result object returned * @param taskResultRequest the request containing the details, including the task ID, * required to retrieve the task result * @return the result of the task execution, of the type {@code O} @@ -55,18 +64,17 @@ protected AsyncOperations(HttpOperations httpOperations, DoclingServeTaskApi tas * to start the task and then repeatedly polls the task status to determine * when the operation is complete. * - * @param the type of the request object being sent - * @param the type of the response object returned upon completion + * @param the type of the request object being sent + * @param the type of the response object returned upon completion * @param request the request object containing the data necessary to initialize the task - * @param uri the endpoint URI to which the request will be sent + * @param uri the endpoint URI to which the request will be sent * @return a {@link CompletionStage} that will be completed with the result of the asynchronous operation */ protected CompletionStage executeAsync(I request, String uri) { ValidationUtils.ensureNotNull(request, "request"); // Start the async conversion and chain the polling logic - return CompletableFuture.supplyAsync(() -> - this.httpOperations.executePost(createAsyncRequestContext(uri, request)) + return supply(() -> this.httpOperations.executePost(createAsyncRequestContext(uri, request)) ).thenCompose(taskResponse -> { LOG.info("Started async conversion with task ID: {}", taskResponse.getTaskId()); @@ -98,7 +106,7 @@ private CompletionStage pollTaskUntilComplete(TaskStatusPollResponse stat .taskId(taskId) .build(); - return CompletableFuture.supplyAsync(() -> this.taskApi.pollTaskStatus(pollRequest)) + return supply(() -> this.taskApi.pollTaskStatus(pollRequest)) .thenCompose(statusResponse -> pollTaskStatus(statusResponse, startTime)); } @@ -116,7 +124,7 @@ private CompletionStage pollTaskStatus(TaskStatusPollResponse statusRespo .taskId(statusResponse.getTaskId()) .build(); - yield CompletableFuture.supplyAsync(() -> getTaskResult(taskResult)); + yield supply(() -> getTaskResult(taskResult)); } case FAILURE -> { @@ -129,11 +137,25 @@ private CompletionStage pollTaskStatus(TaskStatusPollResponse statusRespo } default -> - // Still in progress (PENDING or STARTED), schedule next poll after delay - CompletableFuture.supplyAsync( - () -> null, - CompletableFuture.delayedExecutor(this.asyncPollInterval.toMillis(), TimeUnit.MILLISECONDS) - ).thenCompose(v -> pollTaskUntilComplete(statusResponse, startTime)); + // Still in progress (PENDING or STARTED), schedule next poll after delay + CompletableFuture.supplyAsync( + () -> null, delayed() + ).thenCompose(v -> pollTaskUntilComplete(statusResponse, startTime)); }; } + + // Without a configured executor, defer to CompletableFuture's default async executor, which is not + // always ForkJoinPool.commonPool() (e.g. when the common pool's parallelism is 1 or less). + private CompletableFuture supply(Supplier supplier) { + return (this.asyncExecutor != null) ? + CompletableFuture.supplyAsync(supplier, this.asyncExecutor) : + CompletableFuture.supplyAsync(supplier); + } + + private Executor delayed() { + var millis = this.asyncPollInterval.toMillis(); + return (this.asyncExecutor != null) ? + CompletableFuture.delayedExecutor(millis, TimeUnit.MILLISECONDS, this.asyncExecutor) : + CompletableFuture.delayedExecutor(millis, TimeUnit.MILLISECONDS); + } } diff --git a/docling-serve/docling-serve-client/src/main/java/ai/docling/serve/client/operations/ChunkOperations.java b/docling-serve/docling-serve-client/src/main/java/ai/docling/serve/client/operations/ChunkOperations.java index f2e32919..96375f5b 100644 --- a/docling-serve/docling-serve-client/src/main/java/ai/docling/serve/client/operations/ChunkOperations.java +++ b/docling-serve/docling-serve-client/src/main/java/ai/docling/serve/client/operations/ChunkOperations.java @@ -2,6 +2,9 @@ import java.time.Duration; import java.util.concurrent.CompletionStage; +import java.util.concurrent.Executor; + +import org.jspecify.annotations.Nullable; import ai.docling.serve.api.DoclingServeChunkApi; import ai.docling.serve.api.DoclingServeTaskApi; @@ -19,9 +22,31 @@ public final class ChunkOperations extends AsyncOperations implements DoclingSer private final HttpOperations httpOperations; private final DoclingServeTaskApi taskApi; - public ChunkOperations(HttpOperations httpOperations, DoclingServeTaskApi taskApi, - Duration asyncPollInterval, Duration asyncTimeout) { - super(httpOperations, taskApi, asyncPollInterval, asyncTimeout); + /** + * Creates a new ChunkOperations instance whose async operations run on the default async executor of + * {@link java.util.concurrent.CompletableFuture}. + * + * @param httpOperations the HTTP operations handler for executing requests + * @param taskApi the task operations handler for polling and retrieving results + * @param asyncPollInterval the interval between status polls for async operations + * @param asyncTimeout the maximum time to wait for async operations to complete + */ + public ChunkOperations(HttpOperations httpOperations, DoclingServeTaskApi taskApi, Duration asyncPollInterval, Duration asyncTimeout) { + this(httpOperations, taskApi, asyncPollInterval, asyncTimeout, null); + } + + /** + * Creates a new ChunkOperations instance whose async operations run on the given executor. + * + * @param httpOperations the HTTP operations handler for executing requests + * @param taskApi the task operations handler for polling and retrieving results + * @param asyncPollInterval the interval between status polls for async operations + * @param asyncTimeout the maximum time to wait for async operations to complete + * @param asyncExecutor the executor to run async operations on, or {@code null} to use the + * default async executor of {@link java.util.concurrent.CompletableFuture} + */ + public ChunkOperations(HttpOperations httpOperations, DoclingServeTaskApi taskApi, Duration asyncPollInterval, Duration asyncTimeout, @Nullable Executor asyncExecutor) { + super(httpOperations, taskApi, asyncPollInterval, asyncTimeout, asyncExecutor); this.httpOperations = httpOperations; this.taskApi = taskApi; } diff --git a/docling-serve/docling-serve-client/src/main/java/ai/docling/serve/client/operations/ConvertOperations.java b/docling-serve/docling-serve-client/src/main/java/ai/docling/serve/client/operations/ConvertOperations.java index 53022860..a44fac2c 100644 --- a/docling-serve/docling-serve-client/src/main/java/ai/docling/serve/client/operations/ConvertOperations.java +++ b/docling-serve/docling-serve-client/src/main/java/ai/docling/serve/client/operations/ConvertOperations.java @@ -2,6 +2,9 @@ import java.time.Duration; import java.util.concurrent.CompletionStage; +import java.util.concurrent.Executor; + +import org.jspecify.annotations.Nullable; import ai.docling.serve.api.DoclingServeConvertApi; import ai.docling.serve.api.DoclingServeTaskApi; @@ -27,16 +30,30 @@ public final class ConvertOperations extends AsyncOperations implements DoclingS private final DoclingServeTaskApi taskApi; /** - * Creates a new ConvertOperations instance. + * Creates a new ConvertOperations instance whose async operations run on the default async executor of + * {@link java.util.concurrent.CompletableFuture}. * * @param httpOperations the HTTP operations handler for executing requests * @param taskApi the task operations handler for polling and retrieving results * @param asyncPollInterval the interval between status polls for async operations * @param asyncTimeout the maximum time to wait for async operations to complete */ - public ConvertOperations(HttpOperations httpOperations, DoclingServeTaskApi taskApi, - Duration asyncPollInterval, Duration asyncTimeout) { - super(httpOperations, taskApi, asyncPollInterval, asyncTimeout); + public ConvertOperations(HttpOperations httpOperations, DoclingServeTaskApi taskApi, Duration asyncPollInterval, Duration asyncTimeout) { + this(httpOperations, taskApi, asyncPollInterval, asyncTimeout, null); + } + + /** + * Creates a new ConvertOperations instance whose async operations run on the given executor. + * + * @param httpOperations the HTTP operations handler for executing requests + * @param taskApi the task operations handler for polling and retrieving results + * @param asyncPollInterval the interval between status polls for async operations + * @param asyncTimeout the maximum time to wait for async operations to complete + * @param asyncExecutor the executor to run async operations on, or {@code null} to use the + * default async executor of {@link java.util.concurrent.CompletableFuture} + */ + public ConvertOperations(HttpOperations httpOperations, DoclingServeTaskApi taskApi, Duration asyncPollInterval, Duration asyncTimeout, @Nullable Executor asyncExecutor) { + super(httpOperations, taskApi, asyncPollInterval, asyncTimeout, asyncExecutor); this.httpOperations = httpOperations; this.taskApi = taskApi; } @@ -52,23 +69,21 @@ public ConvertDocumentResponse convertSource(ConvertDocumentRequest request) { final var uri = "/v1/convert/source"; boolean hasMultipleSources = !Utils.isNullOrEmpty(request.getSources()) ? - request.getSources().size() > 1: Boolean.FALSE; - boolean isRemoteTarget = request.getTarget() instanceof S3Target || request.getTarget() instanceof PutTarget - || request.getTarget() instanceof PresignedUrlTarget; + request.getSources().size() > 1 : Boolean.FALSE; + boolean isRemoteTarget = request.getTarget() instanceof S3Target || request.getTarget() instanceof PutTarget || request.getTarget() instanceof PresignedUrlTarget; boolean isZipTarget = request.getTarget() instanceof ZipTarget; - if((hasMultipleSources && !isRemoteTarget) || isZipTarget) { + if ((hasMultipleSources && !isRemoteTarget) || isZipTarget) { StreamResponse response = this.httpOperations - .executePostWithStreamResponse(createRequestContext(uri, request, - StreamResponse.class)); + .executePostWithStreamResponse(createRequestContext(uri, request, StreamResponse.class)); String fileName = response.getHeaders().getFileName().orElse("converted_docs.zip"); return ZipArchiveConvertDocumentResponse .builder().fileName(fileName) .inputStream(response.getBody()) .build(); - } else { - return this.httpOperations.executePost(createRequestContext(uri, request, - ConvertDocumentResponse.class)); + } + else { + return this.httpOperations.executePost(createRequestContext(uri, request, ConvertDocumentResponse.class)); } } diff --git a/docling-serve/docling-serve-client/src/test/java/ai/docling/serve/client/AbstractDoclingServeClientAsyncExecutorTests.java b/docling-serve/docling-serve-client/src/test/java/ai/docling/serve/client/AbstractDoclingServeClientAsyncExecutorTests.java new file mode 100644 index 00000000..826ffe84 --- /dev/null +++ b/docling-serve/docling-serve-client/src/test/java/ai/docling/serve/client/AbstractDoclingServeClientAsyncExecutorTests.java @@ -0,0 +1,247 @@ +package ai.docling.serve.client; + +import static com.github.tomakehurst.wiremock.client.WireMock.get; +import static com.github.tomakehurst.wiremock.client.WireMock.getRequestedFor; +import static com.github.tomakehurst.wiremock.client.WireMock.okJson; +import static com.github.tomakehurst.wiremock.client.WireMock.post; +import static com.github.tomakehurst.wiremock.client.WireMock.urlPathEqualTo; +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import java.net.URI; +import java.time.Duration; +import java.util.Set; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.Executor; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; + +import org.jspecify.annotations.Nullable; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; + +import com.github.tomakehurst.wiremock.junit5.WireMockExtension; +import com.github.tomakehurst.wiremock.stubbing.Scenario; + +import ai.docling.serve.api.DoclingServeApi; +import ai.docling.serve.api.chunk.request.HybridChunkDocumentRequest; +import ai.docling.serve.api.convert.request.ConvertDocumentRequest; +import ai.docling.serve.api.convert.request.source.HttpSource; +import ai.docling.serve.api.convert.response.InBodyConvertDocumentResponse; + +/** + * Tests for the {@code asyncExecutor} option of {@link DoclingServeClient.DoclingServeClientBuilder}. + */ +abstract class AbstractDoclingServeClientAsyncExecutorTests { + private static final String THREAD_PREFIX = "custom-async-"; + private static final String TASK_ID = "task-1"; + private static final String POLL_PATH = "/v1/status/poll/" + TASK_ID; + private static final String RESULT_PATH = "/v1/result/" + TASK_ID; + private static final String POLLED_ONCE = "polled-once"; + + // Tasks submitted to the executor for a task that is polled twice: the task submission, the first poll, the + // delay before the re-poll (CompletableFuture.delayedExecutor submits a task to its base executor once the + // delay has elapsed), the re-poll and, for a successful task, the result retrieval. + private static final int FAILED_TASK_SUBMISSIONS = 4; + private static final int SUCCESSFUL_TASK_SUBMISSIONS = 5; + + private static final String CONVERT_RESULT = """ + { + "document": { + "filename": "dev.html", + "md_content": "# Dev" + }, + "status": "success", + "errors": [], + "processing_time": 0.5, + "timings": {} + } + """; + + private static final String CHUNK_RESULT = """ + { + "chunks": [], + "documents": [], + "processing_time": 0.5 + } + """; + + private final AtomicInteger executions = new AtomicInteger(); + private final Set threadNames = ConcurrentHashMap.newKeySet(); + private ExecutorService customExecutor; + private Executor recordingExecutor; + + protected abstract WireMockExtension getWireMock(); + + protected abstract DoclingServeClient.DoclingServeClientBuilder newClientBuilder(); + + @BeforeEach + void setUp() { + var threadCounter = new AtomicInteger(); + this.customExecutor = Executors.newFixedThreadPool(2, runnable -> new Thread(runnable, THREAD_PREFIX + threadCounter.incrementAndGet())); + + // Counts every task submitted to the executor, so that a step bypassing it shows up as a lower count + this.recordingExecutor = command -> { + this.executions.incrementAndGet(); + this.customExecutor.execute(() -> { + this.threadNames.add(Thread.currentThread().getName()); + command.run(); + }); + }; + } + + @AfterEach + void tearDown() throws InterruptedException { + this.customExecutor.shutdownNow(); + assertThat(this.customExecutor.awaitTermination(5, TimeUnit.SECONDS)).isTrue(); + } + + @Test + void convertSourceAsyncRunsEveryStepOnCustomExecutor() throws Exception { + stubTask("/v1/convert/source/async", "success"); + getWireMock().stubFor(get(urlPathEqualTo(RESULT_PATH)).willReturn(okJson(CONVERT_RESULT))); + + var response = client(this.recordingExecutor) + .convertSourceAsync(convertRequest()) + .toCompletableFuture() + .get(10, TimeUnit.SECONDS); + + assertThat(response).isInstanceOf(InBodyConvertDocumentResponse.class); + assertThat(((InBodyConvertDocumentResponse) response).getDocument().getMarkdownContent()).isEqualTo("# Dev"); + + assertThat(this.executions).hasValue(SUCCESSFUL_TASK_SUBMISSIONS); + assertThat(this.threadNames) + .isNotEmpty() + .allSatisfy(threadName -> assertThat(threadName).startsWith(THREAD_PREFIX)); + } + + @Test + void chunkSourceAsyncRunsEveryStepOnCustomExecutor() throws Exception { + stubTask("/v1/chunk/hybrid/source/async", "success"); + getWireMock().stubFor(get(urlPathEqualTo(RESULT_PATH)).willReturn(okJson(CHUNK_RESULT))); + + var request = HybridChunkDocumentRequest.builder() + .source(HttpSource.builder().url(URI.create("https://docs.arconia.io/arconia-cli/latest/development/dev/")).build()) + .build(); + + var response = client(this.recordingExecutor) + .chunkSourceWithHybridChunkerAsync(request) + .toCompletableFuture() + .get(10, TimeUnit.SECONDS); + + assertThat(response.getProcessingTime()).isEqualTo(0.5); + + assertThat(this.executions).hasValue(SUCCESSFUL_TASK_SUBMISSIONS); + assertThat(this.threadNames) + .isNotEmpty() + .allSatisfy(threadName -> assertThat(threadName).startsWith(THREAD_PREFIX)); + } + + @Test + void failedTaskRunsEveryPollOnCustomExecutor() { + stubTask("/v1/convert/source/async", "failure"); + + var future = client(this.recordingExecutor) + .convertSourceAsync(convertRequest()) + .toCompletableFuture(); + + assertThatThrownBy(() -> future.get(10, TimeUnit.SECONDS)) + .hasRootCauseMessage("Async conversion failed for task %s: Task failed".formatted(TASK_ID)); + + assertThat(this.executions).hasValue(FAILED_TASK_SUBMISSIONS); + getWireMock().verify(0, getRequestedFor(urlPathEqualTo(RESULT_PATH))); + } + + @Test + void asyncOperationsWorkWithoutCustomExecutor() throws Exception { + stubTask("/v1/convert/source/async", "success"); + getWireMock().stubFor(get(urlPathEqualTo(RESULT_PATH)).willReturn(okJson(CONVERT_RESULT))); + + var response = client(null) + .convertSourceAsync(convertRequest()) + .toCompletableFuture() + .get(10, TimeUnit.SECONDS); + + assertThat(((InBodyConvertDocumentResponse) response).getDocument().getMarkdownContent()).isEqualTo("# Dev"); + getWireMock().verify(2, getRequestedFor(urlPathEqualTo(POLL_PATH))); + getWireMock().verify(1, getRequestedFor(urlPathEqualTo(RESULT_PATH))); + assertThat(this.executions).hasValue(0); + } + + @Test + void toBuilderKeepsAsyncExecutor() throws Exception { + stubTask("/v1/convert/source/async", "success"); + getWireMock().stubFor(get(urlPathEqualTo(RESULT_PATH)).willReturn(okJson(CONVERT_RESULT))); + + DoclingServeApi client = client(this.recordingExecutor).toBuilder().build(); + + client.convertSourceAsync(convertRequest()) + .toCompletableFuture() + .get(10, TimeUnit.SECONDS); + + assertThat(this.executions).hasValue(SUCCESSFUL_TASK_SUBMISSIONS); + } + + @Test + void nullAsyncExecutorIsRejected() { + assertThatThrownBy(() -> newClientBuilder().asyncExecutor(null)) + .isInstanceOf(IllegalArgumentException.class) + .hasMessageContaining("asyncExecutor"); + } + + private DoclingServeApi client(@Nullable Executor asyncExecutor) { + var builder = newClientBuilder() + .baseUrl(getWireMock().baseUrl()) + .asyncPollInterval(Duration.ofMillis(50)) + .asyncTimeout(Duration.ofSeconds(10)); + + if (asyncExecutor != null) { + builder.asyncExecutor(asyncExecutor); + } + + return builder.build(); + } + + // The first poll reports the task as still running, which makes the client schedule a delayed re-poll + private void stubTask(String submitPath, String finalStatus) { + var wireMock = getWireMock(); + + wireMock.stubFor( + post(urlPathEqualTo(submitPath)) + .willReturn(okJson(taskStatus("pending"))) + ); + + wireMock.stubFor( + get(urlPathEqualTo(POLL_PATH)) + .inScenario("polling") + .whenScenarioStateIs(Scenario.STARTED) + .willReturn(okJson(taskStatus("started"))) + .willSetStateTo(POLLED_ONCE) + ); + + wireMock.stubFor( + get(urlPathEqualTo(POLL_PATH)) + .inScenario("polling") + .whenScenarioStateIs(POLLED_ONCE) + .willReturn(okJson(taskStatus(finalStatus))) + ); + } + + private static ConvertDocumentRequest convertRequest() { + return ConvertDocumentRequest.builder() + .source(HttpSource.builder().url(URI.create("https://docs.arconia.io/arconia-cli/latest/development/dev/")).build()) + .build(); + } + + private static String taskStatus(String status) { + return """ + { + "task_id": "%s", + "task_status": "%s" + } + """.formatted(TASK_ID, status); + } +} diff --git a/docling-serve/docling-serve-client/src/test/java/ai/docling/serve/client/DoclingServeJackson2ClientAsyncExecutorTests.java b/docling-serve/docling-serve-client/src/test/java/ai/docling/serve/client/DoclingServeJackson2ClientAsyncExecutorTests.java new file mode 100644 index 00000000..c517669f --- /dev/null +++ b/docling-serve/docling-serve-client/src/test/java/ai/docling/serve/client/DoclingServeJackson2ClientAsyncExecutorTests.java @@ -0,0 +1,27 @@ +package ai.docling.serve.client; + +import static com.github.tomakehurst.wiremock.core.WireMockConfiguration.wireMockConfig; + +import org.junit.jupiter.api.extension.RegisterExtension; + +import com.github.tomakehurst.wiremock.junit5.WireMockExtension; + +/** + * Async executor tests for {@link DoclingServeJackson2Client}. + */ +class DoclingServeJackson2ClientAsyncExecutorTests extends AbstractDoclingServeClientAsyncExecutorTests { + @RegisterExtension + static WireMockExtension wireMock = WireMockExtension.newInstance() + .options(wireMockConfig().dynamicPort()) + .build(); + + @Override + protected WireMockExtension getWireMock() { + return wireMock; + } + + @Override + protected DoclingServeClient.DoclingServeClientBuilder newClientBuilder() { + return DoclingServeJackson2Client.builder(); + } +} diff --git a/docling-serve/docling-serve-client/src/test/java/ai/docling/serve/client/DoclingServeJackson3ClientAsyncExecutorTests.java b/docling-serve/docling-serve-client/src/test/java/ai/docling/serve/client/DoclingServeJackson3ClientAsyncExecutorTests.java new file mode 100644 index 00000000..e736fbdf --- /dev/null +++ b/docling-serve/docling-serve-client/src/test/java/ai/docling/serve/client/DoclingServeJackson3ClientAsyncExecutorTests.java @@ -0,0 +1,27 @@ +package ai.docling.serve.client; + +import static com.github.tomakehurst.wiremock.core.WireMockConfiguration.wireMockConfig; + +import org.junit.jupiter.api.extension.RegisterExtension; + +import com.github.tomakehurst.wiremock.junit5.WireMockExtension; + +/** + * Async executor tests for {@link DoclingServeJackson3Client}. + */ +class DoclingServeJackson3ClientAsyncExecutorTests extends AbstractDoclingServeClientAsyncExecutorTests { + @RegisterExtension + static WireMockExtension wireMock = WireMockExtension.newInstance() + .options(wireMockConfig().dynamicPort()) + .build(); + + @Override + protected WireMockExtension getWireMock() { + return wireMock; + } + + @Override + protected DoclingServeClient.DoclingServeClientBuilder newClientBuilder() { + return DoclingServeJackson3Client.builder(); + } +} diff --git a/docs/src/doc/docs/whats-new.md b/docs/src/doc/docs/whats-new.md index d58c48af..a6f13cc2 100644 --- a/docs/src/doc/docs/whats-new.md +++ b/docs/src/doc/docs/whats-new.md @@ -25,6 +25,7 @@ Docling Java {{ gradle.project_version }} includes important breaking changes, a ### {{ gradle.project_version }} +* **Custom `Executor` for async operations** — The async methods (`convertSourceAsync`, `convertSourceBatchAsync`, `convertFilesAsync`, `chunkSourceWith*ChunkerAsync`, ...) used to run on `CompletableFuture`'s default executor (usually `ForkJoinPool.commonPool()`). A new `asyncExecutor(Executor)` builder method lets you run them (task submission, status polling and result retrieval) on your own executor instead, e.g. a virtual-thread executor or one managed by your framework. When not set, the behaviour is unchanged. The client never shuts the executor down. For custom `DoclingApiBuilder` implementations, `asyncExecutor` is a `default` method that throws `UnsupportedOperationException`, so they keep compiling. * **New `DocumentRequest` sealed base class** — `ConvertDocumentRequest`, `BatchConvertDocumentRequest`, and `ChunkDocumentRequest` now extend a common `DocumentRequest` abstract class in the `ai.docling.serve.api.request` package. This enables polymorphism when working with different request types — for example, accepting a `DocumentRequest` and dispatching to the correct endpoint based on the concrete type via pattern matching. * **New `ProcessedDocumentResponse` sealed base class** — `ConvertDocumentResponse` and `ChunkDocumentResponse` now extend a common `ProcessedDocumentResponse` abstract class in the `ai.docling.serve.api.response` package. This enables polymorphic handling of document processing responses — for example, using `ProcessedDocumentResponse` as a type bound in generic APIs that work with both conversion and chunking results. * **`toBuilder()` on the `DocumentRequest` base type** — `DocumentRequest` (and the intermediate `ChunkDocumentRequest`) now expose `toBuilder()`, so a request can be cloned and modified through the base type without first pattern-matching on the concrete subtype. This makes it possible to inject a `source` or `target` once — polymorphically — before dispatching, e.g. `request.toBuilder().source(source).build()`.