mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-07-29 04:08:12 +08:00
feat(api): add unified index API and dataset management endpoints (#14222)
### What problem does this PR solve?
## Summary
Refactor the dataset API layer into a clean service/REST separation
pattern, add a unified `/index` API for graph/raptor/mindmap operations,
and introduce several new dataset management endpoints with full test
coverage.
## Changes
### Service Layer (`dataset_api_service.py`)
- Added `trace_index(dataset_id, tenant_id, index_type)` — unified trace
function for all index types
- Added `run_index`, `delete_index` service functions
- Added `get_dataset`, `get_ingestion_summary`, `list_ingestion_logs`,
`get_ingestion_log`
- Added `run_embedding`, `list_tags`, `aggregate_tags`, `delete_tags`,
`rename_tag`
- Added `get_flattened_metadata`, `get_auto_metadata`,
`update_auto_metadata`
### REST API Layer (`dataset_api.py`)
**New unified routes:**
| Method | Route | Description |
|--------|-------|-------------|
| POST | `/datasets/<id>/index?type=graph\|raptor\|mindmap` | Run index
task |
| GET | `/datasets/<id>/index?type=graph\|raptor\|mindmap` | Trace index
task |
| DELETE | `/datasets/<id>/<index_type>` | Delete index |
| GET | `/datasets/<id>` | Get dataset details |
| GET | `/datasets/<id>/ingestions/summary` | Ingestion summary |
| GET | `/datasets/<id>/ingestions` | List ingestion logs |
| GET | `/datasets/<id>/ingestions/<log_id>` | Get single ingestion log
|
| POST | `/datasets/<id>/embedding` | Run embedding |
| GET | `/datasets/<id>/tags` | List tags |
| GET | `/datasets/tags/aggregation` | Aggregate tags across datasets |
| DELETE | `/datasets/<id>/tags` | Delete tags |
| PUT | `/datasets/<id>/tags` | Rename tag |
| GET | `/datasets/metadata/flattened` | Get flattened metadata |
| GET/PUT | `/datasets/<id>/metadata/config` | New metadata config path
|
**Removed routes (replaced by unified `/index`):**
- `POST /datasets/<id>/mindmap`
- `GET /datasets/<id>/mindmap`
**Preserved legacy routes (backward compatibility):**
- `/run_graphrag`, `/trace_graphrag`, `/run_raptor`, `/trace_raptor`
- `/auto_metadata` GET/PUT
### Test Suite
- Updated `common.py` helpers: added `trace_index`, removed
`run_mindmap`/`trace_mindmap`
- Added 7 new test files with 39 test cases total:
| Test File | Cases |
|-----------|-------|
| `test_get_dataset.py` | 4 |
| `test_ingestion_summary.py` | 2 |
| `test_ingestion_logs.py` | 5 |
| `test_index_api.py` | 14 |
| `test_embedding.py` | 2 |
| `test_tags.py` | 8 |
| `test_flattened_metadata.py` | 4 |
- Deleted `test_mindmap_tasks.py` (covered by unified index tests)
## Design Decisions
1. **Unified `/index?type=...`** — single endpoint replaces 3 separate
route pairs for graph/raptor/mindmap
2. **Backward compatibility** — old routes (`/run_graphrag`,
`/run_raptor`, `/auto_metadata`) preserved alongside new paths
3. **`_VALID_INDEX_TYPES = {"graph", "raptor", "mindmap"}`** — input
validation via constant set
4. **`_INDEX_TYPE_TO_TASK_ID_FIELD`** — maps index type to KB model task
ID field for clean dispatch
## Files Changed
- `api/apps/restful_apis/dataset_api.py`
- `api/apps/services/dataset_api_service.py`
- `sdk/python/ragflow_sdk/modules/dataset.py`
- `test/testcases/test_http_api/common.py`
- `test/testcases/test_http_api/test_dataset_management/` (7 new files)
### Type of change
- [x] New Feature (non-breaking change which adds functionality)
- [x] Refactoring
---------
Signed-off-by: noob <yixiao121314@outlook.com>
This commit is contained in:
@@ -13,38 +13,6 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
#
|
||||
import logging
|
||||
import random
|
||||
import re
|
||||
|
||||
from common.metadata_utils import turn2jsonschema
|
||||
from quart import request
|
||||
import numpy as np
|
||||
|
||||
from api.db.services.connector_service import Connector2KbService
|
||||
from api.db.services.llm_service import LLMBundle
|
||||
from api.db.services.document_service import DocumentService, queue_raptor_o_graphrag_tasks
|
||||
from api.db.services.doc_metadata_service import DocMetadataService
|
||||
from api.db.services.pipeline_operation_log_service import PipelineOperationLogService
|
||||
from api.db.services.task_service import TaskService, GRAPH_RAPTOR_FAKE_DOC_ID
|
||||
from api.db.services.user_service import UserTenantService
|
||||
from api.db.joint_services.tenant_model_service import get_model_config_by_type_and_name, get_model_config_by_id
|
||||
from api.utils.api_utils import (
|
||||
get_error_data_result,
|
||||
server_error_response,
|
||||
get_data_error_result,
|
||||
validate_request,
|
||||
get_request_json,
|
||||
)
|
||||
from api.db import VALID_FILE_TYPES
|
||||
from api.db.services.knowledgebase_service import KnowledgebaseService
|
||||
from api.utils.api_utils import get_json_result
|
||||
from rag.nlp import search
|
||||
from rag.utils.redis_conn import REDIS_CONN
|
||||
from common.constants import RetCode, PipelineTaskType, VALID_TASK_STATUS, LLMType
|
||||
from common import settings
|
||||
from common.doc_store.doc_store_base import OrderByExpr
|
||||
from api.apps import login_required, current_user
|
||||
|
||||
"""
|
||||
Deprecated, todo delete
|
||||
@@ -182,52 +150,6 @@ async def update():
|
||||
return server_error_response(e)
|
||||
"""
|
||||
|
||||
@manager.route('/update_metadata_setting', methods=['post']) # noqa: F821
|
||||
@login_required
|
||||
@validate_request("kb_id", "metadata")
|
||||
async def update_metadata_setting():
|
||||
req = await get_request_json()
|
||||
e, kb = KnowledgebaseService.get_by_id(req["kb_id"])
|
||||
if not e:
|
||||
return get_data_error_result(
|
||||
message="Database error (Knowledgebase rename)!")
|
||||
kb = kb.to_dict()
|
||||
kb["parser_config"]["metadata"] = req["metadata"]
|
||||
kb["parser_config"]["enable_metadata"] = req.get("enable_metadata", True)
|
||||
KnowledgebaseService.update_by_id(kb["id"], kb)
|
||||
return get_json_result(data=kb)
|
||||
|
||||
|
||||
@manager.route('/detail', methods=['GET']) # noqa: F821
|
||||
@login_required
|
||||
def detail():
|
||||
kb_id = request.args["kb_id"]
|
||||
try:
|
||||
tenants = UserTenantService.query(user_id=current_user.id)
|
||||
for tenant in tenants:
|
||||
if KnowledgebaseService.query(
|
||||
tenant_id=tenant.tenant_id, id=kb_id):
|
||||
break
|
||||
else:
|
||||
return get_json_result(
|
||||
data=False, message='Only owner of dataset authorized for this operation.',
|
||||
code=RetCode.OPERATING_ERROR)
|
||||
kb = KnowledgebaseService.get_detail(kb_id)
|
||||
if not kb:
|
||||
return get_data_error_result(
|
||||
message="Can't find this dataset!")
|
||||
kb["size"] = DocumentService.get_total_size_by_kb_id(kb_id=kb["id"],keywords="", run_status=[], types=[])
|
||||
kb["connectors"] = Connector2KbService.list_connectors(kb_id)
|
||||
if kb["parser_config"].get("metadata"):
|
||||
kb["parser_config"]["metadata"] = turn2jsonschema(kb["parser_config"]["metadata"])
|
||||
|
||||
for key in ["graphrag_task_finish_at", "raptor_task_finish_at", "mindmap_task_finish_at"]:
|
||||
if finish_at := kb.get(key):
|
||||
kb[key] = finish_at.strftime("%Y-%m-%d %H:%M:%S")
|
||||
return get_json_result(data=kb)
|
||||
except Exception as e:
|
||||
return server_error_response(e)
|
||||
|
||||
"""
|
||||
Deprecated, todo delete
|
||||
@manager.route('/list', methods=['POST']) # noqa: F821
|
||||
@@ -326,80 +248,6 @@ async def rm():
|
||||
return server_error_response(e)
|
||||
"""
|
||||
|
||||
@manager.route('/<kb_id>/tags', methods=['GET']) # noqa: F821
|
||||
@login_required
|
||||
def list_tags(kb_id):
|
||||
if not KnowledgebaseService.accessible(kb_id, current_user.id):
|
||||
return get_json_result(
|
||||
data=False,
|
||||
message='No authorization.',
|
||||
code=RetCode.AUTHENTICATION_ERROR
|
||||
)
|
||||
|
||||
tenants = UserTenantService.get_tenants_by_user_id(current_user.id)
|
||||
tags = []
|
||||
for tenant in tenants:
|
||||
tags += settings.retriever.all_tags(tenant["tenant_id"], [kb_id])
|
||||
return get_json_result(data=tags)
|
||||
|
||||
|
||||
@manager.route('/tags', methods=['GET']) # noqa: F821
|
||||
@login_required
|
||||
def list_tags_from_kbs():
|
||||
kb_ids = request.args.get("kb_ids", "").split(",")
|
||||
for kb_id in kb_ids:
|
||||
if not KnowledgebaseService.accessible(kb_id, current_user.id):
|
||||
return get_json_result(
|
||||
data=False,
|
||||
message='No authorization.',
|
||||
code=RetCode.AUTHENTICATION_ERROR
|
||||
)
|
||||
|
||||
tenants = UserTenantService.get_tenants_by_user_id(current_user.id)
|
||||
tags = []
|
||||
for tenant in tenants:
|
||||
tags += settings.retriever.all_tags(tenant["tenant_id"], kb_ids)
|
||||
return get_json_result(data=tags)
|
||||
|
||||
|
||||
@manager.route('/<kb_id>/rm_tags', methods=['POST']) # noqa: F821
|
||||
@login_required
|
||||
async def rm_tags(kb_id):
|
||||
req = await get_request_json()
|
||||
if not KnowledgebaseService.accessible(kb_id, current_user.id):
|
||||
return get_json_result(
|
||||
data=False,
|
||||
message='No authorization.',
|
||||
code=RetCode.AUTHENTICATION_ERROR
|
||||
)
|
||||
e, kb = KnowledgebaseService.get_by_id(kb_id)
|
||||
|
||||
for t in req["tags"]:
|
||||
settings.docStoreConn.update({"tag_kwd": t, "kb_id": [kb_id]},
|
||||
{"remove": {"tag_kwd": t}},
|
||||
search.index_name(kb.tenant_id),
|
||||
kb_id)
|
||||
return get_json_result(data=True)
|
||||
|
||||
|
||||
@manager.route('/<kb_id>/rename_tag', methods=['POST']) # noqa: F821
|
||||
@login_required
|
||||
async def rename_tags(kb_id):
|
||||
req = await get_request_json()
|
||||
if not KnowledgebaseService.accessible(kb_id, current_user.id):
|
||||
return get_json_result(
|
||||
data=False,
|
||||
message='No authorization.',
|
||||
code=RetCode.AUTHENTICATION_ERROR
|
||||
)
|
||||
e, kb = KnowledgebaseService.get_by_id(kb_id)
|
||||
|
||||
settings.docStoreConn.update({"tag_kwd": req["from_tag"], "kb_id": [kb_id]},
|
||||
{"remove": {"tag_kwd": req["from_tag"].strip()}, "add": {"tag_kwd": req["to_tag"]}},
|
||||
search.index_name(kb.tenant_id),
|
||||
kb_id)
|
||||
return get_json_result(data=True)
|
||||
|
||||
"""
|
||||
Deprecated, todo delete
|
||||
@manager.route('/<kb_id>/knowledge_graph', methods=['GET']) # noqa: F821
|
||||
@@ -457,143 +305,6 @@ def delete_knowledge_graph(kb_id):
|
||||
return get_json_result(data=True)
|
||||
"""
|
||||
|
||||
@manager.route("/get_meta", methods=["GET"]) # noqa: F821
|
||||
@login_required
|
||||
def get_meta():
|
||||
kb_ids = request.args.get("kb_ids", "").split(",")
|
||||
for kb_id in kb_ids:
|
||||
if not KnowledgebaseService.accessible(kb_id, current_user.id):
|
||||
return get_json_result(
|
||||
data=False,
|
||||
message='No authorization.',
|
||||
code=RetCode.AUTHENTICATION_ERROR
|
||||
)
|
||||
return get_json_result(data=DocMetadataService.get_flatted_meta_by_kbs(kb_ids))
|
||||
|
||||
|
||||
@manager.route("/basic_info", methods=["GET"]) # noqa: F821
|
||||
@login_required
|
||||
def get_basic_info():
|
||||
kb_id = request.args.get("kb_id", "")
|
||||
if not KnowledgebaseService.accessible(kb_id, current_user.id):
|
||||
return get_json_result(
|
||||
data=False,
|
||||
message='No authorization.',
|
||||
code=RetCode.AUTHENTICATION_ERROR
|
||||
)
|
||||
|
||||
basic_info = DocumentService.knowledgebase_basic_info(kb_id)
|
||||
|
||||
return get_json_result(data=basic_info)
|
||||
|
||||
|
||||
@manager.route("/list_pipeline_logs", methods=["POST"]) # noqa: F821
|
||||
@login_required
|
||||
async def list_pipeline_logs():
|
||||
kb_id = request.args.get("kb_id")
|
||||
if not kb_id:
|
||||
return get_json_result(data=False, message='Lack of "KB ID"', code=RetCode.ARGUMENT_ERROR)
|
||||
|
||||
keywords = request.args.get("keywords", "")
|
||||
|
||||
page_number = int(request.args.get("page", 0))
|
||||
items_per_page = int(request.args.get("page_size", 0))
|
||||
orderby = request.args.get("orderby", "create_time")
|
||||
if request.args.get("desc", "true").lower() == "false":
|
||||
desc = False
|
||||
else:
|
||||
desc = True
|
||||
create_date_from = request.args.get("create_date_from", "")
|
||||
create_date_to = request.args.get("create_date_to", "")
|
||||
if create_date_to > create_date_from:
|
||||
return get_data_error_result(message="Create data filter is abnormal.")
|
||||
|
||||
req = await get_request_json()
|
||||
|
||||
operation_status = req.get("operation_status", [])
|
||||
if operation_status:
|
||||
invalid_status = {s for s in operation_status if s not in VALID_TASK_STATUS}
|
||||
if invalid_status:
|
||||
return get_data_error_result(message=f"Invalid filter operation_status status conditions: {', '.join(invalid_status)}")
|
||||
|
||||
types = req.get("types", [])
|
||||
if types:
|
||||
invalid_types = {t for t in types if t not in VALID_FILE_TYPES}
|
||||
if invalid_types:
|
||||
return get_data_error_result(message=f"Invalid filter conditions: {', '.join(invalid_types)} type{'s' if len(invalid_types) > 1 else ''}")
|
||||
|
||||
suffix = req.get("suffix", [])
|
||||
|
||||
try:
|
||||
logs, count = PipelineOperationLogService.get_file_logs_by_kb_id(kb_id, page_number, items_per_page, orderby, desc, keywords, operation_status, types, suffix, create_date_from, create_date_to)
|
||||
return get_json_result(data={"total": count, "logs": logs})
|
||||
except Exception as e:
|
||||
return server_error_response(e)
|
||||
|
||||
|
||||
@manager.route("/list_pipeline_dataset_logs", methods=["POST"]) # noqa: F821
|
||||
@login_required
|
||||
async def list_pipeline_dataset_logs():
|
||||
kb_id = request.args.get("kb_id")
|
||||
if not kb_id:
|
||||
return get_json_result(data=False, message='Lack of "KB ID"', code=RetCode.ARGUMENT_ERROR)
|
||||
|
||||
page_number = int(request.args.get("page", 0))
|
||||
items_per_page = int(request.args.get("page_size", 0))
|
||||
orderby = request.args.get("orderby", "create_time")
|
||||
if request.args.get("desc", "true").lower() == "false":
|
||||
desc = False
|
||||
else:
|
||||
desc = True
|
||||
create_date_from = request.args.get("create_date_from", "")
|
||||
create_date_to = request.args.get("create_date_to", "")
|
||||
if create_date_to > create_date_from:
|
||||
return get_data_error_result(message="Create data filter is abnormal.")
|
||||
|
||||
req = await get_request_json()
|
||||
|
||||
operation_status = req.get("operation_status", [])
|
||||
if operation_status:
|
||||
invalid_status = {s for s in operation_status if s not in VALID_TASK_STATUS}
|
||||
if invalid_status:
|
||||
return get_data_error_result(message=f"Invalid filter operation_status status conditions: {', '.join(invalid_status)}")
|
||||
|
||||
try:
|
||||
logs, tol = PipelineOperationLogService.get_dataset_logs_by_kb_id(kb_id, page_number, items_per_page, orderby, desc, operation_status, create_date_from, create_date_to)
|
||||
return get_json_result(data={"total": tol, "logs": logs})
|
||||
except Exception as e:
|
||||
return server_error_response(e)
|
||||
|
||||
|
||||
@manager.route("/delete_pipeline_logs", methods=["POST"]) # noqa: F821
|
||||
@login_required
|
||||
async def delete_pipeline_logs():
|
||||
kb_id = request.args.get("kb_id")
|
||||
if not kb_id:
|
||||
return get_json_result(data=False, message='Lack of "KB ID"', code=RetCode.ARGUMENT_ERROR)
|
||||
|
||||
req = await get_request_json()
|
||||
log_ids = req.get("log_ids", [])
|
||||
|
||||
PipelineOperationLogService.delete_by_ids(log_ids)
|
||||
|
||||
return get_json_result(data=True)
|
||||
|
||||
|
||||
@manager.route("/pipeline_log_detail", methods=["GET"]) # noqa: F821
|
||||
@login_required
|
||||
def pipeline_log_detail():
|
||||
log_id = request.args.get("log_id")
|
||||
if not log_id:
|
||||
return get_json_result(data=False, message='Lack of "Pipeline log ID"', code=RetCode.ARGUMENT_ERROR)
|
||||
|
||||
ok, log = PipelineOperationLogService.get_by_id(log_id)
|
||||
if not ok:
|
||||
return get_data_error_result(message="Invalid pipeline log ID")
|
||||
|
||||
return get_json_result(data=log.to_dict())
|
||||
|
||||
|
||||
"""
|
||||
Deprecated, todo delete
|
||||
@manager.route("/run_graphrag", methods=["POST"]) # noqa: F821
|
||||
@@ -733,280 +444,3 @@ def trace_raptor():
|
||||
|
||||
return get_json_result(data=task.to_dict())
|
||||
"""
|
||||
|
||||
@manager.route("/run_mindmap", methods=["POST"]) # noqa: F821
|
||||
@login_required
|
||||
async def run_mindmap():
|
||||
req = await get_request_json()
|
||||
|
||||
kb_id = req.get("kb_id", "")
|
||||
if not kb_id:
|
||||
return get_error_data_result(message='Lack of "KB ID"')
|
||||
|
||||
ok, kb = KnowledgebaseService.get_by_id(kb_id)
|
||||
if not ok:
|
||||
return get_error_data_result(message="Invalid Knowledgebase ID")
|
||||
|
||||
task_id = kb.mindmap_task_id
|
||||
if task_id:
|
||||
ok, task = TaskService.get_by_id(task_id)
|
||||
if not ok:
|
||||
logging.warning(f"A valid Mindmap task id is expected for kb {kb_id}")
|
||||
|
||||
if task and task.progress not in [-1, 1]:
|
||||
return get_error_data_result(message=f"Task {task_id} in progress with status {task.progress}. A Mindmap Task is already running.")
|
||||
|
||||
documents, _ = DocumentService.get_by_kb_id(
|
||||
kb_id=kb_id,
|
||||
page_number=0,
|
||||
items_per_page=0,
|
||||
orderby="create_time",
|
||||
desc=False,
|
||||
keywords="",
|
||||
run_status=[],
|
||||
types=[],
|
||||
suffix=[],
|
||||
)
|
||||
if not documents:
|
||||
return get_error_data_result(message=f"No documents in Knowledgebase {kb_id}")
|
||||
|
||||
sample_document = documents[0]
|
||||
document_ids = [document["id"] for document in documents]
|
||||
|
||||
task_id = queue_raptor_o_graphrag_tasks(sample_doc=sample_document, ty="mindmap", priority=0, fake_doc_id=GRAPH_RAPTOR_FAKE_DOC_ID, doc_ids=list(document_ids))
|
||||
|
||||
if not KnowledgebaseService.update_by_id(kb.id, {"mindmap_task_id": task_id}):
|
||||
logging.warning(f"Cannot save mindmap_task_id for kb {kb_id}")
|
||||
|
||||
return get_json_result(data={"mindmap_task_id": task_id})
|
||||
|
||||
|
||||
@manager.route("/trace_mindmap", methods=["GET"]) # noqa: F821
|
||||
@login_required
|
||||
def trace_mindmap():
|
||||
kb_id = request.args.get("kb_id", "")
|
||||
if not kb_id:
|
||||
return get_error_data_result(message='Lack of "KB ID"')
|
||||
|
||||
ok, kb = KnowledgebaseService.get_by_id(kb_id)
|
||||
if not ok:
|
||||
return get_error_data_result(message="Invalid Knowledgebase ID")
|
||||
|
||||
task_id = kb.mindmap_task_id
|
||||
if not task_id:
|
||||
return get_json_result(data={})
|
||||
|
||||
ok, task = TaskService.get_by_id(task_id)
|
||||
if not ok:
|
||||
return get_error_data_result(message="Mindmap Task Not Found or Error Occurred")
|
||||
|
||||
return get_json_result(data=task.to_dict())
|
||||
|
||||
|
||||
@manager.route("/unbind_task", methods=["DELETE"]) # noqa: F821
|
||||
@login_required
|
||||
def delete_kb_task():
|
||||
kb_id = request.args.get("kb_id", "")
|
||||
if not kb_id:
|
||||
return get_error_data_result(message='Lack of "KB ID"')
|
||||
ok, kb = KnowledgebaseService.get_by_id(kb_id)
|
||||
if not ok:
|
||||
return get_json_result(data=True)
|
||||
|
||||
pipeline_task_type = request.args.get("pipeline_task_type", "")
|
||||
if not pipeline_task_type or pipeline_task_type not in [PipelineTaskType.GRAPH_RAG, PipelineTaskType.RAPTOR, PipelineTaskType.MINDMAP]:
|
||||
return get_error_data_result(message="Invalid task type")
|
||||
|
||||
def cancel_task(task_id):
|
||||
REDIS_CONN.set(f"{task_id}-cancel", "x")
|
||||
|
||||
kb_task_id_field: str = ""
|
||||
kb_task_finish_at: str = ""
|
||||
match pipeline_task_type:
|
||||
case PipelineTaskType.GRAPH_RAG:
|
||||
kb_task_id_field = "graphrag_task_id"
|
||||
task_id = kb.graphrag_task_id
|
||||
kb_task_finish_at = "graphrag_task_finish_at"
|
||||
cancel_task(task_id)
|
||||
settings.docStoreConn.delete({"knowledge_graph_kwd": ["graph", "subgraph", "entity", "relation"]}, search.index_name(kb.tenant_id), kb_id)
|
||||
case PipelineTaskType.RAPTOR:
|
||||
kb_task_id_field = "raptor_task_id"
|
||||
task_id = kb.raptor_task_id
|
||||
kb_task_finish_at = "raptor_task_finish_at"
|
||||
cancel_task(task_id)
|
||||
settings.docStoreConn.delete({"raptor_kwd": ["raptor"]}, search.index_name(kb.tenant_id), kb_id)
|
||||
case PipelineTaskType.MINDMAP:
|
||||
kb_task_id_field = "mindmap_task_id"
|
||||
task_id = kb.mindmap_task_id
|
||||
kb_task_finish_at = "mindmap_task_finish_at"
|
||||
cancel_task(task_id)
|
||||
case _:
|
||||
return get_error_data_result(message="Internal Error: Invalid task type")
|
||||
|
||||
|
||||
ok = KnowledgebaseService.update_by_id(kb_id, {kb_task_id_field: "", kb_task_finish_at: None})
|
||||
if not ok:
|
||||
return server_error_response(f"Internal error: cannot delete task {pipeline_task_type}")
|
||||
|
||||
return get_json_result(data=True)
|
||||
|
||||
@manager.route("/check_embedding", methods=["post"]) # noqa: F821
|
||||
@login_required
|
||||
async def check_embedding():
|
||||
|
||||
def _guess_vec_field(src: dict) -> str | None:
|
||||
for k in src or {}:
|
||||
if k.endswith("_vec"):
|
||||
return k
|
||||
return None
|
||||
|
||||
def _as_float_vec(v):
|
||||
if v is None:
|
||||
return []
|
||||
if isinstance(v, str):
|
||||
return [float(x) for x in v.split("\t") if x != ""]
|
||||
if isinstance(v, (list, tuple, np.ndarray)):
|
||||
return [float(x) for x in v]
|
||||
return []
|
||||
|
||||
def _to_1d(x):
|
||||
a = np.asarray(x, dtype=np.float32)
|
||||
return a.reshape(-1)
|
||||
|
||||
def _cos_sim(a, b, eps=1e-12):
|
||||
a = _to_1d(a)
|
||||
b = _to_1d(b)
|
||||
na = np.linalg.norm(a)
|
||||
nb = np.linalg.norm(b)
|
||||
if na < eps or nb < eps:
|
||||
return 0.0
|
||||
return float(np.dot(a, b) / (na * nb))
|
||||
|
||||
def sample_random_chunks_with_vectors(
|
||||
docStoreConn,
|
||||
tenant_id: str,
|
||||
kb_id: str,
|
||||
n: int = 5,
|
||||
base_fields=("docnm_kwd","doc_id","content_with_weight","page_num_int","position_int","top_int"),
|
||||
):
|
||||
index_nm = search.index_name(tenant_id)
|
||||
|
||||
res0 = docStoreConn.search(
|
||||
select_fields=[], highlight_fields=[],
|
||||
condition={"kb_id": kb_id, "available_int": 1},
|
||||
match_expressions=[], order_by=OrderByExpr(),
|
||||
offset=0, limit=1,
|
||||
index_names=index_nm, knowledgebase_ids=[kb_id]
|
||||
)
|
||||
total = docStoreConn.get_total(res0)
|
||||
if total <= 0:
|
||||
return []
|
||||
|
||||
n = min(n, total)
|
||||
offsets = sorted(random.sample(range(min(total,1000)), n))
|
||||
out = []
|
||||
|
||||
for off in offsets:
|
||||
res1 = docStoreConn.search(
|
||||
select_fields=list(base_fields),
|
||||
highlight_fields=[],
|
||||
condition={"kb_id": kb_id, "available_int": 1},
|
||||
match_expressions=[], order_by=OrderByExpr(),
|
||||
offset=off, limit=1,
|
||||
index_names=index_nm, knowledgebase_ids=[kb_id]
|
||||
)
|
||||
ids = docStoreConn.get_doc_ids(res1)
|
||||
if not ids:
|
||||
continue
|
||||
|
||||
cid = ids[0]
|
||||
full_doc = docStoreConn.get(cid, index_nm, [kb_id]) or {}
|
||||
vec_field = _guess_vec_field(full_doc)
|
||||
vec = _as_float_vec(full_doc.get(vec_field))
|
||||
|
||||
out.append({
|
||||
"chunk_id": cid,
|
||||
"kb_id": kb_id,
|
||||
"doc_id": full_doc.get("doc_id"),
|
||||
"doc_name": full_doc.get("docnm_kwd"),
|
||||
"vector_field": vec_field,
|
||||
"vector_dim": len(vec),
|
||||
"vector": vec,
|
||||
"page_num_int": full_doc.get("page_num_int"),
|
||||
"position_int": full_doc.get("position_int"),
|
||||
"top_int": full_doc.get("top_int"),
|
||||
"content_with_weight": full_doc.get("content_with_weight") or "",
|
||||
"question_kwd": full_doc.get("question_kwd") or []
|
||||
})
|
||||
return out
|
||||
|
||||
def _clean(s: str) -> str:
|
||||
s = re.sub(r"</?(table|td|caption|tr|th)( [^<>]{0,12})?>", " ", s or "")
|
||||
return s if s else "None"
|
||||
req = await get_request_json()
|
||||
kb_id = req.get("kb_id", "")
|
||||
tenant_embd_id = req.get("tenant_embd_id")
|
||||
embd_id = req.get("embd_id", "")
|
||||
n = int(req.get("check_num", 5))
|
||||
_, kb = KnowledgebaseService.get_by_id(kb_id)
|
||||
tenant_id = kb.tenant_id
|
||||
if tenant_embd_id:
|
||||
embd_model_config = get_model_config_by_id(tenant_embd_id)
|
||||
elif embd_id:
|
||||
embd_model_config = get_model_config_by_type_and_name(tenant_id, LLMType.EMBEDDING, embd_id)
|
||||
else:
|
||||
return get_error_data_result("`tenant_embd_id` or `embd_id` is required.")
|
||||
emb_mdl = LLMBundle(tenant_id, embd_model_config)
|
||||
samples = sample_random_chunks_with_vectors(settings.docStoreConn, tenant_id=tenant_id, kb_id=kb_id, n=n)
|
||||
|
||||
results, eff_sims = [], []
|
||||
for ck in samples:
|
||||
title = ck.get("doc_name") or "Title"
|
||||
txt_in = "\n".join(ck.get("question_kwd") or []) or ck.get("content_with_weight") or ""
|
||||
txt_in = _clean(txt_in)
|
||||
if not txt_in:
|
||||
results.append({"chunk_id": ck["chunk_id"], "reason": "no_text"})
|
||||
continue
|
||||
|
||||
if not ck.get("vector"):
|
||||
results.append({"chunk_id": ck["chunk_id"], "reason": "no_stored_vector"})
|
||||
continue
|
||||
|
||||
try:
|
||||
v, _ = emb_mdl.encode([title, txt_in])
|
||||
assert len(v[1]) == len(ck["vector"]), f"The dimension ({len(v[1])}) of given embedding model is different from the original ({len(ck['vector'])})"
|
||||
sim_content = _cos_sim(v[1], ck["vector"])
|
||||
title_w = 0.1
|
||||
qv_mix = title_w * v[0] + (1 - title_w) * v[1]
|
||||
sim_mix = _cos_sim(qv_mix, ck["vector"])
|
||||
sim = sim_content
|
||||
mode = "content_only"
|
||||
if sim_mix > sim:
|
||||
sim = sim_mix
|
||||
mode = "title+content"
|
||||
except Exception as e:
|
||||
return get_error_data_result(message=f"Embedding failure. {e}")
|
||||
|
||||
eff_sims.append(sim)
|
||||
results.append({
|
||||
"chunk_id": ck["chunk_id"],
|
||||
"doc_id": ck["doc_id"],
|
||||
"doc_name": ck["doc_name"],
|
||||
"vector_field": ck["vector_field"],
|
||||
"vector_dim": ck["vector_dim"],
|
||||
"cos_sim": round(sim, 6),
|
||||
})
|
||||
|
||||
summary = {
|
||||
"kb_id": kb_id,
|
||||
"model": embd_id,
|
||||
"sampled": len(samples),
|
||||
"valid": len(eff_sims),
|
||||
"avg_cos_sim": round(float(np.mean(eff_sims)) if eff_sims else 0.0, 6),
|
||||
"min_cos_sim": round(float(np.min(eff_sims)) if eff_sims else 0.0, 6),
|
||||
"max_cos_sim": round(float(np.max(eff_sims)) if eff_sims else 0.0, 6),
|
||||
"match_mode": mode,
|
||||
}
|
||||
if summary["avg_cos_sim"] > 0.9:
|
||||
return get_json_result(data={"summary": summary, "results": results})
|
||||
return get_json_result(code=RetCode.NOT_EFFECTIVE, message="Embedding model switch failed: the average similarity between old and new vectors is below 0.9, indicating incompatible vector spaces.", data={"summary": summary, "results": results})
|
||||
|
||||
@@ -31,6 +31,50 @@ from api.utils.validation_utils import (
|
||||
from api.apps.services import dataset_api_service
|
||||
|
||||
|
||||
@manager.route("/datasets/tags/aggregation", methods=["GET"]) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
def aggregate_tags(tenant_id):
|
||||
dataset_ids = request.args.get("dataset_ids", "").split(",")
|
||||
dataset_ids = [d for d in dataset_ids if d]
|
||||
if not dataset_ids:
|
||||
return get_error_data_result(message="Lack of dataset_ids in query parameters")
|
||||
|
||||
try:
|
||||
success, result = dataset_api_service.aggregate_tags(dataset_ids, tenant_id)
|
||||
if success:
|
||||
return get_result(data=result)
|
||||
else:
|
||||
return get_error_data_result(message=result)
|
||||
except ValueError as e:
|
||||
return get_error_argument_result(str(e))
|
||||
except Exception as e:
|
||||
logging.exception(e)
|
||||
return get_error_data_result(message="Internal server error")
|
||||
|
||||
|
||||
@manager.route("/datasets/metadata/flattened", methods=["GET"]) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
def get_flattened_metadata(tenant_id):
|
||||
dataset_ids = request.args.get("dataset_ids", "").split(",")
|
||||
dataset_ids = [d for d in dataset_ids if d]
|
||||
if not dataset_ids:
|
||||
return get_error_data_result(message="Lack of dataset_ids in query parameters")
|
||||
|
||||
try:
|
||||
success, result = dataset_api_service.get_flattened_metadata(dataset_ids, tenant_id)
|
||||
if success:
|
||||
return get_result(data=result)
|
||||
else:
|
||||
return get_error_data_result(message=result)
|
||||
except ValueError as e:
|
||||
return get_error_argument_result(str(e))
|
||||
except Exception as e:
|
||||
logging.exception(e)
|
||||
return get_error_data_result(message="Internal server error")
|
||||
|
||||
|
||||
@manager.route("/datasets", methods=["POST"]) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
@@ -102,6 +146,8 @@ async def create(tenant_id: str=None):
|
||||
return get_result(data=result)
|
||||
else:
|
||||
return get_error_data_result(message=result)
|
||||
except ValueError as e:
|
||||
return get_error_argument_result(str(e))
|
||||
except Exception as e:
|
||||
logging.exception(e)
|
||||
return get_error_data_result(message="Internal server error")
|
||||
@@ -330,7 +376,107 @@ def list_datasets(tenant_id):
|
||||
return get_error_data_result(message="Internal server error")
|
||||
|
||||
|
||||
@manager.route('/datasets/<dataset_id>/knowledge_graph', methods=['GET']) # noqa: F821
|
||||
@manager.route("/datasets/<dataset_id>", methods=["GET"]) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
def get_dataset(tenant_id, dataset_id):
|
||||
try:
|
||||
success, result = dataset_api_service.get_dataset(dataset_id, tenant_id)
|
||||
if success:
|
||||
return get_result(data=result)
|
||||
else:
|
||||
return get_error_data_result(message=result)
|
||||
except ValueError as e:
|
||||
return get_error_argument_result(str(e))
|
||||
except Exception as e:
|
||||
logging.exception(e)
|
||||
return get_error_data_result(message="Internal server error")
|
||||
|
||||
|
||||
@manager.route("/datasets/<dataset_id>/ingestions/summary", methods=["GET"]) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
def get_ingestion_summary(tenant_id, dataset_id):
|
||||
try:
|
||||
success, result = dataset_api_service.get_ingestion_summary(dataset_id, tenant_id)
|
||||
if success:
|
||||
return get_result(data=result)
|
||||
else:
|
||||
return get_error_data_result(message=result)
|
||||
except ValueError as e:
|
||||
return get_error_argument_result(str(e))
|
||||
except Exception as e:
|
||||
logging.exception(e)
|
||||
return get_error_data_result(message="Internal server error")
|
||||
|
||||
|
||||
@manager.route("/datasets/<dataset_id>/tags", methods=["GET"]) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
def list_tags(tenant_id, dataset_id):
|
||||
try:
|
||||
success, result = dataset_api_service.list_tags(dataset_id, tenant_id)
|
||||
if success:
|
||||
return get_result(data=result)
|
||||
else:
|
||||
return get_error_data_result(message=result)
|
||||
except ValueError as e:
|
||||
return get_error_argument_result(str(e))
|
||||
except Exception as e:
|
||||
logging.exception(e)
|
||||
return get_error_data_result(message="Internal server error")
|
||||
|
||||
|
||||
@manager.route("/datasets/<dataset_id>/tags", methods=["DELETE"]) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
async def delete_tags(tenant_id, dataset_id):
|
||||
req = await request.get_json()
|
||||
if not req or "tags" not in req:
|
||||
return get_error_data_result(message="Lack of tags in request body")
|
||||
if not isinstance(req["tags"], list) or not all(isinstance(t, str) for t in req["tags"]):
|
||||
return get_error_argument_result("tags must be a list of strings")
|
||||
|
||||
try:
|
||||
success, result = dataset_api_service.delete_tags(dataset_id, tenant_id, req["tags"])
|
||||
if success:
|
||||
return get_result(data=result)
|
||||
else:
|
||||
return get_error_data_result(message=result)
|
||||
except ValueError as e:
|
||||
return get_error_argument_result(str(e))
|
||||
except Exception as e:
|
||||
logging.exception(e)
|
||||
return get_error_data_result(message="Internal server error")
|
||||
|
||||
|
||||
@manager.route("/datasets/<dataset_id>/tags", methods=["PUT"]) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
async def rename_tag(tenant_id, dataset_id):
|
||||
req = await request.get_json()
|
||||
if not req or "from_tag" not in req or "to_tag" not in req:
|
||||
return get_error_data_result(message="Lack of from_tag or to_tag in request body")
|
||||
if not isinstance(req["from_tag"], str) or not isinstance(req["to_tag"], str):
|
||||
return get_error_argument_result("from_tag and to_tag must be strings")
|
||||
|
||||
if not req["from_tag"].strip() or not req["to_tag"].strip():
|
||||
return get_error_argument_result("from_tag and to_tag must not be empty")
|
||||
|
||||
try:
|
||||
success, result = dataset_api_service.rename_tag(dataset_id, tenant_id, req["from_tag"], req["to_tag"])
|
||||
if success:
|
||||
return get_result(data=result)
|
||||
else:
|
||||
return get_error_data_result(message=result)
|
||||
except ValueError as e:
|
||||
return get_error_argument_result(str(e))
|
||||
except Exception as e:
|
||||
logging.exception(e)
|
||||
return get_error_data_result(message="Internal server error")
|
||||
|
||||
|
||||
@manager.route('/datasets/<dataset_id>/graph/search', methods=['GET']) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
async def knowledge_graph(tenant_id, dataset_id):
|
||||
@@ -349,7 +495,7 @@ async def knowledge_graph(tenant_id, dataset_id):
|
||||
return get_error_data_result(message="Internal server error")
|
||||
|
||||
|
||||
@manager.route('/datasets/<dataset_id>/knowledge_graph', methods=['DELETE']) # noqa: F821
|
||||
@manager.route('/datasets/<dataset_id>/graph', methods=['DELETE']) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
def delete_knowledge_graph(tenant_id, dataset_id):
|
||||
@@ -368,12 +514,67 @@ def delete_knowledge_graph(tenant_id, dataset_id):
|
||||
return get_error_data_result(message="Internal server error")
|
||||
|
||||
|
||||
@manager.route("/datasets/<dataset_id>/run_graphrag", methods=["POST"]) # noqa: F821
|
||||
@manager.route("/datasets/<dataset_id>/index", methods=["POST"]) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
async def run_graphrag(tenant_id, dataset_id):
|
||||
async def run_index(tenant_id, dataset_id):
|
||||
index_type = request.args.get("type", "")
|
||||
try:
|
||||
success, result = dataset_api_service.run_graphrag(dataset_id, tenant_id)
|
||||
success, result = dataset_api_service.run_index(dataset_id, tenant_id, index_type)
|
||||
if success:
|
||||
return get_result(data=result)
|
||||
else:
|
||||
return get_error_data_result(message=result)
|
||||
except ValueError as e:
|
||||
return get_error_argument_result(str(e))
|
||||
except Exception as e:
|
||||
logging.exception(e)
|
||||
return get_error_data_result(message="Internal server error")
|
||||
|
||||
|
||||
@manager.route("/datasets/<dataset_id>/index", methods=["GET"]) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
def trace_index(tenant_id, dataset_id):
|
||||
index_type = request.args.get("type", "")
|
||||
try:
|
||||
success, result = dataset_api_service.trace_index(dataset_id, tenant_id, index_type)
|
||||
if success:
|
||||
return get_result(data=result)
|
||||
else:
|
||||
return get_error_data_result(message=result)
|
||||
except ValueError as e:
|
||||
return get_error_argument_result(str(e))
|
||||
except Exception as e:
|
||||
logging.exception(e)
|
||||
return get_error_data_result(message="Internal server error")
|
||||
|
||||
|
||||
@manager.route("/datasets/<dataset_id>/<index_type>", methods=["DELETE"]) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
def delete_index(tenant_id, dataset_id, index_type):
|
||||
if index_type not in dataset_api_service._VALID_INDEX_TYPES:
|
||||
return get_error_argument_result(f"Invalid index type '{index_type}'")
|
||||
try:
|
||||
success, result = dataset_api_service.delete_index(dataset_id, tenant_id, index_type)
|
||||
if success:
|
||||
return get_result(data=result)
|
||||
else:
|
||||
return get_error_data_result(message=result)
|
||||
except ValueError as e:
|
||||
return get_error_argument_result(str(e))
|
||||
except Exception as e:
|
||||
logging.exception(e)
|
||||
return get_error_data_result(message="Internal server error")
|
||||
|
||||
|
||||
@manager.route("/datasets/<dataset_id>/embedding", methods=["POST"]) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
async def run_embedding(tenant_id, dataset_id):
|
||||
try:
|
||||
success, result = dataset_api_service.run_embedding(dataset_id, tenant_id)
|
||||
if success:
|
||||
return get_result(data=result)
|
||||
else:
|
||||
@@ -383,52 +584,50 @@ async def run_graphrag(tenant_id, dataset_id):
|
||||
return get_error_data_result(message="Internal server error")
|
||||
|
||||
|
||||
@manager.route("/datasets/<dataset_id>/trace_graphrag", methods=["GET"]) # noqa: F821
|
||||
@manager.route("/datasets/<dataset_id>/ingestions", methods=["GET"]) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
def trace_graphrag(tenant_id, dataset_id):
|
||||
def list_ingestion_logs(tenant_id, dataset_id):
|
||||
try:
|
||||
success, result = dataset_api_service.trace_graphrag(dataset_id, tenant_id)
|
||||
page = int(request.args.get("page", 0))
|
||||
page_size = int(request.args.get("page_size", 0))
|
||||
orderby = request.args.get("orderby", "create_time")
|
||||
desc = request.args.get("desc", "true").lower() != "false"
|
||||
operation_status = request.args.getlist("operation_status")
|
||||
create_date_from = request.args.get("create_date_from", None)
|
||||
create_date_to = request.args.get("create_date_to", None)
|
||||
success, result = dataset_api_service.list_ingestion_logs(
|
||||
dataset_id, tenant_id, page, page_size, orderby, desc, operation_status, create_date_from, create_date_to
|
||||
)
|
||||
if success:
|
||||
return get_result(data=result)
|
||||
else:
|
||||
return get_error_data_result(message=result)
|
||||
except ValueError as e:
|
||||
return get_error_argument_result(str(e))
|
||||
except Exception as e:
|
||||
logging.exception(e)
|
||||
return get_error_data_result(message="Internal server error")
|
||||
|
||||
|
||||
@manager.route("/datasets/<dataset_id>/run_raptor", methods=["POST"]) # noqa: F821
|
||||
@manager.route("/datasets/<dataset_id>/ingestions/<log_id>", methods=["GET"]) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
async def run_raptor(tenant_id, dataset_id):
|
||||
def get_ingestion_log(tenant_id, dataset_id, log_id):
|
||||
try:
|
||||
success, result = dataset_api_service.run_raptor(dataset_id, tenant_id)
|
||||
success, result = dataset_api_service.get_ingestion_log(dataset_id, tenant_id, log_id)
|
||||
if success:
|
||||
return get_result(data=result)
|
||||
else:
|
||||
return get_error_data_result(message=result)
|
||||
except ValueError as e:
|
||||
return get_error_argument_result(str(e))
|
||||
except Exception as e:
|
||||
logging.exception(e)
|
||||
return get_error_data_result(message="Internal server error")
|
||||
|
||||
|
||||
@manager.route("/datasets/<dataset_id>/trace_raptor", methods=["GET"]) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
def trace_raptor(tenant_id, dataset_id):
|
||||
try:
|
||||
success, result = dataset_api_service.trace_raptor(dataset_id, tenant_id)
|
||||
if success:
|
||||
return get_result(data=result)
|
||||
else:
|
||||
return get_error_data_result(message=result)
|
||||
except Exception as e:
|
||||
logging.exception(e)
|
||||
return get_error_data_result(message="Internal server error")
|
||||
|
||||
|
||||
@manager.route("/datasets/<dataset_id>/auto_metadata", methods=["GET"]) # noqa: F821
|
||||
@manager.route("/datasets/<dataset_id>/metadata/config", methods=["GET"]) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
def get_auto_metadata(tenant_id, dataset_id):
|
||||
@@ -462,12 +661,14 @@ def get_auto_metadata(tenant_id, dataset_id):
|
||||
return get_result(data=result)
|
||||
else:
|
||||
return get_error_data_result(message=result)
|
||||
except ValueError as e:
|
||||
return get_error_argument_result(str(e))
|
||||
except Exception as e:
|
||||
logging.exception(e)
|
||||
return get_error_data_result(message="Internal server error")
|
||||
|
||||
|
||||
@manager.route("/datasets/<dataset_id>/auto_metadata", methods=["PUT"]) # noqa: F821
|
||||
@manager.route("/datasets/<dataset_id>/metadata/config", methods=["PUT"]) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
async def update_auto_metadata(tenant_id, dataset_id):
|
||||
@@ -512,6 +713,8 @@ async def update_auto_metadata(tenant_id, dataset_id):
|
||||
return get_result(data=result)
|
||||
else:
|
||||
return get_error_data_result(message=result)
|
||||
except ValueError as e:
|
||||
return get_error_argument_result(str(e))
|
||||
except Exception as e:
|
||||
logging.exception(e)
|
||||
return get_error_data_result(message="Internal server error")
|
||||
|
||||
@@ -26,18 +26,22 @@ from api.apps.services.document_api_service import validate_document_update_fiel
|
||||
from api.constants import IMG_BASE64_PREFIX
|
||||
from api.db import VALID_FILE_TYPES
|
||||
from api.db.services.doc_metadata_service import DocMetadataService
|
||||
from api.db.db_models import Task
|
||||
from api.db.services.document_service import DocumentService
|
||||
from api.db.services.file_service import FileService
|
||||
from api.db.services.knowledgebase_service import KnowledgebaseService
|
||||
from api.db.services.task_service import TaskService, cancel_all_task_of
|
||||
from api.common.check_team_permission import check_kb_team_permission
|
||||
from api.utils.api_utils import get_data_error_result, get_error_data_result, get_result, get_json_result, \
|
||||
server_error_response, add_tenant_id_to_kwargs, get_request_json, get_error_argument_result, check_duplicate_ids
|
||||
from api.utils.validation_utils import (
|
||||
UpdateDocumentReq, format_validation_error_message, validate_and_parse_json_request, DeleteDocumentReq,
|
||||
)
|
||||
from common.constants import RetCode
|
||||
from common import settings
|
||||
from common.constants import RetCode, TaskStatus
|
||||
from common.metadata_utils import convert_conditions, meta_filter, turn2jsonschema
|
||||
from common.misc_utils import thread_pool_exec
|
||||
from rag.nlp import search
|
||||
|
||||
@manager.route("/datasets/<dataset_id>/documents/<document_id>", methods=["PATCH"]) # noqa: F821
|
||||
@login_required
|
||||
@@ -192,6 +196,88 @@ async def metadata_summary(dataset_id, tenant_id):
|
||||
return server_error_response(e)
|
||||
|
||||
|
||||
@manager.route("/datasets/<dataset_id>/metadata/update", methods=["POST"]) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
async def metadata_batch_update(dataset_id, tenant_id):
|
||||
"""
|
||||
Batch update metadata for documents in a dataset.
|
||||
---
|
||||
tags:
|
||||
- Documents
|
||||
security:
|
||||
- ApiKeyAuth: []
|
||||
parameters:
|
||||
- in: path
|
||||
name: dataset_id
|
||||
type: string
|
||||
required: true
|
||||
description: ID of the dataset.
|
||||
requestBody:
|
||||
required: true
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
type: object
|
||||
properties:
|
||||
selector:
|
||||
type: object
|
||||
updates:
|
||||
type: array
|
||||
deletes:
|
||||
type: array
|
||||
responses:
|
||||
200:
|
||||
description: Metadata updated successfully.
|
||||
"""
|
||||
if not KnowledgebaseService.accessible(kb_id=dataset_id, user_id=tenant_id):
|
||||
return get_error_data_result(message=f"You don't own the dataset {dataset_id}. ")
|
||||
|
||||
req = await get_request_json()
|
||||
selector = req.get("selector", {}) or {}
|
||||
updates = req.get("updates", []) or []
|
||||
deletes = req.get("deletes", []) or []
|
||||
|
||||
if not isinstance(selector, dict):
|
||||
return get_error_data_result(message="selector must be an object.")
|
||||
if not isinstance(updates, list) or not isinstance(deletes, list):
|
||||
return get_error_data_result(message="updates and deletes must be lists.")
|
||||
|
||||
metadata_condition = selector.get("metadata_condition", {}) or {}
|
||||
if metadata_condition and not isinstance(metadata_condition, dict):
|
||||
return get_error_data_result(message="metadata_condition must be an object.")
|
||||
|
||||
document_ids = selector.get("document_ids", []) or []
|
||||
if document_ids and not isinstance(document_ids, list):
|
||||
return get_error_data_result(message="document_ids must be a list.")
|
||||
|
||||
for upd in updates:
|
||||
if not isinstance(upd, dict) or not upd.get("key") or "value" not in upd:
|
||||
return get_error_data_result(message="Each update requires key and value.")
|
||||
for d in deletes:
|
||||
if not isinstance(d, dict) or not d.get("key"):
|
||||
return get_error_data_result(message="Each delete requires key.")
|
||||
|
||||
target_doc_ids = set()
|
||||
if document_ids:
|
||||
kb_doc_ids = KnowledgebaseService.list_documents_by_ids([dataset_id])
|
||||
invalid_ids = set(document_ids) - set(kb_doc_ids)
|
||||
if invalid_ids:
|
||||
return get_error_data_result(message=f"These documents do not belong to dataset {dataset_id}: {', '.join(invalid_ids)}")
|
||||
target_doc_ids = set(document_ids)
|
||||
|
||||
if metadata_condition:
|
||||
metas = DocMetadataService.get_flatted_meta_by_kbs([dataset_id])
|
||||
filtered_ids = set(meta_filter(metas, convert_conditions(metadata_condition), metadata_condition.get("logic", "and")))
|
||||
target_doc_ids = target_doc_ids & filtered_ids
|
||||
if metadata_condition.get("conditions") and not target_doc_ids:
|
||||
return get_result(data={"updated": 0, "matched_docs": 0})
|
||||
|
||||
target_doc_ids = list(target_doc_ids)
|
||||
updated = DocMetadataService.batch_update_metadata(dataset_id, target_doc_ids, updates, deletes)
|
||||
return get_result(data={"updated": updated, "matched_docs": len(target_doc_ids)})
|
||||
|
||||
|
||||
@manager.route("/datasets/<dataset_id>/documents", methods=["POST"]) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
@@ -1019,3 +1105,217 @@ async def update_metadata(tenant_id, dataset_id):
|
||||
target_doc_ids = list(target_doc_ids)
|
||||
updated = DocMetadataService.batch_update_metadata(dataset_id, target_doc_ids, updates, deletes)
|
||||
return get_result(data={"updated": updated, "matched_docs": len(target_doc_ids)})
|
||||
|
||||
|
||||
@manager.route("/datasets/<dataset_id>/documents/parse", methods=["POST"]) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
async def parse_documents(tenant_id, dataset_id):
|
||||
"""
|
||||
Start parsing documents in a dataset.
|
||||
---
|
||||
tags:
|
||||
- Documents
|
||||
security:
|
||||
- ApiKeyAuth: []
|
||||
parameters:
|
||||
- in: path
|
||||
name: dataset_id
|
||||
type: string
|
||||
required: true
|
||||
description: ID of the dataset.
|
||||
- in: header
|
||||
name: Authorization
|
||||
type: string
|
||||
required: true
|
||||
description: Bearer token for authentication.
|
||||
- in: body
|
||||
name: body
|
||||
description: Document parse parameters.
|
||||
required: true
|
||||
schema:
|
||||
type: object
|
||||
properties:
|
||||
document_ids:
|
||||
type: array
|
||||
items:
|
||||
type: string
|
||||
description: List of document IDs to parse.
|
||||
responses:
|
||||
200:
|
||||
description: Successful operation.
|
||||
"""
|
||||
if not KnowledgebaseService.accessible(kb_id=dataset_id, user_id=tenant_id):
|
||||
return get_error_data_result(message=f"You don't own the dataset {dataset_id}.")
|
||||
|
||||
req = await get_request_json()
|
||||
if req is None:
|
||||
return get_error_data_result(message="Request body is required")
|
||||
|
||||
document_ids = req.get("document_ids")
|
||||
if document_ids is None or not isinstance(document_ids, list):
|
||||
return get_error_data_result(message="`document_ids` is required")
|
||||
if len(document_ids) == 0:
|
||||
return get_error_data_result(message="`document_ids` is required")
|
||||
|
||||
# Check for duplicate document IDs
|
||||
unique_doc_ids, duplicate_messages = check_duplicate_ids(document_ids, "document")
|
||||
errors = duplicate_messages if duplicate_messages else []
|
||||
|
||||
# Validate all document IDs belong to the dataset
|
||||
not_found_ids = []
|
||||
valid_doc_ids = []
|
||||
for doc_id in unique_doc_ids:
|
||||
docs = DocumentService.query(kb_id=dataset_id, id=doc_id)
|
||||
if not docs:
|
||||
not_found_ids.append(doc_id)
|
||||
else:
|
||||
valid_doc_ids.append(doc_id)
|
||||
|
||||
if not_found_ids:
|
||||
errors.append(f"Documents not found: {not_found_ids}")
|
||||
# Still parse valid documents, but return error code
|
||||
if not valid_doc_ids:
|
||||
return get_error_data_result(message=f"Documents not found: {not_found_ids}")
|
||||
|
||||
try:
|
||||
def _run_sync():
|
||||
kb_table_num_map = {}
|
||||
success_count = 0
|
||||
for doc_id in valid_doc_ids:
|
||||
e, doc = DocumentService.get_by_id(doc_id)
|
||||
if not e:
|
||||
errors.append(f"Document not found: {doc_id}")
|
||||
continue
|
||||
|
||||
info = {"run": str(TaskStatus.RUNNING.value), "progress": 0}
|
||||
# If re-running a completed document, clear previous chunks
|
||||
if str(doc.run) == TaskStatus.DONE.value:
|
||||
DocumentService.clear_chunk_num_when_rerun(doc.id)
|
||||
info["progress_msg"] = ""
|
||||
info["chunk_num"] = 0
|
||||
info["token_num"] = 0
|
||||
|
||||
DocumentService.update_by_id(doc_id, info)
|
||||
TaskService.filter_delete([Task.doc_id == doc_id])
|
||||
if settings.docStoreConn.index_exist(search.index_name(tenant_id), doc.kb_id):
|
||||
settings.docStoreConn.delete({"doc_id": doc_id}, search.index_name(tenant_id), doc.kb_id)
|
||||
|
||||
doc_dict = doc.to_dict()
|
||||
DocumentService.run(tenant_id, doc_dict, kb_table_num_map)
|
||||
success_count += 1
|
||||
|
||||
result = {"success_count": success_count}
|
||||
if errors:
|
||||
result["errors"] = errors
|
||||
return result
|
||||
|
||||
result = await thread_pool_exec(_run_sync)
|
||||
if not_found_ids:
|
||||
return get_error_data_result(message=f"Documents not found: {not_found_ids}")
|
||||
return get_result(data=result)
|
||||
except Exception as e:
|
||||
logging.exception(e)
|
||||
return get_error_data_result(message="Internal server error")
|
||||
|
||||
|
||||
@manager.route("/datasets/<dataset_id>/documents/stop", methods=["POST"]) # noqa: F821
|
||||
@login_required
|
||||
@add_tenant_id_to_kwargs
|
||||
async def stop_parse_documents(tenant_id, dataset_id):
|
||||
"""
|
||||
Stop parsing documents in a dataset.
|
||||
---
|
||||
tags:
|
||||
- Documents
|
||||
security:
|
||||
- ApiKeyAuth: []
|
||||
parameters:
|
||||
- in: path
|
||||
name: dataset_id
|
||||
type: string
|
||||
required: true
|
||||
description: ID of the dataset.
|
||||
- in: header
|
||||
name: Authorization
|
||||
type: string
|
||||
required: true
|
||||
description: Bearer token for authentication.
|
||||
- in: body
|
||||
name: body
|
||||
description: Document stop parse parameters.
|
||||
required: true
|
||||
schema:
|
||||
type: object
|
||||
properties:
|
||||
document_ids:
|
||||
type: array
|
||||
items:
|
||||
type: string
|
||||
description: List of document IDs to stop parsing.
|
||||
responses:
|
||||
200:
|
||||
description: Successful operation.
|
||||
"""
|
||||
if not KnowledgebaseService.accessible(kb_id=dataset_id, user_id=tenant_id):
|
||||
return get_error_data_result(message=f"You don't own the dataset {dataset_id}.")
|
||||
|
||||
req = await get_request_json()
|
||||
if req is None:
|
||||
return get_error_data_result(message="Request body is required")
|
||||
|
||||
document_ids = req.get("document_ids")
|
||||
if document_ids is None or not isinstance(document_ids, list):
|
||||
return get_error_data_result(message="`document_ids` is required")
|
||||
if len(document_ids) == 0:
|
||||
return get_error_data_result(message="`document_ids` is required")
|
||||
|
||||
# Check for duplicate document IDs
|
||||
unique_doc_ids, duplicate_messages = check_duplicate_ids(document_ids, "document")
|
||||
errors = duplicate_messages if duplicate_messages else []
|
||||
|
||||
# Validate all document IDs belong to the dataset
|
||||
not_found_ids = []
|
||||
valid_doc_ids = []
|
||||
for doc_id in unique_doc_ids:
|
||||
docs = DocumentService.query(kb_id=dataset_id, id=doc_id)
|
||||
if not docs:
|
||||
not_found_ids.append(doc_id)
|
||||
else:
|
||||
valid_doc_ids.append(doc_id)
|
||||
|
||||
if not_found_ids:
|
||||
return get_error_data_result(message=f"Documents not found: {not_found_ids}")
|
||||
|
||||
try:
|
||||
def _run_sync():
|
||||
success_count = 0
|
||||
for doc_id in valid_doc_ids:
|
||||
e, doc = DocumentService.get_by_id(doc_id)
|
||||
if not e:
|
||||
errors.append(f"Document not found: {doc_id}")
|
||||
continue
|
||||
|
||||
# Check if the document is currently running
|
||||
tasks = list(TaskService.query(doc_id=doc_id))
|
||||
has_unfinished_task = any((task.progress or 0) < 1 for task in tasks)
|
||||
if str(doc.run) not in [TaskStatus.RUNNING.value, TaskStatus.CANCEL.value] and not has_unfinished_task:
|
||||
errors.append("Can't stop parsing document that has not started or already completed")
|
||||
continue
|
||||
|
||||
cancel_all_task_of(doc_id)
|
||||
DocumentService.update_by_id(doc_id, {"run": str(TaskStatus.CANCEL.value)})
|
||||
success_count += 1
|
||||
|
||||
result = {"success_count": success_count}
|
||||
if errors:
|
||||
result["errors"] = errors
|
||||
return result
|
||||
|
||||
result = await thread_pool_exec(_run_sync)
|
||||
if not_found_ids:
|
||||
return get_error_data_result(message=f"Documents not found: {not_found_ids}")
|
||||
return get_result(data=result)
|
||||
except Exception as e:
|
||||
logging.exception(e)
|
||||
return get_error_data_result(message="Internal server error")
|
||||
|
||||
@@ -25,10 +25,30 @@ from api.db.services.file_service import FileService
|
||||
from api.db.services.knowledgebase_service import KnowledgebaseService
|
||||
from api.db.services.connector_service import Connector2KbService
|
||||
from api.db.services.task_service import GRAPH_RAPTOR_FAKE_DOC_ID, TaskService
|
||||
from api.db.services.user_service import TenantService, UserService
|
||||
from api.db.services.user_service import TenantService, UserService, UserTenantService
|
||||
from common.constants import FileSource, StatusEnum
|
||||
from api.utils.api_utils import deep_merge, get_parser_config, remap_dictionary_keys, verify_embedding_availability
|
||||
|
||||
_VALID_INDEX_TYPES = {"graph", "raptor", "mindmap"}
|
||||
|
||||
_INDEX_TYPE_TO_TASK_TYPE = {
|
||||
"graph": "graphrag",
|
||||
"raptor": "raptor",
|
||||
"mindmap": "mindmap",
|
||||
}
|
||||
|
||||
_INDEX_TYPE_TO_TASK_ID_FIELD = {
|
||||
"graph": "graphrag_task_id",
|
||||
"raptor": "raptor_task_id",
|
||||
"mindmap": "mindmap_task_id",
|
||||
}
|
||||
|
||||
_INDEX_TYPE_TO_DISPLAY_NAME = {
|
||||
"graph": "Graph",
|
||||
"raptor": "RAPTOR",
|
||||
"mindmap": "Mindmap",
|
||||
}
|
||||
|
||||
|
||||
async def create_dataset(tenant_id: str, req: dict):
|
||||
"""
|
||||
@@ -158,6 +178,55 @@ async def delete_datasets(tenant_id: str, ids: list = None, delete_all: bool = F
|
||||
return True, {"success_count": success_count, "errors": errors[:5]}
|
||||
|
||||
|
||||
def get_dataset(dataset_id: str, tenant_id: str):
|
||||
"""
|
||||
Get a single dataset.
|
||||
|
||||
:param dataset_id: dataset ID
|
||||
:param tenant_id: tenant ID
|
||||
:return: (success, result) or (success, error_message)
|
||||
"""
|
||||
if not dataset_id:
|
||||
return False, 'Lack of "Dataset ID"'
|
||||
|
||||
if not KnowledgebaseService.accessible(dataset_id, tenant_id):
|
||||
return False, f"User '{tenant_id}' lacks permission for dataset '{dataset_id}'"
|
||||
|
||||
ok, kb = KnowledgebaseService.get_by_id(dataset_id)
|
||||
if not ok:
|
||||
return False, "Invalid Dataset ID"
|
||||
|
||||
response_data = remap_dictionary_keys(kb.to_dict())
|
||||
return True, response_data
|
||||
|
||||
|
||||
def get_ingestion_summary(dataset_id: str, tenant_id: str):
|
||||
"""
|
||||
Get ingestion summary for a dataset.
|
||||
|
||||
:param dataset_id: dataset ID
|
||||
:param tenant_id: tenant ID
|
||||
:return: (success, result) or (success, error_message)
|
||||
"""
|
||||
if not dataset_id:
|
||||
return False, 'Lack of "Dataset ID"'
|
||||
|
||||
if not KnowledgebaseService.accessible(dataset_id, tenant_id):
|
||||
return False, f"User '{tenant_id}' lacks permission for dataset '{dataset_id}'"
|
||||
|
||||
ok, kb = KnowledgebaseService.get_by_id(dataset_id)
|
||||
if not ok:
|
||||
return False, "Invalid Dataset ID"
|
||||
|
||||
status = DocumentService.get_parsing_status_by_kb_ids([dataset_id]).get(dataset_id, {})
|
||||
return True, {
|
||||
"doc_num": kb.doc_num,
|
||||
"chunk_num": kb.chunk_num,
|
||||
"token_num": kb.token_num,
|
||||
"status": status,
|
||||
}
|
||||
|
||||
|
||||
async def update_dataset(tenant_id: str, dataset_id: str, req: dict):
|
||||
"""
|
||||
Update a dataset.
|
||||
@@ -404,14 +473,18 @@ def delete_knowledge_graph(dataset_id: str, tenant_id: str):
|
||||
return True, True
|
||||
|
||||
|
||||
def run_graphrag(dataset_id: str, tenant_id: str):
|
||||
def run_index(dataset_id: str, tenant_id: str, index_type: str):
|
||||
"""
|
||||
Run GraphRAG for a dataset.
|
||||
Run an indexing task (graph/raptor/mindmap) for a dataset.
|
||||
|
||||
:param dataset_id: dataset ID
|
||||
:param tenant_id: tenant ID
|
||||
:param index_type: one of "graph", "raptor", "mindmap"
|
||||
:return: (success, result) or (success, error_message)
|
||||
"""
|
||||
if index_type not in _VALID_INDEX_TYPES:
|
||||
return False, f"Invalid index type '{index_type}'. Must be one of {sorted(_VALID_INDEX_TYPES)}"
|
||||
|
||||
if not dataset_id:
|
||||
return False, 'Lack of "Dataset ID"'
|
||||
if not KnowledgebaseService.accessible(dataset_id, tenant_id):
|
||||
@@ -421,14 +494,18 @@ def run_graphrag(dataset_id: str, tenant_id: str):
|
||||
if not ok:
|
||||
return False, "Invalid Dataset ID"
|
||||
|
||||
task_id = kb.graphrag_task_id
|
||||
if task_id:
|
||||
ok, task = TaskService.get_by_id(task_id)
|
||||
task_type = _INDEX_TYPE_TO_TASK_TYPE[index_type]
|
||||
task_id_field = _INDEX_TYPE_TO_TASK_ID_FIELD[index_type]
|
||||
display_name = _INDEX_TYPE_TO_DISPLAY_NAME[index_type]
|
||||
|
||||
existing_task_id = getattr(kb, task_id_field, None)
|
||||
if existing_task_id:
|
||||
ok, task = TaskService.get_by_id(existing_task_id)
|
||||
if not ok:
|
||||
logging.warning(f"A valid GraphRAG task id is expected for Dataset {dataset_id}")
|
||||
logging.warning(f"A valid {display_name} task id is expected for Dataset {dataset_id}")
|
||||
|
||||
if task and task.progress not in [-1, 1]:
|
||||
return False, f"Task {task_id} in progress with status {task.progress}. A Graph Task is already running."
|
||||
return False, f"Task {existing_task_id} in progress with status {task.progress}. A {display_name} Task is already running."
|
||||
|
||||
documents, _ = DocumentService.get_by_kb_id(
|
||||
kb_id=dataset_id,
|
||||
@@ -447,24 +524,29 @@ def run_graphrag(dataset_id: str, tenant_id: str):
|
||||
sample_document = documents[0]
|
||||
document_ids = [document["id"] for document in documents]
|
||||
|
||||
task_id = queue_raptor_o_graphrag_tasks(sample_doc=sample_document, ty="graphrag", priority=0, fake_doc_id=GRAPH_RAPTOR_FAKE_DOC_ID, doc_ids=list(document_ids))
|
||||
task_id = queue_raptor_o_graphrag_tasks(sample_doc=sample_document, ty=task_type, priority=0, fake_doc_id=GRAPH_RAPTOR_FAKE_DOC_ID, doc_ids=list(document_ids))
|
||||
|
||||
if not KnowledgebaseService.update_by_id(kb.id, {"graphrag_task_id": task_id}):
|
||||
logging.warning(f"Cannot save graphrag_task_id for Dataset {dataset_id}")
|
||||
if not KnowledgebaseService.update_by_id(kb.id, {task_id_field: task_id}):
|
||||
logging.warning(f"Cannot save {task_id_field} for Dataset {dataset_id}")
|
||||
|
||||
return True, {"graphrag_task_id": task_id}
|
||||
return True, {"task_id": task_id}
|
||||
|
||||
|
||||
def trace_graphrag(dataset_id: str, tenant_id: str):
|
||||
def trace_index(dataset_id: str, tenant_id: str, index_type: str):
|
||||
"""
|
||||
Trace GraphRAG task for a dataset.
|
||||
Trace an indexing task (graph/raptor/mindmap) for a dataset.
|
||||
|
||||
:param dataset_id: dataset ID
|
||||
:param tenant_id: tenant ID
|
||||
:param index_type: one of "graph", "raptor", "mindmap"
|
||||
:return: (success, result) or (success, error_message)
|
||||
"""
|
||||
if index_type not in _VALID_INDEX_TYPES:
|
||||
return False, f"Invalid index type '{index_type}'. Must be one of {sorted(_VALID_INDEX_TYPES)}"
|
||||
|
||||
if not dataset_id:
|
||||
return False, 'Lack of "Dataset ID"'
|
||||
|
||||
if not KnowledgebaseService.accessible(dataset_id, tenant_id):
|
||||
return False, "No authorization."
|
||||
|
||||
@@ -472,7 +554,8 @@ def trace_graphrag(dataset_id: str, tenant_id: str):
|
||||
if not ok:
|
||||
return False, "Invalid Dataset ID"
|
||||
|
||||
task_id = kb.graphrag_task_id
|
||||
task_id_field = _INDEX_TYPE_TO_TASK_ID_FIELD[index_type]
|
||||
task_id = getattr(kb, task_id_field, None)
|
||||
if not task_id:
|
||||
return True, {}
|
||||
|
||||
@@ -483,9 +566,9 @@ def trace_graphrag(dataset_id: str, tenant_id: str):
|
||||
return True, task.to_dict()
|
||||
|
||||
|
||||
def run_raptor(dataset_id: str, tenant_id: str):
|
||||
def list_tags(dataset_id: str, tenant_id: str):
|
||||
"""
|
||||
Run RAPTOR for a dataset.
|
||||
List tags for a dataset.
|
||||
|
||||
:param dataset_id: dataset ID
|
||||
:param tenant_id: tenant ID
|
||||
@@ -493,74 +576,65 @@ def run_raptor(dataset_id: str, tenant_id: str):
|
||||
"""
|
||||
if not dataset_id:
|
||||
return False, 'Lack of "Dataset ID"'
|
||||
|
||||
if not KnowledgebaseService.accessible(dataset_id, tenant_id):
|
||||
return False, "No authorization."
|
||||
|
||||
ok, kb = KnowledgebaseService.get_by_id(dataset_id)
|
||||
if not ok:
|
||||
return False, "Invalid Dataset ID"
|
||||
tenants = UserTenantService.get_tenants_by_user_id(tenant_id)
|
||||
tags = []
|
||||
for tenant in tenants:
|
||||
tags += settings.retriever.all_tags(tenant["tenant_id"], [dataset_id])
|
||||
return True, tags
|
||||
|
||||
task_id = kb.raptor_task_id
|
||||
if task_id:
|
||||
ok, task = TaskService.get_by_id(task_id)
|
||||
|
||||
def aggregate_tags(dataset_ids: list[str], tenant_id: str):
|
||||
"""
|
||||
Aggregate tags across multiple datasets.
|
||||
|
||||
:param dataset_ids: list of dataset IDs
|
||||
:param tenant_id: tenant ID
|
||||
:return: (success, result) or (success, error_message)
|
||||
"""
|
||||
if not dataset_ids:
|
||||
return False, 'Lack of "dataset_ids"'
|
||||
|
||||
for dataset_id in dataset_ids:
|
||||
if not KnowledgebaseService.accessible(dataset_id, tenant_id):
|
||||
return False, f"No authorization for dataset '{dataset_id}'"
|
||||
|
||||
dataset_ids_by_tenant = {}
|
||||
for dataset_id in dataset_ids:
|
||||
ok, kb = KnowledgebaseService.get_by_id(dataset_id)
|
||||
if not ok:
|
||||
logging.warning(f"A valid RAPTOR task id is expected for Dataset {dataset_id}")
|
||||
return False, f"Invalid Dataset ID '{dataset_id}'"
|
||||
dataset_ids_by_tenant.setdefault(kb.tenant_id, []).append(dataset_id)
|
||||
|
||||
if task and task.progress not in [-1, 1]:
|
||||
return False, f"Task {task_id} in progress with status {task.progress}. A RAPTOR Task is already running."
|
||||
merged = {}
|
||||
for kb_tenant_id, kb_ids in dataset_ids_by_tenant.items():
|
||||
for bucket in settings.retriever.all_tags(kb_tenant_id, kb_ids):
|
||||
tag = bucket["value"]
|
||||
merged[tag] = merged.get(tag, 0) + bucket["count"]
|
||||
|
||||
documents, _ = DocumentService.get_by_kb_id(
|
||||
kb_id=dataset_id,
|
||||
page_number=0,
|
||||
items_per_page=0,
|
||||
orderby="create_time",
|
||||
desc=False,
|
||||
keywords="",
|
||||
run_status=[],
|
||||
types=[],
|
||||
suffix=[],
|
||||
)
|
||||
if not documents:
|
||||
return False, f"No documents in Dataset {dataset_id}"
|
||||
|
||||
sample_document = documents[0]
|
||||
document_ids = [document["id"] for document in documents]
|
||||
|
||||
task_id = queue_raptor_o_graphrag_tasks(sample_doc=sample_document, ty="raptor", priority=0, fake_doc_id=GRAPH_RAPTOR_FAKE_DOC_ID, doc_ids=list(document_ids))
|
||||
|
||||
if not KnowledgebaseService.update_by_id(kb.id, {"raptor_task_id": task_id}):
|
||||
logging.warning(f"Cannot save raptor_task_id for Dataset {dataset_id}")
|
||||
|
||||
return True, {"raptor_task_id": task_id}
|
||||
return True, [{"value": tag, "count": count} for tag, count in merged.items()]
|
||||
|
||||
|
||||
def trace_raptor(dataset_id: str, tenant_id: str):
|
||||
def get_flattened_metadata(dataset_ids: list[str], tenant_id: str):
|
||||
"""
|
||||
Trace RAPTOR task for a dataset.
|
||||
Get flattened metadata for datasets.
|
||||
|
||||
:param dataset_id: dataset ID
|
||||
:param dataset_ids: list of dataset IDs
|
||||
:param tenant_id: tenant ID
|
||||
:return: (success, result) or (success, error_message)
|
||||
"""
|
||||
if not dataset_id:
|
||||
return False, 'Lack of "Dataset ID"'
|
||||
if not dataset_ids:
|
||||
return False, 'Lack of "dataset_ids"'
|
||||
|
||||
if not KnowledgebaseService.accessible(dataset_id, tenant_id):
|
||||
return False, "No authorization."
|
||||
for dataset_id in dataset_ids:
|
||||
if not KnowledgebaseService.accessible(dataset_id, tenant_id):
|
||||
return False, f"No authorization for dataset '{dataset_id}'"
|
||||
|
||||
ok, kb = KnowledgebaseService.get_by_id(dataset_id)
|
||||
if not ok:
|
||||
return False, "Invalid Dataset ID"
|
||||
|
||||
task_id = kb.raptor_task_id
|
||||
if not task_id:
|
||||
return True, {}
|
||||
|
||||
ok, task = TaskService.get_by_id(task_id)
|
||||
if not ok:
|
||||
return False, "RAPTOR Task Not Found or Error Occurred"
|
||||
|
||||
return True, task.to_dict()
|
||||
from api.db.services.doc_metadata_service import DocMetadataService
|
||||
return True, DocMetadataService.get_flatted_meta_by_kbs(dataset_ids)
|
||||
|
||||
|
||||
def get_auto_metadata(dataset_id: str, tenant_id: str):
|
||||
@@ -627,3 +701,202 @@ async def update_auto_metadata(dataset_id: str, tenant_id: str, cfg: dict):
|
||||
return False, "Update auto-metadata error.(Database error)"
|
||||
|
||||
return True, {"enabled": parser_cfg["enable_metadata"], "fields": fields}
|
||||
|
||||
|
||||
def delete_tags(dataset_id: str, tenant_id: str, tags: list[str]):
|
||||
"""
|
||||
Delete tags from a dataset.
|
||||
|
||||
:param dataset_id: dataset ID
|
||||
:param tenant_id: tenant ID
|
||||
:param tags: list of tags to delete
|
||||
:return: (success, result) or (success, error_message)
|
||||
"""
|
||||
if not dataset_id:
|
||||
return False, 'Lack of "Dataset ID"'
|
||||
|
||||
if not KnowledgebaseService.accessible(dataset_id, tenant_id):
|
||||
return False, "No authorization."
|
||||
|
||||
ok, kb = KnowledgebaseService.get_by_id(dataset_id)
|
||||
if not ok:
|
||||
return False, "Invalid Dataset ID"
|
||||
|
||||
from rag.nlp import search
|
||||
for t in tags:
|
||||
settings.docStoreConn.update({"tag_kwd": t, "kb_id": [dataset_id]},
|
||||
{"remove": {"tag_kwd": t}},
|
||||
search.index_name(kb.tenant_id),
|
||||
dataset_id)
|
||||
|
||||
return True, {}
|
||||
|
||||
def list_ingestion_logs(dataset_id: str, tenant_id: str, page: int, page_size: int, orderby: str, desc: bool, operation_status: list = None, create_date_from: str = None, create_date_to: str = None):
|
||||
"""
|
||||
List ingestion logs for a dataset.
|
||||
|
||||
:param dataset_id: dataset ID
|
||||
:param tenant_id: tenant ID
|
||||
:param page: page number
|
||||
:param page_size: items per page
|
||||
:param orderby: order by field
|
||||
:param desc: descending order
|
||||
:param operation_status: filter by operation status
|
||||
:param create_date_from: filter start date
|
||||
:param create_date_to: filter end date
|
||||
:return: (success, result) or (success, error_message)
|
||||
"""
|
||||
if not dataset_id:
|
||||
return False, 'Lack of "Dataset ID"'
|
||||
|
||||
if not KnowledgebaseService.accessible(dataset_id, tenant_id):
|
||||
return False, "No authorization."
|
||||
|
||||
from api.db.services.pipeline_operation_log_service import PipelineOperationLogService
|
||||
logs, total = PipelineOperationLogService.get_dataset_logs_by_kb_id(
|
||||
dataset_id, page, page_size, orderby, desc, operation_status or [], create_date_from, create_date_to
|
||||
)
|
||||
return True, {"total": total, "logs": logs}
|
||||
|
||||
|
||||
def get_ingestion_log(dataset_id: str, tenant_id: str, log_id: str):
|
||||
"""
|
||||
Get a single ingestion log.
|
||||
|
||||
:param dataset_id: dataset ID
|
||||
:param tenant_id: tenant ID
|
||||
:param log_id: log ID
|
||||
:return: (success, result) or (success, error_message)
|
||||
"""
|
||||
if not dataset_id:
|
||||
return False, 'Lack of "Dataset ID"'
|
||||
|
||||
if not KnowledgebaseService.accessible(dataset_id, tenant_id):
|
||||
return False, "No authorization."
|
||||
|
||||
from api.db.services.pipeline_operation_log_service import PipelineOperationLogService
|
||||
fields = PipelineOperationLogService.get_dataset_logs_fields()
|
||||
log = PipelineOperationLogService.model.select(*fields).where(
|
||||
(PipelineOperationLogService.model.id == log_id) & (PipelineOperationLogService.model.kb_id == dataset_id)
|
||||
).first()
|
||||
if not log:
|
||||
return False, "Log not found"
|
||||
|
||||
return True, log.to_dict()
|
||||
|
||||
|
||||
def delete_index(dataset_id: str, tenant_id: str, index_type: str):
|
||||
"""
|
||||
Delete an indexing task (graph/raptor/mindmap) for a dataset.
|
||||
|
||||
:param dataset_id: dataset ID
|
||||
:param tenant_id: tenant ID
|
||||
:param index_type: one of "graph", "raptor", "mindmap"
|
||||
:return: (success, result) or (success, error_message)
|
||||
"""
|
||||
if index_type not in _VALID_INDEX_TYPES:
|
||||
return False, f"Invalid index type '{index_type}'. Must be one of {sorted(_VALID_INDEX_TYPES)}"
|
||||
|
||||
if not dataset_id:
|
||||
return False, 'Lack of "Dataset ID"'
|
||||
|
||||
if not KnowledgebaseService.accessible(dataset_id, tenant_id):
|
||||
return False, "No authorization."
|
||||
|
||||
ok, kb = KnowledgebaseService.get_by_id(dataset_id)
|
||||
if not ok:
|
||||
return False, "Invalid Dataset ID"
|
||||
|
||||
task_id_field = _INDEX_TYPE_TO_TASK_ID_FIELD[index_type]
|
||||
task_finish_at_field = f"{task_id_field.replace('_task_id', '_task_finish_at')}"
|
||||
task_id = getattr(kb, task_id_field, None)
|
||||
|
||||
if task_id:
|
||||
from rag.utils.redis_conn import REDIS_CONN
|
||||
try:
|
||||
REDIS_CONN.set(f"{task_id}-cancel", "x")
|
||||
except Exception as e:
|
||||
logging.exception(e)
|
||||
TaskService.delete_by_id(task_id)
|
||||
|
||||
if index_type == "graph":
|
||||
from rag.nlp import search
|
||||
settings.docStoreConn.delete({"knowledge_graph_kwd": ["graph", "subgraph", "entity", "relation"]},
|
||||
search.index_name(kb.tenant_id), dataset_id)
|
||||
elif index_type == "raptor":
|
||||
from rag.nlp import search
|
||||
settings.docStoreConn.delete({"raptor_kwd": ["raptor"]},
|
||||
search.index_name(kb.tenant_id), dataset_id)
|
||||
|
||||
KnowledgebaseService.update_by_id(kb.id, {task_id_field: "", task_finish_at_field: None})
|
||||
return True, {}
|
||||
|
||||
|
||||
def run_embedding(dataset_id: str, tenant_id: str):
|
||||
"""
|
||||
Run embedding for all documents in a dataset.
|
||||
|
||||
:param dataset_id: dataset ID
|
||||
:param tenant_id: tenant ID
|
||||
:return: (success, result) or (success, error_message)
|
||||
"""
|
||||
if not dataset_id:
|
||||
return False, 'Lack of "Dataset ID"'
|
||||
|
||||
if not KnowledgebaseService.accessible(dataset_id, tenant_id):
|
||||
return False, "No authorization."
|
||||
|
||||
ok, kb = KnowledgebaseService.get_by_id(dataset_id)
|
||||
if not ok:
|
||||
return False, "Invalid Dataset ID"
|
||||
|
||||
documents, _ = DocumentService.get_by_kb_id(
|
||||
kb_id=dataset_id,
|
||||
page_number=0,
|
||||
items_per_page=0,
|
||||
orderby="create_time",
|
||||
desc=False,
|
||||
keywords="",
|
||||
run_status=[],
|
||||
types=[],
|
||||
suffix=[],
|
||||
)
|
||||
if not documents:
|
||||
return False, f"No documents in Dataset {dataset_id}"
|
||||
|
||||
kb_table_num_map = {}
|
||||
for doc in documents:
|
||||
doc["tenant_id"] = tenant_id
|
||||
DocumentService.run(tenant_id, doc, kb_table_num_map)
|
||||
|
||||
return True, {"scheduled_count": len(documents)}
|
||||
|
||||
|
||||
def rename_tag(dataset_id: str, tenant_id: str, from_tag: str, to_tag: str):
|
||||
"""
|
||||
Rename a tag in a dataset.
|
||||
|
||||
:param dataset_id: dataset ID
|
||||
:param tenant_id: tenant ID
|
||||
:param from_tag: original tag name
|
||||
:param to_tag: new tag name
|
||||
:return: (success, result) or (success, error_message)
|
||||
"""
|
||||
if not dataset_id:
|
||||
return False, 'Lack of "Dataset ID"'
|
||||
|
||||
if not KnowledgebaseService.accessible(dataset_id, tenant_id):
|
||||
return False, "No authorization."
|
||||
|
||||
ok, kb = KnowledgebaseService.get_by_id(dataset_id)
|
||||
if not ok:
|
||||
return False, "Invalid Dataset ID"
|
||||
|
||||
from rag.nlp import search
|
||||
settings.docStoreConn.update({"tag_kwd": from_tag, "kb_id": [dataset_id]},
|
||||
{"remove": {"tag_kwd": from_tag.strip()}, "add": {"tag_kwd": to_tag}},
|
||||
search.index_name(kb.tenant_id),
|
||||
dataset_id)
|
||||
|
||||
return True, {"from": from_tag, "to": to_tag}
|
||||
|
||||
|
||||
@@ -454,19 +454,27 @@ class DocMetadataService:
|
||||
# Index exists - check if document exists
|
||||
try:
|
||||
doc_exists = settings.docStoreConn.get(
|
||||
index_name=index_name,
|
||||
id=doc_id,
|
||||
kb_id=kb_id
|
||||
doc_id,
|
||||
index_name,
|
||||
[kb_id]
|
||||
)
|
||||
if doc_exists:
|
||||
# Document exists - use partial update
|
||||
# Document exists - replace meta_fields entirely
|
||||
# Use upsert to fully replace the meta_fields field
|
||||
# (ES update with doc parameter does deep merge on object fields,
|
||||
# which would retain old keys that should be removed)
|
||||
settings.docStoreConn.es.update(
|
||||
index=index_name,
|
||||
id=doc_id,
|
||||
refresh=True,
|
||||
doc={"meta_fields": processed_meta}
|
||||
body={
|
||||
"script": {
|
||||
"source": "ctx._source.meta_fields = params.meta_fields",
|
||||
"params": {"meta_fields": processed_meta}
|
||||
}
|
||||
}
|
||||
)
|
||||
logging.debug(f"Successfully updated metadata for document {doc_id} using ES partial update")
|
||||
logging.debug(f"Successfully updated metadata for document {doc_id} using ES script update")
|
||||
return True
|
||||
except Exception as e:
|
||||
logging.debug(f"Document {doc_id} not found in index, will insert: {e}")
|
||||
|
||||
Reference in New Issue
Block a user