diff --git a/.github/workflows/_test-units.yml b/.github/workflows/_test-units.yml index 545df7e7..770f17e0 100644 --- a/.github/workflows/_test-units.yml +++ b/.github/workflows/_test-units.yml @@ -59,9 +59,6 @@ jobs: - "windows-2022" python-version: - "3.10" - - "3.11" - - "3.12" - - "3.13" - "3.14" runs-on: ${{ matrix.os }} steps: diff --git a/mindee/v2/client.py b/mindee/v2/client.py index b5b89fc1..d25351c1 100644 --- a/mindee/v2/client.py +++ b/mindee/v2/client.py @@ -1,6 +1,5 @@ import warnings from time import sleep -from typing import TypeVar import httpx @@ -11,20 +10,25 @@ from mindee.logger import logger from mindee.mindee_http.cancellation_token import CancellationToken from mindee.parsing.common.common_response import CommonStatus +from mindee.v2.client_options.base_annotation_parameters import BaseAnnotationParameters from mindee.v2.client_options.base_product_parameters import BaseProductParameters +from mindee.v2.client_options.base_rag_document_upload_parameters import ( + BaseRagDocumentUploadParameters, +) from mindee.v2.client_options.base_search_parameters import ( BaseSearchParameters, TypeSearchResponse, ) from mindee.v2.mindee_http.mindee_api_v2 import MindeeAPIV2 -from mindee.v2.parsing.inference.base_inference_response import BaseInferenceResponse +from mindee.v2.parsing.base_rag_annotation_response import ( + TypeRagAnnotationResponse, +) +from mindee.v2.parsing.inference.base_inference_response import ( + TypeBaseInferenceResponse, +) from mindee.v2.parsing.job.job_response import JobResponse from mindee.v2.parsing.search.search_response import SearchResponse -TypeBaseInferenceResponse = TypeVar( - "TypeBaseInferenceResponse", bound=BaseInferenceResponse -) - class Client: """ @@ -169,6 +173,143 @@ def enqueue_and_get_result( raise MindeeError(f"Couldn't retrieve document after {try_counter + 1} tries.") + def upload_rag_document( + self, + input_source: LocalInputSource, + parameters: BaseRagDocumentUploadParameters[TypeRagAnnotationResponse], + ) -> TypeRagAnnotationResponse: + """ + Not recommended for general use, prefer ``upload_and_get_rag_document``. + You will need to poll until the document is ready for use. + Add a document to the RAG database. + """ + return self.mindee_api.req_post_rag_document(input_source, parameters) + + def upload_and_get_rag_document( + self, + input_source: LocalInputSource, + parameters: BaseRagDocumentUploadParameters[TypeRagAnnotationResponse], + polling_options: PollingOptions | None = None, + cancellation_token: CancellationToken | None = None, + ) -> TypeRagAnnotationResponse: + """ + Add a document to the RAG database and return the initial annotation. + """ + initial_response = self.upload_rag_document(input_source, parameters) + if initial_response.status != "Processing": + return initial_response + if polling_options is None: + polling_options = PollingOptions() + return self._poll_for_rag_document( + initial_response, polling_options, cancellation_token + ) + + def get_rag_document( + self, response_type: type[TypeRagAnnotationResponse], document_id: str + ) -> TypeRagAnnotationResponse: + """ + Not recommended for general use, prefer ``get_ready_rag_document``. + You will need to poll until the document is ready for use. + Get a document's info and annotations from the RAG database. + """ + return self.mindee_api.req_get_rag_annotation(response_type, document_id) + + def get_ready_rag_document( + self, + response_type: type[TypeRagAnnotationResponse], + document_id: str, + polling_options: PollingOptions | None = None, + cancellation_token: CancellationToken | None = None, + ): + """ + Get a document's info and annotations from the RAG database. + """ + initial_response = self.get_rag_document(response_type, document_id) + if initial_response.status != "Processing": + return initial_response + if polling_options is None: + polling_options = PollingOptions() + return self._poll_for_rag_document( + initial_response, polling_options, cancellation_token + ) + + def update_rag_annotations( + self, parameters: BaseAnnotationParameters[TypeRagAnnotationResponse] + ) -> TypeRagAnnotationResponse: + """ + Not recommended for general use, prefer ``update_and_get_rag_annotations``. + You will need to poll until the document is ready for use. + Update a document's annotations in the RAG database. + """ + return self.mindee_api.req_patch_rag_annotation(parameters) + + def update_and_get_rag_annotations( + self, + parameters: BaseAnnotationParameters[TypeRagAnnotationResponse], + polling_options: PollingOptions | None = None, + cancellation_token: CancellationToken | None = None, + ) -> TypeRagAnnotationResponse: + """ + Update a document's annotations in the RAG database. + """ + initial_response = self.update_rag_annotations(parameters) + if initial_response.status != "Processing": + return initial_response + if polling_options is None: + polling_options = PollingOptions() + return self._poll_for_rag_document( + initial_response, polling_options, cancellation_token + ) + + def delete_extraction_rag_document(self, document_id: str) -> bool: + """ + Delete a document from the RAG database. + For extraction models only. + """ + return self.mindee_api.req_delete_extraction_rag_document(document_id) + + def _poll_for_rag_document( + self, + initial_response: TypeRagAnnotationResponse, + polling_options: PollingOptions, + cancellation_token: CancellationToken | None = None, + ) -> TypeRagAnnotationResponse: + """ + Poll until the document is finished processing or the max number of attempts is reached. + """ + logger.info("Polling for RAG document ID: %s", initial_response.id) + polling_options.validate_settings() + max_retries = polling_options.max_retries + 1 + + logger.debug( + "Waiting %s seconds before attempting to retrieve the result...", + polling_options.initial_delay_sec, + ) + + if cancellation_token and cancellation_token.is_canceled: + raise MindeeError("Request canceled through cancellation token.") + + sleep(polling_options.initial_delay_sec) + document_id = initial_response.id + retry_count = 1 + + while retry_count < max_retries: + if cancellation_token and cancellation_token.is_canceled: + raise MindeeError("Request canceled through cancellation token.") + logger.info("Poll attempt %s of %s", retry_count, max_retries) + + response = self.get_rag_document(type(initial_response), document_id) + retry_count += 1 + + if response.status == "Processing": + sleep(polling_options.delay_sec) + continue + if response.status == "Failed": + raise MindeeError("Job failed without an error payload.") + return response + + raise MindeeError(f"RAG polling not complete after {retry_count - 1} attempts.") + def search( self, params: BaseSearchParameters[TypeSearchResponse] ) -> TypeSearchResponse: @@ -183,7 +324,7 @@ def search_models( self, name: str | None = None, model_type: str | None = None ) -> SearchResponse: """ - Deprecated. Use `search` instead. + Deprecated: use `search` instead. """ warnings.warn( "search_models is deprecated, use search instead.", diff --git a/mindee/v2/client_options/base_annotation_parameters.py b/mindee/v2/client_options/base_annotation_parameters.py new file mode 100644 index 00000000..75ceed64 --- /dev/null +++ b/mindee/v2/client_options/base_annotation_parameters.py @@ -0,0 +1,26 @@ +from abc import ABC, abstractmethod +from dataclasses import dataclass +from typing import Generic + +from mindee.parsing.common.string_dict import StringDict +from mindee.v2.parsing.base_rag_annotation_response import TypeRagAnnotationResponse + + +@dataclass(kw_only=True) +class BaseAnnotationParameters(ABC, Generic[TypeRagAnnotationResponse]): + """Base parameters for document annotations.""" + + # Note: DocumentId is included in the request URL path, it is not a parameter. + document_id: str + """UID of the annotated document.""" + + _response_class: type[TypeRagAnnotationResponse] + """Response class for the annotation.""" + + @abstractmethod + def get_request_parameters(self) -> dict[str, str | StringDict]: + """Gets the request parameters for the upload request.""" + + def get_response_class(self) -> type[TypeRagAnnotationResponse]: + """Gets the response class for the search.""" + return self._response_class diff --git a/mindee/v2/client_options/base_product_parameters.py b/mindee/v2/client_options/base_product_parameters.py index 2a00f08f..ac030dd7 100644 --- a/mindee/v2/client_options/base_product_parameters.py +++ b/mindee/v2/client_options/base_product_parameters.py @@ -29,7 +29,7 @@ class BaseProductParameters(ABC): """Options for polling. Set only if having timeout issues.""" close_file: bool = True - """Whether to close the file after product.""" + """Whether to close the file after uploading. Default: True.""" _slug: ClassVar[str] """Slug of the product.""" diff --git a/mindee/v2/client_options/base_rag_document_upload_parameters.py b/mindee/v2/client_options/base_rag_document_upload_parameters.py new file mode 100644 index 00000000..4377647a --- /dev/null +++ b/mindee/v2/client_options/base_rag_document_upload_parameters.py @@ -0,0 +1,27 @@ +from abc import ABC +from dataclasses import dataclass +from typing import Generic + +from mindee.v2.parsing.base_rag_annotation_response import TypeRagAnnotationResponse + + +@dataclass(kw_only=True) +class BaseRagDocumentUploadParameters(ABC, Generic[TypeRagAnnotationResponse]): + """Base parameters for document upload operations.""" + + model_id: str + """UUID of the model that the uploaded RAG document is linked to.""" + + close_file: bool = True + """Whether to close the file after uploading. Default: True.""" + + _response_class: type[TypeRagAnnotationResponse] + """Response class for the annotation.""" + + def get_request_parameters(self) -> dict[str, str]: + """Gets the request parameters for the upload request.""" + return {"model_id": self.model_id} + + def get_response_class(self) -> type[TypeRagAnnotationResponse]: + """Gets the response class for the search.""" + return self._response_class diff --git a/mindee/v2/client_options/base_search_parameters.py b/mindee/v2/client_options/base_search_parameters.py index e0919888..25079558 100644 --- a/mindee/v2/client_options/base_search_parameters.py +++ b/mindee/v2/client_options/base_search_parameters.py @@ -1,10 +1,8 @@ from abc import ABC from dataclasses import dataclass -from typing import ClassVar, Generic, TypeVar +from typing import ClassVar, Generic -from mindee.v2.parsing.search.base_search_response import BaseSearchResponse - -TypeSearchResponse = TypeVar("TypeSearchResponse", bound=BaseSearchResponse) +from mindee.v2.parsing.search.base_search_response import TypeSearchResponse @dataclass(kw_only=True) diff --git a/mindee/v2/error/mindee_http_error_v2.py b/mindee/v2/error/mindee_http_error_v2.py index 0cc7af8a..8dd80d51 100644 --- a/mindee/v2/error/mindee_http_error_v2.py +++ b/mindee/v2/error/mindee_http_error_v2.py @@ -1,5 +1,6 @@ import json +from mindee.logger import logger from mindee.parsing.common.string_dict import StringDict from mindee.v2.parsing.error import ErrorItem, ErrorResponse, IErrorResponse @@ -48,4 +49,6 @@ def handle_error_v2(raw_response: StringDict) -> None: """ if "status" not in raw_response or "detail" not in raw_response: raise MindeeHTTPUnknownErrorV2(json.dumps(raw_response, indent=2)) - raise MindeeHTTPErrorV2(ErrorResponse(raw_response)) + error_response = ErrorResponse(raw_response) + logger.error("\n%s", error_response) + raise MindeeHTTPErrorV2(error_response) diff --git a/mindee/v2/mindee_http/mindee_api_v2.py b/mindee/v2/mindee_http/mindee_api_v2.py index 0bce8226..6f2ed4d6 100644 --- a/mindee/v2/mindee_http/mindee_api_v2.py +++ b/mindee/v2/mindee_http/mindee_api_v2.py @@ -1,7 +1,6 @@ import json import os from collections.abc import Callable -from typing import TypeVar import httpx @@ -12,7 +11,11 @@ from mindee.mindee_http.settings_mixin import SettingsMixin from mindee.parsing.common.string_dict import StringDict from mindee.v1.mindee_http.base_settings import USER_AGENT +from mindee.v2.client_options.base_annotation_parameters import BaseAnnotationParameters from mindee.v2.client_options.base_product_parameters import BaseProductParameters +from mindee.v2.client_options.base_rag_document_upload_parameters import ( + BaseRagDocumentUploadParameters, +) from mindee.v2.client_options.base_search_parameters import ( BaseSearchParameters, TypeSearchResponse, @@ -22,7 +25,10 @@ MindeeHTTPUnknownErrorV2, handle_error_v2, ) -from mindee.v2.parsing import BaseInferenceResponse +from mindee.v2.parsing.base_rag_annotation_response import TypeRagAnnotationResponse +from mindee.v2.parsing.inference.base_inference_response import ( + TypeBaseInferenceResponse, +) from mindee.v2.parsing.job.job_response import JobResponse from mindee.v2.parsing.search.search_response import SearchResponse @@ -35,8 +41,6 @@ REQUEST_TIMEOUT_ENV_NAME = "MINDEE_REQUEST_TIMEOUT" TIMEOUT_DEFAULT = 120 -ResponseT = TypeVar("ResponseT", bound=BaseInferenceResponse) - class MindeeAPIV2(SettingsMixin): """Settings class relating to API V2 requests.""" @@ -45,7 +49,7 @@ class MindeeAPIV2(SettingsMixin): """Root of the URL to use for polling.""" api_key: str | None """API Key for the client.""" - http_client: httpx.Client | None + _http_client: httpx.Client | None """HTTP client for making requests.""" request_timeout: float @@ -65,7 +69,7 @@ def __init__(self, api_key: str | None, http_client: httpx.Client | None = None) f"'{API_KEY_V2_ENV_NAME}' environment variable." ) self.url_root = f"{self.base_url.rstrip('/')}" - self.http_client = http_client + self._http_client = http_client self.request_timeout = float( os.environ.get(REQUEST_TIMEOUT_ENV_NAME, TIMEOUT_DEFAULT) ) @@ -113,13 +117,7 @@ def req_post_product_enqueue( elif isinstance(input_source, URLInputSource): data["url"] = input_source.url - post_caller: Callable - if self.http_client is None or self.http_client.is_closed: - post_caller = httpx.post - else: - post_caller = self.http_client.post - - response = post_caller( + response = self.post_caller( url, headers=self.base_headers, data=data, @@ -138,13 +136,7 @@ def req_get_job_by_id(self, job_id: str) -> JobResponse: :param job_id: Job ID, returned by the enqueue request. """ - get_caller: Callable - if self.http_client is None or self.http_client.is_closed: - get_caller = httpx.get - else: - get_caller = self.http_client.get - - response = get_caller( + response = self.get_caller( url=f"{self.url_root}/v2/jobs/{job_id}", headers=self.base_headers, follow_redirects=False, @@ -156,22 +148,16 @@ def req_get_job_by_id(self, job_id: str) -> JobResponse: return JobResponse(dict_response) def req_get_product_result_by_url( - self, response_type: type[ResponseT], url: str - ) -> ResponseT: + self, response_class: type[TypeBaseInferenceResponse], url: str + ) -> TypeBaseInferenceResponse: """ Get the result of an inference that was previously enqueued. :param url: URL to use for the request. - :param response_type: Type of the response to return. + :param response_class: Type of the response to return. :return: Response object from the request. """ - get_caller: Callable - if self.http_client is None or self.http_client.is_closed: - get_caller = httpx.get - else: - get_caller = self.http_client.get - - response = get_caller( + response = self.get_caller( url=url, headers=self.base_headers, follow_redirects=False, @@ -180,20 +166,20 @@ def req_get_product_result_by_url( dict_response = self._response_json(response) if not is_valid_sync_response(response): handle_error_v2(dict_response) - return response_type(dict_response) + return response_class(dict_response) def req_get_product_result_by_id( - self, response_type: type[ResponseT], inference_id: str - ) -> ResponseT: + self, response_class: type[TypeBaseInferenceResponse], inference_id: str + ) -> TypeBaseInferenceResponse: """ Sends a request matching a given queue_id. Returns either a Job or a Document. :param inference_id: Inference ID, returned by the job request. - :param response_type: Type of the response to return. + :param response_class: Type of the response to return. """ - slug = response_type.get_result_slug() + slug = response_class.get_product_slug() return self.req_get_product_result_by_url( - response_type=response_type, + response_class=response_class, url=f"{self.url_root}/v2/products/{slug}/results/{inference_id}", ) @@ -205,14 +191,9 @@ def req_search( :param params: Search parameters :return: A search response containing the matching resources """ - get_caller: Callable - if self.http_client is None or self.http_client.is_closed: - get_caller = httpx.get - else: - get_caller = self.http_client.get slug = params.get_slug() response_class = params.get_response_class() - response = get_caller( + response = self.get_caller( url=f"{self.url_root}/v2/search/{slug}", headers=self.base_headers, params=params.get_request_parameters(), @@ -228,20 +209,15 @@ def req_get_search_models( self, name: str | None, model_type: str | None ) -> SearchResponse: """ - Deprecated. Use `req_search` instead. + Deprecated: use `req_search` instead. """ - get_caller: Callable - if self.http_client is None or self.http_client.is_closed: - get_caller = httpx.get - else: - get_caller = self.http_client.get params = {} if name: params["name"] = name if model_type: params["model_type"] = model_type - response = get_caller( + response = self.get_caller( url=f"{self.url_root}/v2/search/models", headers=self.base_headers, params=params, @@ -253,6 +229,102 @@ def req_get_search_models( handle_error_v2(dict_response) return SearchResponse(dict_response) + def req_post_rag_document( + self, + input_source: LocalInputSource, + params: BaseRagDocumentUploadParameters[TypeRagAnnotationResponse], + ) -> TypeRagAnnotationResponse: + """Add a document to the RAG database.""" + response_class = params.get_response_class() + slug = response_class.get_product_slug() + response = self.post_caller( + url=f"{self.url_root}/v2/products/{slug}/rag-documents", + headers=self.base_headers, + files={"file": input_source.read_contents(params.close_file)}, + data=params.get_request_parameters(), + follow_redirects=False, + timeout=self.request_timeout, + ) + dict_response = self._response_json(response) + if not is_valid_sync_response(response): + handle_error_v2(dict_response) + return response_class(dict_response) + + def req_get_rag_annotation( + self, response_class: type[TypeRagAnnotationResponse], document_id: str + ) -> TypeRagAnnotationResponse: + """Get a document's info and annotations from the RAG database.""" + slug = response_class.get_product_slug() + response = self.get_caller( + url=f"{self.url_root}/v2/products/{slug}/rag-documents/{document_id}", + headers=self.base_headers, + follow_redirects=False, + timeout=self.request_timeout, + ) + dict_response = self._response_json(response) + if not is_valid_sync_response(response): + handle_error_v2(dict_response) + return response_class(dict_response) + + def req_patch_rag_annotation( + self, params: BaseAnnotationParameters[TypeRagAnnotationResponse] + ) -> TypeRagAnnotationResponse: + """Update a document's annotations in the RAG database.""" + response_class = params.get_response_class() + slug = response_class.get_product_slug() + response = self.patch_caller( + url=f"{self.url_root}/v2/products/{slug}/rag-documents/{params.document_id}", + headers=self.base_headers, + timeout=self.request_timeout, + follow_redirects=False, + json=params.get_request_parameters(), + ) + dict_response = self._response_json(response) + if not is_valid_sync_response(response): + handle_error_v2(dict_response) + return response_class(dict_response) + + def req_delete_extraction_rag_document(self, document_id: str) -> bool: + """ + Deletes a document from the RAG database. + For extraction models only. + """ + response = self.delete_caller( + url=f"{self.url_root}/v2/products/extraction/rag-documents/{document_id}", + headers=self.base_headers, + timeout=self.request_timeout, + follow_redirects=False, + ) + return response.is_success + + @property + def get_caller(self) -> Callable: + """Caller for GET requests.""" + if self._http_client is None or self._http_client.is_closed: + return httpx.get + return self._http_client.get + + @property + def post_caller(self) -> Callable: + """Caller for POST requests.""" + if self._http_client is None or self._http_client.is_closed: + return httpx.post + return self._http_client.post + + @property + def patch_caller(self) -> Callable: + """Caller for PATCH requests.""" + if self._http_client is None or self._http_client.is_closed: + return httpx.patch + return self._http_client.patch + + @property + def delete_caller(self) -> Callable: + """Caller for DELETE requests.""" + if self._http_client is None or self._http_client.is_closed: + return httpx.delete + return self._http_client.delete + @staticmethod def _response_json(response: httpx.Response) -> StringDict: try: @@ -265,11 +337,11 @@ def _response_json(response: httpx.Response) -> StringDict: def close(self) -> None: """Closes the underlying HTTP client.""" - if self.http_client and not self.http_client.is_closed: - self.http_client.close() + if self._http_client and not self._http_client.is_closed: + self._http_client.close() def __enter__(self): - self.http_client = httpx.Client() + self._http_client = httpx.Client() return self def __exit__(self, exc_type, exc_val, exc_tb): @@ -277,8 +349,8 @@ def __exit__(self, exc_type, exc_val, exc_tb): def delete_http_client(self): """Delete the underlying HTTP client.""" - httpx_client = getattr(self, "http_client", None) - if httpx_client and not self.http_client.is_closed: + httpx_client = getattr(self, "_http_client", None) + if httpx_client and not self._http_client.is_closed: logger.info("Force-closing unclosed Mindee Client (V2) %s.", str(self)) self.close() diff --git a/mindee/v2/parsing/base_rag_annotation_response.py b/mindee/v2/parsing/base_rag_annotation_response.py new file mode 100644 index 00000000..219006dd --- /dev/null +++ b/mindee/v2/parsing/base_rag_annotation_response.py @@ -0,0 +1,40 @@ +from datetime import datetime +from typing import ClassVar, TypeVar + +from mindee.parsing.common import CommonResponse +from mindee.parsing.common.string_dict import StringDict + + +class BaseRagAnnotationResponse(CommonResponse): + """Base class for all RAG document responses from the V2 API.""" + + id: str + """Unique identifier of the RAG document.""" + filename: str + """Original filename of the uploaded document.""" + created_at: datetime + """Date and time of the document creation.""" + status: str + """Current status of the RAG document.""" + + _slug: ClassVar[str] + """Slug of the product.""" + + def __init__(self, raw_response: StringDict) -> None: + super().__init__(raw_response) + self.id = raw_response["id"] + self.filename = raw_response["filename"] + self.created_at = datetime.fromisoformat( + raw_response["created_at"].replace("Z", "+00:00") + ) + self.status = raw_response["status"] + + @classmethod + def get_product_slug(cls) -> str: + """Get the product's slug.""" + return cls._slug + + +TypeRagAnnotationResponse = TypeVar( + "TypeRagAnnotationResponse", bound=BaseRagAnnotationResponse +) diff --git a/mindee/v2/parsing/inference/base_inference_response.py b/mindee/v2/parsing/inference/base_inference_response.py index 1c46fc28..d9d6027d 100644 --- a/mindee/v2/parsing/inference/base_inference_response.py +++ b/mindee/v2/parsing/inference/base_inference_response.py @@ -1,5 +1,5 @@ from abc import ABC -from typing import ClassVar +from typing import ClassVar, TypeVar from mindee.parsing.common.common_response import CommonResponse from mindee.v2.parsing.inference.base_inference import BaseInference @@ -18,6 +18,11 @@ def __str__(self) -> str: return str(self.inference) @classmethod - def get_result_slug(cls) -> str: - """Getter for the inference slug.""" + def get_product_slug(cls) -> str: + """Get the product's slug.""" return cls._slug + + +TypeBaseInferenceResponse = TypeVar( + "TypeBaseInferenceResponse", bound=BaseInferenceResponse +) diff --git a/mindee/v2/parsing/inference/field/base_field.py b/mindee/v2/parsing/inference/field/base_field.py index b8a12450..c9984cf5 100644 --- a/mindee/v2/parsing/inference/field/base_field.py +++ b/mindee/v2/parsing/inference/field/base_field.py @@ -1,4 +1,5 @@ from enum import Enum +from typing import TYPE_CHECKING, ClassVar, TypeAlias, Union from mindee.parsing.common.string_dict import StringDict from mindee.v2.parsing.inference.field.field_confidence import FieldConfidence @@ -13,18 +14,35 @@ class FieldType(str, Enum): SIMPLE = "SimpleField" +if TYPE_CHECKING: + from mindee.v2.parsing.inference.field.list_field import ListField + from mindee.v2.parsing.inference.field.object_field import ObjectField + from mindee.v2.parsing.inference.field.simple_field import SimpleField + + +ResultFieldsType: TypeAlias = Union["SimpleField", "ObjectField", "ListField"] + + class BaseField: - """Field with base information.""" + """Base class for V2 fields.""" - field_type: FieldType - _indent_level: int locations: list[FieldLocation] + """List of the location candidates for the value.""" + confidence: FieldConfidence | None + """Confidence associated with the field.""" + + _indent_level: int + """For pretty printing.""" + + _field_type: FieldType + _registry: ClassVar[dict[str, type[ResultFieldsType]]] = {} def __init__( self, field_type: FieldType, raw_response: StringDict, indent_level: int = 0 ) -> None: - self.field_type = field_type + + self._field_type = field_type self._indent_level = indent_level self.confidence = None @@ -41,6 +59,32 @@ def __init__( for location in raw_response["locations"]: self.locations.append(FieldLocation(location)) + @property + def field_type(self) -> FieldType: + """The field type.""" + return self._field_type + + @classmethod + def register(cls, discriminator_key: str): + """Class decorator: subclasses declare which JSON key identifies them.""" + + def decorator(subclass): + cls._registry[discriminator_key] = subclass + return subclass + + return decorator + + @classmethod + def build(cls, raw_response: dict, indent_level: int) -> ResultFieldsType: + """Build an instance of the appropriate subclass.""" + + if not isinstance(raw_response, dict): + raise ValueError("Field must be a dict") + for key, subclass in cls._registry.items(): + if key in raw_response: + return subclass(raw_response, indent_level) + raise ValueError("Invalid structure for field") + def multi_str(self) -> str: """String representation of the field in a list.""" return str(self) diff --git a/mindee/v2/parsing/inference/field/factory.py b/mindee/v2/parsing/inference/field/factory.py deleted file mode 100644 index f980fa5e..00000000 --- a/mindee/v2/parsing/inference/field/factory.py +++ /dev/null @@ -1,17 +0,0 @@ -from mindee.parsing.common.string_dict import StringDict -from mindee.v2.error.mindee_api_v2_error import MindeeAPIV2Error -from mindee.v2.parsing.inference.field.list_field import ListField -from mindee.v2.parsing.inference.field.object_field import ObjectField -from mindee.v2.parsing.inference.field.simple_field import SimpleField - - -def create_field(raw_response: StringDict, indent_level: int = 0): - """The central parser function to be injected down the tree.""" - if "value" in raw_response: - return SimpleField(raw_response, indent_level) - if "items" in raw_response: - return ListField(raw_response, create_field, indent_level) - if "fields" in raw_response: - return ObjectField(raw_response, create_field, indent_level) - - raise MindeeAPIV2Error(f"Unrecognized field type in {raw_response}.") diff --git a/mindee/v2/parsing/inference/field/inference_fields.py b/mindee/v2/parsing/inference/field/inference_fields.py index 9b73043c..51636819 100644 --- a/mindee/v2/parsing/inference/field/inference_fields.py +++ b/mindee/v2/parsing/inference/field/inference_fields.py @@ -1,8 +1,11 @@ -from collections.abc import Callable from typing import TYPE_CHECKING, cast from mindee.parsing.common.string_dict import StringDict -from mindee.v2.parsing.inference.field.base_field import BaseField, FieldType +from mindee.v2.parsing.inference.field.base_field import ( + BaseField, + FieldType, + ResultFieldsType, +) if TYPE_CHECKING: from mindee.v2.parsing.inference.field.list_field import ListField @@ -10,18 +13,17 @@ from mindee.v2.parsing.inference.field.simple_field import SimpleField -class InferenceFields(dict[str, BaseField]): +class InferenceFields(dict[str, ResultFieldsType]): """Inference fields dict.""" def __init__( self, raw_response: StringDict, - parser_func: Callable[[StringDict, int], BaseField], indent_level: int = 0, ) -> None: super().__init__() for key, value in raw_response.items(): - self[key] = parser_func(value, indent_level) + self[key] = BaseField.build(value, indent_level) def __getattr__(self, item): try: diff --git a/mindee/v2/parsing/inference/field/list_field.py b/mindee/v2/parsing/inference/field/list_field.py index af844aae..79969f0e 100644 --- a/mindee/v2/parsing/inference/field/list_field.py +++ b/mindee/v2/parsing/inference/field/list_field.py @@ -1,11 +1,10 @@ -from collections.abc import Callable - from mindee.parsing.common.string_dict import StringDict from mindee.v2.parsing.inference.field.base_field import BaseField, FieldType from mindee.v2.parsing.inference.field.object_field import ObjectField from mindee.v2.parsing.inference.field.simple_field import SimpleField +@BaseField.register("items") class ListField(BaseField): """List field containing multiple fields.""" @@ -15,14 +14,13 @@ class ListField(BaseField): def __init__( self, raw_response: StringDict, - parser_func: Callable[[StringDict, int], BaseField], indent_level: int = 0, ): super().__init__(FieldType.LIST, raw_response, indent_level) self.items = [] for item in raw_response["items"]: - self.items.append(parser_func(item, indent_level)) + self.items.append(BaseField.build(item, indent_level)) @property def simple_items(self) -> list[SimpleField]: diff --git a/mindee/v2/parsing/inference/field/object_field.py b/mindee/v2/parsing/inference/field/object_field.py index fb0d4514..e14e088a 100644 --- a/mindee/v2/parsing/inference/field/object_field.py +++ b/mindee/v2/parsing/inference/field/object_field.py @@ -1,4 +1,3 @@ -from collections.abc import Callable from typing import TYPE_CHECKING, cast from mindee.parsing.common.string_dict import StringDict @@ -10,19 +9,18 @@ from mindee.v2.parsing.inference.field.simple_field import SimpleField +@BaseField.register("fields") class ObjectField(BaseField): """Object field containing multiple fields.""" fields: InferenceFields """Fields contained in the object.""" - def __init__( - self, raw_response: StringDict, parser_func: Callable, indent_level: int = 0 - ): + def __init__(self, raw_response: StringDict, indent_level: int = 0): super().__init__(FieldType.OBJECT, raw_response, indent_level) inner_fields = raw_response.get("fields", raw_response) - self.fields = InferenceFields(inner_fields, parser_func, self._indent_level + 1) + self.fields = InferenceFields(inner_fields, self._indent_level + 1) def single_str(self) -> str: """String representation of a single object field.""" diff --git a/mindee/v2/parsing/inference/field/simple_field.py b/mindee/v2/parsing/inference/field/simple_field.py index f6e6c809..f2b5fd97 100644 --- a/mindee/v2/parsing/inference/field/simple_field.py +++ b/mindee/v2/parsing/inference/field/simple_field.py @@ -2,6 +2,7 @@ from mindee.v2.parsing.inference.field.base_field import BaseField, FieldType +@BaseField.register("value") class SimpleField(BaseField): """Simple field containing a single value.""" diff --git a/mindee/v2/parsing/search/base_search_response.py b/mindee/v2/parsing/search/base_search_response.py index f4446943..81498d2c 100644 --- a/mindee/v2/parsing/search/base_search_response.py +++ b/mindee/v2/parsing/search/base_search_response.py @@ -1,4 +1,5 @@ from abc import ABC, abstractmethod +from typing import TypeVar from mindee.parsing.common import StringDict from mindee.parsing.common.common_response import CommonResponse @@ -26,3 +27,6 @@ def __str__(self) -> str: lines: list[str] = self.body_lines() lines += ["Pagination Metadata", "###################", str(self.pagination)] return "\n".join(lines) + + +TypeSearchResponse = TypeVar("TypeSearchResponse", bound=BaseSearchResponse) diff --git a/mindee/v2/product/extraction/params/string_data_class.py b/mindee/v2/parsing/string_data_class.py similarity index 73% rename from mindee/v2/product/extraction/params/string_data_class.py rename to mindee/v2/parsing/string_data_class.py index b47ab4c1..1f4342de 100644 --- a/mindee/v2/product/extraction/params/string_data_class.py +++ b/mindee/v2/parsing/string_data_class.py @@ -13,5 +13,8 @@ def _no_none_values(items) -> dict: def __str__(self) -> str: return json.dumps( - asdict(self, dict_factory=self._no_none_values), indent=None, sort_keys=True + asdict(self, dict_factory=self._no_none_values), + indent=None, + sort_keys=True, + separators=(",", ":"), ) diff --git a/mindee/v2/product/extraction/__init__.py b/mindee/v2/product/extraction/__init__.py index 5a7f41b2..b2a95116 100644 --- a/mindee/v2/product/extraction/__init__.py +++ b/mindee/v2/product/extraction/__init__.py @@ -7,7 +7,6 @@ from mindee.v2.product.extraction.params.extraction_parameters import ( ExtractionParameters, ) -from mindee.v2.product.extraction.params.string_data_class import StringDataClass __all__ = [ "DataSchema", @@ -17,5 +16,4 @@ "ExtractionParameters", "ExtractionResponse", "ExtractionResult", - "StringDataClass", ] diff --git a/mindee/v2/product/extraction/extraction_response.py b/mindee/v2/product/extraction/extraction_response.py index 2bb744ce..e06d7e7c 100644 --- a/mindee/v2/product/extraction/extraction_response.py +++ b/mindee/v2/product/extraction/extraction_response.py @@ -10,6 +10,7 @@ class ExtractionResponse(BaseInferenceResponse): inference: ExtractionInference """Inference result.""" + _slug: ClassVar[str] = "extraction" def __init__(self, raw_response: StringDict) -> None: @@ -18,8 +19,3 @@ def __init__(self, raw_response: StringDict) -> None: def __str__(self) -> str: return str(self.inference) - - @classmethod - def get_result_slug(cls) -> str: - """Getter for the inference slug.""" - return cls._slug diff --git a/mindee/v2/product/extraction/extraction_result.py b/mindee/v2/product/extraction/extraction_result.py index 2f0c46ea..821aae00 100644 --- a/mindee/v2/product/extraction/extraction_result.py +++ b/mindee/v2/product/extraction/extraction_result.py @@ -1,6 +1,5 @@ from mindee.parsing.common.string_dict import StringDict from mindee.v2.parsing.inference.field import InferenceFields -from mindee.v2.parsing.inference.field.factory import create_field from mindee.v2.parsing.inference.rag_metadata import RAGMetadata from mindee.v2.parsing.inference.raw_text import RawText @@ -16,7 +15,7 @@ class ExtractionResult: """RAG metadata.""" def __init__(self, raw_response: StringDict) -> None: - self.fields = InferenceFields(raw_response["fields"], create_field) + self.fields = InferenceFields(raw_response["fields"]) if raw_response.get("raw_text"): self.raw_text = RawText(raw_response["raw_text"]) if raw_response.get("rag"): diff --git a/mindee/v2/product/extraction/params/__init__.py b/mindee/v2/product/extraction/params/__init__.py index e50473e8..3c88e962 100644 --- a/mindee/v2/product/extraction/params/__init__.py +++ b/mindee/v2/product/extraction/params/__init__.py @@ -4,12 +4,10 @@ from mindee.v2.product.extraction.params.extraction_parameters import ( ExtractionParameters, ) -from mindee.v2.product.extraction.params.string_data_class import StringDataClass __all__ = [ "DataSchema", "DataSchemaField", "DataSchemaReplace", "ExtractionParameters", - "StringDataClass", ] diff --git a/mindee/v2/product/extraction/params/data_schema.py b/mindee/v2/product/extraction/params/data_schema.py index 5392f49b..302e0889 100644 --- a/mindee/v2/product/extraction/params/data_schema.py +++ b/mindee/v2/product/extraction/params/data_schema.py @@ -1,15 +1,16 @@ import json from dataclasses import dataclass +from mindee.parsing.common import StringDict +from mindee.v2.parsing.string_data_class import StringDataClass from mindee.v2.product.extraction.params.data_schema_replace import DataSchemaReplace -from mindee.v2.product.extraction.params.string_data_class import StringDataClass @dataclass class DataSchema(StringDataClass): """Modify the Data Schema.""" - replace: DataSchemaReplace | dict | str | None = None + replace: DataSchemaReplace | StringDict | str | None = None """If set, completely replaces the data schema of the model.""" def __post_init__(self) -> None: diff --git a/mindee/v2/product/extraction/params/data_schema_field.py b/mindee/v2/product/extraction/params/data_schema_field.py index a986354a..c0619106 100644 --- a/mindee/v2/product/extraction/params/data_schema_field.py +++ b/mindee/v2/product/extraction/params/data_schema_field.py @@ -1,6 +1,6 @@ from dataclasses import dataclass -from mindee.v2.product.extraction.params.string_data_class import StringDataClass +from mindee.v2.parsing.string_data_class import StringDataClass @dataclass diff --git a/mindee/v2/product/extraction/params/data_schema_replace.py b/mindee/v2/product/extraction/params/data_schema_replace.py index 74620d0b..f54fe16f 100644 --- a/mindee/v2/product/extraction/params/data_schema_replace.py +++ b/mindee/v2/product/extraction/params/data_schema_replace.py @@ -1,7 +1,7 @@ from dataclasses import dataclass +from mindee.v2.parsing.string_data_class import StringDataClass from mindee.v2.product.extraction.params.data_schema_field import DataSchemaField -from mindee.v2.product.extraction.params.string_data_class import StringDataClass @dataclass diff --git a/mindee/v2/product/extraction/rag_documents/__init__.py b/mindee/v2/product/extraction/rag_documents/__init__.py new file mode 100644 index 00000000..b98fadf3 --- /dev/null +++ b/mindee/v2/product/extraction/rag_documents/__init__.py @@ -0,0 +1,29 @@ +from mindee.v2.product.extraction.rag_documents.annotated_list_field import ( + AnnotatedListField, +) +from mindee.v2.product.extraction.rag_documents.annotated_object_field import ( + AnnotatedObjectField, +) +from mindee.v2.product.extraction.rag_documents.annotated_simple_field import ( + AnnotatedSimpleField, +) +from mindee.v2.product.extraction.rag_documents.extraction_rag_annotation_response import ( + ExtractionRagAnnotationResponse, +) +from mindee.v2.product.extraction.rag_documents.params.rag_document_annotation_parameters import ( + RagDocumentAnnotationParameters, +) +from mindee.v2.product.extraction.rag_documents.params.rag_document_upload_parameters import ( + RagDocumentUploadParameters, +) +from mindee.v2.product.extraction.rag_documents.rag_annotation import RagAnnotation + +__all__ = [ + "AnnotatedListField", + "AnnotatedObjectField", + "AnnotatedSimpleField", + "ExtractionRagAnnotationResponse", + "RagAnnotation", + "RagDocumentAnnotationParameters", + "RagDocumentUploadParameters", +] diff --git a/mindee/v2/product/extraction/rag_documents/annotated_base_field.py b/mindee/v2/product/extraction/rag_documents/annotated_base_field.py new file mode 100644 index 00000000..b0313496 --- /dev/null +++ b/mindee/v2/product/extraction/rag_documents/annotated_base_field.py @@ -0,0 +1,70 @@ +from dataclasses import InitVar, dataclass +from typing import TYPE_CHECKING, ClassVar, TypeAlias, Union + +from mindee.parsing.common.string_dict import StringDict +from mindee.v2.parsing.inference.field.base_field import FieldType +from mindee.v2.parsing.string_data_class import StringDataClass + +if TYPE_CHECKING: + from mindee.v2.product.extraction.rag_documents.annotated_list_field import ( + AnnotatedListField, + ) + from mindee.v2.product.extraction.rag_documents.annotated_object_field import ( + AnnotatedObjectField, + ) + from mindee.v2.product.extraction.rag_documents.annotated_simple_field import ( + AnnotatedSimpleField, + ) + +AnnotatedFieldsType: TypeAlias = Union[ + "AnnotatedSimpleField", "AnnotatedObjectField", "AnnotatedListField" +] + + +@dataclass +class AnnotatedBaseField(StringDataClass): + """Base class for annotated fields.""" + + _field_type: ClassVar[FieldType] + raw_response: InitVar[StringDict] + + selected: bool = False + """When true, use the RAG information for the final result. When false, use the Data Schema information.""" + + guidelines: str | None = None + """Guidelines or instructions for processing this field.""" + + _registry: ClassVar[dict[str, type[AnnotatedFieldsType]]] = {} + + def __post_init__(self, raw_response: StringDict): + if "selected" in raw_response and raw_response["selected"] is not None: + self.selected = raw_response["selected"] + + if "guidelines" in raw_response and raw_response["guidelines"] is not None: + self.guidelines = raw_response.get("guidelines") + + @property + def field_type(self) -> FieldType: + """The field type.""" + return self._field_type + + @classmethod + def register(cls, discriminator_key: str): + """Class decorator: subclasses declare which JSON key identifies them.""" + + def decorator(subclass): + cls._registry[discriminator_key] = subclass + return subclass + + return decorator + + @classmethod + def build(cls, raw_response: dict) -> AnnotatedFieldsType: + """Build an instance of the appropriate subclass.""" + + if not isinstance(raw_response, dict): + raise ValueError("Field must be a dict") + for key, subclass in cls._registry.items(): + if key in raw_response: + return subclass(raw_response) + raise ValueError(f"Invalid structure for field: '{raw_response}'") diff --git a/mindee/v2/product/extraction/rag_documents/annotated_fields.py b/mindee/v2/product/extraction/rag_documents/annotated_fields.py new file mode 100644 index 00000000..13ac66bf --- /dev/null +++ b/mindee/v2/product/extraction/rag_documents/annotated_fields.py @@ -0,0 +1,60 @@ +from collections.abc import Iterable +from typing import TYPE_CHECKING, Any, cast + +from mindee.parsing.common.string_dict import StringDict +from mindee.v2.parsing.inference.field.base_field import FieldType +from mindee.v2.product.extraction.rag_documents.annotated_base_field import ( + AnnotatedBaseField, + AnnotatedFieldsType, +) + +if TYPE_CHECKING: + from mindee.v2.product.extraction.rag_documents.annotated_list_field import ( + AnnotatedListField, + ) + from mindee.v2.product.extraction.rag_documents.annotated_object_field import ( + AnnotatedObjectField, + ) + from mindee.v2.product.extraction.rag_documents.annotated_simple_field import ( + AnnotatedSimpleField, + ) + + +class AnnotatedFields(dict[str, AnnotatedFieldsType]): + """A dictionary of field names and their corresponding annotation.""" + + def __init__( + self, raw_response: StringDict | Iterable[tuple[str, Any]] | None = None + ): + super().__init__() + if raw_response is None: + return + + # Handle standard instantiation with a raw JSON dict + if hasattr(raw_response, "items"): + for key, value in raw_response.items(): + self[key] = AnnotatedBaseField.build(value) + # Handle the (key, value) generator passed recursively by dataclasses.asdict() + else: + self.update(raw_response) + + def get_simple_field(self, field_name: str) -> "AnnotatedSimpleField": + """Retrieve a simple field by its name.""" + field = self.get(field_name) + if field and field.field_type == FieldType.SIMPLE: + return cast("AnnotatedSimpleField", field) + raise ValueError(f"Field {field_name} is not an AnnotatedSimpleField.") + + def get_list_field(self, field_name: str) -> "AnnotatedListField": + """Retrieve a list field by its name.""" + field = self.get(field_name) + if field and field.field_type == FieldType.LIST: + return cast("AnnotatedListField", field) + raise ValueError(f"Field {field_name} is not an AnnotatedListField.") + + def get_object_field(self, field_name: str) -> "AnnotatedObjectField": + """Retrieve an object field by its name.""" + field = self.get(field_name) + if field and field.field_type == FieldType.OBJECT: + return cast("AnnotatedObjectField", field) + raise ValueError(f"Field {field_name} is not an AnnotatedObjectField.") diff --git a/mindee/v2/product/extraction/rag_documents/annotated_list_field.py b/mindee/v2/product/extraction/rag_documents/annotated_list_field.py new file mode 100644 index 00000000..b1a634b2 --- /dev/null +++ b/mindee/v2/product/extraction/rag_documents/annotated_list_field.py @@ -0,0 +1,49 @@ +from dataclasses import dataclass, field +from typing import ClassVar + +from mindee.parsing.common.string_dict import StringDict +from mindee.v2.parsing.inference.field.base_field import FieldType +from mindee.v2.product.extraction.rag_documents.annotated_base_field import ( + AnnotatedBaseField, +) +from mindee.v2.product.extraction.rag_documents.annotated_object_field import ( + AnnotatedObjectField, +) +from mindee.v2.product.extraction.rag_documents.annotated_simple_field import ( + AnnotatedSimpleField, +) + + +@AnnotatedBaseField.register("items") +@dataclass +class AnnotatedListField(AnnotatedBaseField): + """A ListField with additional configuration for annotation.""" + + _field_type: ClassVar[FieldType] = field(init=False, default=FieldType.LIST) + items: list[AnnotatedBaseField] = field(init=False) + + def __post_init__(self, raw_response: StringDict): + super().__post_init__(raw_response) + self.items = [AnnotatedBaseField.build(item) for item in raw_response["items"]] + + @property + def simple_items(self) -> list[AnnotatedSimpleField]: + """List of items as ``AnnotatedSimpleField``.""" + simple_items = [] + for item in self.items: + if isinstance(item, AnnotatedSimpleField): + simple_items.append(item) + else: + raise ValueError("List item is not an AnnotatedSimpleField field.") + return simple_items + + @property + def object_items(self) -> list[AnnotatedObjectField]: + """List of items as ``AnnotatedObjectField``.""" + object_items = [] + for item in self.items: + if isinstance(item, AnnotatedObjectField): + object_items.append(item) + else: + raise ValueError("List item is not an AnnotatedObjectField field.") + return object_items diff --git a/mindee/v2/product/extraction/rag_documents/annotated_object_field.py b/mindee/v2/product/extraction/rag_documents/annotated_object_field.py new file mode 100644 index 00000000..fd215a59 --- /dev/null +++ b/mindee/v2/product/extraction/rag_documents/annotated_object_field.py @@ -0,0 +1,33 @@ +from dataclasses import dataclass, field +from typing import ClassVar + +from mindee.parsing.common.string_dict import StringDict +from mindee.v2.parsing.inference.field.base_field import FieldType +from mindee.v2.product.extraction.rag_documents.annotated_base_field import ( + AnnotatedBaseField, +) +from mindee.v2.product.extraction.rag_documents.annotated_fields import AnnotatedFields +from mindee.v2.product.extraction.rag_documents.annotated_simple_field import ( + AnnotatedSimpleField, +) + + +@AnnotatedBaseField.register("fields") +@dataclass +class AnnotatedObjectField(AnnotatedBaseField): + """An ObjectField with additional configuration for annotation.""" + + _field_type: ClassVar[FieldType] = field(init=False, default=FieldType.OBJECT) + fields: AnnotatedFields = field(init=False) + + def __post_init__(self, raw_response: StringDict): + super().__post_init__(raw_response) + self.fields = AnnotatedFields(raw_response["fields"]) + + def get_simple_field(self, field_name: str) -> AnnotatedSimpleField: + """Retrieve a Simple field by its name.""" + return self.fields.get_simple_field(field_name) + + def get_object_field(self, field_name: str) -> "AnnotatedObjectField": + """Retrieve an Object field by its name.""" + return self.fields.get_object_field(field_name) diff --git a/mindee/v2/product/extraction/rag_documents/annotated_simple_field.py b/mindee/v2/product/extraction/rag_documents/annotated_simple_field.py new file mode 100644 index 00000000..dcee0bac --- /dev/null +++ b/mindee/v2/product/extraction/rag_documents/annotated_simple_field.py @@ -0,0 +1,22 @@ +from dataclasses import dataclass, field +from typing import ClassVar + +from mindee.parsing.common.string_dict import StringDict +from mindee.v2.parsing.inference.field.base_field import FieldType +from mindee.v2.product.extraction.rag_documents.annotated_base_field import ( + AnnotatedBaseField, +) + + +@AnnotatedBaseField.register("value") +@dataclass +class AnnotatedSimpleField(AnnotatedBaseField): + """A SimpleField with additional configuration for annotation.""" + + _field_type: ClassVar[FieldType] = field(init=False, default=FieldType.SIMPLE) + value: str | float | bool | None = field(init=False) + + def __post_init__(self, raw_response: StringDict): + super().__post_init__(raw_response) + + self.value = raw_response["value"] diff --git a/mindee/v2/product/extraction/rag_documents/extraction_rag_annotation_response.py b/mindee/v2/product/extraction/rag_documents/extraction_rag_annotation_response.py new file mode 100644 index 00000000..b542a4d1 --- /dev/null +++ b/mindee/v2/product/extraction/rag_documents/extraction_rag_annotation_response.py @@ -0,0 +1,35 @@ +from datetime import datetime +from typing import ClassVar + +from mindee.parsing.common.string_dict import StringDict +from mindee.v2.parsing.base_rag_annotation_response import BaseRagAnnotationResponse +from mindee.v2.product.extraction.rag_documents.rag_annotation import RagAnnotation + + +class ExtractionRagAnnotationResponse(BaseRagAnnotationResponse): + """Response for a RAG document.""" + + model_id: str + """Model identifier linked to the RAG document.""" + + total_matches: int + """Number of times this document was used in an inference.""" + + last_match_at: datetime | None = None + """Date and time of the latest matching inference, if any.""" + + annotation: RagAnnotation | None = None + """Annotation metadata associated with the document.""" + + _slug: ClassVar[str] = "extraction" + + def __init__(self, raw_response: StringDict) -> None: + super().__init__(raw_response) + self.model_id = raw_response["model_id"] + self.total_matches = raw_response["total_matches"] + if raw_response.get("last_match_at"): + self.last_match_at = datetime.fromisoformat( + raw_response["last_match_at"].replace("Z", "+00:00") + ) + if raw_response.get("annotation"): + self.annotation = RagAnnotation(raw_response["annotation"]) diff --git a/mindee/v2/product/extraction/rag_documents/params/__init__.py b/mindee/v2/product/extraction/rag_documents/params/__init__.py new file mode 100644 index 00000000..49a3463e --- /dev/null +++ b/mindee/v2/product/extraction/rag_documents/params/__init__.py @@ -0,0 +1,11 @@ +from mindee.v2.product.extraction.rag_documents.params.rag_document_annotation_parameters import ( + RagDocumentAnnotationParameters, +) +from mindee.v2.product.extraction.rag_documents.params.rag_document_upload_parameters import ( + RagDocumentUploadParameters, +) + +__all__ = [ + "RagDocumentAnnotationParameters", + "RagDocumentUploadParameters", +] diff --git a/mindee/v2/product/extraction/rag_documents/params/rag_document_annotation_parameters.py b/mindee/v2/product/extraction/rag_documents/params/rag_document_annotation_parameters.py new file mode 100644 index 00000000..26383080 --- /dev/null +++ b/mindee/v2/product/extraction/rag_documents/params/rag_document_annotation_parameters.py @@ -0,0 +1,46 @@ +import json +from dataclasses import asdict, dataclass +from typing import Any, ClassVar + +from mindee.parsing.common.string_dict import StringDict +from mindee.v2.client_options.base_annotation_parameters import BaseAnnotationParameters +from mindee.v2.product.extraction.rag_documents.extraction_rag_annotation_response import ( + ExtractionRagAnnotationResponse, +) +from mindee.v2.product.extraction.rag_documents.rag_annotation import RagAnnotation + + +@dataclass(kw_only=True) +class RagDocumentAnnotationParameters( + BaseAnnotationParameters[ExtractionRagAnnotationResponse] +): + """Annotation parameters for RAG documents.""" + + status: str | None = None + """New public status to apply to the document (for example, to deactivate it).""" + + annotation: RagAnnotation | StringDict | str | None = None + """Field-level RAG annotation and guidelines configuration for the document.""" + + _slug: ClassVar[str] = "extraction" + _response_class: type[ExtractionRagAnnotationResponse] = ( + ExtractionRagAnnotationResponse + ) + + def __post_init__(self) -> None: + if isinstance(self.annotation, str): + self.annotation = RagAnnotation(json.loads(self.annotation)) + elif isinstance(self.annotation, dict): + self.annotation = RagAnnotation(self.annotation) + + def get_request_parameters(self) -> dict[str, str | StringDict]: + """Gets the request parameters for the upload request.""" + parameters: dict[str, Any] = {} + + if self.status: + parameters["status"] = self.status + + if self.annotation is not None and isinstance(self.annotation, RagAnnotation): + parameters["annotation"] = asdict(self.annotation) + + return parameters diff --git a/mindee/v2/product/extraction/rag_documents/params/rag_document_upload_parameters.py b/mindee/v2/product/extraction/rag_documents/params/rag_document_upload_parameters.py new file mode 100644 index 00000000..07aa1d56 --- /dev/null +++ b/mindee/v2/product/extraction/rag_documents/params/rag_document_upload_parameters.py @@ -0,0 +1,21 @@ +from dataclasses import dataclass +from typing import ClassVar + +from mindee.v2.client_options.base_rag_document_upload_parameters import ( + BaseRagDocumentUploadParameters, +) +from mindee.v2.product.extraction.rag_documents.extraction_rag_annotation_response import ( + ExtractionRagAnnotationResponse, +) + + +@dataclass(kw_only=True) +class RagDocumentUploadParameters( + BaseRagDocumentUploadParameters[ExtractionRagAnnotationResponse] +): + """Upload parameters for RAG documents.""" + + _slug: ClassVar[str] = "extraction" + _response_class: type[ExtractionRagAnnotationResponse] = ( + ExtractionRagAnnotationResponse + ) diff --git a/mindee/v2/product/extraction/rag_documents/rag_annotation.py b/mindee/v2/product/extraction/rag_documents/rag_annotation.py new file mode 100644 index 00000000..9d41bc6a --- /dev/null +++ b/mindee/v2/product/extraction/rag_documents/rag_annotation.py @@ -0,0 +1,18 @@ +from dataclasses import InitVar, dataclass, field + +from mindee.parsing.common.string_dict import StringDict +from mindee.v2.parsing.string_data_class import StringDataClass +from mindee.v2.product.extraction.rag_documents.annotated_fields import AnnotatedFields + + +@dataclass +class RagAnnotation(StringDataClass): + """A RAG annotation enriched with field-level configuration.""" + + raw_response: InitVar[StringDict] + + fields: AnnotatedFields = field(init=False) + """Annotated fields.""" + + def __post_init__(self, raw_response: StringDict): + self.fields = AnnotatedFields(raw_response["fields"]) diff --git a/tests/v2/input/test_inference_parameters.py b/tests/v2/input/test_inference_parameters.py index 9c965d87..d4c41b01 100644 --- a/tests/v2/input/test_inference_parameters.py +++ b/tests/v2/input/test_inference_parameters.py @@ -12,7 +12,7 @@ (V2_PRODUCT_PATH / "extraction" / "data_schema_replace_param.json").read_text() ) expected_data_schema_str = json.dumps( - expected_data_schema_dict, indent=None, sort_keys=True + expected_data_schema_dict, indent=None, sort_keys=True, separators=(",", ":") ) diff --git a/tests/v2/product/extraction/test_extraction_response.py b/tests/v2/product/extraction/test_extraction_response.py index 8da8aec2..33d05223 100644 --- a/tests/v2/product/extraction/test_extraction_response.py +++ b/tests/v2/product/extraction/test_extraction_response.py @@ -7,14 +7,14 @@ FieldConfidence, InferenceFields, ListField, + ObjectField, + SimpleField, ) -from mindee.v2.parsing.inference.field.object_field import ObjectField -from mindee.v2.parsing.inference.field.simple_field import SimpleField from mindee.v2.parsing.inference.inference_active_options import InferenceActiveOptions from mindee.v2.parsing.inference.inference_file import InferenceFile from mindee.v2.parsing.inference.inference_model import InferenceModel from mindee.v2.parsing.inference.rag_metadata import RAGMetadata -from mindee.v2.product.extraction.extraction_inference import ExtractionInference +from mindee.v2.product.extraction import ExtractionInference from tests.utils import V2_PRODUCT_PATH from tests.v2.product.utils import get_product_samples diff --git a/tests/v2/product/extraction/test_rag_documents.py b/tests/v2/product/extraction/test_rag_documents.py new file mode 100644 index 00000000..5e75dc98 --- /dev/null +++ b/tests/v2/product/extraction/test_rag_documents.py @@ -0,0 +1,199 @@ +import json + +import pytest + +from mindee.v2.product.extraction.rag_documents import ( + ExtractionRagAnnotationResponse, + RagAnnotation, + RagDocumentAnnotationParameters, + RagDocumentUploadParameters, +) +from mindee.v2.product.extraction.rag_documents.annotated_list_field import ( + AnnotatedListField, +) +from mindee.v2.product.extraction.rag_documents.annotated_object_field import ( + AnnotatedObjectField, +) +from mindee.v2.product.extraction.rag_documents.annotated_simple_field import ( + AnnotatedSimpleField, +) +from tests.utils import V2_PRODUCT_PATH + +expected_annotation = { + "fields": { + "simple": {"selected": False, "guidelines": None, "value": True}, + "list": {"selected": False, "guidelines": None, "items": []}, + "object": {"selected": False, "guidelines": None, "fields": {}}, + } +} + + +@pytest.mark.v2 +def test_post_parameters_must_init(): + """Should init POST parameters.""" + parameters = RagDocumentUploadParameters(model_id="invalid-model-id") + req_params = parameters.get_request_parameters() + assert req_params["model_id"] == "invalid-model-id" + + +@pytest.mark.v2 +def test_patch_parameters_must_init_from_dict(): + """Should init PATCH parameters from an annotation dict.""" + annotation = { + "fields": { + "simple": {"value": True}, + "list": {"items": []}, + "object": {"fields": {}}, + } + } + parameters = RagDocumentAnnotationParameters( + document_id="invalid-document-id", status="Active", annotation=annotation + ) + req_params = parameters.get_request_parameters() + assert parameters.document_id == "invalid-document-id" + assert req_params["status"] == "Active" + assert req_params["annotation"] == expected_annotation + + +@pytest.mark.v2 +def test_patch_parameters_must_init_from_object(): + """Should init PATCH parameters from an annotation instance.""" + annotation = RagAnnotation( + { + "fields": { + "simple": {"value": True}, + "list": {"items": []}, + "object": {"fields": {}}, + } + } + ) + parameters = RagDocumentAnnotationParameters( + document_id="invalid-document-id", status="Active", annotation=annotation + ) + req_params = parameters.get_request_parameters() + assert parameters.document_id == "invalid-document-id" + assert req_params["status"] == "Active" + assert req_params["annotation"] == expected_annotation + + +@pytest.mark.v2 +def test_patch_parameters_must_init_from_string(): + """Should init PATCH parameters from an annotation string.""" + parameters = RagDocumentAnnotationParameters( + document_id="invalid-document-id", + status="Active", + annotation='{"fields": {"simple":{"value": true}, "list":{"items": []}, "object":{"fields": {}}}}', + ) + req_params = parameters.get_request_parameters() + assert parameters.document_id == "invalid-document-id" + assert req_params["status"] == "Active" + assert req_params["annotation"] == expected_annotation + + +@pytest.mark.v2 +def test_rag_documents_post_must_have_valid_properties(): + """Should load a POST response from a JSON string.""" + + response = _get_response("extraction/rag_documents/post_response.json") + assert response is not None + assert response.id == "cc831599-c545-48b7-aa27-6d7ccd5b8d32" + assert response.status == "Processing" + assert response.annotation is None + + +@pytest.mark.v2 +def test_rag_documents_get_draft_must_have_valid_properties(): + """Should load a GET response from a JSON string.""" + + response = _get_response("extraction/rag_documents/get_response_draft.json") + assert response is not None + assert response.id == "cc831599-c545-48b7-aa27-6d7ccd5b8d32" + assert response.status == "Draft" + assert response.annotation is not None + fields = response.annotation.fields + assert fields is not None + + # null simple field + tip_field = fields.get_simple_field("tip") + assert isinstance(tip_field, AnnotatedSimpleField) + assert tip_field.selected is False + assert tip_field.guidelines is None + assert tip_field.value is None + + # filled simple field + date_field = fields.get_simple_field("date") + assert isinstance(date_field, AnnotatedSimpleField) + assert date_field.selected is False + assert date_field.guidelines is None + assert date_field.value == "2019-11-02" + + # filled object field + locale_field = fields.get_object_field("locale") + assert isinstance(locale_field, AnnotatedObjectField) + assert locale_field.selected is False + assert locale_field.guidelines is None + assert locale_field.fields is not None + assert len(locale_field.fields) == 3 + assert locale_field.get_simple_field("country").value == "US" + assert locale_field.get_simple_field("currency").value == "USD" + assert locale_field.get_simple_field("language").value is None + + # list of simple fields + reference_numbers_field = fields.get_list_field("reference_numbers") + assert isinstance(reference_numbers_field, AnnotatedListField) + assert reference_numbers_field.selected is False + assert reference_numbers_field.guidelines is None + assert reference_numbers_field.items is not None + assert len(reference_numbers_field.items) == 1 + assert reference_numbers_field.simple_items[0].value == "2412/2019" + + # list of object fields + line_items_field = fields.get_list_field("line_items") + assert isinstance(line_items_field, AnnotatedListField) + assert line_items_field.selected is False + assert line_items_field.guidelines is None + assert line_items_field.items is not None + assert len(line_items_field.items) == 3 + + line_item_0 = line_items_field.object_items[0] + assert line_item_0.fields is not None + assert len(line_item_0.fields) == 8 + assert line_item_0.fields["description"].value == "Front and rear brake cables" + assert line_item_0.fields["quantity"].value == 1 + assert line_item_0.fields["unit_price"].value == 100 + assert line_item_0.fields["total_price"].value == 100 + assert line_item_0.fields["tax_rate"].value is None + assert line_item_0.fields["tax_amount"].value is None + assert line_item_0.fields["product_code"].value is None + assert line_item_0.fields["unit_measure"].value is None + + line_item_1 = line_items_field.object_items[1] + assert line_item_1.fields is not None + assert len(line_item_1.fields) == 8 + assert line_item_1.fields["description"].value == "New set of pedal arms" + assert line_item_1.fields["quantity"].value == 2 + assert line_item_1.fields["unit_price"].value == 25 + assert line_item_1.fields["total_price"].value == 50 + assert line_item_1.fields["tax_rate"].value is None + assert line_item_1.fields["tax_amount"].value is None + assert line_item_1.fields["product_code"].value is None + assert line_item_1.fields["unit_measure"].value is None + + line_item_2 = line_items_field.object_items[2] + assert line_item_2.fields is not None + assert len(line_item_2.fields) == 8 + assert line_item_2.fields["description"].value == "Labor 3hrs" + assert line_item_2.fields["quantity"].value == 3 + assert line_item_2.fields["unit_price"].value == 15 + assert line_item_2.fields["total_price"].value == 45 + assert line_item_2.fields["tax_rate"].value is None + assert line_item_2.fields["tax_amount"].value is None + assert line_item_2.fields["product_code"].value is None + assert line_item_2.fields["unit_measure"].value is None + + +def _get_response(path: str) -> ExtractionRagAnnotationResponse: + file_path = V2_PRODUCT_PATH / path + with open(file_path, encoding="utf-8") as f: + response_dict = json.load(f) + return ExtractionRagAnnotationResponse(response_dict) diff --git a/tests/v2/product/extraction/test_rag_documents_integration.py b/tests/v2/product/extraction/test_rag_documents_integration.py new file mode 100644 index 00000000..648f50c6 --- /dev/null +++ b/tests/v2/product/extraction/test_rag_documents_integration.py @@ -0,0 +1,104 @@ +import os + +import pytest + +from mindee import PathInput +from mindee.v2 import Client +from mindee.v2.error import MindeeHTTPErrorV2 +from mindee.v2.product.extraction.rag_documents import ( + ExtractionRagAnnotationResponse, + RagDocumentAnnotationParameters, + RagDocumentUploadParameters, +) +from tests.utils import V2_PRODUCT_PATH + + +@pytest.mark.integration +@pytest.mark.v2 +def test_rag_document_lifecycle_must_succeed(): + """Should perform the entire lifecycle of a RAG document.""" + extraction_model_id = os.getenv("MINDEE_V2_SE_TESTS_FINDOC_MODEL_ID") + if extraction_model_id is None: + raise RuntimeError( + "MINDEE_V2_SE_TESTS_FINDOC_MODEL_ID environment variable is not set" + ) + + client = Client() + + input_source = PathInput( + V2_PRODUCT_PATH / "extraction" / "financial_document" / "default_sample.jpg" + ) + post_parameters = RagDocumentUploadParameters(model_id=extraction_model_id) + post_response = client.upload_and_get_rag_document(input_source, post_parameters) + assert post_response is not None + + post_annotation = post_response.annotation + assert post_annotation is not None + + document_id = post_response.id + assert document_id is not None + + assert post_response.status == "Draft" + + post_annotation.fields["supplier_name"].selected = True + post_annotation.fields["supplier_name"].guidelines = "I am the walrus!" + post_annotation.fields.get_simple_field("invoice_number").selected = True + post_annotation.fields.get_simple_field( + "invoice_number" + ).guidelines = "koo koo katchoo!" + + patch_annotation_response = client.update_rag_annotations( + RagDocumentAnnotationParameters( + document_id=document_id, + annotation=post_annotation, + ) + ) + assert patch_annotation_response is not None + patch_annotation = patch_annotation_response.annotation + assert patch_annotation is not None + + assert ( + patch_annotation.fields.get_simple_field("supplier_name").guidelines + == "I am the walrus!" + ) + assert patch_annotation.fields.get_simple_field("supplier_name").selected is True + assert ( + patch_annotation.fields.get_simple_field("invoice_number").guidelines + == "koo koo katchoo!" + ) + assert patch_annotation.fields.get_simple_field("invoice_number").selected is True + + get_response = client.get_ready_rag_document( + ExtractionRagAnnotationResponse, document_id + ) + assert get_response is not None + get_annotation = get_response.annotation + assert get_annotation is not None + + assert get_response.status == "Draft" + + assert ( + get_annotation.fields.get_simple_field("supplier_name").guidelines + == "I am the walrus!" + ) + assert get_annotation.fields.get_simple_field("supplier_name").selected is True + assert ( + get_annotation.fields.get_simple_field("invoice_number").guidelines + == "koo koo katchoo!" + ) + assert get_annotation.fields.get_simple_field("invoice_number").selected is True + + patch_status_response = client.update_and_get_rag_annotations( + RagDocumentAnnotationParameters( + document_id=document_id, + status="Active", + ) + ) + assert patch_status_response is not None + assert patch_status_response.status == "Active" + + delete_response = client.delete_extraction_rag_document(document_id) + assert delete_response is True + + with pytest.raises(MindeeHTTPErrorV2): + client.get_rag_document(ExtractionRagAnnotationResponse, document_id) diff --git a/tests/v2/test_base_classes.py b/tests/v2/test_base_classes.py index 5bb46b2b..4f648594 100644 --- a/tests/v2/test_base_classes.py +++ b/tests/v2/test_base_classes.py @@ -38,7 +38,7 @@ def test_base_inference(): def test_base_response(): class DummyResponse(BaseInferenceResponse): - _slug = "dummy/results" + _slug = "dummy" def __init__(self): self.inference = "dummy inference" @@ -46,4 +46,4 @@ def __init__(self): response = DummyResponse() assert str(response) == "dummy inference" - assert response.get_result_slug() == "dummy/results" + assert response.get_product_slug() == "dummy" diff --git a/tests/v2/test_client.py b/tests/v2/test_client.py index edb1dfd3..8dd63b25 100644 --- a/tests/v2/test_client.py +++ b/tests/v2/test_client.py @@ -254,7 +254,7 @@ def test_client_closes_httpx_connections() -> None: with pytest.raises( AttributeError, match=r"'NoneType' object has no attribute 'get'" ): - client.mindee_api.http_client.get("https://google.com") + client.mindee_api._http_client.get("https://google.com") @pytest.mark.v2 diff --git a/tests/v2/test_client_integration.py b/tests/v2/test_client_integration.py index fd477ba5..35d96e37 100644 --- a/tests/v2/test_client_integration.py +++ b/tests/v2/test_client_integration.py @@ -302,7 +302,10 @@ def test_data_schema_must_succeed( _basic_assert_success(response=response, page_count=1, model_id=findoc_model_id) assert response.inference.active_options.data_schema.replace is True assert response.inference.result.fields["test_replace"] is not None - assert response.inference.result.fields["test_replace"].value == "a test value" + assert ( + response.inference.result.fields.get_simple_field("test_replace").value + == "a test value" + ) @pytest.mark.integration