diff --git a/pom.xml b/pom.xml index b01a5f9ff..fd5fcea65 100644 --- a/pom.xml +++ b/pom.xml @@ -59,7 +59,40 @@ maven-install-plugin ${org.apache.maven.install.version} - + + org.projectlombok + lombok-maven-plugin + 1.18.20.0 + + + generate-sources + + delombok + + + + + ${project.basedir}/src/main/java + ${project.build.directory}/delombok + false + + + + org.apache.maven.plugins + maven-javadoc-plugin + ${org.apache.maven.javadoc.version} + + ${project.build.directory}/delombok + + + + attach-javadocs + + jar + + + + maven-site-plugin ${org.apache.maven.site.version} diff --git a/src/main/java/com/mindee/exceptions/MindeeInputException.java b/src/main/java/com/mindee/exceptions/MindeeInputException.java new file mode 100644 index 000000000..e7412a995 --- /dev/null +++ b/src/main/java/com/mindee/exceptions/MindeeInputException.java @@ -0,0 +1,23 @@ +package com.mindee.exceptions; + +import com.mindee.MindeeException; + +/** + * Represent an invalid or malformed input. + */ +public class MindeeInputException extends MindeeException { + + /** + * {@link Exception} + */ + public MindeeInputException(String message) { + super(message); + } + + /** + * {@link Exception} + */ + public MindeeInputException(String message, Exception innerException) { + super(message, innerException); + } +} diff --git a/src/main/java/com/mindee/v2/MindeeClient.java b/src/main/java/com/mindee/v2/MindeeClient.java index 670588323..705f30317 100644 --- a/src/main/java/com/mindee/v2/MindeeClient.java +++ b/src/main/java/com/mindee/v2/MindeeClient.java @@ -3,12 +3,15 @@ import com.mindee.MindeeException; import com.mindee.input.LocalInputSource; import com.mindee.input.URLInputSource; +import com.mindee.v2.clientoptions.BaseAnnotationParameters; import com.mindee.v2.clientoptions.BaseProductParameters; +import com.mindee.v2.clientoptions.BaseRagDocumentUploadParameters; import com.mindee.v2.clientoptions.BaseSearchParameters; import com.mindee.v2.clientoptions.PollingOptions; import com.mindee.v2.http.MindeeApiV2; import com.mindee.v2.http.MindeeHttpApiV2; import com.mindee.v2.http.MindeeHttpExceptionV2; +import com.mindee.v2.parsing.BaseRagAnnotationResponse; import com.mindee.v2.parsing.BaseResponse; import com.mindee.v2.parsing.JobResponse; import com.mindee.v2.parsing.error.ErrorResponse; @@ -78,10 +81,10 @@ public JobResponse enqueue( * Can be used for polling. */ public JobResponse getJobFromUrl(String pollingUrl) { - logger.log(System.Logger.Level.INFO, "Getting Job at: {0}", pollingUrl); if (pollingUrl == null || pollingUrl.isBlank()) { throw new IllegalArgumentException("Job URL cannot be null or blank."); } + logger.log(System.Logger.Level.INFO, "Getting Job at: {0}", pollingUrl); return mindeeApi.reqGetJobByUrl(pollingUrl); } @@ -90,10 +93,10 @@ public JobResponse getJobFromUrl(String pollingUrl) { * Can be used for polling. */ public JobResponse getJob(String jobId) { - logger.log(System.Logger.Level.INFO, "Getting job ID: {0}", jobId); if (jobId == null || jobId.isBlank()) { throw new IllegalArgumentException("jobId must not be null or blank."); } + logger.log(System.Logger.Level.INFO, "Getting job ID: {0}", jobId); return mindeeApi.reqGetJobById(jobId); } @@ -105,10 +108,10 @@ public TResponse getResult( Class responseClass, String inferenceId ) { - logger.log(System.Logger.Level.INFO, "Getting result with ID: {0}", inferenceId); - if (inferenceId == null || inferenceId.trim().isEmpty()) { + if (inferenceId == null || inferenceId.isBlank()) { throw new IllegalArgumentException("inferenceId must not be null or blank."); } + logger.log(System.Logger.Level.INFO, "Getting result with ID: {0}", inferenceId); return mindeeApi.reqGetResultById(responseClass, inferenceId); } @@ -120,10 +123,10 @@ public TResponse getResultFromUrl( Class responseClass, String inferenceUrl ) { - logger.log(System.Logger.Level.INFO, "Getting result at: {0}", inferenceUrl); - if (inferenceUrl == null || inferenceUrl.trim().isEmpty()) { + if (inferenceUrl == null || inferenceUrl.isBlank()) { throw new IllegalArgumentException("inferenceUrl must not be null or blank."); } + logger.log(System.Logger.Level.INFO, "Getting result at: {0}", inferenceUrl); return mindeeApi.reqGetResultByUrl(responseClass, inferenceUrl); } @@ -277,6 +280,241 @@ public SearchResponse searchModels(String modelName, String modelType) { .reqGetSearch(ModelSearchParameters.builder().name(modelName).modelType(modelType).build()); } + /** + * Not recommended for general use, prefer {@link #uploadAndGetRagDocument}. + * You will need to poll until the document is ready for use. + *

+ * Add a document to the RAG database. + *

+ * + * @param inputSource The file to upload. + * @param parameters RAG document upload parameters. + * @return an instance of {@link BaseRagAnnotationResponse}. + */ + public TAnnotationResponse uploadRagDocument( + LocalInputSource inputSource, + BaseRagDocumentUploadParameters parameters + ) throws IOException { + logger.log(System.Logger.Level.INFO, "Adding a document to the RAG database"); + return mindeeApi.reqPostRagDocument(parameters, inputSource); + } + + /** + * Add a document to the RAG database, poll, and return the initial annotation. + * + * @param inputSource The file to upload. + * @param parameters RAG document upload parameters. + * @return an instance of {@link BaseRagAnnotationResponse}. + */ + public TAnnotationResponse uploadAndGetRagDocument( + LocalInputSource inputSource, + BaseRagDocumentUploadParameters parameters + ) throws IOException, InterruptedException { + return uploadAndGetRagDocument(inputSource, parameters, null); + } + + /** + * Add a document to the RAG database, poll, and return the initial annotation. + * + * @param inputSource The file to upload. + * @param parameters RAG document upload parameters. + * @param pollingOptions Polling options (if null, default options are used). + * @return an instance of {@link BaseRagAnnotationResponse}. + */ + public TAnnotationResponse uploadAndGetRagDocument( + LocalInputSource inputSource, + BaseRagDocumentUploadParameters parameters, + PollingOptions pollingOptions + ) throws IOException, InterruptedException { + if (pollingOptions == null) { + pollingOptions = PollingOptions.builder().build(); + } + TAnnotationResponse initialResponse = uploadRagDocument(inputSource, parameters); + if (!"Processing".equals(initialResponse.getStatus())) { + return initialResponse; + } + return pollForRagDocument(parameters.getResponseClass(), initialResponse, pollingOptions); + } + + /** + * Not recommended for general use, prefer {@link #getReadyRagDocument}. + * You will need to poll until the document is ready for use. + *

+ * Get a document's info and annotations from the RAG database. + *

+ * + * @param responseClass The class of the response. + * @param documentId The document's ID. + * @return an instance of {@link BaseRagAnnotationResponse}. + */ + public TAnnotationResponse getRagDocument( + Class responseClass, + String documentId + ) { + if (documentId == null || documentId.isBlank()) { + throw new IllegalArgumentException("documentId must not be null or blank."); + } + logger.log(System.Logger.Level.INFO, "Getting RAG document ID: {0}", documentId); + return mindeeApi.reqGetRagAnnotation(responseClass, documentId); + } + + /** + * Get a document's info and annotations from the RAG database, polling if it is still processing. + * + * @param responseClass The class of the response. + * @param documentId The document's ID. + * @return an instance of {@link BaseRagAnnotationResponse}. + */ + public TAnnotationResponse getReadyRagDocument( + Class responseClass, + String documentId + ) throws InterruptedException { + return getReadyRagDocument(responseClass, documentId, null); + } + + /** + * Get a document's info and annotations from the RAG database, polling if it is still processing. + * + * @param responseClass The class of the response. + * @param documentId The document's ID. + * @param pollingOptions Polling options (if null, default options are used). + * @return an instance of {@link BaseRagAnnotationResponse}. + */ + public TAnnotationResponse getReadyRagDocument( + Class responseClass, + String documentId, + PollingOptions pollingOptions + ) throws InterruptedException { + TAnnotationResponse initialResponse = getRagDocument(responseClass, documentId); + if (!"Processing".equals(initialResponse.getStatus())) { + return initialResponse; + } + + if (pollingOptions == null) { + pollingOptions = PollingOptions.builder().build(); + } + return pollForRagDocument(responseClass, initialResponse, pollingOptions); + } + + /** + * Not recommended for general use, prefer {@link #updateAndGetRagAnnotation}. + * You will need to poll until the document is ready for use. + *

+ * Update a document's annotations in the RAG database. + *

+ * + * @param parameters Annotation parameters. + * @return an instance of {@link BaseRagAnnotationResponse}. + */ + public TAnnotationResponse updateRagAnnotation( + BaseAnnotationParameters parameters + ) { + logger + .log(System.Logger.Level.INFO, "Updating RAG document ID: {0}", parameters.getDocumentId()); + return mindeeApi.reqPatchRagAnnotation(parameters); + } + + /** + * Update a document's annotations in the RAG database and poll until complete. + * + * @param parameters Annotation parameters. + * @return an instance of {@link BaseRagAnnotationResponse}. + */ + public TAnnotationResponse updateAndGetRagAnnotation( + BaseAnnotationParameters parameters + ) throws InterruptedException { + return updateAndGetRagAnnotation(parameters, null); + } + + /** + * Update a document's annotations in the RAG database and poll until complete. + * + * @param parameters Annotation parameters. + * @param pollingOptions Polling options (if null, default options are used). + * @return an instance of {@link BaseRagAnnotationResponse}. + */ + public TAnnotationResponse updateAndGetRagAnnotation( + BaseAnnotationParameters parameters, + PollingOptions pollingOptions + ) throws InterruptedException { + TAnnotationResponse initialResponse = updateRagAnnotation(parameters); + if (!"Processing".equals(initialResponse.getStatus())) { + return initialResponse; + } + + if (pollingOptions == null) { + pollingOptions = PollingOptions.builder().build(); + } + return pollForRagDocument(parameters.getResponseClass(), initialResponse, pollingOptions); + } + + /** + * Delete a document from the RAG database. + * For extraction models only. + * + * @param documentId The document's ID. + * @return true if successful. + */ + public boolean deleteExtractionRagDocument(String documentId) { + if (documentId == null || documentId.isBlank()) { + throw new IllegalArgumentException("documentId must not be null or blank."); + } + logger.log(System.Logger.Level.INFO, "Deleting RAG document ID: {0}", documentId); + return mindeeApi.reqDeleteExtractionRagDocument(documentId); + } + + /** + * Poll until the RAG document is finished processing or the max number of attempts is reached. + * + * @param responseClass The class of the response. + * @param initialResponse The initial annotation response. + * @param pollingOptions Polling options. + * @return an instance of {@link BaseRagAnnotationResponse}. + * @throws InterruptedException Throws if the thread is interrupted. + */ + private TAnnotationResponse pollForRagDocument( + Class responseClass, + BaseRagAnnotationResponse initialResponse, + PollingOptions pollingOptions + ) throws InterruptedException { + logger + .log(System.Logger.Level.INFO, "Polling for RAG document ID: {0}", initialResponse.getId()); + int maxRetries = pollingOptions.getMaxRetries() + 1; + + logger + .log( + System.Logger.Level.DEBUG, + "Waiting {0} seconds before attempting to retrieve the result...", + pollingOptions.getInitialDelaySec() + ); + + interruptibleSleep( + (long) (pollingOptions.getInitialDelaySec() * 1000), + pollingOptions.getCancelToken() + ); + + String documentId = initialResponse.getId(); + long intervalMillis = (long) (pollingOptions.getIntervalSec() * 1000); + int retryCount = 1; + + while (retryCount < maxRetries) { + logger.log(System.Logger.Level.DEBUG, "Poll attempt {0} of {1}", retryCount, maxRetries); + + TAnnotationResponse response = getRagDocument(responseClass, documentId); + retryCount++; + + String status = response.getStatus(); + if ("Processing".equals(status)) { + interruptibleSleep(intervalMillis, pollingOptions.getCancelToken()); + } else if ("Failed".equals(status)) { + throw new MindeeException("RAG failed without an error payload."); + } else { + return response; + } + } + throw new MindeeException("RAG polling not complete after " + retryCount + " attempts."); + } + /** * Common logic for polling an asynchronous job for local & url files. * diff --git a/src/main/java/com/mindee/v2/clientoptions/BaseAnnotationParameters.java b/src/main/java/com/mindee/v2/clientoptions/BaseAnnotationParameters.java new file mode 100644 index 000000000..b2c63ba4b --- /dev/null +++ b/src/main/java/com/mindee/v2/clientoptions/BaseAnnotationParameters.java @@ -0,0 +1,53 @@ +package com.mindee.v2.clientoptions; + +import com.mindee.v2.parsing.BaseRagAnnotationResponse; +import java.util.Map; +import java.util.Objects; +import lombok.Getter; + +/** + * Base parameters for document annotations. + */ +@Getter +public abstract class BaseAnnotationParameters { + private final Class responseClass; + + /** + * UUID of the annotated document. + */ + private final String documentId; + + /** + * Base constructor. + * + * @param documentId {@link #documentId} + */ + protected BaseAnnotationParameters(Class responseClass, String documentId) { + this.responseClass = Objects.requireNonNull(responseClass, "responseClass cannot be null"); + + if (documentId == null || documentId.trim().isEmpty()) { + throw new IllegalArgumentException("DocumentId cannot be null or whitespace."); + } + + // Note: DocumentId is included in the request URL path, it is not a parameter. + this.documentId = documentId.trim(); + } + + /** + * Gets the request parameters for the upload request. + */ + public abstract Map getRequestParameters(); + + protected abstract static class BaseBuilder> { + protected String documentId; + + @SuppressWarnings("unchecked") + protected T self() { + return (T) this; + } + + protected BaseBuilder(String documentId) { + this.documentId = documentId; + } + } +} diff --git a/src/main/java/com/mindee/v2/clientoptions/BaseProductParameters.java b/src/main/java/com/mindee/v2/clientoptions/BaseProductParameters.java index f411cda7d..73dd37926 100644 --- a/src/main/java/com/mindee/v2/clientoptions/BaseProductParameters.java +++ b/src/main/java/com/mindee/v2/clientoptions/BaseProductParameters.java @@ -26,6 +26,13 @@ public abstract class BaseProductParameters { */ protected final String[] webhookIds; + /** + * Base constructor. + * + * @param modelId {@link #modelId} + * @param alias {@link #alias} + * @param webhookIds {@link #webhookIds} + */ protected BaseProductParameters(String modelId, String alias, String[] webhookIds) { if (modelId == null || modelId.trim().isBlank()) { throw new IllegalArgumentException("modelId cannot be null or whitespace."); diff --git a/src/main/java/com/mindee/v2/clientoptions/BaseRagDocumentUploadParameters.java b/src/main/java/com/mindee/v2/clientoptions/BaseRagDocumentUploadParameters.java new file mode 100644 index 000000000..d2073ec41 --- /dev/null +++ b/src/main/java/com/mindee/v2/clientoptions/BaseRagDocumentUploadParameters.java @@ -0,0 +1,59 @@ +package com.mindee.v2.clientoptions; + +import com.mindee.v2.parsing.BaseRagAnnotationResponse; +import java.util.HashMap; +import java.util.Map; +import java.util.Objects; +import lombok.Getter; + +/** + * Base parameters for document upload operations. + */ +@Getter +public abstract class BaseRagDocumentUploadParameters { + private final Class responseClass; + + /** + * UUID of the model that the uploaded RAG document is linked to. + */ + private final String modelId; + + /** + * Base constructor. + * + * @param modelId {@link #modelId} + */ + protected BaseRagDocumentUploadParameters( + Class responseClass, + String modelId + ) { + this.responseClass = Objects.requireNonNull(responseClass, "responseClass cannot be null"); + + if (modelId == null || modelId.isBlank()) { + throw new IllegalArgumentException("ModelId cannot be null or whitespace."); + } + this.modelId = modelId.trim(); + } + + /** + * Gets the request parameters for the upload request. + */ + public Map getRequestParameters() { + Map parameters = new HashMap<>(); + parameters.put("model_id", modelId); + return parameters; + } + + protected abstract static class BaseBuilder> { + protected String modelId; + + @SuppressWarnings("unchecked") + protected T self() { + return (T) this; + } + + protected BaseBuilder(String modelId) { + this.modelId = modelId; + } + } +} diff --git a/src/main/java/com/mindee/v2/clientoptions/BaseSearchParameters.java b/src/main/java/com/mindee/v2/clientoptions/BaseSearchParameters.java index 75020137c..c858aecb7 100644 --- a/src/main/java/com/mindee/v2/clientoptions/BaseSearchParameters.java +++ b/src/main/java/com/mindee/v2/clientoptions/BaseSearchParameters.java @@ -24,6 +24,10 @@ public abstract class BaseSearchParameters responseClass, diff --git a/src/main/java/com/mindee/v2/http/MindeeApiV2.java b/src/main/java/com/mindee/v2/http/MindeeApiV2.java index fb00ce6ff..204600288 100644 --- a/src/main/java/com/mindee/v2/http/MindeeApiV2.java +++ b/src/main/java/com/mindee/v2/http/MindeeApiV2.java @@ -1,10 +1,16 @@ package com.mindee.v2.http; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.fasterxml.jackson.databind.json.JsonMapper; import com.mindee.MindeeException; import com.mindee.http.MindeeApiCommon; import com.mindee.input.InputSource; +import com.mindee.input.LocalInputSource; +import com.mindee.v2.clientoptions.BaseAnnotationParameters; import com.mindee.v2.clientoptions.BaseProductParameters; +import com.mindee.v2.clientoptions.BaseRagDocumentUploadParameters; import com.mindee.v2.clientoptions.BaseSearchParameters; +import com.mindee.v2.parsing.BaseRagAnnotationResponse; import com.mindee.v2.parsing.BaseResponse; import com.mindee.v2.parsing.JobResponse; import com.mindee.v2.parsing.error.ErrorResponse; @@ -13,6 +19,9 @@ import com.mindee.v2.product.ProductAttributes; import com.mindee.v2.search.models.ModelSearchParameters; import java.io.IOException; +import java.nio.charset.StandardCharsets; +import org.apache.hc.core5.http.ClassicHttpResponse; +import org.apache.hc.core5.http.io.entity.EntityUtils; /** * Communicate with the Mindee HTTP API V2. @@ -22,10 +31,14 @@ *

*/ public abstract class MindeeApiV2 extends MindeeApiCommon { + + protected static final ObjectMapper mapper = JsonMapper.builder().findAndAddModules().build(); + protected final System.Logger logger = System.getLogger(getClass().getName()); + /** * Send a file to the asynchronous processing queue for a product. * - * @param inputSource Local input source from URL. + * @param inputSource Local input source or URL input source. * @param parameters parameters. */ public abstract JobResponse reqPostProductEnqueue( @@ -36,21 +49,22 @@ public abstract JobResponse reqPostProductEnqueue( /** * Get the status of an inference that was previously enqueued. * - * @param pollingUrl The job URL as returned by the predict_async route. + * @param pollingUrl The job ID as returned by the enqueue call. */ public abstract JobResponse reqGetJobByUrl(String pollingUrl); /** - * Attempts to poll the queue. + * Get the status of an inference that was previously enqueued. * - * @param jobId id of the job to get. + * @param jobId The job ID as returned by the enqueue call. */ public abstract JobResponse reqGetJobById(String jobId); /** - * Retrieves the inference from a 302 redirect. + * Get the result of an inference that was previously enqueued. * - * @param inferenceId ID of the inference to poll. + * @param responseClass The class of the response. + * @param inferenceId URL to poll. */ public abstract TResponse reqGetResultById( Class responseClass, @@ -58,8 +72,10 @@ public abstract TResponse reqGetResultById( ); /** - * Retrieves the inference from a given URL. - * The inference will only be available after it has finished processing. + * Get the result of an inference that was previously enqueued. + * + * @param responseClass The class of the response. + * @param inferenceUrl URL to poll. */ public abstract TResponse reqGetResultByUrl( Class responseClass, @@ -68,14 +84,116 @@ public abstract TResponse reqGetResultByUrl( /** * Retrieves a list of resources with the given criteria. + * + * @param parameters Search parameters. */ public abstract TSearchResponse reqGetSearch( BaseSearchParameters parameters ); + /** + * Add a document to the RAG database. + * + * @param parameters RAG document upload parameters. + * @param localInputSource Local input source. + */ + public abstract TAnnotationResponse reqPostRagDocument( + BaseRagDocumentUploadParameters parameters, + LocalInputSource localInputSource + ) throws IOException; + + /** + * Get a document's info and annotations from the RAG database. + * + * @param responseClass The class of the response. + * @param documentId The ID of the document. + */ + public abstract TAnnotationResponse reqGetRagAnnotation( + Class responseClass, + String documentId + ); + + /** + * Update a document's annotations in the RAG database. + * + * @param parameters Annotation parameters. + */ + public abstract TAnnotationResponse reqPatchRagAnnotation( + BaseAnnotationParameters parameters + ); + + /** + * Deletes a document from the RAG database. + * For extraction models only. + * + * @param documentId The ID of the document to delete. + */ + public abstract boolean reqDeleteExtractionRagDocument(String documentId); + + /** + * Retrieves a list of models available for a given API key. + * + * @param parameters Model search parameters. + */ @Deprecated public abstract SearchResponse reqGetSearch(ModelSearchParameters parameters); + /** + * Get the error from the server response. + */ + protected MindeeHttpExceptionV2 getErrorFromResponse(ClassicHttpResponse response) { + logger.log(System.Logger.Level.INFO, "Parsing error response ..."); + + String rawBody; + try { + rawBody = response.getEntity() == null + ? "" + : EntityUtils.toString(response.getEntity(), StandardCharsets.UTF_8); + + logger.log(System.Logger.Level.DEBUG, "HTTP response: {0}", rawBody); + + var errorResponse = mapper.readValue(rawBody, ErrorResponse.class); + + if (errorResponse.getDetail() == null) { + errorResponse = makeUnknownError(response.getCode()); + } + return new MindeeHttpExceptionV2(errorResponse); + + } catch (Exception exception) { + return new MindeeHttpExceptionV2(makeUnknownError(response.getCode()), exception); + } + } + + protected R deserializeResponse( + String body, + Class clazz, + int httpStatus + ) throws MindeeHttpExceptionV2 { + + if (httpStatus >= 200 && httpStatus < 400) { + try { + var model = mapper.readerFor(clazz).readValue(body); + model.setRawResponse(body); + return model; + } catch (Exception exception) { + throw new MindeeException( + "Couldn't deserialize server response:\n" + exception.getMessage() + ); + } + } + + ErrorResponse errorResponse; + try { + errorResponse = mapper.readValue(body, ErrorResponse.class); + if (errorResponse.getDetail() == null) { + errorResponse = makeUnknownError(httpStatus); + } + } catch (Exception ignored) { + errorResponse = makeUnknownError(httpStatus); + } + throw new MindeeHttpExceptionV2(errorResponse); + } + /** * Creates an "unknown error" response from an HTTP status code. */ diff --git a/src/main/java/com/mindee/v2/http/MindeeHttpApiV2.java b/src/main/java/com/mindee/v2/http/MindeeHttpApiV2.java index e7c827248..ce53f0ab7 100644 --- a/src/main/java/com/mindee/v2/http/MindeeHttpApiV2.java +++ b/src/main/java/com/mindee/v2/http/MindeeHttpApiV2.java @@ -1,17 +1,17 @@ package com.mindee.v2.http; -import com.fasterxml.jackson.databind.ObjectMapper; -import com.fasterxml.jackson.databind.json.JsonMapper; import com.mindee.MindeeException; import com.mindee.input.InputSource; import com.mindee.input.LocalInputSource; import com.mindee.input.URLInputSource; import com.mindee.v2.MindeeSettings; +import com.mindee.v2.clientoptions.BaseAnnotationParameters; import com.mindee.v2.clientoptions.BaseProductParameters; +import com.mindee.v2.clientoptions.BaseRagDocumentUploadParameters; import com.mindee.v2.clientoptions.BaseSearchParameters; +import com.mindee.v2.parsing.BaseRagAnnotationResponse; import com.mindee.v2.parsing.BaseResponse; import com.mindee.v2.parsing.JobResponse; -import com.mindee.v2.parsing.error.ErrorResponse; import com.mindee.v2.parsing.search.BaseSearchResponse; import com.mindee.v2.parsing.search.SearchResponse; import com.mindee.v2.search.models.ModelSearchParameters; @@ -19,17 +19,19 @@ import java.net.URISyntaxException; import java.nio.charset.StandardCharsets; import lombok.Builder; +import org.apache.hc.client5.http.classic.methods.HttpDelete; import org.apache.hc.client5.http.classic.methods.HttpGet; +import org.apache.hc.client5.http.classic.methods.HttpPatch; import org.apache.hc.client5.http.classic.methods.HttpPost; import org.apache.hc.client5.http.classic.methods.HttpUriRequestBase; import org.apache.hc.client5.http.config.RequestConfig; import org.apache.hc.client5.http.entity.mime.HttpMultipartMode; import org.apache.hc.client5.http.entity.mime.MultipartEntityBuilder; import org.apache.hc.client5.http.impl.classic.HttpClientBuilder; -import org.apache.hc.core5.http.ClassicHttpResponse; import org.apache.hc.core5.http.ContentType; import org.apache.hc.core5.http.HttpHeaders; import org.apache.hc.core5.http.io.entity.EntityUtils; +import org.apache.hc.core5.http.io.entity.StringEntity; import org.apache.hc.core5.net.URIBuilder; /** @@ -37,9 +39,6 @@ */ public final class MindeeHttpApiV2 extends MindeeApiV2 { - private static final System.Logger logger = System.getLogger(MindeeHttpApiV2.class.getName()); - private static final ObjectMapper mapper = JsonMapper.builder().findAndAddModules().build(); - /** * The MindeeSetting needed to make the api call. */ @@ -154,6 +153,100 @@ public TSearchResponse reqGetSearch return this.executeAPIRequest(get, parameters.getResponseClass()); } + @Override + public TAnnotationResponse reqPostRagDocument( + BaseRagDocumentUploadParameters parameters, + LocalInputSource localInputSource + ) { + var productInfo = getResponseProductAttributes(parameters.getResponseClass()); + var url = String + .format("%s/products/%s/rag-documents", this.mindeeSettings.getBaseUrl(), productInfo.slug()); + var post = buildHttpPost(url); + + var builder = MultipartEntityBuilder.create(); + builder.setMode(HttpMultipartMode.EXTENDED); + builder + .addBinaryBody( + "file", + localInputSource.getFile(), + ContentType.DEFAULT_BINARY, + localInputSource.getFilename() + ); + + parameters.getRequestParameters().forEach(builder::addTextBody); + post.setEntity(builder.build()); + + logger.log(System.Logger.Level.DEBUG, "HTTP POST to {0} ...", url); + return executeAPIRequest(post, parameters.getResponseClass()); + } + + @Override + public TAnnotationResponse reqGetRagAnnotation( + Class responseClass, + String documentId + ) { + var productInfo = getResponseProductAttributes(responseClass); + var url = String + .format( + "%s/products/%s/rag-documents/%s", + this.mindeeSettings.getBaseUrl(), + productInfo.slug(), + documentId + ); + var get = new HttpGet(url); + + logger.log(System.Logger.Level.DEBUG, "HTTP GET to {0} ...", url); + return executeAPIRequest(get, responseClass); + } + + @Override + public TAnnotationResponse reqPatchRagAnnotation( + BaseAnnotationParameters parameters + ) { + var productInfo = getResponseProductAttributes(parameters.getResponseClass()); + var url = String + .format( + "%s/products/%s/rag-documents/%s", + this.mindeeSettings.getBaseUrl(), + productInfo.slug(), + parameters.getDocumentId() + ); + var patch = new HttpPatch(url); + + try { + var json = mapper.writeValueAsString(parameters.getRequestParameters()); + patch.setEntity(new StringEntity(json, ContentType.APPLICATION_JSON)); + } catch (com.fasterxml.jackson.core.JsonProcessingException e) { + throw new com.mindee.MindeeException("Failed to serialize patch parameters", e); + } + + logger.log(System.Logger.Level.DEBUG, "HTTP PATCH to {0} ...", url); + return executeAPIRequest(patch, parameters.getResponseClass()); + } + + @Override + public boolean reqDeleteExtractionRagDocument(String documentId) { + var url = this.mindeeSettings.getBaseUrl() + "/products/extraction/rag-documents/" + documentId; + var delete = new HttpDelete(url); + + logger.log(System.Logger.Level.DEBUG, "HTTP DELETE to {0} ...", url); + + if (this.mindeeSettings.getApiKey().isPresent()) { + delete.setHeader(HttpHeaders.AUTHORIZATION, this.mindeeSettings.getApiKey().get()); + } + delete.setHeader(HttpHeaders.USER_AGENT, getUserAgent()); + + try (var httpClient = httpClientBuilder.build()) { + return httpClient.execute(delete, response -> { + int statusCode = response.getCode(); + EntityUtils.consumeQuietly(response.getEntity()); + return statusCode >= 200 && statusCode < 300; + }); + } catch (IOException err) { + throw new MindeeException(err.getMessage(), err); + } + } + @Override @Deprecated public SearchResponse reqGetSearch(ModelSearchParameters parameters) { @@ -221,12 +314,12 @@ private TResponse executeAPIRequest( var responseEntity = response.getEntity(); var statusCode = response.getCode(); if (isInvalidStatusCode(statusCode)) { - throw getHttpError(response); + throw getErrorFromResponse(response); } try { var raw = EntityUtils.toString(response.getEntity(), StandardCharsets.UTF_8); logger.log(System.Logger.Level.DEBUG, "HTTP response: {0}", raw); - return deserializeOrThrow(raw, responseClass, response.getCode()); + return deserializeResponse(raw, responseClass, response.getCode()); } finally { EntityUtils.consumeQuietly(responseEntity); } @@ -236,27 +329,6 @@ private TResponse executeAPIRequest( } } - private MindeeHttpExceptionV2 getHttpError(ClassicHttpResponse response) { - String rawBody; - try { - rawBody = response.getEntity() == null - ? "" - : EntityUtils.toString(response.getEntity(), StandardCharsets.UTF_8); - - logger.log(System.Logger.Level.DEBUG, "HTTP response: {0}", rawBody); - - var errorResponse = mapper.readValue(rawBody, ErrorResponse.class); - - if (errorResponse.getDetail() == null) { - errorResponse = makeUnknownError(response.getCode()); - } - return new MindeeHttpExceptionV2(errorResponse); - - } catch (Exception exception) { - return new MindeeHttpExceptionV2(makeUnknownError(response.getCode()), exception); - } - } - private HttpPost buildHttpPost(String url) { HttpPost post; try { @@ -270,34 +342,4 @@ private HttpPost buildHttpPost(String url) { } return post; } - - private R deserializeOrThrow( - String body, - Class clazz, - int httpStatus - ) throws MindeeHttpExceptionV2 { - - if (httpStatus >= 200 && httpStatus < 400) { - try { - var model = mapper.readerFor(clazz).readValue(body); - model.setRawResponse(body); - return model; - } catch (Exception exception) { - throw new MindeeException( - "Couldn't deserialize server response:\n" + exception.getMessage() - ); - } - } - - ErrorResponse errorResponse; - try { - errorResponse = mapper.readValue(body, ErrorResponse.class); - if (errorResponse.getDetail() == null) { - errorResponse = makeUnknownError(httpStatus); - } - } catch (Exception ignored) { - errorResponse = makeUnknownError(httpStatus); - } - throw new MindeeHttpExceptionV2(errorResponse); - } } diff --git a/src/main/java/com/mindee/v2/parsing/BaseRagAnnotationResponse.java b/src/main/java/com/mindee/v2/parsing/BaseRagAnnotationResponse.java new file mode 100644 index 000000000..d47516851 --- /dev/null +++ b/src/main/java/com/mindee/v2/parsing/BaseRagAnnotationResponse.java @@ -0,0 +1,41 @@ +package com.mindee.v2.parsing; + +import com.fasterxml.jackson.annotation.JsonIgnoreProperties; +import com.fasterxml.jackson.annotation.JsonProperty; +import lombok.EqualsAndHashCode; +import lombok.Getter; +import lombok.NoArgsConstructor; + +/** + * Base class for all RAG document responses from the V2 API. + */ +@Getter +@EqualsAndHashCode(callSuper = true) +@JsonIgnoreProperties(ignoreUnknown = true) +@NoArgsConstructor +public class BaseRagAnnotationResponse extends BaseResponse { + + /** + * Unique identifier of the RAG document. + */ + @JsonProperty("id") + private String id; + + /** + * Original filename of the uploaded document. + */ + @JsonProperty("filename") + private String filename; + + /** + * Date and time of the document creation. + */ + @JsonProperty("created_at") + private String createdAt; + + /** + * Current status of the RAG document. + */ + @JsonProperty("status") + private String status; +} diff --git a/src/main/java/com/mindee/v2/parsing/inference/field/InferenceFields.java b/src/main/java/com/mindee/v2/parsing/inference/field/InferenceFields.java index 51c650ea7..3473ee117 100644 --- a/src/main/java/com/mindee/v2/parsing/inference/field/InferenceFields.java +++ b/src/main/java/com/mindee/v2/parsing/inference/field/InferenceFields.java @@ -14,7 +14,7 @@ public final class InferenceFields extends LinkedHashMap { /** - * Retrieves the field as a `SimpleField`. + * Retrieves the field as a {@link SimpleField}. * * @param fieldName the name of the field * @throws IllegalStateException if the field is not a SimpleField @@ -24,7 +24,7 @@ public SimpleField getSimpleField(String fieldName) throws IllegalStateException } /** - * Retrieves the field as a `ListField`. + * Retrieves the field as a {@link ListField}. * * @param fieldName the name of the field * @throws IllegalStateException if the field is not a ListField @@ -34,7 +34,7 @@ public ListField getListField(String fieldName) throws IllegalStateException { } /** - * Retrieves the field as an `ObjectField`. + * Retrieves the field as an {@link ObjectField}. * * @param fieldName the name of the field * @throws IllegalStateException if the field is not a ObjectField diff --git a/src/main/java/com/mindee/v2/product/extraction/params/ExtractionParameters.java b/src/main/java/com/mindee/v2/product/extraction/params/ExtractionParameters.java index 26468dfee..aae6ed23c 100644 --- a/src/main/java/com/mindee/v2/product/extraction/params/ExtractionParameters.java +++ b/src/main/java/com/mindee/v2/product/extraction/params/ExtractionParameters.java @@ -90,7 +90,7 @@ public Map getRequestParameters() { /** * Create a new builder. * - * @param modelId the mandatory model identifier + * @param modelId {@link #modelId} * @return a fresh {@link Builder} */ public static Builder builder(String modelId) { @@ -112,40 +112,37 @@ public static final class Builder extends BaseProductParameters.BaseBuilder { + /** + * Retrieves the field as an {@link AnnotatedSimpleField}. + * + * @param fieldName the name of the field + * @throws IllegalStateException if the field is not a SimpleField + */ + public AnnotatedSimpleField getSimpleField(String fieldName) throws IllegalStateException { + return this.get(fieldName).getSimpleField(); + } + + /** + * Retrieves the field as a {@link AnnotatedListField}. + * + * @param fieldName the name of the field + * @throws IllegalStateException if the field is not a ListField + */ + public AnnotatedListField getListField(String fieldName) throws IllegalStateException { + return this.get(fieldName).getListField(); + } + + /** + * Retrieves the field as an {@link AnnotatedObjectField}. + * + * @param fieldName the name of the field + * @throws IllegalStateException if the field is not a ObjectField + */ + public AnnotatedObjectField getObjectField(String fieldName) throws IllegalStateException { + return this.get(fieldName).getObjectField(); + } +} diff --git a/src/main/java/com/mindee/v2/product/extraction/ragdocuments/AnnotatedListField.java b/src/main/java/com/mindee/v2/product/extraction/ragdocuments/AnnotatedListField.java new file mode 100644 index 000000000..7291b5745 --- /dev/null +++ b/src/main/java/com/mindee/v2/product/extraction/ragdocuments/AnnotatedListField.java @@ -0,0 +1,86 @@ +package com.mindee.v2.product.extraction.ragdocuments; + +import com.fasterxml.jackson.annotation.JsonIgnore; +import com.fasterxml.jackson.annotation.JsonIgnoreProperties; +import com.fasterxml.jackson.annotation.JsonProperty; +import java.util.ArrayList; +import java.util.List; +import java.util.stream.Collectors; +import lombok.EqualsAndHashCode; +import lombok.Getter; +import lombok.NoArgsConstructor; + +/** + * A ListField with additional configuration for annotation. + */ +@EqualsAndHashCode(callSuper = true) +@JsonIgnoreProperties(ignoreUnknown = true) +@NoArgsConstructor +public class AnnotatedListField extends AnnotatedBaseField { + + /** + * List of dynamic fields, prefer SimpleItems or ObjectItems. + */ + @Getter + @JsonProperty("items") + private List items = new ArrayList<>(); + + private List simpleItems; + private List objectItems; + + /** + * Default constructor. + */ + public AnnotatedListField( + List items, + boolean selected, + String guidelines + ) { + super(selected, guidelines); + this.items = items; + } + + /** + * List of simple fields. + */ + @JsonIgnore + public List getSimpleItems() { + if (simpleItems != null) { + return simpleItems; + } + + if (items == null) { + return new ArrayList<>(); + } + + simpleItems = items + .stream() + .filter(item -> item.getSimpleField() != null) + .map(AnnotatedDynamicField::getSimpleField) + .collect(Collectors.toList()); + + return simpleItems; + } + + /** + * List of object fields. + */ + @JsonIgnore + public List getObjectItems() { + if (objectItems != null) { + return objectItems; + } + + if (items == null) { + return new ArrayList<>(); + } + + objectItems = items + .stream() + .filter(item -> item.getObjectField() != null) + .map(AnnotatedDynamicField::getObjectField) + .collect(Collectors.toList()); + + return objectItems; + } +} diff --git a/src/main/java/com/mindee/v2/product/extraction/ragdocuments/AnnotatedObjectField.java b/src/main/java/com/mindee/v2/product/extraction/ragdocuments/AnnotatedObjectField.java new file mode 100644 index 000000000..06f6770ce --- /dev/null +++ b/src/main/java/com/mindee/v2/product/extraction/ragdocuments/AnnotatedObjectField.java @@ -0,0 +1,31 @@ +package com.mindee.v2.product.extraction.ragdocuments; + +import com.fasterxml.jackson.annotation.JsonIgnoreProperties; +import com.fasterxml.jackson.annotation.JsonProperty; +import lombok.EqualsAndHashCode; +import lombok.Getter; +import lombok.NoArgsConstructor; + +/** + * An ObjectField with additional configuration for annotation. + */ +@Getter +@EqualsAndHashCode(callSuper = true) +@JsonIgnoreProperties(ignoreUnknown = true) +@NoArgsConstructor +public class AnnotatedObjectField extends AnnotatedBaseField { + + /** + * Sub-fields of the field. + */ + @JsonProperty("fields") + private AnnotatedFields fields; + + /** + * Default constructor. + */ + public AnnotatedObjectField(AnnotatedFields fields, boolean selected, String guidelines) { + super(selected, guidelines); + this.fields = fields; + } +} diff --git a/src/main/java/com/mindee/v2/product/extraction/ragdocuments/AnnotatedSimpleField.java b/src/main/java/com/mindee/v2/product/extraction/ragdocuments/AnnotatedSimpleField.java new file mode 100644 index 000000000..2f9a123a4 --- /dev/null +++ b/src/main/java/com/mindee/v2/product/extraction/ragdocuments/AnnotatedSimpleField.java @@ -0,0 +1,33 @@ +package com.mindee.v2.product.extraction.ragdocuments; + +import com.fasterxml.jackson.annotation.JsonIgnoreProperties; +import com.fasterxml.jackson.annotation.JsonProperty; +import com.fasterxml.jackson.databind.annotation.JsonDeserialize; +import lombok.EqualsAndHashCode; +import lombok.Getter; +import lombok.NoArgsConstructor; + +/** + * A SimpleField with additional configuration for annotation. + */ +@Getter +@NoArgsConstructor +@EqualsAndHashCode(callSuper = true) +@JsonIgnoreProperties(ignoreUnknown = true) +@JsonDeserialize(using = AnnotatedSimpleFieldDeserializer.class) +public class AnnotatedSimpleField extends AnnotatedBaseField { + + /** + * Field value, one of: string, bool, int, double, null. + */ + @JsonProperty("value") + private Object value; + + /** + * Default constructor. + */ + public AnnotatedSimpleField(Object value, boolean selected, String guidelines) { + super(selected, guidelines); + this.value = value; + } +} diff --git a/src/main/java/com/mindee/v2/product/extraction/ragdocuments/AnnotatedSimpleFieldDeserializer.java b/src/main/java/com/mindee/v2/product/extraction/ragdocuments/AnnotatedSimpleFieldDeserializer.java new file mode 100644 index 000000000..0ba480319 --- /dev/null +++ b/src/main/java/com/mindee/v2/product/extraction/ragdocuments/AnnotatedSimpleFieldDeserializer.java @@ -0,0 +1,56 @@ +package com.mindee.v2.product.extraction.ragdocuments; + +import com.fasterxml.jackson.core.JsonParser; +import com.fasterxml.jackson.core.ObjectCodec; +import com.fasterxml.jackson.databind.DeserializationContext; +import com.fasterxml.jackson.databind.JsonDeserializer; +import com.fasterxml.jackson.databind.JsonNode; +import java.io.IOException; + +/** + * Custom deserializer for {@link AnnotatedSimpleField}. + */ +public final class AnnotatedSimpleFieldDeserializer extends JsonDeserializer { + + @Override + public AnnotatedSimpleField deserialize( + JsonParser jp, + DeserializationContext ctxt + ) throws IOException { + ObjectCodec codec = jp.getCodec(); + JsonNode root = codec.readTree(jp); + + JsonNode valueNode = root.get("value"); + Object value = null; + + if (valueNode != null && !valueNode.isNull()) { + switch (valueNode.getNodeType()) { + case BOOLEAN: + value = valueNode.booleanValue(); + break; + case NUMBER: + value = valueNode.doubleValue(); + break; + case STRING: + value = valueNode.textValue(); + break; + default: + value = codec.treeToValue(valueNode, Object.class); + } + } + + boolean selected = false; + JsonNode selectedNode = root.get("selected"); + if (selectedNode != null && !selectedNode.isNull()) { + selected = selectedNode.booleanValue(); + } + + String guidelines = null; + JsonNode guidelinesNode = root.get("guidelines"); + if (guidelinesNode != null && !guidelinesNode.isNull()) { + guidelines = guidelinesNode.textValue(); + } + + return new AnnotatedSimpleField(value, selected, guidelines); + } +} diff --git a/src/main/java/com/mindee/v2/product/extraction/ragdocuments/DynamicAnnotationFieldDeserializer.java b/src/main/java/com/mindee/v2/product/extraction/ragdocuments/DynamicAnnotationFieldDeserializer.java new file mode 100644 index 000000000..e8d84371b --- /dev/null +++ b/src/main/java/com/mindee/v2/product/extraction/ragdocuments/DynamicAnnotationFieldDeserializer.java @@ -0,0 +1,68 @@ +package com.mindee.v2.product.extraction.ragdocuments; + +import com.fasterxml.jackson.core.JsonParser; +import com.fasterxml.jackson.core.ObjectCodec; +import com.fasterxml.jackson.databind.DeserializationContext; +import com.fasterxml.jackson.databind.JsonDeserializer; +import com.fasterxml.jackson.databind.JsonNode; +import java.io.IOException; +import java.util.ArrayList; + +/** + * Custom deserializer for {@link AnnotatedDynamicField}. + */ +public final class DynamicAnnotationFieldDeserializer + extends JsonDeserializer { + + @Override + public AnnotatedDynamicField deserialize( + JsonParser jp, + DeserializationContext ctxt + ) throws IOException { + ObjectCodec codec = jp.getCodec(); + JsonNode root = codec.readTree(jp); + + if (root == null || root.isNull()) { + return null; + } + + // -------- LIST OF FIELDS -------- + JsonNode itemsNode = root.get("items"); + if (itemsNode != null && itemsNode.isArray()) { + String guidelines = null; + JsonNode guidelinesNode = root.get("guidelines"); + if (guidelinesNode != null && !guidelinesNode.isNull()) { + guidelines = guidelinesNode.textValue(); + } + + boolean selected = false; + JsonNode selectedNode = root.get("selected"); + if (selectedNode != null && !selectedNode.isNull()) { + selected = selectedNode.booleanValue(); + } + + var items = new ArrayList(); + for (JsonNode item : itemsNode) { + items.add(codec.treeToValue(item, AnnotatedDynamicField.class)); + } + AnnotatedListField listField = new AnnotatedListField(items, selected, guidelines); + + return new AnnotatedDynamicField(listField); + } + + // -------- OBJECT FIELD -------- + JsonNode fieldsNode = root.get("fields"); + if (fieldsNode != null && fieldsNode.isObject()) { + AnnotatedObjectField objectField = codec.treeToValue(root, AnnotatedObjectField.class); + return new AnnotatedDynamicField(objectField); + } + + // -------- SIMPLE FIELD -------- + if (root.has("value")) { + AnnotatedSimpleField simpleField = codec.treeToValue(root, AnnotatedSimpleField.class); + return new AnnotatedDynamicField(simpleField); + } + + throw new IOException("Unknown field: " + root.toString()); + } +} diff --git a/src/main/java/com/mindee/v2/product/extraction/ragdocuments/ExtractionRagAnnotationResponse.java b/src/main/java/com/mindee/v2/product/extraction/ragdocuments/ExtractionRagAnnotationResponse.java new file mode 100644 index 000000000..ea26e31c1 --- /dev/null +++ b/src/main/java/com/mindee/v2/product/extraction/ragdocuments/ExtractionRagAnnotationResponse.java @@ -0,0 +1,44 @@ +package com.mindee.v2.product.extraction.ragdocuments; + +import com.fasterxml.jackson.annotation.JsonIgnoreProperties; +import com.fasterxml.jackson.annotation.JsonProperty; +import com.mindee.v2.parsing.BaseRagAnnotationResponse; +import com.mindee.v2.product.ProductAttributes; +import lombok.EqualsAndHashCode; +import lombok.Getter; +import lombok.NoArgsConstructor; + +/** + * Response for a RAG document. + */ +@Getter +@EqualsAndHashCode(callSuper = true) +@JsonIgnoreProperties(ignoreUnknown = true) +@NoArgsConstructor +@ProductAttributes(slug = "extraction") +public class ExtractionRagAnnotationResponse extends BaseRagAnnotationResponse { + + /** + * Model identifier linked to the RAG document. + */ + @JsonProperty("model_id") + private String modelId; + + /** + * Number of times this document was used in an inference. + */ + @JsonProperty("total_matches") + private int totalMatches; + + /** + * Date and time of the latest matching inference, if any. + */ + @JsonProperty("last_match_at") + private String lastMatchAt; + + /** + * Annotation metadata associated with the document. + */ + @JsonProperty("annotation") + private RagAnnotation annotation; +} diff --git a/src/main/java/com/mindee/v2/product/extraction/ragdocuments/RagAnnotation.java b/src/main/java/com/mindee/v2/product/extraction/ragdocuments/RagAnnotation.java new file mode 100644 index 000000000..c99a4ce1d --- /dev/null +++ b/src/main/java/com/mindee/v2/product/extraction/ragdocuments/RagAnnotation.java @@ -0,0 +1,25 @@ +package com.mindee.v2.product.extraction.ragdocuments; + +import com.fasterxml.jackson.annotation.JsonIgnoreProperties; +import com.fasterxml.jackson.annotation.JsonProperty; +import lombok.AllArgsConstructor; +import lombok.EqualsAndHashCode; +import lombok.Getter; +import lombok.NoArgsConstructor; + +/** + * A RAG annotation enriched with field-level configuration. + */ +@Getter +@EqualsAndHashCode +@JsonIgnoreProperties(ignoreUnknown = true) +@NoArgsConstructor +@AllArgsConstructor +public class RagAnnotation { + + /** + * Annotated fields. + */ + @JsonProperty("fields") + private AnnotatedFields fields; +} diff --git a/src/main/java/com/mindee/v2/product/extraction/ragdocuments/params/RagDocumentAnnotationParameters.java b/src/main/java/com/mindee/v2/product/extraction/ragdocuments/params/RagDocumentAnnotationParameters.java new file mode 100644 index 000000000..ee242e206 --- /dev/null +++ b/src/main/java/com/mindee/v2/product/extraction/ragdocuments/params/RagDocumentAnnotationParameters.java @@ -0,0 +1,134 @@ +package com.mindee.v2.product.extraction.ragdocuments.params; + +import com.fasterxml.jackson.core.JsonProcessingException; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.fasterxml.jackson.databind.json.JsonMapper; +import com.mindee.exceptions.MindeeInputException; +import com.mindee.v2.clientoptions.BaseAnnotationParameters; +import com.mindee.v2.product.extraction.ragdocuments.ExtractionRagAnnotationResponse; +import com.mindee.v2.product.extraction.ragdocuments.RagAnnotation; +import java.util.HashMap; +import java.util.Map; +import lombok.Getter; + +/** + * Annotation parameters for RAG documents. + */ +@Getter +public class RagDocumentAnnotationParameters + extends BaseAnnotationParameters { + + private static final ObjectMapper mapper = JsonMapper.builder().findAndAddModules().build(); + + /** + * New public status to apply to the document (for example, to deactivate it). + */ + private final String status; + + /** + * Field-level RAG annotation and guidelines configuration for the document. + */ + private final RagAnnotation annotation; + + /** + * Constructor with only document ID. + * + * @param documentId {@link BaseAnnotationParameters#getDocumentId()} + */ + public RagDocumentAnnotationParameters(String documentId) { + this(documentId, null, null); + } + + /** + * Constructor with document ID and status. + * + * @param documentId {@link BaseAnnotationParameters#getDocumentId()} + * @param status {@link #status} + */ + public RagDocumentAnnotationParameters(String documentId, String status) { + this(documentId, status, null); + } + + /** + * Default constructor. + * + * @param documentId {@link BaseAnnotationParameters#getDocumentId()} + * @param status {@link #status} + * @param annotation {@link #annotation} + */ + public RagDocumentAnnotationParameters(String documentId, String status, Object annotation) { + super(ExtractionRagAnnotationResponse.class, documentId); + this.status = status; + + if (annotation instanceof RagAnnotation) { + this.annotation = (RagAnnotation) annotation; + } else if (annotation instanceof String) { + try { + this.annotation = mapper.readValue((String) annotation, RagAnnotation.class); + } catch (JsonProcessingException e) { + throw new MindeeInputException("Invalid RAG Annotation format.", e); + } + } else if (annotation == null) { + this.annotation = null; + } else { + throw new MindeeInputException("Invalid RAG Annotation format."); + } + } + + /** + * {@inheritDoc} + */ + @Override + public Map getRequestParameters() { + Map parameters = new HashMap<>(); + + if (status != null && !status.isEmpty()) { + parameters.put("status", status); + } + + if (annotation != null) { + parameters.put("annotation", annotation); + } + + return parameters; + } + + /** + * Create a new builder. + * + * @param documentId {@link BaseAnnotationParameters#getDocumentId()} + * @return a fresh {@link Builder} + */ + public static Builder builder(String documentId) { + return new Builder(documentId); + } + + /** + * Fluent builder for {@link RagDocumentAnnotationParameters}. + */ + public static final class Builder extends BaseAnnotationParameters.BaseBuilder { + private String status; + private Object annotation; + + Builder(String documentId) { + super(documentId); + } + + /** @param status {@link #status} */ + public Builder status(String status) { + this.status = status; + return this; + } + + /** @param annotation {@link #annotation} */ + public Builder annotation(Object annotation) { + this.annotation = annotation; + return this; + } + + /** Build an immutable {@link RagDocumentAnnotationParameters} instance. */ + public RagDocumentAnnotationParameters build() { + return new RagDocumentAnnotationParameters(documentId, status, annotation); + } + } +} diff --git a/src/main/java/com/mindee/v2/product/extraction/ragdocuments/params/RagDocumentUploadParameters.java b/src/main/java/com/mindee/v2/product/extraction/ragdocuments/params/RagDocumentUploadParameters.java new file mode 100644 index 000000000..ceedbd82c --- /dev/null +++ b/src/main/java/com/mindee/v2/product/extraction/ragdocuments/params/RagDocumentUploadParameters.java @@ -0,0 +1,45 @@ +package com.mindee.v2.product.extraction.ragdocuments.params; + +import com.mindee.v2.clientoptions.BaseRagDocumentUploadParameters; +import com.mindee.v2.product.extraction.ragdocuments.ExtractionRagAnnotationResponse; + +/** + * Upload parameters for RAG documents. + */ +public class RagDocumentUploadParameters + extends BaseRagDocumentUploadParameters { + + /** + * {@inheritDoc} + * + * @param modelId {@inheritDoc} + */ + public RagDocumentUploadParameters(String modelId) { + super(ExtractionRagAnnotationResponse.class, modelId); + } + + /** + * Create a new builder. + * + * @param modelId {@link BaseRagDocumentUploadParameters#getModelId()} + * @return a fresh {@link Builder} + */ + public static Builder builder(String modelId) { + return new Builder(modelId); + } + + /** + * Fluent builder for {@link RagDocumentUploadParameters}. + */ + public static final class Builder extends BaseRagDocumentUploadParameters.BaseBuilder { + + Builder(String modelId) { + super(modelId); + } + + /** Build an immutable {@link RagDocumentUploadParameters} instance. */ + public RagDocumentUploadParameters build() { + return new RagDocumentUploadParameters(modelId); + } + } +} diff --git a/src/test/java/com/mindee/v2/MindeeClientIT.java b/src/test/java/com/mindee/v2/MindeeClientIT.java index f2beb8690..48d2d1014 100644 --- a/src/test/java/com/mindee/v2/MindeeClientIT.java +++ b/src/test/java/com/mindee/v2/MindeeClientIT.java @@ -18,17 +18,17 @@ @TestInstance(TestInstance.Lifecycle.PER_CLASS) @Tag("integration") -@DisplayName("MindeeV2 – Integration Tests") +@DisplayName("MindeeV2 – Integration") class MindeeClientIT { - private MindeeClient mindeeClient; + private MindeeClient client; private String modelId; @BeforeAll void setUp() { String apiKey = System.getenv("MINDEE_V2_API_KEY"); modelId = System.getenv("MINDEE_V2_SE_TESTS_FINDOC_MODEL_ID"); - mindeeClient = new MindeeClient(apiKey); + client = new MindeeClient(apiKey); } @Test @@ -52,7 +52,7 @@ void parseFile_emptyMultiPage_mustSucceed() throws IOException, InterruptedExcep .maxRetries(80) .build(); - var response = mindeeClient + var response = client .enqueueAndGetResult(ExtractionResponse.class, source, params, pollingOptions); assertNotNull(response); @@ -98,7 +98,7 @@ void parseFile_filledSinglePage_mustSucceed() throws IOException, InterruptedExc .textContext("this is an invoice") .build(); - var response = mindeeClient.enqueueAndGetResult(ExtractionResponse.class, source, params); + var response = client.enqueueAndGetResult(ExtractionResponse.class, source, params); assertNotNull(response); var inference = response.getInference(); @@ -147,7 +147,7 @@ void parseFile_dataSchemaReplace_mustSucceed() throws IOException, InterruptedEx .dataSchema(Files.readString(getV2ProductPath("extraction/data_schema_replace_param.json"))) .build(); - var response = mindeeClient.enqueueAndGetResult(ExtractionResponse.class, source, params); + var response = client.enqueueAndGetResult(ExtractionResponse.class, source, params); assertNotNull(response); ExtractionInference inference = response.getInference(); assertNotNull(inference); @@ -177,7 +177,7 @@ void invalidModel_mustThrowError() throws IOException { MindeeHttpExceptionV2 err = assertThrows( MindeeHttpExceptionV2.class, - () -> mindeeClient.enqueue(source, params) + () -> client.enqueue(source, params) ); assertEquals(422, err.getStatus()); } @@ -193,7 +193,7 @@ void invalidWebhook_mustThrowError() throws IOException { MindeeHttpExceptionV2 err = assertThrows( MindeeHttpExceptionV2.class, - () -> mindeeClient.enqueue(source, params) + () -> client.enqueue(source, params) ); assertEquals(422, err.getStatus()); } @@ -203,7 +203,7 @@ void invalidWebhook_mustThrowError() throws IOException { void invalidJob_mustThrowError() { MindeeHttpExceptionV2 err = assertThrows( MindeeHttpExceptionV2.class, - () -> mindeeClient.getResult(ExtractionResponse.class, "INVALID_JOB_ID") + () -> client.getResult(ExtractionResponse.class, "INVALID_JOB_ID") ); assertEquals(422, err.getStatus()); assertNotNull(err); @@ -218,7 +218,7 @@ void urlInputSource_mustNotRaiseErrors() throws IOException, InterruptedExceptio var options = ExtractionParameters.builder(modelId).build(); - var response = mindeeClient.enqueueAndGetResult(ExtractionResponse.class, urlSource, options); + var response = client.enqueueAndGetResult(ExtractionResponse.class, urlSource, options); assertNotNull(response); assertNotNull(response.getInference()); @@ -227,7 +227,7 @@ void urlInputSource_mustNotRaiseErrors() throws IOException, InterruptedExceptio @Test @DisplayName("Search for models by name") void searchModelsByName_mustSucceed() { - SearchResponse response = mindeeClient.searchModels("crop"); + SearchResponse response = client.searchModels("crop"); assertNotNull(response); assertFalse(response.getModels().isEmpty()); } @@ -241,21 +241,21 @@ void getResultFromUrl_mustSucceed() throws IOException, InterruptedException { .alias("java-integration-test_get-result-from-url") .build(); - var enqueueResp = mindeeClient.enqueue(source, params); + var enqueueResp = client.enqueue(source, params); assertNotNull(enqueueResp); var jobId = enqueueResp.getJob().getId(); String resultUrl = null; for (int i = 0; i < 80 && resultUrl == null; i++) { Thread.sleep(1500); - var poll = mindeeClient.getJob(jobId); + var poll = client.getJob(jobId); if (poll.getJob().getStatus().equals("Processed")) { resultUrl = poll.getJob().getResultUrl(); } } assertNotNull(resultUrl, "Job must expose a result_url once processed"); - var response = mindeeClient.getResultFromUrl(ExtractionResponse.class, resultUrl); + var response = client.getResultFromUrl(ExtractionResponse.class, resultUrl); assertNotNull(response); assertNotNull(response.getInference()); assertNotNull(response.getInference().getId()); diff --git a/src/test/java/com/mindee/v2/MindeeClientTest.java b/src/test/java/com/mindee/v2/MindeeClientTest.java index d03979442..fe2b911b6 100644 --- a/src/test/java/com/mindee/v2/MindeeClientTest.java +++ b/src/test/java/com/mindee/v2/MindeeClientTest.java @@ -10,16 +10,20 @@ import com.fasterxml.jackson.databind.ObjectMapper; import com.mindee.input.InputSource; import com.mindee.input.LocalInputSource; +import com.mindee.v2.clientoptions.BaseAnnotationParameters; import com.mindee.v2.clientoptions.BaseProductParameters; +import com.mindee.v2.clientoptions.BaseRagDocumentUploadParameters; import com.mindee.v2.clientoptions.BaseSearchParameters; import com.mindee.v2.clientoptions.PollingOptions; import com.mindee.v2.http.MindeeApiV2; +import com.mindee.v2.parsing.BaseRagAnnotationResponse; import com.mindee.v2.parsing.BaseResponse; import com.mindee.v2.parsing.JobResponse; import com.mindee.v2.parsing.search.BaseSearchResponse; import com.mindee.v2.parsing.search.SearchResponse; import com.mindee.v2.product.extraction.ExtractionResponse; import com.mindee.v2.product.extraction.params.ExtractionParameters; +import com.mindee.v2.product.extraction.ragdocuments.ExtractionRagAnnotationResponse; import com.mindee.v2.search.models.ModelSearchParameters; import com.mindee.v2.search.models.ModelSearchResponse; import java.io.IOException; @@ -57,18 +61,49 @@ public JobResponse reqGetJobByUrl(String jobId) { return jobResponse; } - @Override public JobResponse reqGetJobById(String jobId) { return jobResponse; } @Override + @SuppressWarnings("unchecked") public TSearchResponse reqGetSearch( BaseSearchParameters parameters ) { return (TSearchResponse) new ModelSearchResponse(); } + @Override + @SuppressWarnings("unchecked") + public TAnnotationResponse reqPostRagDocument( + BaseRagDocumentUploadParameters parameters, + LocalInputSource localInputSource + ) { + return (TAnnotationResponse) new ExtractionRagAnnotationResponse(); + } + + @Override + @SuppressWarnings("unchecked") + public TAnnotationResponse reqGetRagAnnotation( + Class responseClass, + String documentId + ) { + return (TAnnotationResponse) new ExtractionRagAnnotationResponse(); + } + + @Override + @SuppressWarnings("unchecked") + public TAnnotationResponse reqPatchRagAnnotation( + BaseAnnotationParameters parameters + ) { + return (TAnnotationResponse) new ExtractionRagAnnotationResponse(); + } + + @Override + public boolean reqDeleteExtractionRagDocument(String documentId) { + return true; + } + @Override @Deprecated public SearchResponse reqGetSearch(ModelSearchParameters parameters) { @@ -76,6 +111,7 @@ public SearchResponse reqGetSearch(ModelSearchParameters parameters) { } @Override + @SuppressWarnings("unchecked") public TResponse reqGetResultById( Class tResponseClass, String inferenceId @@ -84,6 +120,7 @@ public TResponse reqGetResultById( } @Override + @SuppressWarnings("unchecked") public TResponse reqGetResultByUrl( Class tResponseClass, String inferenceUrl diff --git a/src/test/java/com/mindee/v2/product/ClassificationTest.java b/src/test/java/com/mindee/v2/product/ClassificationTest.java index 08b83b2c2..2582af45b 100644 --- a/src/test/java/com/mindee/v2/product/ClassificationTest.java +++ b/src/test/java/com/mindee/v2/product/ClassificationTest.java @@ -13,7 +13,7 @@ import org.junit.jupiter.api.Nested; import org.junit.jupiter.api.Test; -@DisplayName("MindeeV2 - Classification Model Tests") +@DisplayName("MindeeV2 - Classification Response") public class ClassificationTest { private ClassificationResponse loadResponse(String filePath) throws IOException { var localResponse = new LocalResponse(getV2ProductPath(filePath)); @@ -26,7 +26,7 @@ class SinglePredictionTest { @Test @DisplayName("classification properties must be valid") void singleMustHaveValidProperties() throws IOException { - ClassificationResponse response = loadResponse("classification/default_sample.json"); + var response = loadResponse("classification/default_sample.json"); assertNotNull(response.getInference()); assertEquals( "invoice", diff --git a/src/test/java/com/mindee/v2/product/CropTest.java b/src/test/java/com/mindee/v2/product/CropTest.java index e2405cf02..4f302577c 100644 --- a/src/test/java/com/mindee/v2/product/CropTest.java +++ b/src/test/java/com/mindee/v2/product/CropTest.java @@ -21,7 +21,7 @@ import org.junit.jupiter.api.Nested; import org.junit.jupiter.api.Test; -@DisplayName("MindeeV2 - Crop Model Tests") +@DisplayName("MindeeV2 - Crop Response") public class CropTest { private static final Path outputPath = getResourcePath("output/v2/product/crop"); diff --git a/src/test/java/com/mindee/v2/product/OcrTest.java b/src/test/java/com/mindee/v2/product/OcrTest.java index 3f1ae485f..9b8571f5f 100644 --- a/src/test/java/com/mindee/v2/product/OcrTest.java +++ b/src/test/java/com/mindee/v2/product/OcrTest.java @@ -11,7 +11,7 @@ import org.junit.jupiter.api.Nested; import org.junit.jupiter.api.Test; -@DisplayName("MindeeV2 - OCR Model Tests") +@DisplayName("MindeeV2 - OCR Response") public class OcrTest { private OcrResponse loadResponse(String filePath) throws IOException { var localResponse = new LocalResponse(getV2ProductPath(filePath)); diff --git a/src/test/java/com/mindee/v2/product/SplitTest.java b/src/test/java/com/mindee/v2/product/SplitTest.java index 0b9b43bbd..426d6f955 100644 --- a/src/test/java/com/mindee/v2/product/SplitTest.java +++ b/src/test/java/com/mindee/v2/product/SplitTest.java @@ -21,7 +21,7 @@ import org.junit.jupiter.api.Nested; import org.junit.jupiter.api.Test; -@DisplayName("MindeeV2 - Split Model Tests") +@DisplayName("MindeeV2 - Split Response") public class SplitTest { private static final Path outputPath = getResourcePath("output/v2/product/split"); diff --git a/src/test/java/com/mindee/v2/product/extraction/ExtractionParametersTest.java b/src/test/java/com/mindee/v2/product/extraction/ExtractionParametersTest.java new file mode 100644 index 000000000..ccb9febe2 --- /dev/null +++ b/src/test/java/com/mindee/v2/product/extraction/ExtractionParametersTest.java @@ -0,0 +1,46 @@ +package com.mindee.v2.product.extraction; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNull; + +import com.mindee.v2.product.extraction.params.ExtractionParameters; +import java.util.Map; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; + +@DisplayName("MindeeV2 - Extraction Parameters") +public class ExtractionParametersTest { + private static final String MODEL_ID = "test-model-id"; + + @Test + @DisplayName("should init with minimum values") + void parameters_mustInit() { + ExtractionParameters productParams = ExtractionParameters.builder(MODEL_ID).build(); + assertEquals(MODEL_ID, productParams.getModelId()); + } + + @Nested + @DisplayName("Data Schema") + class DataSchemaTests { + private Map dataSchemaDict; + private String dataSchemaString; + + @Test + @DisplayName("should leave unset when not provided") + void dataSchema_shouldLeaveUnsetWhenNotProvided() { + ExtractionParameters inferenceParameters = ExtractionParameters.builder(MODEL_ID).build(); + assertNull(inferenceParameters.getDataSchema()); + } + + @Test + @DisplayName("should initialize from a string") + void dataSchemaString_shouldInitialize() { + ExtractionParameters inferenceParameters = ExtractionParameters + .builder(MODEL_ID) + .dataSchema(dataSchemaString) + .build(); + assertEquals(dataSchemaString, inferenceParameters.getDataSchema()); + } + } +} diff --git a/src/test/java/com/mindee/v2/product/ExtractionTest.java b/src/test/java/com/mindee/v2/product/extraction/ExtractionResponseTest.java similarity index 98% rename from src/test/java/com/mindee/v2/product/ExtractionTest.java rename to src/test/java/com/mindee/v2/product/extraction/ExtractionResponseTest.java index 0f6ae84d2..2ec79354f 100644 --- a/src/test/java/com/mindee/v2/product/ExtractionTest.java +++ b/src/test/java/com/mindee/v2/product/extraction/ExtractionResponseTest.java @@ -1,4 +1,4 @@ -package com.mindee.v2.product; +package com.mindee.v2.product.extraction; import static com.mindee.TestingUtilities.getV2ProductPath; import static org.junit.jupiter.api.Assertions.assertEquals; @@ -27,8 +27,6 @@ import com.mindee.v2.parsing.inference.field.ListField; import com.mindee.v2.parsing.inference.field.ObjectField; import com.mindee.v2.parsing.inference.field.SimpleField; -import com.mindee.v2.product.extraction.ExtractionInference; -import com.mindee.v2.product.extraction.ExtractionResponse; import java.io.IOException; import java.math.BigDecimal; import java.nio.file.Files; @@ -40,8 +38,8 @@ import org.junit.jupiter.api.Nested; import org.junit.jupiter.api.Test; -@DisplayName("MindeeV2 - Extraction Model Tests") -class ExtractionTest { +@DisplayName("MindeeV2 - Extraction Response") +class ExtractionResponseTest { private ExtractionResponse loadResponse(String filePath) throws IOException { var localResponse = new LocalResponse(getV2ProductPath(filePath)); @@ -180,11 +178,11 @@ class DeepNestedFieldsTest { @Test @DisplayName("all nested structures must be typed correctly") void deepNestedFields_mustExposeCorrectTypes() throws IOException { - ExtractionResponse resp = loadResponse("extraction/deep_nested_fields.json"); - ExtractionInference inf = resp.getInference(); - assertNotNull(inf); + var response = loadResponse("extraction/deep_nested_fields.json"); + ExtractionInference inference = response.getInference(); + assertNotNull(inference); - var root = inf.getResult().getFields(); + var root = inference.getResult().getFields(); assertNotNull(root.get("field_simple").getSimpleField()); assertNotNull(root.get("field_object").getObjectField()); diff --git a/src/test/java/com/mindee/v2/product/extraction/RagDocumentsIT.java b/src/test/java/com/mindee/v2/product/extraction/RagDocumentsIT.java new file mode 100644 index 000000000..b17166cc8 --- /dev/null +++ b/src/test/java/com/mindee/v2/product/extraction/RagDocumentsIT.java @@ -0,0 +1,121 @@ +package com.mindee.v2.product.extraction; + +import static com.mindee.TestingUtilities.getV2ProductPath; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import com.mindee.input.LocalInputSource; +import com.mindee.v2.MindeeClient; +import com.mindee.v2.http.MindeeHttpExceptionV2; +import com.mindee.v2.product.extraction.ragdocuments.ExtractionRagAnnotationResponse; +import com.mindee.v2.product.extraction.ragdocuments.params.RagDocumentAnnotationParameters; +import com.mindee.v2.product.extraction.ragdocuments.params.RagDocumentUploadParameters; +import java.io.IOException; +import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Tag; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.TestInstance; + +@TestInstance(TestInstance.Lifecycle.PER_CLASS) +@Tag("integration") +@DisplayName("MindeeV2 – Integration") +public class RagDocumentsIT { + private MindeeClient client; + private String modelId; + + @BeforeAll + void setUp() { + String apiKey = System.getenv("MINDEE_V2_API_KEY"); + modelId = System.getenv("MINDEE_V2_SE_TESTS_FINDOC_MODEL_ID"); + client = new MindeeClient(apiKey); + } + + @Test + @DisplayName("Should perform the entire lifecycle of a RAG document.") + void ragDocument_lifecycle_mustSucceed() throws IOException, InterruptedException { + var inputSource = new LocalInputSource( + getV2ProductPath("extraction/financial_document/default_sample.jpg") + ); + var parameters = RagDocumentUploadParameters.builder(modelId).build(); + + var postResponse = client.uploadAndGetRagDocument(inputSource, parameters); + assertNotNull(postResponse); + + var postAnnotation = postResponse.getAnnotation(); + assertNotNull(postAnnotation.getFields()); + + var documentId = postResponse.getId(); + assertNotNull(documentId); + + assertEquals("Draft", postResponse.getStatus()); + + postAnnotation.getFields().get("supplier_name").getSimpleField().setSelected(true); + postAnnotation + .getFields() + .get("supplier_name") + .getSimpleField() + .setGuidelines("I am the walrus!"); + postAnnotation.getFields().get("invoice_number").getSimpleField().setSelected(true); + postAnnotation + .getFields() + .get("invoice_number") + .getSimpleField() + .setGuidelines("koo koo katchoo!"); + + var patchAnnotationParameters = RagDocumentAnnotationParameters + .builder(documentId) + .annotation(postAnnotation) + .build(); + + var patchAnnotationResponse = client.updateRagAnnotation(patchAnnotationParameters); + assertNotNull(patchAnnotationResponse); + var patchAnnotation = patchAnnotationResponse.getAnnotation(); + assertEquals( + "I am the walrus!", + patchAnnotation.getFields().get("supplier_name").getSimpleField().getGuidelines() + ); + assertTrue(patchAnnotation.getFields().get("supplier_name").getSimpleField().isSelected()); + assertEquals( + "koo koo katchoo!", + patchAnnotation.getFields().get("invoice_number").getSimpleField().getGuidelines() + ); + assertTrue(patchAnnotation.getFields().get("invoice_number").getSimpleField().isSelected()); + + var getResponse = client.getReadyRagDocument(ExtractionRagAnnotationResponse.class, documentId); + assertNotNull(getResponse); + var getAnnotation = getResponse.getAnnotation(); + assertNotNull(getAnnotation); + + assertEquals("Draft", getResponse.getStatus()); + + assertEquals( + "I am the walrus!", + getAnnotation.getFields().get("supplier_name").getSimpleField().getGuidelines() + ); + assertTrue(getAnnotation.getFields().get("supplier_name").getSimpleField().isSelected()); + assertEquals( + "koo koo katchoo!", + getAnnotation.getFields().get("invoice_number").getSimpleField().getGuidelines() + ); + assertTrue(getAnnotation.getFields().get("invoice_number").getSimpleField().isSelected()); + + var patchStatusParameters = RagDocumentAnnotationParameters + .builder(documentId) + .status("Active") + .build(); + + var patchStatusResponse = client.updateAndGetRagAnnotation(patchStatusParameters); + assertNotNull(patchStatusResponse); + assertEquals("Active", patchStatusResponse.getStatus()); + + boolean deleteResponse = client.deleteExtractionRagDocument(documentId); + assertTrue(deleteResponse); + + assertThrows(MindeeHttpExceptionV2.class, () -> { + client.getRagDocument(ExtractionRagAnnotationResponse.class, documentId); + }); + } +} diff --git a/src/test/java/com/mindee/v2/product/extraction/RagDocumentsTest.java b/src/test/java/com/mindee/v2/product/extraction/RagDocumentsTest.java new file mode 100644 index 000000000..72df9da5e --- /dev/null +++ b/src/test/java/com/mindee/v2/product/extraction/RagDocumentsTest.java @@ -0,0 +1,195 @@ +package com.mindee.v2.product.extraction; + +import static com.mindee.TestingUtilities.getV2ProductPath; +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.assertNull; + +import com.fasterxml.jackson.core.JsonProcessingException; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.mindee.v2.parsing.LocalResponse; +import com.mindee.v2.product.extraction.ragdocuments.AnnotatedDynamicField; +import com.mindee.v2.product.extraction.ragdocuments.AnnotatedFields; +import com.mindee.v2.product.extraction.ragdocuments.AnnotatedListField; +import com.mindee.v2.product.extraction.ragdocuments.AnnotatedObjectField; +import com.mindee.v2.product.extraction.ragdocuments.AnnotatedSimpleField; +import com.mindee.v2.product.extraction.ragdocuments.ExtractionRagAnnotationResponse; +import com.mindee.v2.product.extraction.ragdocuments.RagAnnotation; +import com.mindee.v2.product.extraction.ragdocuments.params.RagDocumentAnnotationParameters; +import com.mindee.v2.product.extraction.ragdocuments.params.RagDocumentUploadParameters; +import java.io.IOException; +import java.nio.file.Files; +import java.util.ArrayList; +import java.util.Map; +import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; + +@DisplayName("MindeeV2 - Extraction Rag Documents") +public class RagDocumentsTest { + + static ObjectMapper mapper = new ObjectMapper(); + private static String expectedAnnotation; + + @BeforeAll + static void init() throws IOException { + String rawJson = Files + .readString(getV2ProductPath("extraction/rag_documents/test_annotation.json")); + expectedAnnotation = mapper.readTree(rawJson).toString(); + } + + @Test + @DisplayName("should init POST parameters") + void postParameters_mustInit() { + RagDocumentUploadParameters parameters = RagDocumentUploadParameters + .builder("invalid-model-id") + .build(); + + Map reqParams = parameters.getRequestParameters(); + assertEquals("invalid-model-id", reqParams.get("model_id")); + } + + @Test + @DisplayName("Should init PATCH parameters from an annotation instance.") + void patchParameters_mustInitFromObject() throws JsonProcessingException { + var fields = new AnnotatedFields(); + fields.put("simple", new AnnotatedDynamicField(new AnnotatedSimpleField(true, false, null))); + fields + .put( + "list", + new AnnotatedDynamicField(new AnnotatedListField(new ArrayList<>(), false, null)) + ); + fields + .put( + "object", + new AnnotatedDynamicField(new AnnotatedObjectField(new AnnotatedFields(), false, null)) + ); + + var annotation = new RagAnnotation(fields); + RagDocumentAnnotationParameters parameters = RagDocumentAnnotationParameters + .builder("invalid-document-id") + .status("Active") + .annotation(annotation) + .build(); + + var reqParams = parameters.getRequestParameters(); + assertEquals("invalid-document-id", parameters.getDocumentId()); + assertEquals("Active", reqParams.get("status")); + assertEquals(expectedAnnotation, mapper.writeValueAsString(reqParams.get("annotation"))); + } + + @Test + @DisplayName("should load a POST response from a JSON string") + void ragDocumentsPost_mustHaveValidProperties() throws IOException { + var response = loadResponse("extraction/rag_documents/post_response.json"); + + assertNotNull(response); + assertEquals("cc831599-c545-48b7-aa27-6d7ccd5b8d32", response.getId()); + assertEquals("Processing", response.getStatus()); + assertNull(response.getAnnotation()); + } + + @Test + @DisplayName("should load a GET response from a JSON string") + void ragDocumentsGetDraft_mustHaveValidProperties() throws IOException { + var response = loadResponse("extraction/rag_documents/get_response_draft.json"); + + assertNotNull(response); + assertEquals("cc831599-c545-48b7-aa27-6d7ccd5b8d32", response.getId()); + assertEquals("Draft", response.getStatus()); + assertNotNull(response.getAnnotation()); + + var fields = response.getAnnotation().getFields(); + assertNotNull(fields); + + // null simple field + var tipField = fields.getSimpleField("tip"); + assertNotNull(tipField); + assertFalse(tipField.isSelected()); + assertNull(tipField.getGuidelines()); + assertNull(tipField.getValue()); + + // filled simple field + var dateField = fields.getSimpleField("date"); + assertNotNull(dateField); + assertFalse(dateField.isSelected()); + assertNull(dateField.getGuidelines()); + assertEquals("2019-11-02", dateField.getValue()); + + // filled object field + var localeField = fields.getObjectField("locale"); + assertNotNull(localeField); + assertFalse(localeField.isSelected()); + assertNull(localeField.getGuidelines()); + assertNotNull(localeField.getFields()); + assertEquals(3, localeField.getFields().size()); + assertEquals("US", localeField.getFields().getSimpleField("country").getValue()); + assertEquals("USD", localeField.getFields().getSimpleField("currency").getValue()); + assertNull(localeField.getFields().getSimpleField("language").getValue()); + + // list of simple fields + var referenceNumbersField = fields.getListField("reference_numbers"); + assertNotNull(referenceNumbersField); + assertFalse(referenceNumbersField.isSelected()); + assertNull(referenceNumbersField.getGuidelines()); + assertNotNull(referenceNumbersField.getSimpleItems()); + assertEquals(1, referenceNumbersField.getSimpleItems().size()); + assertEquals("2412/2019", referenceNumbersField.getSimpleItems().get(0).getValue()); + + // list of object fields + var lineItemsField = fields.getListField("line_items"); + assertNotNull(lineItemsField); + assertFalse(lineItemsField.isSelected()); + assertNull(lineItemsField.getGuidelines()); + assertNotNull(lineItemsField.getObjectItems()); + assertEquals(3, lineItemsField.getObjectItems().size()); + + var lineItem0 = lineItemsField.getObjectItems().get(0); + assertNotNull(lineItem0.getFields()); + assertEquals(8, lineItem0.getFields().size()); + assertEquals( + "Front and rear brake cables", + lineItem0.getFields().getSimpleField("description").getValue() + ); + assertEquals(1.0, lineItem0.getFields().getSimpleField("quantity").getValue()); + assertEquals(100.0, lineItem0.getFields().getSimpleField("unit_price").getValue()); + assertEquals(100.0, lineItem0.getFields().getSimpleField("total_price").getValue()); + assertNull(lineItem0.getFields().getSimpleField("tax_rate").getValue()); + assertNull(lineItem0.getFields().getSimpleField("tax_amount").getValue()); + assertNull(lineItem0.getFields().getSimpleField("product_code").getValue()); + assertNull(lineItem0.getFields().getSimpleField("unit_measure").getValue()); + + var lineItem1 = lineItemsField.getObjectItems().get(1); + assertNotNull(lineItem1.getFields()); + assertEquals(8, lineItem1.getFields().size()); + assertEquals( + "New set of pedal arms", + lineItem1.getFields().getSimpleField("description").getValue() + ); + assertEquals(2.0, lineItem1.getFields().getSimpleField("quantity").getValue()); + assertEquals(25.0, lineItem1.getFields().getSimpleField("unit_price").getValue()); + assertEquals(50.0, lineItem1.getFields().getSimpleField("total_price").getValue()); + assertNull(lineItem1.getFields().getSimpleField("tax_rate").getValue()); + assertNull(lineItem1.getFields().getSimpleField("tax_amount").getValue()); + assertNull(lineItem1.getFields().getSimpleField("product_code").getValue()); + assertNull(lineItem1.getFields().getSimpleField("unit_measure").getValue()); + + var lineItem2 = lineItemsField.getObjectItems().get(2); + assertNotNull(lineItem2.getFields()); + assertEquals(8, lineItem2.getFields().size()); + assertEquals("Labor 3hrs", lineItem2.getFields().getSimpleField("description").getValue()); + assertEquals(3.0, lineItem2.getFields().getSimpleField("quantity").getValue()); + assertEquals(15.0, lineItem2.getFields().getSimpleField("unit_price").getValue()); + assertEquals(45.0, lineItem2.getFields().getSimpleField("total_price").getValue()); + assertNull(lineItem2.getFields().getSimpleField("tax_rate").getValue()); + assertNull(lineItem2.getFields().getSimpleField("tax_amount").getValue()); + assertNull(lineItem2.getFields().getSimpleField("product_code").getValue()); + assertNull(lineItem2.getFields().getSimpleField("unit_measure").getValue()); + } + + private ExtractionRagAnnotationResponse loadResponse(String filePath) throws IOException { + var localResponse = new LocalResponse(getV2ProductPath(filePath)); + return localResponse.deserializeResponse(ExtractionRagAnnotationResponse.class); + } +} diff --git a/src/test/resources b/src/test/resources index 8b8423d23..bf5893c5c 160000 --- a/src/test/resources +++ b/src/test/resources @@ -1 +1 @@ -Subproject commit 8b8423d239a360f7a8ac72d77e53fb644a76e345 +Subproject commit bf5893c5c631f37f01419b4a9c71c8a5ce036939