mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-07-31 21:13:49 +08:00
Feat: Add knowledge compilation workflows (#16515)
## Summary - Add knowledge compilation template APIs, services, and builtin template seed data - Add advanced knowledge compile structure/artifact/RAPTOR workflow support - Update parsing, dataset/document APIs, and supporting services for compilation workflows
This commit is contained in:
@@ -37,7 +37,95 @@ from rag.nlp import search
|
||||
|
||||
CANVAS_DEBUG_DOC_ID = "dataflow_x"
|
||||
GRAPH_RAPTOR_FAKE_DOC_ID = "graph_raptor_x"
|
||||
TASK_MAX_LOG_LENGTH = int(os.environ.get("TASK_MAX_LOG_LENGTH", 3000)) # TEXT MAX is 64 KiB bytes!
|
||||
TASK_MAX_LOG_LENGTH = int(os.environ.get("TASK_MAX_LOG_LENGTH", 3000)) # TEXT MAX is 64 KiB bytes!
|
||||
DOC_CHUNKING_COUNTER_TTL_SECONDS = 7 * 24 * 3600
|
||||
|
||||
|
||||
def _doc_chunking_pending_key(doc_id: str) -> str:
|
||||
return f"doc:chunking_pending:{doc_id}"
|
||||
|
||||
|
||||
def _doc_chunking_aborted_key(doc_id: str) -> str:
|
||||
return f"doc:chunking_aborted:{doc_id}"
|
||||
|
||||
|
||||
def _doc_chunking_done_key(task_id: str) -> str:
|
||||
return f"doc:chunking_done:{task_id}"
|
||||
|
||||
|
||||
def seed_doc_chunking_counter(doc_id: str, pending_count: int) -> bool:
|
||||
if not doc_id or pending_count <= 0:
|
||||
return False
|
||||
try:
|
||||
REDIS_CONN.delete(_doc_chunking_aborted_key(doc_id))
|
||||
return REDIS_CONN.set(
|
||||
_doc_chunking_pending_key(doc_id),
|
||||
str(pending_count),
|
||||
exp=DOC_CHUNKING_COUNTER_TTL_SECONDS,
|
||||
)
|
||||
except Exception:
|
||||
logging.exception("Failed to seed chunking counter for doc %s", doc_id)
|
||||
return False
|
||||
|
||||
|
||||
def clear_doc_chunking_counter(doc_id: str) -> None:
|
||||
if not doc_id:
|
||||
return
|
||||
try:
|
||||
REDIS_CONN.delete(_doc_chunking_pending_key(doc_id))
|
||||
except Exception:
|
||||
logging.exception("Failed to clear chunking counter for doc %s", doc_id)
|
||||
|
||||
|
||||
def abort_doc_chunking_counter(doc_id: str) -> None:
|
||||
if not doc_id:
|
||||
return
|
||||
try:
|
||||
REDIS_CONN.delete(_doc_chunking_pending_key(doc_id))
|
||||
REDIS_CONN.set(
|
||||
_doc_chunking_aborted_key(doc_id),
|
||||
"1",
|
||||
exp=DOC_CHUNKING_COUNTER_TTL_SECONDS,
|
||||
)
|
||||
except Exception:
|
||||
logging.exception("Failed to abort chunking counter for doc %s", doc_id)
|
||||
|
||||
|
||||
def is_doc_chunking_aborted(doc_id: str) -> bool:
|
||||
if not doc_id:
|
||||
return False
|
||||
try:
|
||||
return bool(REDIS_CONN.get(_doc_chunking_aborted_key(doc_id)))
|
||||
except Exception:
|
||||
logging.exception("Failed to read chunking abort marker for doc %s", doc_id)
|
||||
return False
|
||||
|
||||
|
||||
def credit_doc_chunking_task(doc_id: str, task_id: str) -> int | None:
|
||||
"""Credit one completed standard chunking task.
|
||||
|
||||
Returns the post-decrement pending count when this task was credited for
|
||||
the first time. Returns a positive value when this task was already
|
||||
credited, so callers treat retries as not-last.
|
||||
"""
|
||||
if not doc_id or not task_id:
|
||||
return None
|
||||
try:
|
||||
first_credit = REDIS_CONN.set_if_absent(
|
||||
_doc_chunking_done_key(task_id),
|
||||
"1",
|
||||
exp=DOC_CHUNKING_COUNTER_TTL_SECONDS,
|
||||
)
|
||||
if not first_credit:
|
||||
return 1
|
||||
pending_key = _doc_chunking_pending_key(doc_id)
|
||||
if REDIS_CONN.get(pending_key) is None:
|
||||
return -1
|
||||
return REDIS_CONN.decrby(pending_key, 1)
|
||||
except Exception:
|
||||
logging.exception("Failed to credit chunking task %s for doc %s", task_id, doc_id)
|
||||
return None
|
||||
|
||||
|
||||
def trim_header_by_lines(text: str, max_length) -> str:
|
||||
# Trim header text to maximum length while preserving line breaks
|
||||
@@ -50,8 +138,8 @@ def trim_header_by_lines(text: str, max_length) -> str:
|
||||
if len_text <= max_length:
|
||||
return text
|
||||
for i in range(len_text):
|
||||
if text[i] == '\n' and len_text - i <= max_length:
|
||||
return text[i + 1:]
|
||||
if text[i] == "\n" and len_text - i <= max_length:
|
||||
return text[i + 1 :]
|
||||
return text
|
||||
|
||||
|
||||
@@ -69,6 +157,7 @@ class TaskService(CommonService):
|
||||
Attributes:
|
||||
model: The Task model class for database operations.
|
||||
"""
|
||||
|
||||
model = Task
|
||||
|
||||
@classmethod
|
||||
@@ -96,6 +185,7 @@ class TaskService(CommonService):
|
||||
cls.model.doc_id,
|
||||
cls.model.from_page,
|
||||
cls.model.to_page,
|
||||
cls.model.task_type,
|
||||
cls.model.retry_count,
|
||||
Document.kb_id,
|
||||
Document.parser_id,
|
||||
@@ -116,32 +206,34 @@ class TaskService(CommonService):
|
||||
]
|
||||
docs = (
|
||||
cls.model.select(*fields)
|
||||
.join(Document, on=(doc_id == Document.id))
|
||||
.join(Knowledgebase, on=(Document.kb_id == Knowledgebase.id))
|
||||
.join(Tenant, on=(Knowledgebase.tenant_id == Tenant.id))
|
||||
.where(cls.model.id == task_id)
|
||||
.join(Document, on=(doc_id == Document.id))
|
||||
.join(Knowledgebase, on=(Document.kb_id == Knowledgebase.id))
|
||||
.join(Tenant, on=(Knowledgebase.tenant_id == Tenant.id))
|
||||
.where(cls.model.id == task_id)
|
||||
)
|
||||
docs = list(docs.dicts())
|
||||
if not docs:
|
||||
return None
|
||||
doc = docs[0]
|
||||
|
||||
msg = f"\n{datetime.now().strftime('%H:%M:%S')} Task has been received."
|
||||
prog = random.random() / 10.0
|
||||
if docs[0]["retry_count"] >= 3:
|
||||
if doc["retry_count"] >= 3:
|
||||
msg = "\nERROR: Task is abandoned after 3 times attempts."
|
||||
prog = -1
|
||||
|
||||
cls.model.update(
|
||||
progress_msg=cls.model.progress_msg + msg,
|
||||
progress=prog,
|
||||
retry_count=docs[0]["retry_count"] + 1,
|
||||
).where(cls.model.id == docs[0]["id"]).execute()
|
||||
retry_count=doc["retry_count"] + 1,
|
||||
).where(cls.model.id == doc["id"]).execute()
|
||||
|
||||
if docs[0]["retry_count"] >= 3:
|
||||
abort_doc_chunking_counter(docs[0]["doc_id"])
|
||||
DocumentService.update_by_id(docs[0]["doc_id"], {"progress": -1, "run": TaskStatus.FAIL.value, "update_time": current_timestamp(), "update_date": get_format_time()})
|
||||
return None
|
||||
|
||||
return docs[0]
|
||||
return doc
|
||||
|
||||
@classmethod
|
||||
@DB.connection_context()
|
||||
@@ -165,10 +257,7 @@ class TaskService(CommonService):
|
||||
cls.model.digest,
|
||||
cls.model.chunk_ids,
|
||||
]
|
||||
tasks = (
|
||||
cls.model.select(*fields).order_by(cls.model.from_page.asc(), cls.model.create_time.desc())
|
||||
.where(cls.model.doc_id == doc_id)
|
||||
)
|
||||
tasks = cls.model.select(*fields).order_by(cls.model.from_page.asc(), cls.model.create_time.desc()).where(cls.model.doc_id == doc_id)
|
||||
tasks = list(tasks.dicts())
|
||||
if not tasks:
|
||||
return None
|
||||
@@ -189,20 +278,8 @@ class TaskService(CommonService):
|
||||
list[dict]: List of task dictionaries containing task details.
|
||||
Returns None if no tasks are found.
|
||||
"""
|
||||
fields = [
|
||||
cls.model.id,
|
||||
cls.model.doc_id,
|
||||
cls.model.from_page,
|
||||
cls.model.progress,
|
||||
cls.model.progress_msg,
|
||||
cls.model.digest,
|
||||
cls.model.chunk_ids,
|
||||
cls.model.create_time
|
||||
]
|
||||
tasks = (
|
||||
cls.model.select(*fields).order_by(cls.model.create_time.desc())
|
||||
.where(cls.model.doc_id.in_(doc_ids))
|
||||
)
|
||||
fields = [cls.model.id, cls.model.doc_id, cls.model.from_page, cls.model.progress, cls.model.progress_msg, cls.model.digest, cls.model.chunk_ids, cls.model.create_time]
|
||||
tasks = cls.model.select(*fields).order_by(cls.model.create_time.desc()).where(cls.model.doc_id.in_(doc_ids))
|
||||
tasks = list(tasks.dicts())
|
||||
if not tasks:
|
||||
return None
|
||||
@@ -238,9 +315,7 @@ class TaskService(CommonService):
|
||||
"""
|
||||
with DB.lock("get_task", -1):
|
||||
docs = (
|
||||
cls.model.select(
|
||||
*[Document.id, Document.kb_id, Document.location, File.parent_id]
|
||||
)
|
||||
cls.model.select(*[Document.id, Document.kb_id, Document.location, File.parent_id])
|
||||
.join(Document, on=(cls.model.doc_id == Document.id))
|
||||
.join(
|
||||
File2Document,
|
||||
@@ -326,11 +401,7 @@ class TaskService(CommonService):
|
||||
cls.model.update(progress_msg=progress_msg).where(cls.model.id == id).execute()
|
||||
if "progress" in info:
|
||||
prog = info["progress"]
|
||||
cls.model.update(progress=prog).where(
|
||||
(cls.model.id == id) &
|
||||
((prog >= 1) | ((cls.model.progress != -1) &
|
||||
((prog == -1) | (prog > cls.model.progress))))
|
||||
).execute()
|
||||
cls.model.update(progress=prog).where((cls.model.id == id) & ((prog >= 1) | ((cls.model.progress != -1) & ((prog == -1) | (prog > cls.model.progress))))).execute()
|
||||
else:
|
||||
with DB.lock("update_progress", -1):
|
||||
if info["progress_msg"]:
|
||||
@@ -338,11 +409,7 @@ class TaskService(CommonService):
|
||||
cls.model.update(progress_msg=progress_msg).where(cls.model.id == id).execute()
|
||||
if "progress" in info:
|
||||
prog = info["progress"]
|
||||
cls.model.update(progress=prog).where(
|
||||
(cls.model.id == id) &
|
||||
((prog >= 1) | ((cls.model.progress != -1) &
|
||||
((prog == -1) | (prog > cls.model.progress))))
|
||||
).execute()
|
||||
cls.model.update(progress=prog).where((cls.model.id == id) & ((prog >= 1) | ((cls.model.progress != -1) & ((prog == -1) | (prog > cls.model.progress))))).execute()
|
||||
|
||||
begin_at = task.begin_at
|
||||
if begin_at is not None:
|
||||
@@ -456,18 +523,22 @@ def queue_tasks(doc: dict, bucket: str, name: str, priority: int):
|
||||
if pre_task["chunk_ids"]:
|
||||
pre_chunk_ids.extend(pre_task["chunk_ids"].split())
|
||||
if pre_chunk_ids:
|
||||
settings.docStoreConn.delete({"id": pre_chunk_ids}, search.index_name(chunking_config["tenant_id"]),
|
||||
chunking_config["kb_id"])
|
||||
settings.docStoreConn.delete({"id": pre_chunk_ids}, search.index_name(chunking_config["tenant_id"]), chunking_config["kb_id"])
|
||||
DocumentService.update_by_id(doc["id"], {"chunk_num": ck_num})
|
||||
|
||||
bulk_insert_into_db(Task, parse_task_array, True)
|
||||
DocumentService.begin2parse(doc["id"])
|
||||
|
||||
unfinished_task_array = [task for task in parse_task_array if task["progress"] < 1.0]
|
||||
for unfinished_task in unfinished_task_array:
|
||||
assert REDIS_CONN.queue_product(
|
||||
settings.get_svr_queue_name(priority, suffix), message=unfinished_task
|
||||
), "Can't access Redis. Please check the Redis' status."
|
||||
chunking_n = sum(1 for task in unfinished_task_array if not task.get("task_type"))
|
||||
if chunking_n > 0:
|
||||
assert seed_doc_chunking_counter(doc["id"], chunking_n), "Can't access Redis. Please check the Redis' status."
|
||||
try:
|
||||
for unfinished_task in unfinished_task_array:
|
||||
assert REDIS_CONN.queue_product(settings.get_svr_queue_name(priority, suffix), message=unfinished_task), "Can't access Redis. Please check the Redis' status."
|
||||
except Exception:
|
||||
abort_doc_chunking_counter(doc["id"])
|
||||
raise
|
||||
|
||||
|
||||
def reuse_prev_task_chunks(task: dict, prev_tasks: list[dict], chunking_config: dict):
|
||||
@@ -494,8 +565,7 @@ def reuse_prev_task_chunks(task: dict, prev_tasks: list[dict], chunking_config:
|
||||
idx = 0
|
||||
while idx < len(prev_tasks):
|
||||
prev_task = prev_tasks[idx]
|
||||
if prev_task.get("from_page", 0) == task.get("from_page", 0) \
|
||||
and prev_task.get("digest", 0) == task.get("digest", ""):
|
||||
if prev_task.get("from_page", 0) == task.get("from_page", 0) and prev_task.get("digest", 0) == task.get("digest", ""):
|
||||
break
|
||||
idx += 1
|
||||
|
||||
@@ -506,18 +576,22 @@ def reuse_prev_task_chunks(task: dict, prev_tasks: list[dict], chunking_config:
|
||||
return 0
|
||||
task["chunk_ids"] = prev_task["chunk_ids"]
|
||||
task["progress"] = 1.0
|
||||
if "from_page" in task and "to_page" in task and (int(task['to_page']) - int(task['from_page']) >= 10 ** 6 or (int(task['from_page']) == MAXIMUM_TASK_PAGE_NUMBER and int(task['to_page']) == MAXIMUM_TASK_PAGE_NUMBER)):
|
||||
if (
|
||||
"from_page" in task
|
||||
and "to_page" in task
|
||||
and (int(task["to_page"]) - int(task["from_page"]) >= 10**6 or (int(task["from_page"]) == MAXIMUM_TASK_PAGE_NUMBER and int(task["to_page"]) == MAXIMUM_TASK_PAGE_NUMBER))
|
||||
):
|
||||
task["progress_msg"] = f"Page({task['from_page']}~{task['to_page']}): "
|
||||
else:
|
||||
task["progress_msg"] = ""
|
||||
task["progress_msg"] = " ".join(
|
||||
[datetime.now().strftime("%H:%M:%S"), task["progress_msg"], "Reused previous task's chunks."])
|
||||
task["progress_msg"] = " ".join([datetime.now().strftime("%H:%M:%S"), task["progress_msg"], "Reused previous task's chunks."])
|
||||
prev_task["chunk_ids"] = ""
|
||||
|
||||
return len(task["chunk_ids"].split())
|
||||
|
||||
|
||||
def cancel_all_task_of(doc_id):
|
||||
abort_doc_chunking_counter(doc_id)
|
||||
for t in TaskService.query(doc_id=doc_id):
|
||||
try:
|
||||
REDIS_CONN.set(f"{t.id}-cancel", "x")
|
||||
@@ -535,7 +609,7 @@ def has_canceled(task_id):
|
||||
return False
|
||||
|
||||
|
||||
def queue_dataflow(tenant_id:str, flow_id:str, task_id:str, doc_id:str=CANVAS_DEBUG_DOC_ID, file:dict=None, priority: int=0, rerun:bool=False) -> tuple[bool, str]:
|
||||
def queue_dataflow(tenant_id: str, flow_id: str, task_id: str, doc_id: str = CANVAS_DEBUG_DOC_ID, file: dict = None, priority: int = 0, rerun: bool = False) -> tuple[bool, str]:
|
||||
|
||||
task = dict(
|
||||
id=task_id,
|
||||
@@ -544,7 +618,7 @@ def queue_dataflow(tenant_id:str, flow_id:str, task_id:str, doc_id:str=CANVAS_DE
|
||||
to_page=MAXIMUM_TASK_PAGE_NUMBER,
|
||||
task_type="dataflow" if not rerun else "dataflow_rerun",
|
||||
priority=priority,
|
||||
begin_at= datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
|
||||
begin_at=datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
|
||||
)
|
||||
if doc_id not in [CANVAS_DEBUG_DOC_ID, GRAPH_RAPTOR_FAKE_DOC_ID]:
|
||||
TaskService.model.delete().where(TaskService.model.doc_id == doc_id).execute()
|
||||
@@ -556,9 +630,7 @@ def queue_dataflow(tenant_id:str, flow_id:str, task_id:str, doc_id:str=CANVAS_DE
|
||||
task["dataflow_id"] = flow_id
|
||||
task["file"] = file
|
||||
|
||||
if not REDIS_CONN.queue_product(
|
||||
settings.get_svr_queue_name(priority, "common"), message=task
|
||||
):
|
||||
if not REDIS_CONN.queue_product(settings.get_svr_queue_name(priority, "common"), message=task):
|
||||
return False, "Can't access Redis. Please check the Redis' status."
|
||||
|
||||
return True, ""
|
||||
|
||||
Reference in New Issue
Block a user