mirror of
https://github.com/langgenius/dify.git
synced 2026-09-05 08:48:10 +08:00
chore: dep inject for model in dataset document actions (#40787)
Co-authored-by: geneachina <geneachina@gmail.com>
This commit is contained in:
parent
b6694d462f
commit
b6319c819c
@ -1487,8 +1487,10 @@ class DocumentRetryApi(DocumentResource):
|
||||
@with_current_tenant_id
|
||||
@rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT)
|
||||
@with_session
|
||||
@model_validate(DocumentRetryPayload)
|
||||
def post(
|
||||
self,
|
||||
req_data: DocumentRetryPayload,
|
||||
session: Session,
|
||||
current_tenant_id: str,
|
||||
current_user: Account,
|
||||
@ -1509,13 +1511,12 @@ class DocumentRetryApi(DocumentResource):
|
||||
except services.errors.account.NoPermissionError as e:
|
||||
raise Forbidden(str(e))
|
||||
|
||||
payload = DocumentRetryPayload.model_validate(console_ns.payload or {})
|
||||
documents = DocumentService.get_documents_by_ids(
|
||||
DatasetRefService.create_dataset_ref(dataset), payload.document_ids, session
|
||||
DatasetRefService.create_dataset_ref(dataset), req_data.document_ids, session
|
||||
)
|
||||
documents_by_id = {document.id: document for document in documents}
|
||||
retry_documents = []
|
||||
for document_id in payload.document_ids:
|
||||
for document_id in req_data.document_ids:
|
||||
try:
|
||||
document = documents_by_id.get(document_id)
|
||||
|
||||
@ -1551,7 +1552,15 @@ class DocumentRenameApi(DocumentResource):
|
||||
@with_current_user
|
||||
@rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT)
|
||||
@with_session
|
||||
def post(self, session: Session, current_user: Account, dataset_id: UUID, document_id: UUID):
|
||||
@model_validate(DocumentRenamePayload)
|
||||
def post(
|
||||
self,
|
||||
req_data: DocumentRenamePayload,
|
||||
session: Session,
|
||||
current_user: Account,
|
||||
dataset_id: UUID,
|
||||
document_id: UUID,
|
||||
):
|
||||
# The role of the current user in the ta table must be admin, owner, editor, or dataset_operator
|
||||
if not current_user.is_dataset_editor:
|
||||
raise Forbidden()
|
||||
@ -1559,10 +1568,9 @@ class DocumentRenameApi(DocumentResource):
|
||||
if not dataset:
|
||||
raise NotFound("Dataset not found.")
|
||||
DatasetService.check_dataset_operator_permission(current_user, dataset, session=session)
|
||||
payload = DocumentRenamePayload.model_validate(console_ns.payload or {})
|
||||
|
||||
try:
|
||||
document = DocumentService.rename_document(str(dataset_id), str(document_id), payload.name, session)
|
||||
document = DocumentService.rename_document(str(dataset_id), str(document_id), req_data.name, session)
|
||||
except services.errors.document.DocumentIndexingError:
|
||||
raise DocumentIndexingError("Cannot delete document during indexing.")
|
||||
|
||||
@ -1673,7 +1681,8 @@ class DocumentGenerateSummaryApi(Resource):
|
||||
@with_current_user
|
||||
@rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT)
|
||||
@with_session
|
||||
def post(self, session: Session, current_user: Account, dataset_id: UUID):
|
||||
@model_validate(GenerateSummaryPayload)
|
||||
def post(self, req_data: GenerateSummaryPayload, session: Session, current_user: Account, dataset_id: UUID):
|
||||
"""
|
||||
Generate summary index for specified documents.
|
||||
|
||||
@ -1697,9 +1706,7 @@ class DocumentGenerateSummaryApi(Resource):
|
||||
except services.errors.account.NoPermissionError as e:
|
||||
raise Forbidden(str(e))
|
||||
|
||||
# Validate request payload
|
||||
payload = GenerateSummaryPayload.model_validate(console_ns.payload or {})
|
||||
document_list = payload.document_list
|
||||
document_list = req_data.document_list
|
||||
|
||||
if not document_list:
|
||||
from werkzeug.exceptions import BadRequest
|
||||
|
||||
@ -28,10 +28,13 @@ from controllers.console.datasets.datasets_document import (
|
||||
DocumentProcessingApi,
|
||||
DocumentRecoverApi,
|
||||
DocumentRenameApi,
|
||||
DocumentRenamePayload,
|
||||
DocumentResource,
|
||||
DocumentRetryApi,
|
||||
DocumentRetryPayload,
|
||||
DocumentStatusApi,
|
||||
DocumentSummaryStatusApi,
|
||||
GenerateSummaryPayload,
|
||||
GetProcessRuleApi,
|
||||
WebsiteDocumentSyncApi,
|
||||
)
|
||||
@ -829,16 +832,16 @@ class TestDocumentRetryApi:
|
||||
method = unwrap(api.post)
|
||||
user, tenant_id = patch_tenant
|
||||
payload = {"document_ids": ["doc-1"]}
|
||||
req_data = DocumentRetryPayload.model_validate(payload)
|
||||
doc = MagicMock(id="doc-1", indexing_status="indexing")
|
||||
session = MagicMock()
|
||||
session.scalars.return_value.all.return_value = [doc]
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
patch.object(type(console_ns), "payload", payload),
|
||||
app.test_request_context("/"),
|
||||
patch("controllers.console.datasets.datasets_document.DocumentService.check_archived", return_value=True),
|
||||
patch("controllers.console.datasets.datasets_document.DocumentService.retry_document") as retry_mock,
|
||||
):
|
||||
resp, status = method(api, session, tenant_id, user, "ds-1")
|
||||
resp, status = method(api, req_data, session, tenant_id, user, "ds-1")
|
||||
assert status == 204
|
||||
retry_mock.assert_called_once_with("ds-1", [], session)
|
||||
|
||||
@ -847,18 +850,18 @@ class TestDocumentRetryApi:
|
||||
method = unwrap(api.post)
|
||||
user, tenant_id = patch_tenant
|
||||
payload = {"document_ids": ["doc-1"]}
|
||||
req_data = DocumentRetryPayload.model_validate(payload)
|
||||
document = MagicMock(id="doc-1", indexing_status=IndexingStatus.INDEXING, archived=False)
|
||||
session = MagicMock()
|
||||
session.scalars.return_value.all.return_value = [document]
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
patch.object(type(console_ns), "payload", payload),
|
||||
app.test_request_context("/"),
|
||||
patch("controllers.console.datasets.datasets_document.DocumentService.check_archived", return_value=False),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DocumentService.retry_document", return_value=None
|
||||
) as retry_mock,
|
||||
):
|
||||
response, status = method(api, session, tenant_id, user, "ds-1")
|
||||
response, status = method(api, req_data, session, tenant_id, user, "ds-1")
|
||||
assert status == 204
|
||||
retry_mock.assert_called_once_with("ds-1", [document], session)
|
||||
|
||||
@ -869,20 +872,20 @@ class TestDocumentRetryApi:
|
||||
method = unwrap(api.post)
|
||||
user, tenant_id = patch_tenant
|
||||
payload = {"document_ids": ["doc-1", "doc-2"]}
|
||||
req_data = DocumentRetryPayload.model_validate(payload)
|
||||
first_document = MagicMock(id="doc-1", indexing_status=IndexingStatus.ERROR, archived=False)
|
||||
second_document = MagicMock(id="doc-2", indexing_status=IndexingStatus.ERROR, archived=False)
|
||||
session = MagicMock()
|
||||
session.scalars.return_value.all.return_value = [first_document, second_document]
|
||||
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
patch.object(type(console_ns), "payload", payload),
|
||||
app.test_request_context("/"),
|
||||
patch("controllers.console.datasets.datasets_document.DocumentService.check_archived", return_value=False),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DocumentService.retry_document", return_value=None
|
||||
) as retry_mock,
|
||||
):
|
||||
response, status = method(api, session, tenant_id, user, "ds-1")
|
||||
response, status = method(api, req_data, session, tenant_id, user, "ds-1")
|
||||
|
||||
assert status == 204
|
||||
statement = session.scalars.call_args.args[0]
|
||||
@ -900,17 +903,17 @@ class TestDocumentRetryApi:
|
||||
method = unwrap(api.post)
|
||||
user, tenant_id = patch_tenant
|
||||
payload = {"document_ids": ["doc-1"]}
|
||||
req_data = DocumentRetryPayload.model_validate(payload)
|
||||
document = MagicMock(id="doc-1", indexing_status=IndexingStatus.COMPLETED, archived=False)
|
||||
session = MagicMock()
|
||||
session.scalars.return_value.all.return_value = [document]
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
patch.object(type(console_ns), "payload", payload),
|
||||
app.test_request_context("/"),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DocumentService.retry_document", return_value=None
|
||||
) as retry_mock,
|
||||
):
|
||||
response, status = method(api, session, tenant_id, user, "ds-1")
|
||||
response, status = method(api, req_data, session, tenant_id, user, "ds-1")
|
||||
assert status == 204
|
||||
retry_mock.assert_called_once_with("ds-1", [], session)
|
||||
|
||||
@ -920,10 +923,10 @@ class TestDocumentRetryApi:
|
||||
user, tenant_id = patch_tenant
|
||||
session = MagicMock()
|
||||
payload = {"document_ids": ["doc-1"]}
|
||||
req_data = DocumentRetryPayload.model_validate(payload)
|
||||
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
patch.object(type(console_ns), "payload", payload),
|
||||
app.test_request_context("/"),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DatasetService.get_dataset_for_tenant",
|
||||
return_value=None,
|
||||
@ -931,7 +934,7 @@ class TestDocumentRetryApi:
|
||||
patch("controllers.console.datasets.datasets_document.DocumentService.retry_document") as retry_document,
|
||||
):
|
||||
with pytest.raises(NotFound):
|
||||
method(api, session, tenant_id, user, "foreign-dataset")
|
||||
method(api, req_data, session, tenant_id, user, "foreign-dataset")
|
||||
|
||||
session.scalars.assert_not_called()
|
||||
bypass_knowledge_rate_limit.assert_not_called()
|
||||
@ -1046,9 +1049,9 @@ class TestDocumentGenerateSummaryApi:
|
||||
user, _ = patch_tenant
|
||||
dataset = MagicMock(indexing_technique="high_quality", summary_index_setting={"enable": True})
|
||||
payload = {"document_list": ["doc-1", "doc-2"]}
|
||||
req_data = GenerateSummaryPayload.model_validate(payload)
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
patch.object(type(console_ns), "payload", payload),
|
||||
app.test_request_context("/"),
|
||||
patch("controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=dataset),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DocumentService.get_documents_by_ids",
|
||||
@ -1056,7 +1059,7 @@ class TestDocumentGenerateSummaryApi:
|
||||
),
|
||||
):
|
||||
with pytest.raises(NotFound):
|
||||
method(api, MagicMock(), user, "ds-1")
|
||||
method(api, req_data, MagicMock(), user, "ds-1")
|
||||
|
||||
def test_generate_not_enabled(self, app: Flask, patch_tenant, patch_permission):
|
||||
api = DocumentGenerateSummaryApi()
|
||||
@ -1064,13 +1067,13 @@ class TestDocumentGenerateSummaryApi:
|
||||
user, _ = patch_tenant
|
||||
dataset = MagicMock(indexing_technique="high_quality", summary_index_setting={"enable": False})
|
||||
payload = {"document_list": ["doc-1"]}
|
||||
req_data = GenerateSummaryPayload.model_validate(payload)
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
patch.object(type(console_ns), "payload", payload),
|
||||
app.test_request_context("/"),
|
||||
patch("controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=dataset),
|
||||
):
|
||||
with pytest.raises(ValueError):
|
||||
method(api, MagicMock(), user, "ds-1")
|
||||
method(api, req_data, MagicMock(), user, "ds-1")
|
||||
|
||||
def test_generate_summary_success_with_qa_skip(self, app: Flask, patch_tenant, patch_permission):
|
||||
api = DocumentGenerateSummaryApi()
|
||||
@ -1080,9 +1083,9 @@ class TestDocumentGenerateSummaryApi:
|
||||
doc1 = MagicMock(id="doc-1", doc_form=IndexStructureType.QA_INDEX)
|
||||
doc2 = MagicMock(id="doc-2", doc_form=IndexStructureType.PARAGRAPH_INDEX)
|
||||
payload = {"document_list": ["doc-1", "doc-2"]}
|
||||
req_data = GenerateSummaryPayload.model_validate(payload)
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
patch.object(type(console_ns), "payload", payload),
|
||||
app.test_request_context("/"),
|
||||
patch("controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=dataset),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DocumentService.get_documents_by_ids",
|
||||
@ -1092,7 +1095,7 @@ class TestDocumentGenerateSummaryApi:
|
||||
"controllers.console.datasets.datasets_document.generate_summary_index_task.delay", return_value=None
|
||||
),
|
||||
):
|
||||
response, status = method(api, MagicMock(), user, "ds-1")
|
||||
response, status = method(api, req_data, MagicMock(), user, "ds-1")
|
||||
assert status == 200
|
||||
|
||||
|
||||
@ -1420,12 +1423,12 @@ class TestDocumentRenameApi:
|
||||
method = unwrap(api.post)
|
||||
user, _ = patch_tenant
|
||||
payload = {"name": "Renamed Document"}
|
||||
req_data = DocumentRenamePayload.model_validate(payload)
|
||||
renamed_document = make_document(id="doc-renamed", name="Renamed Document")
|
||||
session = MagicMock()
|
||||
session.scalar.return_value = 0
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
patch.object(type(console_ns), "payload", payload),
|
||||
app.test_request_context("/"),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=make_dataset()
|
||||
),
|
||||
@ -1438,7 +1441,7 @@ class TestDocumentRenameApi:
|
||||
return_value=renamed_document,
|
||||
),
|
||||
):
|
||||
response = method(api, session, user, "ds-1", "doc-1")
|
||||
response = method(api, req_data, session, user, "ds-1", "doc-1")
|
||||
assert response["id"] == "doc-renamed"
|
||||
assert response["name"] == "Renamed Document"
|
||||
assert response["data_source_info"] == {}
|
||||
@ -1484,13 +1487,13 @@ class TestDocumentGenerateSummaryApiSuccess:
|
||||
user, _ = patch_tenant
|
||||
dataset = MagicMock(indexing_technique="economy", summary_index_setting={"enable": True})
|
||||
payload = {"document_list": ["doc-1"]}
|
||||
req_data = GenerateSummaryPayload.model_validate(payload)
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
patch.object(type(console_ns), "payload", payload),
|
||||
app.test_request_context("/"),
|
||||
patch("controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=dataset),
|
||||
):
|
||||
with pytest.raises(ValueError):
|
||||
method(api, MagicMock(), user, "ds-1")
|
||||
method(api, req_data, MagicMock(), user, "ds-1")
|
||||
|
||||
|
||||
class TestDocumentProcessingApiResume:
|
||||
|
||||
Loading…
Reference in New Issue
Block a user