Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
word_extractor.py271 linesDownload Raw Back to extractor
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"![image]({dify_config.CONSOLE_API_URL}/files/{upload_file.id}/file-preview)"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