base.py 17 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459
  1. import asyncio
  2. import base64
  3. import logging
  4. import os
  5. import time
  6. from copy import copy
  7. from io import BytesIO
  8. from typing import Any, AsyncGenerator
  9. import httpx
  10. from unstructured_client import UnstructuredClient
  11. from unstructured_client.models import operations, shared
  12. from core import parsers
  13. from core.base import (
  14. AsyncParser,
  15. ChunkingStrategy,
  16. Document,
  17. DocumentChunk,
  18. DocumentType,
  19. RecursiveCharacterTextSplitter,
  20. )
  21. from core.base.abstractions import R2RSerializable
  22. from core.base.providers.ingestion import IngestionConfig, IngestionProvider
  23. from core.providers.ocr import MistralOCRProvider
  24. from core.utils import generate_extraction_id
  25. from ...database import PostgresDatabaseProvider
  26. from ...llm import (
  27. LiteLLMCompletionProvider,
  28. OpenAICompletionProvider,
  29. R2RCompletionProvider,
  30. )
  31. logger = logging.getLogger()
  32. class FallbackElement(R2RSerializable):
  33. text: str
  34. metadata: dict[str, Any]
  35. class UnstructuredIngestionConfig(IngestionConfig):
  36. combine_under_n_chars: int = 128
  37. max_characters: int = 500
  38. new_after_n_chars: int = 1500
  39. overlap: int = 64
  40. coordinates: bool | None = None
  41. encoding: str | None = None # utf-8
  42. extract_image_block_types: list[str] | None = None
  43. gz_uncompressed_content_type: str | None = None
  44. hi_res_model_name: str | None = None
  45. include_orig_elements: bool | None = None
  46. include_page_breaks: bool | None = None
  47. languages: list[str] | None = None
  48. multipage_sections: bool | None = None
  49. ocr_languages: list[str] | None = None
  50. # output_format: Optional[str] = "application/json"
  51. overlap_all: bool | None = None
  52. pdf_infer_table_structure: bool | None = None
  53. similarity_threshold: float | None = None
  54. skip_infer_table_types: list[str] | None = None
  55. split_pdf_concurrency_level: int | None = None
  56. split_pdf_page: bool | None = None
  57. starting_page_number: int | None = None
  58. strategy: str | None = None
  59. chunking_strategy: str | ChunkingStrategy | None = None # type: ignore
  60. unique_element_ids: bool | None = None
  61. xml_keep_tags: bool | None = None
  62. def to_ingestion_request(self):
  63. import json
  64. x = json.loads(self.json())
  65. x.pop("extra_fields", None)
  66. x.pop("provider", None)
  67. x.pop("excluded_parsers", None)
  68. x = {k: v for k, v in x.items() if v is not None}
  69. return x
  70. class UnstructuredIngestionProvider(IngestionProvider):
  71. R2R_FALLBACK_PARSERS = {
  72. DocumentType.GIF: [parsers.ImageParser], # type: ignore
  73. DocumentType.JPEG: [parsers.ImageParser], # type: ignore
  74. DocumentType.JPG: [parsers.ImageParser], # type: ignore
  75. DocumentType.PNG: [parsers.ImageParser], # type: ignore
  76. DocumentType.SVG: [parsers.ImageParser], # type: ignore
  77. DocumentType.HEIC: [parsers.ImageParser], # type: ignore
  78. DocumentType.MP3: [parsers.AudioParser], # type: ignore
  79. DocumentType.JSON: [parsers.JSONParser], # type: ignore
  80. DocumentType.HTML: [parsers.HTMLParser], # type: ignore
  81. DocumentType.XLS: [parsers.XLSParser], # type: ignore
  82. DocumentType.XLSX: [parsers.XLSXParser], # type: ignore
  83. #DocumentType.DOC: [parsers.DOCParser], # type: ignore
  84. DocumentType.PPT: [parsers.PPTParser], # type: ignore
  85. DocumentType.CSV: [parsers.CSVParserAdvanced], # type: ignore
  86. }
  87. EXTRA_PARSERS = {
  88. #DocumentType.CSV: {"advanced": parsers.CSVParserAdvanced}, # type: ignore
  89. DocumentType.PDF: {
  90. "ocr": parsers.OCRPDFParser, # type: ignore
  91. "unstructured": parsers.PDFParserUnstructured, # type: ignore
  92. "zerox": parsers.VLMPDFParser, # type: ignore
  93. }
  94. #DocumentType.XLSX: {"advanced": parsers.XLSXParserAdvanced}, # type: ignore
  95. }
  96. IMAGE_TYPES = {
  97. DocumentType.GIF,
  98. DocumentType.HEIC,
  99. DocumentType.JPG,
  100. DocumentType.JPEG,
  101. DocumentType.PNG,
  102. DocumentType.SVG,
  103. }
  104. def __init__(
  105. self,
  106. config: UnstructuredIngestionConfig,
  107. database_provider: PostgresDatabaseProvider,
  108. llm_provider: (
  109. LiteLLMCompletionProvider
  110. | OpenAICompletionProvider
  111. | R2RCompletionProvider
  112. ),
  113. ocr_provider: MistralOCRProvider,
  114. ):
  115. super().__init__(config, database_provider, llm_provider)
  116. self.config: UnstructuredIngestionConfig = config
  117. self.database_provider: PostgresDatabaseProvider = database_provider
  118. self.llm_provider: (
  119. LiteLLMCompletionProvider
  120. | OpenAICompletionProvider
  121. | R2RCompletionProvider
  122. ) = llm_provider
  123. self.ocr_provider: MistralOCRProvider = ocr_provider
  124. self.client: UnstructuredClient | httpx.AsyncClient
  125. #config.provider = "unstructured_api"
  126. if config.provider == "unstructured_api":
  127. try:
  128. self.unstructured_api_auth = os.environ["UNSTRUCTURED_API_KEY"]
  129. except KeyError as e:
  130. raise ValueError(
  131. "UNSTRUCTURED_API_KEY environment variable is not set"
  132. ) from e
  133. self.unstructured_api_url = os.environ.get(
  134. "UNSTRUCTURED_API_URL",
  135. "https://api.unstructuredapp.io/general/v0/general",
  136. )
  137. self.client = UnstructuredClient(
  138. api_key_auth=self.unstructured_api_auth,
  139. server_url=self.unstructured_api_url,
  140. )
  141. self.shared = shared
  142. self.operations = operations
  143. else:
  144. try:
  145. self.local_unstructured_url = os.environ[
  146. "UNSTRUCTURED_SERVICE_URL"
  147. ]
  148. except KeyError as e:
  149. raise ValueError(
  150. "UNSTRUCTURED_SERVICE_URL environment variable is not set"
  151. ) from e
  152. self.client = httpx.AsyncClient()
  153. self.parsers: dict[DocumentType, AsyncParser] = {}
  154. self._initialize_parsers()
  155. def _initialize_parsers(self):
  156. for doc_type, parsers in self.R2R_FALLBACK_PARSERS.items():
  157. for parser in parsers:
  158. if (
  159. doc_type not in self.config.excluded_parsers
  160. and doc_type not in self.parsers
  161. ):
  162. # will choose the first parser in the list
  163. self.parsers[doc_type] = parser(
  164. config=self.config,
  165. database_provider=self.database_provider,
  166. llm_provider=self.llm_provider,
  167. )
  168. # TODO - Reduce code duplication between Unstructured & R2R
  169. for doc_type, parser_names in self.config.extra_parsers.items():
  170. if not isinstance(parser_names, list):
  171. parser_names = [parser_names]
  172. for parser_name in parser_names:
  173. parser_key = f"{parser_name}_{str(doc_type)}"
  174. try:
  175. self.parsers[parser_key] = self.EXTRA_PARSERS[doc_type][
  176. parser_name
  177. ](
  178. config=self.config,
  179. database_provider=self.database_provider,
  180. llm_provider=self.llm_provider,
  181. ocr_provider=self.ocr_provider,
  182. )
  183. logger.info(
  184. f"Initialized extra parser {parser_name} for {doc_type}"
  185. )
  186. except KeyError as e:
  187. logger.error(
  188. f"Parser {parser_name} for document type {doc_type} not found: {e}"
  189. )
  190. async def parse_fallback(
  191. self,
  192. file_content: bytes,
  193. ingestion_config: dict,
  194. parser_name: str,
  195. ) -> AsyncGenerator[FallbackElement, None]:
  196. contents = []
  197. async for chunk in self.parsers[parser_name].ingest( # type: ignore
  198. file_content, **ingestion_config
  199. ): # type: ignore
  200. if isinstance(chunk, dict) and chunk.get("content"):
  201. contents.append(chunk)
  202. elif chunk: # Handle string output for backward compatibility
  203. contents.append({"content": chunk})
  204. if not contents:
  205. logging.warning(
  206. "No valid text content was extracted during parsing"
  207. )
  208. return
  209. logging.info(f"Fallback ingestion with config = {ingestion_config}")
  210. vlm_ocr_one_page_per_chunk = ingestion_config.get(
  211. "vlm_ocr_one_page_per_chunk", True
  212. )
  213. iteration = 0
  214. for content_item in contents:
  215. text = content_item["content"]
  216. if vlm_ocr_one_page_per_chunk and parser_name.startswith(
  217. ("zerox_", "ocr_")
  218. ):
  219. # Use one page per chunk for OCR/VLM
  220. metadata = {"chunk_id": iteration}
  221. if "page_number" in content_item:
  222. metadata["page_number"] = content_item["page_number"]
  223. yield FallbackElement(
  224. text=text or "No content extracted.",
  225. metadata=metadata,
  226. )
  227. iteration += 1
  228. await asyncio.sleep(0)
  229. else:
  230. # Use regular text splitting
  231. loop = asyncio.get_event_loop()
  232. splitter = RecursiveCharacterTextSplitter(
  233. chunk_size=ingestion_config["new_after_n_chars"],
  234. chunk_overlap=ingestion_config["overlap"],
  235. )
  236. chunks = await loop.run_in_executor(
  237. None, splitter.create_documents, [text]
  238. )
  239. for text_chunk in chunks:
  240. metadata = {"chunk_id": iteration}
  241. if "page_number" in content_item:
  242. metadata["page_number"] = content_item["page_number"]
  243. yield FallbackElement(
  244. text=text_chunk.page_content,
  245. metadata=metadata,
  246. )
  247. iteration += 1
  248. await asyncio.sleep(0)
  249. async def parse(
  250. self,
  251. file_content: bytes,
  252. document: Document,
  253. ingestion_config_override: dict,
  254. ) -> AsyncGenerator[DocumentChunk, None]:
  255. ingestion_config = copy(
  256. {
  257. **self.config.to_ingestion_request(),
  258. **(ingestion_config_override or {}),
  259. }
  260. )
  261. # cleanup extra fields
  262. ingestion_config.pop("provider", None)
  263. ingestion_config.pop("excluded_parsers", None)
  264. t0 = time.time()
  265. parser_overrides = ingestion_config_override.get(
  266. "parser_overrides", {}
  267. )
  268. elements = []
  269. #parser_overrides = {"pdf": "unstructured"}
  270. # TODO - Cleanup this approach to be less hardcoded
  271. # TODO - Remove code duplication between Unstructured & R2R
  272. logger.info(f"Parser overrides: {parser_overrides}")
  273. logger.info(f"R2R fallback parsers is: {document.document_type.value}")
  274. logger.info(f"R2R fallback parsers is: {self.EXTRA_PARSERS.keys()}")
  275. logger.info(f"R2R fallback parsers is: {document.document_type.value in self.EXTRA_PARSERS.keys()}")
  276. logger.info(f"R2R fallback parsers is: {document.document_type.value in parser_overrides}")
  277. #if document.document_type.value in parser_overrides:
  278. if document.document_type.value in parser_overrides:
  279. logger.info(
  280. f"Using parser_override for {document.document_type} with input value {parser_overrides[document.document_type.value]}"
  281. )
  282. if parser_overrides[document.document_type.value] == "zerox":
  283. async for element in self.parse_fallback(
  284. file_content,
  285. ingestion_config=ingestion_config,
  286. parser_name=f"zerox_{DocumentType.PDF.value}",
  287. ):
  288. logger.warning(
  289. f"Using parser_override for {document.document_type}"
  290. )
  291. elements.append(element)
  292. elif parser_overrides[document.document_type.value] == "ocr":
  293. async for element in self.parse_fallback(
  294. file_content,
  295. ingestion_config=ingestion_config,
  296. parser_name=f"ocr_{DocumentType.PDF.value}",
  297. ):
  298. logger.warning(
  299. f"Using OCR parser_override for {document.document_type}"
  300. )
  301. elements.append(element)
  302. elif document.document_type in self.R2R_FALLBACK_PARSERS.keys():
  303. logger.info(
  304. f"Parsing {document.document_type}: {document.id} with fallback parser"
  305. )
  306. try:
  307. async for element in self.parse_fallback(
  308. file_content,
  309. ingestion_config=ingestion_config,
  310. parser_name=document.document_type,
  311. ):
  312. elements.append(element)
  313. except Exception as e:
  314. logger.error(f"Error parsing {document.document_type}: {e}")
  315. raise e
  316. else:
  317. logger.info(
  318. f"Parsing {document.document_type}: {document.id} with unstructured"
  319. )
  320. file_io = BytesIO(file_content)
  321. logger.info(f"Provider is: {self.config.provider}")
  322. # TODO - Include check on excluded parsers here.
  323. if self.config.provider == "unstructured_api":
  324. logger.info(f"Using API to parse document {document.id}")
  325. files = self.shared.Files(
  326. content=file_io.read(),
  327. file_name=document.metadata.get("title", "unknown_file"),
  328. )
  329. ingestion_config.pop("app", None)
  330. ingestion_config.pop("extra_parsers", None)
  331. req = self.operations.PartitionRequest(
  332. partition_parameters=self.shared.PartitionParameters(
  333. files=files,
  334. **ingestion_config,
  335. )
  336. )
  337. elements = await self.client.general.partition_async( # type: ignore
  338. request=req
  339. )
  340. elements = list(elements.elements) # type: ignore
  341. else:
  342. logger.info(
  343. f"Using local unstructured fastapi server to parse document {document.id}"
  344. )
  345. # Base64 encode the file content
  346. encoded_content = base64.b64encode(file_io.read()).decode(
  347. "utf-8"
  348. )
  349. logger.info(
  350. f"Sending a request to {self.local_unstructured_url}/partition"
  351. )
  352. #ingestion_config["strategy"] = "hi_res"
  353. print(ingestion_config)
  354. response = await self.client.post(
  355. f"{self.local_unstructured_url}/partition",
  356. json={
  357. "file_content": encoded_content, # Use encoded string
  358. "ingestion_config": ingestion_config,
  359. "filename": document.metadata.get("title", None),
  360. },
  361. timeout=3600, # Adjust timeout as needed
  362. )
  363. if response.status_code != 200:
  364. logger.error(f"Error partitioning file: {response.text}")
  365. raise ValueError(
  366. f"Error partitioning file: {response.text}"
  367. )
  368. elements = response.json().get("elements", [])
  369. iteration = 0 # if there are no chunks
  370. for iteration, element in enumerate(elements):
  371. if isinstance(element, FallbackElement):
  372. text = element.text
  373. metadata = copy(document.metadata)
  374. metadata.update(element.metadata)
  375. else:
  376. element_dict = (
  377. element if isinstance(element, dict) else element.to_dict()
  378. )
  379. text = element_dict.get("text", "")
  380. if text == "":
  381. continue
  382. metadata = copy(document.metadata)
  383. for key, value in element_dict.items():
  384. if key == "metadata":
  385. for k, v in value.items():
  386. if k not in metadata and k != "orig_elements":
  387. metadata[f"unstructured_{k}"] = v
  388. # indicate that the document was chunked using unstructured
  389. # nullifies the need for chunking in the pipeline
  390. metadata["partitioned_by_unstructured"] = True
  391. metadata["chunk_order"] = iteration
  392. # creating the text extraction
  393. yield DocumentChunk(
  394. id=generate_extraction_id(document.id, iteration),
  395. document_id=document.id,
  396. owner_id=document.owner_id,
  397. collection_ids=document.collection_ids,
  398. data=text,
  399. metadata=metadata,
  400. )
  401. logger.debug(
  402. f"Parsed document with id={document.id}, title={document.metadata.get('title', None)}, "
  403. f"user_id={document.metadata.get('user_id', None)}, metadata={document.metadata} "
  404. f"into {iteration + 1} extractions in t={time.time() - t0:.2f} seconds."
  405. )
  406. def get_parser_for_document_type(self, doc_type: DocumentType) -> str:
  407. return "unstructured_api"