diff --git a/docling-serve/docling-serve-client/src/main/java/ai/docling/serve/client/operations/TaskOperations.java b/docling-serve/docling-serve-client/src/main/java/ai/docling/serve/client/operations/TaskOperations.java index 1b3c392c..7021202b 100644 --- a/docling-serve/docling-serve-client/src/main/java/ai/docling/serve/client/operations/TaskOperations.java +++ b/docling-serve/docling-serve-client/src/main/java/ai/docling/serve/client/operations/TaskOperations.java @@ -1,6 +1,7 @@ package ai.docling.serve.client.operations; import java.io.IOException; +import java.io.InputStream; import java.nio.charset.StandardCharsets; import ai.docling.serve.api.DoclingServeTaskApi; @@ -36,7 +37,7 @@ public TaskOperations(HttpOperations httpOperations) { * unique task identifier and optional wait time for polling. * Must not be null. * @return a {@link TaskStatusPollResponse} containing the current status of - * the task, its position in the queue, and any associated metadata. + * the task, its position in the queue, and any associated metadata. * @throws IllegalArgumentException if the {@code request} is null. */ public TaskStatusPollResponse pollTaskStatus(TaskStatusPollRequest request) { @@ -44,9 +45,7 @@ public TaskStatusPollResponse pollTaskStatus(TaskStatusPollRequest request) { return this.httpOperations.executeGet(createRequestContext( "/v1/status/poll/%s?wait=%d".formatted( - request.getTaskId(), - request.getWaitTime().toSeconds()), - TaskStatusPollResponse.class) + request.getTaskId(), request.getWaitTime().toSeconds()), TaskStatusPollResponse.class) ); } @@ -58,7 +57,7 @@ public TaskStatusPollResponse pollTaskStatus(TaskStatusPollRequest request) { * @param request an instance of {@link TaskResultRequest} containing the unique task * identifier. Must not be null. * @return a {@link ConvertDocumentResponse} containing details about the converted - * document, processing time, status, and any associated errors or metadata. + * document, processing time, status, and any associated errors or metadata. * @throws IllegalArgumentException if {@code request} is null. */ public ConvertDocumentResponse convertTaskResult(TaskResultRequest request) { @@ -69,9 +68,9 @@ public ConvertDocumentResponse convertTaskResult(TaskResultRequest request) { case HttpOperations.CONTENT_TYPE_JSON -> { try (var is = response.getBody()) { return httpOperations - .readValue(new String(is.readAllBytes(), StandardCharsets.UTF_8) - , ConvertDocumentResponse.class); - } catch (IOException e) { + .readValue(new String(is.readAllBytes(), StandardCharsets.UTF_8), ConvertDocumentResponse.class); + } + catch (IOException e) { throw new DoclingServeClientException(e); } } @@ -82,7 +81,10 @@ public ConvertDocumentResponse convertTaskResult(TaskResultRequest request) { .inputStream(response.getBody()) .build(); } - default -> throw new DoclingServeClientException(null, "Invalid Content-Type in Task API response"); + default -> { + closeQuietly(response.getBody()); + throw new DoclingServeClientException(null, "Invalid Content-Type in Task API response"); + } } } @@ -96,7 +98,7 @@ public ConvertDocumentResponse convertTaskResult(TaskResultRequest request) { * @param request an instance of {@link TaskResultRequest} containing the unique task * identifier. Must not be null. * @return a {@link ChunkDocumentResponse} containing details about the chunks, - * documents, processing time, and any associated metadata. + * documents, processing time, and any associated metadata. * @throws IllegalArgumentException if {@code request} is null. */ public ChunkDocumentResponse chunkTaskResult(TaskResultRequest request) { @@ -104,6 +106,15 @@ public ChunkDocumentResponse chunkTaskResult(TaskResultRequest request) { return this.httpOperations.executeGet(createRequestContext("/v1/result/%s".formatted(request.getTaskId()), ChunkDocumentResponse.class)); } + private static void closeQuietly(InputStream inputStream) { + try { + inputStream.close(); + } + catch (IOException ignored) { + // the Content-Type error is the failure worth reporting + } + } + private RequestContext createRequestContext(String uri, Class responseType) { return RequestContext.builder() .responseType(responseType) diff --git a/docling-serve/docling-serve-client/src/test/java/ai/docling/serve/client/operations/TaskOperationsTests.java b/docling-serve/docling-serve-client/src/test/java/ai/docling/serve/client/operations/TaskOperationsTests.java new file mode 100644 index 00000000..91cce7a8 --- /dev/null +++ b/docling-serve/docling-serve-client/src/test/java/ai/docling/serve/client/operations/TaskOperationsTests.java @@ -0,0 +1,76 @@ +package ai.docling.serve.client.operations; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import java.io.ByteArrayInputStream; +import java.io.InputStream; +import java.util.Optional; +import java.util.concurrent.atomic.AtomicBoolean; + +import org.junit.jupiter.api.Test; + +import ai.docling.serve.api.task.request.TaskResultRequest; +import ai.docling.serve.client.DoclingServeClientException; + +class TaskOperationsTests { + @Test + void convertTaskResultClosesBodyOnUnexpectedContentType() { + var closed = new AtomicBoolean(); + var body = new ByteArrayInputStream(new byte[]{ + 1, + 2, + 3 + }) { + @Override + public void close() { + closed.set(true); + } + }; + var taskOperations = new TaskOperations(new StubHttpOperations(body, "text/html")); + var request = TaskResultRequest.builder().taskId("task-1").build(); + + assertThatThrownBy(() -> taskOperations.convertTaskResult(request)) + .isInstanceOf(DoclingServeClientException.class) + .hasMessageContaining("Invalid Content-Type"); + assertThat(closed).isTrue(); + } + + private static final class StubHttpOperations extends HttpOperations { + private final InputStream body; + private final String contentType; + + StubHttpOperations(InputStream body, String contentType) { + this.body = body; + this.contentType = contentType; + } + + @Override + protected StreamResponse executeGetWithStreamResponse(RequestContext requestContext) { + return StreamResponse.builder() + .body(body) + .headers(name -> CONTENT_TYPE_HEADER.equals(name) ? Optional.of(contentType) : Optional.empty()) + .build(); + } + + @Override + protected O executeGet(RequestContext requestContext) { + throw new UnsupportedOperationException(); + } + + @Override + protected O executePost(RequestContext requestContext) { + throw new UnsupportedOperationException(); + } + + @Override + protected StreamResponse executePostWithStreamResponse(RequestContext requestContext) { + throw new UnsupportedOperationException(); + } + + @Override + protected T readValue(String json, Class valueType) { + throw new UnsupportedOperationException(); + } + } +}