|
10 | 10 | from mindee.logger import logger |
11 | 11 | from mindee.mindee_http.cancellation_token import CancellationToken |
12 | 12 | from mindee.parsing.common.common_response import CommonStatus |
| 13 | +from mindee.v2.client_options.base_annotation_parameters import BaseAnnotationParameters |
13 | 14 | from mindee.v2.client_options.base_product_parameters import BaseProductParameters |
| 15 | +from mindee.v2.client_options.base_rag_document_upload_parameters import ( |
| 16 | + BaseRagDocumentUploadParameters, |
| 17 | +) |
14 | 18 | from mindee.v2.client_options.base_search_parameters import ( |
15 | 19 | BaseSearchParameters, |
16 | 20 | TypeSearchResponse, |
17 | 21 | ) |
18 | 22 | from mindee.v2.mindee_http.mindee_api_v2 import MindeeAPIV2 |
| 23 | +from mindee.v2.parsing.base_rag_annotation_response import ( |
| 24 | + TypeRagAnnotationResponse, |
| 25 | +) |
19 | 26 | from mindee.v2.parsing.inference.base_inference_response import ( |
20 | 27 | TypeBaseInferenceResponse, |
21 | 28 | ) |
@@ -166,6 +173,143 @@ def enqueue_and_get_result( |
166 | 173 |
|
167 | 174 | raise MindeeError(f"Couldn't retrieve document after {try_counter + 1} tries.") |
168 | 175 |
|
| 176 | + def upload_rag_document( |
| 177 | + self, |
| 178 | + input_source: LocalInputSource, |
| 179 | + parameters: BaseRagDocumentUploadParameters[TypeRagAnnotationResponse], |
| 180 | + ) -> TypeRagAnnotationResponse: |
| 181 | + """ |
| 182 | + Not recommended for general use, prefer ``upload_and_get_rag_document``. |
| 183 | + You will need to poll until the document is ready for use. |
| 184 | + Add a document to the RAG database. |
| 185 | + """ |
| 186 | + return self.mindee_api.req_post_rag_document(input_source, parameters) |
| 187 | + |
| 188 | + def upload_and_get_rag_document( |
| 189 | + self, |
| 190 | + input_source: LocalInputSource, |
| 191 | + parameters: BaseRagDocumentUploadParameters[TypeRagAnnotationResponse], |
| 192 | + polling_options: PollingOptions | None = None, |
| 193 | + cancellation_token: CancellationToken | None = None, |
| 194 | + ) -> TypeRagAnnotationResponse: |
| 195 | + """ |
| 196 | + Add a document to the RAG database and return the initial annotation. |
| 197 | + """ |
| 198 | + initial_response = self.upload_rag_document(input_source, parameters) |
| 199 | + if initial_response.status != "Processing": |
| 200 | + return initial_response |
| 201 | + if polling_options is None: |
| 202 | + polling_options = PollingOptions() |
| 203 | + return self._poll_for_rag_document( |
| 204 | + initial_response, polling_options, cancellation_token |
| 205 | + ) |
| 206 | + |
| 207 | + def get_rag_document( |
| 208 | + self, response_type: type[TypeRagAnnotationResponse], document_id: str |
| 209 | + ) -> TypeRagAnnotationResponse: |
| 210 | + """ |
| 211 | + Not recommended for general use, prefer ``get_ready_rag_document``. |
| 212 | + You will need to poll until the document is ready for use. |
| 213 | + Get a document's info and annotations from the RAG database. |
| 214 | + """ |
| 215 | + return self.mindee_api.req_get_rag_annotation(response_type, document_id) |
| 216 | + |
| 217 | + def get_ready_rag_document( |
| 218 | + self, |
| 219 | + response_type: type[TypeRagAnnotationResponse], |
| 220 | + document_id: str, |
| 221 | + polling_options: PollingOptions | None = None, |
| 222 | + cancellation_token: CancellationToken | None = None, |
| 223 | + ): |
| 224 | + """ |
| 225 | + Get a document's info and annotations from the RAG database. |
| 226 | + """ |
| 227 | + initial_response = self.get_rag_document(response_type, document_id) |
| 228 | + if initial_response.status != "Processing": |
| 229 | + return initial_response |
| 230 | + if polling_options is None: |
| 231 | + polling_options = PollingOptions() |
| 232 | + return self._poll_for_rag_document( |
| 233 | + initial_response, polling_options, cancellation_token |
| 234 | + ) |
| 235 | + |
| 236 | + def update_rag_annotations( |
| 237 | + self, parameters: BaseAnnotationParameters[TypeRagAnnotationResponse] |
| 238 | + ) -> TypeRagAnnotationResponse: |
| 239 | + """ |
| 240 | + Not recommended for general use, prefer ``update_and_get_rag_annotations``. |
| 241 | + You will need to poll until the document is ready for use. |
| 242 | + Update a document's annotations in the RAG database. |
| 243 | + """ |
| 244 | + return self.mindee_api.req_patch_rag_annotation(parameters) |
| 245 | + |
| 246 | + def update_and_get_rag_annotations( |
| 247 | + self, |
| 248 | + parameters: BaseAnnotationParameters[TypeRagAnnotationResponse], |
| 249 | + polling_options: PollingOptions | None = None, |
| 250 | + cancellation_token: CancellationToken | None = None, |
| 251 | + ) -> TypeRagAnnotationResponse: |
| 252 | + """ |
| 253 | + Update a document's annotations in the RAG database. |
| 254 | + """ |
| 255 | + initial_response = self.update_rag_annotations(parameters) |
| 256 | + if initial_response.status != "Processing": |
| 257 | + return initial_response |
| 258 | + if polling_options is None: |
| 259 | + polling_options = PollingOptions() |
| 260 | + return self._poll_for_rag_document( |
| 261 | + initial_response, polling_options, cancellation_token |
| 262 | + ) |
| 263 | + |
| 264 | + def delete_extraction_rag_document(self, document_id: str) -> bool: |
| 265 | + """ |
| 266 | + Delete a document from the RAG database. |
| 267 | + For extraction models only. |
| 268 | + """ |
| 269 | + return self.mindee_api.req_delete_extraction_rag_document(document_id) |
| 270 | + |
| 271 | + def _poll_for_rag_document( |
| 272 | + self, |
| 273 | + initial_response: TypeRagAnnotationResponse, |
| 274 | + polling_options: PollingOptions, |
| 275 | + cancellation_token: CancellationToken | None = None, |
| 276 | + ) -> TypeRagAnnotationResponse: |
| 277 | + """ |
| 278 | + Poll until the document is finished processing or the max number of attempts is reached. |
| 279 | + """ |
| 280 | + logger.info("Polling for RAG document ID: %s", initial_response.id) |
| 281 | + max_retries = polling_options.max_retries + 1 |
| 282 | + |
| 283 | + logger.debug( |
| 284 | + "Waiting %s seconds before attempting to retrieve the result...", |
| 285 | + polling_options.initial_delay_sec, |
| 286 | + ) |
| 287 | + |
| 288 | + if cancellation_token and cancellation_token.is_canceled: |
| 289 | + raise MindeeError("Request canceled through cancellation token.") |
| 290 | + |
| 291 | + sleep(polling_options.initial_delay_sec) |
| 292 | + document_id = initial_response.id |
| 293 | + retry_count = 1 |
| 294 | + |
| 295 | + while retry_count < max_retries: |
| 296 | + if cancellation_token and cancellation_token.is_canceled: |
| 297 | + raise MindeeError("Request canceled through cancellation token.") |
| 298 | + |
| 299 | + sleep(polling_options.delay_sec) |
| 300 | + logger.info("Poll attempt %s of %s", retry_count, max_retries) |
| 301 | + |
| 302 | + response = self.get_rag_document(type(initial_response), document_id) |
| 303 | + retry_count += 1 |
| 304 | + |
| 305 | + if response.status == "Processing": |
| 306 | + continue |
| 307 | + if response.status == "Failed": |
| 308 | + raise MindeeError("Job failed without an error payload.") |
| 309 | + return response |
| 310 | + |
| 311 | + raise MindeeError(f"RAG polling not complete after {retry_count - 1} attempts.") |
| 312 | + |
169 | 313 | def search( |
170 | 314 | self, params: BaseSearchParameters[TypeSearchResponse] |
171 | 315 | ) -> TypeSearchResponse: |
|
0 commit comments