Underground-Digital/Workflow-Engine
0
1from collections.abc import Sequence2from typing import Any, Optional3 4from sqlalchemy import func5 6from core.model_manager import ModelManager7from core.model_runtime.entities.model_entities import ModelType8from core.rag.models.document import Document9from extensions.ext_database import db10from models.dataset import Dataset, DocumentSegment11 12 13class DatasetDocumentStore:14 def __init__(15 self,16 dataset: Dataset,17 user_id: str,18 document_id: Optional[str] = None,19 ):20 self._dataset = dataset21 self._user_id = user_id22 self._document_id = document_id23 24 @classmethod25 def from_dict(cls, config_dict: dict[str, Any]) -> "DatasetDocumentStore":26 return cls(**config_dict)27 28 def to_dict(self) -> dict[str, Any]:29 """Serialize to dict."""30 return {31 "dataset_id": self._dataset.id,32 }33 34 @property35 def dateset_id(self) -> Any:36 return self._dataset.id37 38 @property39 def user_id(self) -> Any:40 return self._user_id41 42 @property43 def docs(self) -> dict[str, Document]:44 document_segments = (45 db.session.query(DocumentSegment).filter(DocumentSegment.dataset_id == self._dataset.id).all()46 )47 48 output = {}49 for document_segment in document_segments:50 doc_id = document_segment.index_node_id51 output[doc_id] = Document(52 page_content=document_segment.content,53 metadata={54 "doc_id": document_segment.index_node_id,55 "doc_hash": document_segment.index_node_hash,56 "document_id": document_segment.document_id,57 "dataset_id": document_segment.dataset_id,58 },59 )60 61 return output62 63 def add_documents(self, docs: Sequence[Document], allow_update: bool = True) -> None:64 max_position = (65 db.session.query(func.max(DocumentSegment.position))66 .filter(DocumentSegment.document_id == self._document_id)67 .scalar()68 )69 70 if max_position is None:71 max_position = 072 embedding_model = None73 if self._dataset.indexing_technique == "high_quality":74 model_manager = ModelManager()75 embedding_model = model_manager.get_model_instance(76 tenant_id=self._dataset.tenant_id,77 provider=self._dataset.embedding_model_provider,78 model_type=ModelType.TEXT_EMBEDDING,79 model=self._dataset.embedding_model,80 )81 82 for doc in docs:83 if not isinstance(doc, Document):84 raise ValueError("doc must be a Document")85 86 segment_document = self.get_document_segment(doc_id=doc.metadata["doc_id"])87 88 # NOTE: doc could already exist in the store, but we overwrite it89 if not allow_update and segment_document:90 raise ValueError(91 f"doc_id {doc.metadata['doc_id']} already exists. Set allow_update to True to overwrite."92 )93 94 # calc embedding use tokens95 if embedding_model:96 tokens = embedding_model.get_text_embedding_num_tokens(texts=[doc.page_content])97 else:98 tokens = 099 100 if not segment_document:101 max_position += 1102 103 segment_document = DocumentSegment(104 tenant_id=self._dataset.tenant_id,105 dataset_id=self._dataset.id,106 document_id=self._document_id,107 index_node_id=doc.metadata["doc_id"],108 index_node_hash=doc.metadata["doc_hash"],109 position=max_position,110 content=doc.page_content,111 word_count=len(doc.page_content),112 tokens=tokens,113 enabled=False,114 created_by=self._user_id,115 )116 if doc.metadata.get("answer"):117 segment_document.answer = doc.metadata.pop("answer", "")118 119 db.session.add(segment_document)120 else:121 segment_document.content = doc.page_content122 if doc.metadata.get("answer"):123 segment_document.answer = doc.metadata.pop("answer", "")124 segment_document.index_node_hash = doc.metadata["doc_hash"]125 segment_document.word_count = len(doc.page_content)126 segment_document.tokens = tokens127 128 db.session.commit()129 130 def document_exists(self, doc_id: str) -> bool:131 """Check if document exists."""132 result = self.get_document_segment(doc_id)133 return result is not None134 135 def get_document(self, doc_id: str, raise_error: bool = True) -> Optional[Document]:136 document_segment = self.get_document_segment(doc_id)137 138 if document_segment is None:139 if raise_error:140 raise ValueError(f"doc_id {doc_id} not found.")141 else:142 return None143 144 return Document(145 page_content=document_segment.content,146 metadata={147 "doc_id": document_segment.index_node_id,148 "doc_hash": document_segment.index_node_hash,149 "document_id": document_segment.document_id,150 "dataset_id": document_segment.dataset_id,151 },152 )153 154 def delete_document(self, doc_id: str, raise_error: bool = True) -> None:155 document_segment = self.get_document_segment(doc_id)156 157 if document_segment is None:158 if raise_error:159 raise ValueError(f"doc_id {doc_id} not found.")160 else:161 return None162 163 db.session.delete(document_segment)164 db.session.commit()165 166 def set_document_hash(self, doc_id: str, doc_hash: str) -> None:167 """Set the hash for a given doc_id."""168 document_segment = self.get_document_segment(doc_id)169 170 if document_segment is None:171 return None172 173 document_segment.index_node_hash = doc_hash174 db.session.commit()175 176 def get_document_hash(self, doc_id: str) -> Optional[str]:177 """Get the stored hash for a document, if it exists."""178 document_segment = self.get_document_segment(doc_id)179 180 if document_segment is None:181 return None182 183 return document_segment.index_node_hash184 185 def get_document_segment(self, doc_id: str) -> DocumentSegment:186 document_segment = (187 db.session.query(DocumentSegment)188 .filter(DocumentSegment.dataset_id == self._dataset.id, DocumentSegment.index_node_id == doc_id)189 .first()190 )191 192 return document_segment193 