indexing_runner.py 36 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869
  1. import concurrent.futures
  2. import datetime
  3. import json
  4. import logging
  5. import re
  6. import threading
  7. import time
  8. import uuid
  9. from typing import Optional, cast
  10. from flask import Flask, current_app
  11. from flask_login import current_user
  12. from sqlalchemy.orm.exc import ObjectDeletedError
  13. from configs import dify_config
  14. from core.errors.error import ProviderTokenNotInitError
  15. from core.llm_generator.llm_generator import LLMGenerator
  16. from core.model_manager import ModelInstance, ModelManager
  17. from core.model_runtime.entities.model_entities import ModelType
  18. from core.rag.cleaner.clean_processor import CleanProcessor
  19. from core.rag.datasource.keyword.keyword_factory import Keyword
  20. from core.rag.docstore.dataset_docstore import DatasetDocumentStore
  21. from core.rag.extractor.entity.extract_setting import ExtractSetting
  22. from core.rag.index_processor.index_processor_base import BaseIndexProcessor
  23. from core.rag.index_processor.index_processor_factory import IndexProcessorFactory
  24. from core.rag.models.document import Document
  25. from core.rag.splitter.fixed_text_splitter import (
  26. EnhanceRecursiveCharacterTextSplitter,
  27. FixedRecursiveCharacterTextSplitter,
  28. )
  29. from core.rag.splitter.text_splitter import TextSplitter
  30. from extensions.ext_database import db
  31. from extensions.ext_redis import redis_client
  32. from extensions.ext_storage import storage
  33. from libs import helper
  34. from models.dataset import Dataset, DatasetProcessRule, DocumentSegment
  35. from models.dataset import Document as DatasetDocument
  36. from models.model import UploadFile
  37. from services.feature_service import FeatureService
  38. class IndexingRunner:
  39. def __init__(self):
  40. self.storage = storage
  41. self.model_manager = ModelManager()
  42. def run(self, tenant_id, dataset_documents: list[DatasetDocument]):
  43. """Run the indexing process."""
  44. for dataset_document in dataset_documents:
  45. try:
  46. # get dataset
  47. dataset = Dataset.query.filter_by(id=dataset_document.dataset_id).first()
  48. if not dataset:
  49. raise ValueError("no dataset found")
  50. # get the process rule
  51. processing_rule = (
  52. db.session.query(DatasetProcessRule)
  53. .filter(DatasetProcessRule.id == dataset_document.dataset_process_rule_id)
  54. .first()
  55. )
  56. index_type = dataset_document.doc_form
  57. index_processor = IndexProcessorFactory(index_type).init_index_processor()
  58. # extract
  59. text_docs = self._extract(index_processor, dataset_document, processing_rule.to_dict())
  60. # transform
  61. documents = self._transform(
  62. index_processor, dataset, text_docs, dataset_document.doc_language, processing_rule.to_dict()
  63. )
  64. # save segment
  65. self._load_segments(dataset, dataset_document, documents)
  66. # load
  67. self._load(
  68. tenant_id=tenant_id,
  69. index_processor=index_processor,
  70. dataset=dataset,
  71. dataset_document=dataset_document,
  72. documents=documents,
  73. )
  74. except DocumentIsPausedError:
  75. raise DocumentIsPausedError("Document paused, document id: {}".format(dataset_document.id))
  76. except ProviderTokenNotInitError as e:
  77. dataset_document.indexing_status = "error"
  78. dataset_document.error = str(e.description)
  79. dataset_document.stopped_at = datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None)
  80. db.session.commit()
  81. except ObjectDeletedError:
  82. logging.warning("Document deleted, document id: {}".format(dataset_document.id))
  83. except Exception as e:
  84. logging.exception("consume document failed")
  85. dataset_document.indexing_status = "error"
  86. dataset_document.error = str(e)
  87. dataset_document.stopped_at = datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None)
  88. db.session.commit()
  89. def run_in_splitting_status(self, dataset_document: DatasetDocument):
  90. """Run the indexing process when the index_status is splitting."""
  91. try:
  92. # get dataset
  93. dataset = Dataset.query.filter_by(id=dataset_document.dataset_id).first()
  94. if not dataset:
  95. raise ValueError("no dataset found")
  96. # get exist document_segment list and delete
  97. document_segments = DocumentSegment.query.filter_by(
  98. dataset_id=dataset.id, document_id=dataset_document.id
  99. ).all()
  100. for document_segment in document_segments:
  101. db.session.delete(document_segment)
  102. db.session.commit()
  103. # get the process rule
  104. processing_rule = (
  105. db.session.query(DatasetProcessRule)
  106. .filter(DatasetProcessRule.id == dataset_document.dataset_process_rule_id)
  107. .first()
  108. )
  109. index_type = dataset_document.doc_form
  110. index_processor = IndexProcessorFactory(index_type).init_index_processor()
  111. # extract
  112. text_docs = self._extract(index_processor, dataset_document, processing_rule.to_dict())
  113. # transform
  114. documents = self._transform(
  115. index_processor, dataset, text_docs, dataset_document.doc_language, processing_rule.to_dict()
  116. )
  117. # save segment
  118. self._load_segments(dataset, dataset_document, documents)
  119. # load
  120. self._load(
  121. index_processor=index_processor, dataset=dataset, dataset_document=dataset_document, documents=documents
  122. )
  123. except DocumentIsPausedError:
  124. raise DocumentIsPausedError("Document paused, document id: {}".format(dataset_document.id))
  125. except ProviderTokenNotInitError as e:
  126. dataset_document.indexing_status = "error"
  127. dataset_document.error = str(e.description)
  128. dataset_document.stopped_at = datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None)
  129. db.session.commit()
  130. except Exception as e:
  131. logging.exception("consume document failed")
  132. dataset_document.indexing_status = "error"
  133. dataset_document.error = str(e)
  134. dataset_document.stopped_at = datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None)
  135. db.session.commit()
  136. def run_in_indexing_status(self, dataset_document: DatasetDocument):
  137. """Run the indexing process when the index_status is indexing."""
  138. try:
  139. # get dataset
  140. dataset = Dataset.query.filter_by(id=dataset_document.dataset_id).first()
  141. if not dataset:
  142. raise ValueError("no dataset found")
  143. # get exist document_segment list and delete
  144. document_segments = DocumentSegment.query.filter_by(
  145. dataset_id=dataset.id, document_id=dataset_document.id
  146. ).all()
  147. documents = []
  148. if document_segments:
  149. for document_segment in document_segments:
  150. # transform segment to node
  151. if document_segment.status != "completed":
  152. document = Document(
  153. page_content=document_segment.content,
  154. metadata={
  155. "doc_id": document_segment.index_node_id,
  156. "doc_hash": document_segment.index_node_hash,
  157. "document_id": document_segment.document_id,
  158. "dataset_id": document_segment.dataset_id,
  159. },
  160. )
  161. documents.append(document)
  162. # build index
  163. # get the process rule
  164. processing_rule = (
  165. db.session.query(DatasetProcessRule)
  166. .filter(DatasetProcessRule.id == dataset_document.dataset_process_rule_id)
  167. .first()
  168. )
  169. index_type = dataset_document.doc_form
  170. index_processor = IndexProcessorFactory(index_type).init_index_processor()
  171. self._load(
  172. index_processor=index_processor, dataset=dataset, dataset_document=dataset_document, documents=documents
  173. )
  174. except DocumentIsPausedError:
  175. raise DocumentIsPausedError("Document paused, document id: {}".format(dataset_document.id))
  176. except ProviderTokenNotInitError as e:
  177. dataset_document.indexing_status = "error"
  178. dataset_document.error = str(e.description)
  179. dataset_document.stopped_at = datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None)
  180. db.session.commit()
  181. except Exception as e:
  182. logging.exception("consume document failed")
  183. dataset_document.indexing_status = "error"
  184. dataset_document.error = str(e)
  185. dataset_document.stopped_at = datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None)
  186. db.session.commit()
  187. def indexing_estimate(
  188. self,
  189. tenant_id: str,
  190. extract_settings: list[ExtractSetting],
  191. tmp_processing_rule: dict,
  192. doc_form: Optional[str] = None,
  193. doc_language: str = "English",
  194. dataset_id: Optional[str] = None,
  195. indexing_technique: str = "economy",
  196. ) -> dict:
  197. """
  198. Estimate the indexing for the document.
  199. """
  200. # check document limit
  201. features = FeatureService.get_features(tenant_id)
  202. if features.billing.enabled:
  203. count = len(extract_settings)
  204. batch_upload_limit = dify_config.BATCH_UPLOAD_LIMIT
  205. if count > batch_upload_limit:
  206. raise ValueError(f"You have reached the batch upload limit of {batch_upload_limit}.")
  207. embedding_model_instance = None
  208. if dataset_id:
  209. dataset = Dataset.query.filter_by(id=dataset_id).first()
  210. if not dataset:
  211. raise ValueError("Dataset not found.")
  212. if dataset.indexing_technique == "high_quality" or indexing_technique == "high_quality":
  213. if dataset.embedding_model_provider:
  214. embedding_model_instance = self.model_manager.get_model_instance(
  215. tenant_id=tenant_id,
  216. provider=dataset.embedding_model_provider,
  217. model_type=ModelType.TEXT_EMBEDDING,
  218. model=dataset.embedding_model,
  219. )
  220. else:
  221. embedding_model_instance = self.model_manager.get_default_model_instance(
  222. tenant_id=tenant_id,
  223. model_type=ModelType.TEXT_EMBEDDING,
  224. )
  225. else:
  226. if indexing_technique == "high_quality":
  227. embedding_model_instance = self.model_manager.get_default_model_instance(
  228. tenant_id=tenant_id,
  229. model_type=ModelType.TEXT_EMBEDDING,
  230. )
  231. preview_texts = []
  232. total_segments = 0
  233. index_type = doc_form
  234. index_processor = IndexProcessorFactory(index_type).init_index_processor()
  235. all_text_docs = []
  236. for extract_setting in extract_settings:
  237. # extract
  238. text_docs = index_processor.extract(extract_setting, process_rule_mode=tmp_processing_rule["mode"])
  239. all_text_docs.extend(text_docs)
  240. processing_rule = DatasetProcessRule(
  241. mode=tmp_processing_rule["mode"], rules=json.dumps(tmp_processing_rule["rules"])
  242. )
  243. # get splitter
  244. splitter = self._get_splitter(processing_rule, embedding_model_instance)
  245. # split to documents
  246. documents = self._split_to_documents_for_estimate(
  247. text_docs=text_docs, splitter=splitter, processing_rule=processing_rule
  248. )
  249. total_segments += len(documents)
  250. for document in documents:
  251. if len(preview_texts) < 5:
  252. preview_texts.append(document.page_content)
  253. if doc_form and doc_form == "qa_model":
  254. if len(preview_texts) > 0:
  255. # qa model document
  256. response = LLMGenerator.generate_qa_document(
  257. current_user.current_tenant_id, preview_texts[0], doc_language
  258. )
  259. document_qa_list = self.format_split_text(response)
  260. return {"total_segments": total_segments * 20, "qa_preview": document_qa_list, "preview": preview_texts}
  261. return {"total_segments": total_segments, "preview": preview_texts}
  262. def _extract(
  263. self, index_processor: BaseIndexProcessor, dataset_document: DatasetDocument, process_rule: dict
  264. ) -> list[Document]:
  265. # load file
  266. if dataset_document.data_source_type not in {"upload_file", "notion_import", "website_crawl"}:
  267. return []
  268. data_source_info = dataset_document.data_source_info_dict
  269. text_docs = []
  270. if dataset_document.data_source_type == "upload_file":
  271. if not data_source_info or "upload_file_id" not in data_source_info:
  272. raise ValueError("no upload file found")
  273. file_detail = (
  274. db.session.query(UploadFile).filter(UploadFile.id == data_source_info["upload_file_id"]).one_or_none()
  275. )
  276. if file_detail:
  277. extract_setting = ExtractSetting(
  278. datasource_type="upload_file", upload_file=file_detail, document_model=dataset_document.doc_form
  279. )
  280. text_docs = index_processor.extract(extract_setting, process_rule_mode=process_rule["mode"])
  281. elif dataset_document.data_source_type == "notion_import":
  282. if (
  283. not data_source_info
  284. or "notion_workspace_id" not in data_source_info
  285. or "notion_page_id" not in data_source_info
  286. ):
  287. raise ValueError("no notion import info found")
  288. extract_setting = ExtractSetting(
  289. datasource_type="notion_import",
  290. notion_info={
  291. "notion_workspace_id": data_source_info["notion_workspace_id"],
  292. "notion_obj_id": data_source_info["notion_page_id"],
  293. "notion_page_type": data_source_info["type"],
  294. "document": dataset_document,
  295. "tenant_id": dataset_document.tenant_id,
  296. },
  297. document_model=dataset_document.doc_form,
  298. )
  299. text_docs = index_processor.extract(extract_setting, process_rule_mode=process_rule["mode"])
  300. elif dataset_document.data_source_type == "website_crawl":
  301. if (
  302. not data_source_info
  303. or "provider" not in data_source_info
  304. or "url" not in data_source_info
  305. or "job_id" not in data_source_info
  306. ):
  307. raise ValueError("no website import info found")
  308. extract_setting = ExtractSetting(
  309. datasource_type="website_crawl",
  310. website_info={
  311. "provider": data_source_info["provider"],
  312. "job_id": data_source_info["job_id"],
  313. "tenant_id": dataset_document.tenant_id,
  314. "url": data_source_info["url"],
  315. "mode": data_source_info["mode"],
  316. "only_main_content": data_source_info["only_main_content"],
  317. },
  318. document_model=dataset_document.doc_form,
  319. )
  320. text_docs = index_processor.extract(extract_setting, process_rule_mode=process_rule["mode"])
  321. # update document status to splitting
  322. self._update_document_index_status(
  323. document_id=dataset_document.id,
  324. after_indexing_status="splitting",
  325. extra_update_params={
  326. DatasetDocument.word_count: sum(len(text_doc.page_content) for text_doc in text_docs),
  327. DatasetDocument.parsing_completed_at: datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None),
  328. },
  329. )
  330. # replace doc id to document model id
  331. text_docs = cast(list[Document], text_docs)
  332. for text_doc in text_docs:
  333. text_doc.metadata["document_id"] = dataset_document.id
  334. text_doc.metadata["dataset_id"] = dataset_document.dataset_id
  335. return text_docs
  336. @staticmethod
  337. def filter_string(text):
  338. text = re.sub(r"<\|", "<", text)
  339. text = re.sub(r"\|>", ">", text)
  340. text = re.sub(r"[\x00-\x08\x0B\x0C\x0E-\x1F\x7F\xEF\xBF\xBE]", "", text)
  341. # Unicode U+FFFE
  342. text = re.sub("\ufffe", "", text)
  343. return text
  344. @staticmethod
  345. def _get_splitter(
  346. processing_rule: DatasetProcessRule, embedding_model_instance: Optional[ModelInstance]
  347. ) -> TextSplitter:
  348. """
  349. Get the NodeParser object according to the processing rule.
  350. """
  351. if processing_rule.mode == "custom":
  352. # The user-defined segmentation rule
  353. rules = json.loads(processing_rule.rules)
  354. segmentation = rules["segmentation"]
  355. max_segmentation_tokens_length = dify_config.INDEXING_MAX_SEGMENTATION_TOKENS_LENGTH
  356. if segmentation["max_tokens"] < 50 or segmentation["max_tokens"] > max_segmentation_tokens_length:
  357. raise ValueError(f"Custom segment length should be between 50 and {max_segmentation_tokens_length}.")
  358. separator = segmentation["separator"]
  359. if separator:
  360. separator = separator.replace("\\n", "\n")
  361. if segmentation.get("chunk_overlap"):
  362. chunk_overlap = segmentation["chunk_overlap"]
  363. else:
  364. chunk_overlap = 0
  365. character_splitter = FixedRecursiveCharacterTextSplitter.from_encoder(
  366. chunk_size=segmentation["max_tokens"],
  367. chunk_overlap=chunk_overlap,
  368. fixed_separator=separator,
  369. separators=["\n\n", "。", ". ", " ", ""],
  370. embedding_model_instance=embedding_model_instance,
  371. )
  372. else:
  373. # Automatic segmentation
  374. character_splitter = EnhanceRecursiveCharacterTextSplitter.from_encoder(
  375. chunk_size=DatasetProcessRule.AUTOMATIC_RULES["segmentation"]["max_tokens"],
  376. chunk_overlap=DatasetProcessRule.AUTOMATIC_RULES["segmentation"]["chunk_overlap"],
  377. separators=["\n\n", "。", ". ", " ", ""],
  378. embedding_model_instance=embedding_model_instance,
  379. )
  380. return character_splitter
  381. def _step_split(
  382. self,
  383. text_docs: list[Document],
  384. splitter: TextSplitter,
  385. dataset: Dataset,
  386. dataset_document: DatasetDocument,
  387. processing_rule: DatasetProcessRule,
  388. ) -> list[Document]:
  389. """
  390. Split the text documents into documents and save them to the document segment.
  391. """
  392. documents = self._split_to_documents(
  393. text_docs=text_docs,
  394. splitter=splitter,
  395. processing_rule=processing_rule,
  396. tenant_id=dataset.tenant_id,
  397. document_form=dataset_document.doc_form,
  398. document_language=dataset_document.doc_language,
  399. )
  400. # save node to document segment
  401. doc_store = DatasetDocumentStore(
  402. dataset=dataset, user_id=dataset_document.created_by, document_id=dataset_document.id
  403. )
  404. # add document segments
  405. doc_store.add_documents(documents)
  406. # update document status to indexing
  407. cur_time = datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None)
  408. self._update_document_index_status(
  409. document_id=dataset_document.id,
  410. after_indexing_status="indexing",
  411. extra_update_params={
  412. DatasetDocument.cleaning_completed_at: cur_time,
  413. DatasetDocument.splitting_completed_at: cur_time,
  414. },
  415. )
  416. # update segment status to indexing
  417. self._update_segments_by_document(
  418. dataset_document_id=dataset_document.id,
  419. update_params={
  420. DocumentSegment.status: "indexing",
  421. DocumentSegment.indexing_at: datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None),
  422. },
  423. )
  424. return documents
  425. def _split_to_documents(
  426. self,
  427. text_docs: list[Document],
  428. splitter: TextSplitter,
  429. processing_rule: DatasetProcessRule,
  430. tenant_id: str,
  431. document_form: str,
  432. document_language: str,
  433. ) -> list[Document]:
  434. """
  435. Split the text documents into nodes.
  436. """
  437. all_documents = []
  438. all_qa_documents = []
  439. for text_doc in text_docs:
  440. # document clean
  441. document_text = self._document_clean(text_doc.page_content, processing_rule)
  442. text_doc.page_content = document_text
  443. # parse document to nodes
  444. documents = splitter.split_documents([text_doc])
  445. split_documents = []
  446. for document_node in documents:
  447. if document_node.page_content.strip():
  448. doc_id = str(uuid.uuid4())
  449. hash = helper.generate_text_hash(document_node.page_content)
  450. document_node.metadata["doc_id"] = doc_id
  451. document_node.metadata["doc_hash"] = hash
  452. # delete Splitter character
  453. page_content = document_node.page_content
  454. if page_content.startswith(".") or page_content.startswith("。"):
  455. page_content = page_content[1:]
  456. else:
  457. page_content = page_content
  458. document_node.page_content = page_content
  459. if document_node.page_content:
  460. split_documents.append(document_node)
  461. all_documents.extend(split_documents)
  462. # processing qa document
  463. if document_form == "qa_model":
  464. for i in range(0, len(all_documents), 10):
  465. threads = []
  466. sub_documents = all_documents[i : i + 10]
  467. for doc in sub_documents:
  468. document_format_thread = threading.Thread(
  469. target=self.format_qa_document,
  470. kwargs={
  471. "flask_app": current_app._get_current_object(),
  472. "tenant_id": tenant_id,
  473. "document_node": doc,
  474. "all_qa_documents": all_qa_documents,
  475. "document_language": document_language,
  476. },
  477. )
  478. threads.append(document_format_thread)
  479. document_format_thread.start()
  480. for thread in threads:
  481. thread.join()
  482. return all_qa_documents
  483. return all_documents
  484. def format_qa_document(self, flask_app: Flask, tenant_id: str, document_node, all_qa_documents, document_language):
  485. format_documents = []
  486. if document_node.page_content is None or not document_node.page_content.strip():
  487. return
  488. with flask_app.app_context():
  489. try:
  490. # qa model document
  491. response = LLMGenerator.generate_qa_document(tenant_id, document_node.page_content, document_language)
  492. document_qa_list = self.format_split_text(response)
  493. qa_documents = []
  494. for result in document_qa_list:
  495. qa_document = Document(
  496. page_content=result["question"], metadata=document_node.metadata.model_copy()
  497. )
  498. doc_id = str(uuid.uuid4())
  499. hash = helper.generate_text_hash(result["question"])
  500. qa_document.metadata["answer"] = result["answer"]
  501. qa_document.metadata["doc_id"] = doc_id
  502. qa_document.metadata["doc_hash"] = hash
  503. qa_documents.append(qa_document)
  504. format_documents.extend(qa_documents)
  505. except Exception as e:
  506. logging.exception("Failed to format qa document")
  507. all_qa_documents.extend(format_documents)
  508. def _split_to_documents_for_estimate(
  509. self, text_docs: list[Document], splitter: TextSplitter, processing_rule: DatasetProcessRule
  510. ) -> list[Document]:
  511. """
  512. Split the text documents into nodes.
  513. """
  514. all_documents = []
  515. for text_doc in text_docs:
  516. # document clean
  517. document_text = self._document_clean(text_doc.page_content, processing_rule)
  518. text_doc.page_content = document_text
  519. # parse document to nodes
  520. documents = splitter.split_documents([text_doc])
  521. split_documents = []
  522. for document in documents:
  523. if document.page_content is None or not document.page_content.strip():
  524. continue
  525. doc_id = str(uuid.uuid4())
  526. hash = helper.generate_text_hash(document.page_content)
  527. document.metadata["doc_id"] = doc_id
  528. document.metadata["doc_hash"] = hash
  529. split_documents.append(document)
  530. all_documents.extend(split_documents)
  531. return all_documents
  532. @staticmethod
  533. def _document_clean(text: str, processing_rule: DatasetProcessRule) -> str:
  534. """
  535. Clean the document text according to the processing rules.
  536. """
  537. if processing_rule.mode == "automatic":
  538. rules = DatasetProcessRule.AUTOMATIC_RULES
  539. else:
  540. rules = json.loads(processing_rule.rules) if processing_rule.rules else {}
  541. document_text = CleanProcessor.clean(text, {"rules": rules})
  542. return document_text
  543. @staticmethod
  544. def format_split_text(text):
  545. regex = r"Q\d+:\s*(.*?)\s*A\d+:\s*([\s\S]*?)(?=Q\d+:|$)"
  546. matches = re.findall(regex, text, re.UNICODE)
  547. return [{"question": q, "answer": re.sub(r"\n\s*", "\n", a.strip())} for q, a in matches if q and a]
  548. def _load(
  549. self,
  550. tenant_id :str,
  551. index_processor: BaseIndexProcessor,
  552. dataset: Dataset,
  553. dataset_document: DatasetDocument,
  554. documents: list[Document],
  555. ) -> None:
  556. """
  557. insert index and update document/segment status to completed
  558. """
  559. embedding_model_instance = None
  560. if dataset.indexing_technique == "high_quality":
  561. embedding_model_instance = self.model_manager.get_model_instance(
  562. tenant_id=dataset.tenant_id,
  563. provider=dataset.embedding_model_provider,
  564. model_type=ModelType.TEXT_EMBEDDING,
  565. model=dataset.embedding_model,
  566. )
  567. # chunk nodes by chunk size
  568. indexing_start_at = time.perf_counter()
  569. tokens = 0
  570. chunk_size = 10
  571. # create keyword index
  572. create_keyword_thread = threading.Thread(
  573. target=self._process_keyword_index,
  574. args=(current_app._get_current_object(), tenant_id, dataset.id, dataset_document.id, documents),
  575. )
  576. create_keyword_thread.start()
  577. if dataset.indexing_technique == "high_quality":
  578. with concurrent.futures.ThreadPoolExecutor(max_workers=10) as executor:
  579. futures = []
  580. for i in range(0, len(documents), chunk_size):
  581. chunk_documents = documents[i : i + chunk_size]
  582. futures.append(
  583. executor.submit(
  584. self._process_chunk,
  585. current_app._get_current_object(),
  586. index_processor,
  587. chunk_documents,
  588. dataset,
  589. dataset_document,
  590. embedding_model_instance,
  591. )
  592. )
  593. for future in futures:
  594. tokens += future.result()
  595. create_keyword_thread.join()
  596. indexing_end_at = time.perf_counter()
  597. # update document status to completed
  598. self._update_document_index_status(
  599. document_id=dataset_document.id,
  600. after_indexing_status="completed",
  601. extra_update_params={
  602. DatasetDocument.tokens: tokens,
  603. DatasetDocument.completed_at: datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None),
  604. DatasetDocument.indexing_latency: indexing_end_at - indexing_start_at,
  605. DatasetDocument.error: None,
  606. },
  607. )
  608. @staticmethod
  609. def _process_keyword_index(flask_app, tenant_id, dataset_id, document_id, documents):
  610. with flask_app.app_context():
  611. dataset = Dataset.query.filter_by(id=dataset_id).first()
  612. if not dataset:
  613. raise ValueError("no dataset found")
  614. keyword = Keyword(dataset)
  615. keyword.create(tenant_id, documents)
  616. if dataset.indexing_technique != "high_quality":
  617. document_ids = [document.metadata["doc_id"] for document in documents]
  618. db.session.query(DocumentSegment).filter(
  619. DocumentSegment.document_id == document_id,
  620. DocumentSegment.dataset_id == dataset_id,
  621. DocumentSegment.index_node_id.in_(document_ids),
  622. DocumentSegment.status == "indexing",
  623. ).update(
  624. {
  625. DocumentSegment.status: "completed",
  626. DocumentSegment.enabled: True,
  627. DocumentSegment.completed_at: datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None),
  628. }
  629. )
  630. db.session.commit()
  631. def _process_chunk(
  632. self, flask_app, index_processor, chunk_documents, dataset, dataset_document, embedding_model_instance
  633. ):
  634. with flask_app.app_context():
  635. # check document is paused
  636. self._check_document_paused_status(dataset_document.id)
  637. tokens = 0
  638. if embedding_model_instance:
  639. tokens += sum(
  640. embedding_model_instance.get_text_embedding_num_tokens([document.page_content])
  641. for document in chunk_documents
  642. )
  643. # load index
  644. index_processor.load(dataset, chunk_documents, with_keywords=False)
  645. document_ids = [document.metadata["doc_id"] for document in chunk_documents]
  646. db.session.query(DocumentSegment).filter(
  647. DocumentSegment.document_id == dataset_document.id,
  648. DocumentSegment.dataset_id == dataset.id,
  649. DocumentSegment.index_node_id.in_(document_ids),
  650. DocumentSegment.status == "indexing",
  651. ).update(
  652. {
  653. DocumentSegment.status: "completed",
  654. DocumentSegment.enabled: True,
  655. DocumentSegment.completed_at: datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None),
  656. }
  657. )
  658. db.session.commit()
  659. return tokens
  660. @staticmethod
  661. def _check_document_paused_status(document_id: str):
  662. indexing_cache_key = "document_{}_is_paused".format(document_id)
  663. result = redis_client.get(indexing_cache_key)
  664. if result:
  665. raise DocumentIsPausedError()
  666. @staticmethod
  667. def _update_document_index_status(
  668. document_id: str, after_indexing_status: str, extra_update_params: Optional[dict] = None
  669. ) -> None:
  670. """
  671. Update the document indexing status.
  672. """
  673. count = DatasetDocument.query.filter_by(id=document_id, is_paused=True).count()
  674. if count > 0:
  675. raise DocumentIsPausedError()
  676. document = DatasetDocument.query.filter_by(id=document_id).first()
  677. if not document:
  678. raise DocumentIsDeletedPausedError()
  679. update_params = {DatasetDocument.indexing_status: after_indexing_status}
  680. if extra_update_params:
  681. update_params.update(extra_update_params)
  682. DatasetDocument.query.filter_by(id=document_id).update(update_params)
  683. db.session.commit()
  684. @staticmethod
  685. def _update_segments_by_document(dataset_document_id: str, update_params: dict) -> None:
  686. """
  687. Update the document segment by document id.
  688. """
  689. DocumentSegment.query.filter_by(document_id=dataset_document_id).update(update_params)
  690. db.session.commit()
  691. @staticmethod
  692. def batch_add_segments(segments: list[DocumentSegment], dataset: Dataset):
  693. """
  694. Batch add segments index processing
  695. """
  696. documents = []
  697. for segment in segments:
  698. document = Document(
  699. page_content=segment.content,
  700. metadata={
  701. "doc_id": segment.index_node_id,
  702. "doc_hash": segment.index_node_hash,
  703. "document_id": segment.document_id,
  704. "dataset_id": segment.dataset_id,
  705. },
  706. )
  707. documents.append(document)
  708. # save vector index
  709. index_type = dataset.doc_form
  710. index_processor = IndexProcessorFactory(index_type).init_index_processor()
  711. index_processor.load(dataset, documents)
  712. def _transform(
  713. self,
  714. index_processor: BaseIndexProcessor,
  715. dataset: Dataset,
  716. text_docs: list[Document],
  717. doc_language: str,
  718. process_rule: dict,
  719. ) -> list[Document]:
  720. # get embedding model instance
  721. embedding_model_instance = None
  722. if dataset.indexing_technique == "high_quality":
  723. if dataset.embedding_model_provider:
  724. embedding_model_instance = self.model_manager.get_model_instance(
  725. tenant_id=dataset.tenant_id,
  726. provider=dataset.embedding_model_provider,
  727. model_type=ModelType.TEXT_EMBEDDING,
  728. model=dataset.embedding_model,
  729. )
  730. else:
  731. embedding_model_instance = self.model_manager.get_default_model_instance(
  732. tenant_id=dataset.tenant_id,
  733. model_type=ModelType.TEXT_EMBEDDING,
  734. )
  735. documents = index_processor.transform(
  736. text_docs,
  737. embedding_model_instance=embedding_model_instance,
  738. process_rule=process_rule,
  739. tenant_id=dataset.tenant_id,
  740. doc_language=doc_language,
  741. )
  742. return documents
  743. def _load_segments(self, dataset, dataset_document, documents):
  744. # save node to document segment
  745. doc_store = DatasetDocumentStore(
  746. dataset=dataset, user_id=dataset_document.created_by, document_id=dataset_document.id
  747. )
  748. # add document segments
  749. doc_store.add_documents(documents)
  750. # update document status to indexing
  751. cur_time = datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None)
  752. self._update_document_index_status(
  753. document_id=dataset_document.id,
  754. after_indexing_status="indexing",
  755. extra_update_params={
  756. DatasetDocument.cleaning_completed_at: cur_time,
  757. DatasetDocument.splitting_completed_at: cur_time,
  758. },
  759. )
  760. # update segment status to indexing
  761. self._update_segments_by_document(
  762. dataset_document_id=dataset_document.id,
  763. update_params={
  764. DocumentSegment.status: "indexing",
  765. DocumentSegment.indexing_at: datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None),
  766. },
  767. )
  768. pass
  769. class DocumentIsPausedError(Exception):
  770. pass
  771. class DocumentIsDeletedPausedError(Exception):
  772. pass