Underground-Digital/Workflow-Engine
0
1"""Abstract interface for document loader implementations."""2 3import datetime4import logging5import mimetypes6import os7import re8import tempfile9import uuid10from urllib.parse import urlparse11from xml.etree import ElementTree12 13import requests14from docx import Document as DocxDocument15 16from configs import dify_config17from core.helper import ssrf_proxy18from core.rag.extractor.extractor_base import BaseExtractor19from core.rag.models.document import Document20from extensions.ext_database import db21from extensions.ext_storage import storage22from models.enums import CreatedByRole23from models.model import UploadFile24 25logger = logging.getLogger(__name__)26 27 28class WordExtractor(BaseExtractor):29 """Load docx files.30 31 32 Args:33 file_path: Path to the file to load.34 """35 36 def __init__(self, file_path: str, tenant_id: str, user_id: str):37 """Initialize with file path."""38 self.file_path = file_path39 self.tenant_id = tenant_id40 self.user_id = user_id41 42 if "~" in self.file_path:43 self.file_path = os.path.expanduser(self.file_path)44 45 # If the file is a web path, download it to a temporary file, and use that46 if not os.path.isfile(self.file_path) and self._is_valid_url(self.file_path):47 r = requests.get(self.file_path)48 49 if r.status_code != 200:50 raise ValueError(f"Check the url of your file; returned status code {r.status_code}")51 52 self.web_path = self.file_path53 # TODO: use a better way to handle the file54 self.temp_file = tempfile.NamedTemporaryFile() # noqa: SIM11555 self.temp_file.write(r.content)56 self.file_path = self.temp_file.name57 elif not os.path.isfile(self.file_path):58 raise ValueError(f"File path {self.file_path} is not a valid file or url")59 60 def __del__(self) -> None:61 if hasattr(self, "temp_file"):62 self.temp_file.close()63 64 def extract(self) -> list[Document]:65 """Load given path as single page."""66 content = self.parse_docx(self.file_path, "storage")67 return [68 Document(69 page_content=content,70 metadata={"source": self.file_path},71 )72 ]73 74 @staticmethod75 def _is_valid_url(url: str) -> bool:76 """Check if the url is valid."""77 parsed = urlparse(url)78 return bool(parsed.netloc) and bool(parsed.scheme)79 80 def _extract_images_from_docx(self, doc, image_folder):81 os.makedirs(image_folder, exist_ok=True)82 image_count = 083 image_map = {}84 85 for rel in doc.part.rels.values():86 if "image" in rel.target_ref:87 image_count += 188 if rel.is_external:89 url = rel.reltype90 response = ssrf_proxy.get(url, stream=True)91 if response.status_code == 200:92 image_ext = mimetypes.guess_extension(response.headers["Content-Type"])93 file_uuid = str(uuid.uuid4())94 file_key = "image_files/" + self.tenant_id + "/" + file_uuid + "." + image_ext95 mime_type, _ = mimetypes.guess_type(file_key)96 storage.save(file_key, response.content)97 else:98 continue99 else:100 image_ext = rel.target_ref.split(".")[-1]101 # user uuid as file name102 file_uuid = str(uuid.uuid4())103 file_key = "image_files/" + self.tenant_id + "/" + file_uuid + "." + image_ext104 mime_type, _ = mimetypes.guess_type(file_key)105 106 storage.save(file_key, rel.target_part.blob)107 # save file to db108 upload_file = UploadFile(109 tenant_id=self.tenant_id,110 storage_type=dify_config.STORAGE_TYPE,111 key=file_key,112 name=file_key,113 size=0,114 extension=str(image_ext),115 mime_type=mime_type or "",116 created_by=self.user_id,117 created_by_role=CreatedByRole.ACCOUNT,118 created_at=datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None),119 used=True,120 used_by=self.user_id,121 used_at=datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None),122 )123 124 db.session.add(upload_file)125 db.session.commit()126 image_map[rel.target_part] = (127 f""128 )129 130 return image_map131 132 def _table_to_markdown(self, table, image_map):133 markdown = []134 # calculate the total number of columns135 total_cols = max(len(row.cells) for row in table.rows)136 137 header_row = table.rows[0]138 headers = self._parse_row(header_row, image_map, total_cols)139 markdown.append("| " + " | ".join(headers) + " |")140 markdown.append("| " + " | ".join(["---"] * total_cols) + " |")141 142 for row in table.rows[1:]:143 row_cells = self._parse_row(row, image_map, total_cols)144 markdown.append("| " + " | ".join(row_cells) + " |")145 return "\n".join(markdown)146 147 def _parse_row(self, row, image_map, total_cols):148 # Initialize a row, all of which are empty by default149 row_cells = [""] * total_cols150 col_index = 0151 for cell in row.cells:152 # make sure the col_index is not out of range153 while col_index < total_cols and row_cells[col_index] != "":154 col_index += 1155 # if col_index is out of range the loop is jumped156 if col_index >= total_cols:157 break158 cell_content = self._parse_cell(cell, image_map).strip()159 cell_colspan = cell.grid_span or 1160 for i in range(cell_colspan):161 if col_index + i < total_cols:162 row_cells[col_index + i] = cell_content if i == 0 else ""163 col_index += cell_colspan164 return row_cells165 166 def _parse_cell(self, cell, image_map):167 cell_content = []168 for paragraph in cell.paragraphs:169 parsed_paragraph = self._parse_cell_paragraph(paragraph, image_map)170 if parsed_paragraph:171 cell_content.append(parsed_paragraph)172 unique_content = list(dict.fromkeys(cell_content))173 return " ".join(unique_content)174 175 def _parse_cell_paragraph(self, paragraph, image_map):176 paragraph_content = []177 for run in paragraph.runs:178 if run.element.xpath(".//a:blip"):179 for blip in run.element.xpath(".//a:blip"):180 image_id = blip.get("{http://schemas.openxmlformats.org/officeDocument/2006/relationships}embed")181 if not image_id:182 continue183 image_part = paragraph.part.rels[image_id].target_part184 185 if image_part in image_map:186 image_link = image_map[image_part]187 paragraph_content.append(image_link)188 else:189 paragraph_content.append(run.text)190 return "".join(paragraph_content).strip()191 192 def _parse_paragraph(self, paragraph, image_map):193 paragraph_content = []194 for run in paragraph.runs:195 if run.element.xpath(".//a:blip"):196 for blip in run.element.xpath(".//a:blip"):197 embed_id = blip.get("{http://schemas.openxmlformats.org/officeDocument/2006/relationships}embed")198 if embed_id:199 rel_target = run.part.rels[embed_id].target_ref200 if rel_target in image_map:201 paragraph_content.append(image_map[rel_target])202 if run.text.strip():203 paragraph_content.append(run.text.strip())204 return " ".join(paragraph_content) if paragraph_content else ""205 206 def parse_docx(self, docx_path, image_folder):207 doc = DocxDocument(docx_path)208 os.makedirs(image_folder, exist_ok=True)209 210 content = []211 212 image_map = self._extract_images_from_docx(doc, image_folder)213 214 hyperlinks_url = None215 url_pattern = re.compile(r"http://[^\s+]+//|https://[^\s+]+")216 for para in doc.paragraphs:217 for run in para.runs:218 if run.text and hyperlinks_url:219 result = f" [{run.text}]({hyperlinks_url}) "220 run.text = result221 hyperlinks_url = None222 if "HYPERLINK" in run.element.xml:223 try:224 xml = ElementTree.XML(run.element.xml)225 x_child = [c for c in xml.iter() if c is not None]226 for x in x_child:227 if x_child is None:228 continue229 if x.tag.endswith("instrText"):230 for i in url_pattern.findall(x.text):231 hyperlinks_url = str(i)232 except Exception as e:233 logger.error(e)234 235 def parse_paragraph(paragraph):236 paragraph_content = []237 for run in paragraph.runs:238 if hasattr(run.element, "tag") and isinstance(run.element.tag, str) and run.element.tag.endswith("r"):239 drawing_elements = run.element.findall(240 ".//{http://schemas.openxmlformats.org/wordprocessingml/2006/main}drawing"241 )242 for drawing in drawing_elements:243 blip_elements = drawing.findall(244 ".//{http://schemas.openxmlformats.org/drawingml/2006/main}blip"245 )246 for blip in blip_elements:247 embed_id = blip.get(248 "{http://schemas.openxmlformats.org/officeDocument/2006/relationships}embed"249 )250 if embed_id:251 image_part = doc.part.related_parts.get(embed_id)252 if image_part in image_map:253 paragraph_content.append(image_map[image_part])254 if run.text.strip():255 paragraph_content.append(run.text.strip())256 return "".join(paragraph_content) if paragraph_content else ""257 258 paragraphs = doc.paragraphs.copy()259 tables = doc.tables.copy()260 for element in doc.element.body:261 if hasattr(element, "tag"):262 if isinstance(element.tag, str) and element.tag.endswith("p"): # paragraph263 para = paragraphs.pop(0)264 parsed_paragraph = parse_paragraph(para)265 if parsed_paragraph:266 content.append(parsed_paragraph)267 elif isinstance(element.tag, str) and element.tag.endswith("tbl"): # table268 table = tables.pop(0)269 content.append(self._table_to_markdown(table, image_map))270 return "\n".join(content)271 