mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
chore: more upload file size for paid user (#39967)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
@@ -95,6 +95,7 @@ HOLOGRES_EF_CONSTRUCTION=400
|
||||
|
||||
# Upload configuration
|
||||
UPLOAD_FILE_SIZE_LIMIT=15
|
||||
KNOWLEDGE_UPLOAD_FILE_SIZE_LIMIT_FOR_PAID_PLAN=15
|
||||
UPLOAD_FILE_BATCH_LIMIT=5
|
||||
UPLOAD_IMAGE_FILE_SIZE_LIMIT=10
|
||||
UPLOAD_VIDEO_FILE_SIZE_LIMIT=100
|
||||
|
||||
@@ -32,6 +32,7 @@ def test_file_upload_config_returns_console_limits(
|
||||
assert response.status_code == 200
|
||||
assert response.json == {
|
||||
"file_size_limit": dify_config.UPLOAD_FILE_SIZE_LIMIT,
|
||||
"knowledge_file_size_limit": dify_config.UPLOAD_FILE_SIZE_LIMIT,
|
||||
"batch_count_limit": dify_config.UPLOAD_FILE_BATCH_LIMIT,
|
||||
"file_upload_limit": dify_config.BATCH_UPLOAD_LIMIT,
|
||||
"image_file_size_limit": dify_config.UPLOAD_IMAGE_FILE_SIZE_LIMIT,
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
import pytest
|
||||
|
||||
from configs.feature import FileUploadConfig
|
||||
|
||||
|
||||
def test_paid_plan_file_size_limit_uses_its_default(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("UPLOAD_FILE_SIZE_LIMIT", "23")
|
||||
monkeypatch.delenv("KNOWLEDGE_UPLOAD_FILE_SIZE_LIMIT_FOR_PAID_PLAN", raising=False)
|
||||
|
||||
config = FileUploadConfig()
|
||||
|
||||
assert config.UPLOAD_FILE_SIZE_LIMIT == 23
|
||||
assert config.KNOWLEDGE_UPLOAD_FILE_SIZE_LIMIT_FOR_PAID_PLAN == 15
|
||||
|
||||
|
||||
def test_paid_plan_file_size_limit_can_be_configured_separately(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("UPLOAD_FILE_SIZE_LIMIT", "23")
|
||||
monkeypatch.setenv("KNOWLEDGE_UPLOAD_FILE_SIZE_LIMIT_FOR_PAID_PLAN", "50")
|
||||
|
||||
config = FileUploadConfig()
|
||||
|
||||
assert config.UPLOAD_FILE_SIZE_LIMIT == 23
|
||||
assert config.KNOWLEDGE_UPLOAD_FILE_SIZE_LIMIT_FOR_PAID_PLAN == 50
|
||||
@@ -0,0 +1,19 @@
|
||||
import pytest
|
||||
|
||||
from configs.middleware.vdb.tidb_on_qdrant_config import TidbOnQdrantConfig
|
||||
|
||||
|
||||
def test_estimated_storage_limits_default(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.delenv("TIDB_ON_QDRANT_ESTIMATED_STORAGE_LIMITS_MB", raising=False)
|
||||
|
||||
config = TidbOnQdrantConfig()
|
||||
|
||||
assert config.TIDB_ON_QDRANT_ESTIMATED_STORAGE_LIMITS_MB == "sandbox:60,professional:6400,team:25600"
|
||||
|
||||
|
||||
def test_estimated_storage_limits_custom(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("TIDB_ON_QDRANT_ESTIMATED_STORAGE_LIMITS_MB", "sandbox:61,professional:6500,team:26000")
|
||||
|
||||
config = TidbOnQdrantConfig()
|
||||
|
||||
assert config.TIDB_ON_QDRANT_ESTIMATED_STORAGE_LIMITS_MB == "sandbox:61,professional:6500,team:26000"
|
||||
@@ -41,6 +41,10 @@ from core.rag.index_processor.constant.index_type import IndexStructureType
|
||||
from models.dataset import Dataset
|
||||
from models.dataset import Document as DatasetDocument
|
||||
from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus
|
||||
from services.vector_space_admission_service import (
|
||||
VECTOR_SPACE_ADMISSION_ERROR_CODE,
|
||||
format_vector_space_admission_error,
|
||||
)
|
||||
|
||||
|
||||
def make_serializable_document(**overrides):
|
||||
@@ -1115,9 +1119,10 @@ class TestDocumentBatchIndexingStatusApi:
|
||||
api = DocumentBatchIndexingStatusApi()
|
||||
method = unwrap(api.get)
|
||||
user, _ = patch_tenant
|
||||
error = format_vector_space_admission_error(61, 50)
|
||||
document = MagicMock(
|
||||
id="doc-1",
|
||||
indexing_status=IndexingStatus.COMPLETED,
|
||||
indexing_status=IndexingStatus.ERROR,
|
||||
is_paused=False,
|
||||
processing_started_at=None,
|
||||
parsing_completed_at=None,
|
||||
@@ -1125,7 +1130,7 @@ class TestDocumentBatchIndexingStatusApi:
|
||||
splitting_completed_at=None,
|
||||
completed_at=None,
|
||||
paused_at=None,
|
||||
error=None,
|
||||
error=error,
|
||||
stopped_at=None,
|
||||
)
|
||||
session = MagicMock()
|
||||
@@ -1136,14 +1141,17 @@ class TestDocumentBatchIndexingStatusApi:
|
||||
"data": [
|
||||
{
|
||||
"id": "doc-1",
|
||||
"indexing_status": "completed",
|
||||
"indexing_status": "error",
|
||||
"processing_started_at": None,
|
||||
"parsing_completed_at": None,
|
||||
"cleaning_completed_at": None,
|
||||
"splitting_completed_at": None,
|
||||
"completed_at": None,
|
||||
"paused_at": None,
|
||||
"error": None,
|
||||
"error": error,
|
||||
"error_code": VECTOR_SPACE_ADMISSION_ERROR_CODE,
|
||||
"estimated_vector_space_mb": 61,
|
||||
"vector_space_limit_mb": 50,
|
||||
"stopped_at": None,
|
||||
"completed_segments": 2,
|
||||
"total_segments": 3,
|
||||
|
||||
@@ -10,6 +10,7 @@ from services.feature_service import (
|
||||
LicenseStatus,
|
||||
LimitationModel,
|
||||
SystemFeatureModel,
|
||||
VectorSpaceLimitationModel,
|
||||
)
|
||||
|
||||
|
||||
@@ -40,7 +41,7 @@ class TestFeatureVectorSpaceApi:
|
||||
from controllers.console.feature import FeatureVectorSpaceApi
|
||||
|
||||
get_vector_space = mocker.patch("controllers.console.feature.FeatureService.get_vector_space")
|
||||
get_vector_space.return_value = LimitationModel(size=5120, limit=20480)
|
||||
get_vector_space.return_value = VectorSpaceLimitationModel(size=5120, limit=20480)
|
||||
|
||||
api = FeatureVectorSpaceApi()
|
||||
|
||||
@@ -50,6 +51,24 @@ class TestFeatureVectorSpaceApi:
|
||||
assert result == {"size": 5120, "limit": 20480}
|
||||
get_vector_space.assert_called_once_with("tenant_123")
|
||||
|
||||
def test_get_vector_space_preserves_unknown_usage(self, mocker: MockerFixture):
|
||||
from controllers.console.feature import FeatureVectorSpaceApi
|
||||
|
||||
get_vector_space = mocker.patch("controllers.console.feature.FeatureService.get_vector_space")
|
||||
get_vector_space.return_value = VectorSpaceLimitationModel(size=0, limit=50, usage_unknown=True)
|
||||
|
||||
result = unwrap(FeatureVectorSpaceApi.get)(FeatureVectorSpaceApi(), "tenant_123")
|
||||
|
||||
assert result == {"size": 0, "limit": 50, "usage_unknown": True}
|
||||
get_vector_space.assert_called_once_with("tenant_123")
|
||||
|
||||
def test_vector_space_response_schema_marks_usage_unknown_optional(self):
|
||||
schema = VectorSpaceLimitationModel.model_json_schema(mode="serialization")
|
||||
|
||||
assert schema["required"] == ["size", "limit"]
|
||||
assert schema["properties"]["usage_unknown"]["type"] == "boolean"
|
||||
assert "usage_unknown" not in schema["required"]
|
||||
|
||||
|
||||
class TestTrialModelsApi:
|
||||
def test_get_trial_models_success(self, mocker: MockerFixture):
|
||||
|
||||
@@ -87,12 +87,20 @@ class TestFileApiGet:
|
||||
api = FileApi()
|
||||
get_method = unwrap(api.get)
|
||||
|
||||
with app.test_request_context():
|
||||
data, status = get_method(api)
|
||||
with (
|
||||
app.test_request_context(),
|
||||
patch(
|
||||
"controllers.console.files.FeatureService.get_knowledge_file_size_limit",
|
||||
return_value=50,
|
||||
) as get_knowledge_file_size_limit,
|
||||
):
|
||||
data, status = get_method(api, "tenant-1")
|
||||
|
||||
assert status == 200
|
||||
assert "file_size_limit" in data
|
||||
assert data["knowledge_file_size_limit"] == 50
|
||||
assert "batch_count_limit" in data
|
||||
get_knowledge_file_size_limit.assert_called_once_with("tenant-1")
|
||||
assert data["skill_file_size_limit"] == dify_config.UPLOAD_SKILL_FILE_SIZE_LIMIT
|
||||
|
||||
|
||||
@@ -200,6 +208,33 @@ class TestFileApiPost:
|
||||
assert result is upload_file
|
||||
assert mock_file_service.upload_file.call_args.kwargs["tenant_id"] == "app-tenant-id"
|
||||
|
||||
def test_dataset_source_from_query_uses_knowledge_limit(
|
||||
self,
|
||||
app: Flask,
|
||||
mock_account_context,
|
||||
mock_file_service,
|
||||
):
|
||||
upload_file = MagicMock()
|
||||
mock_file_service.upload_file.return_value = upload_file
|
||||
|
||||
with (
|
||||
app.test_request_context(
|
||||
"/?source=datasets",
|
||||
method="POST",
|
||||
data={"file": (io.BytesIO(b"hello"), "test.txt")},
|
||||
),
|
||||
patch(
|
||||
"controllers.console.files.FeatureService.get_knowledge_file_size_limit",
|
||||
return_value=50,
|
||||
) as get_knowledge_file_size_limit,
|
||||
):
|
||||
result = upload_file_from_request(current_user=mock_account_context)
|
||||
|
||||
assert result is upload_file
|
||||
assert mock_file_service.upload_file.call_args.kwargs["source"] == "datasets"
|
||||
assert mock_file_service.upload_file.call_args.kwargs["default_file_size_limit"] == 50
|
||||
get_knowledge_file_size_limit.assert_called_once_with(mock_account_context.current_tenant_id)
|
||||
|
||||
def test_upload_with_invalid_source(self, app: Flask, mock_account_context, mock_file_service):
|
||||
"""Test that invalid source parameter gets normalized to None"""
|
||||
api = FileApi()
|
||||
|
||||
@@ -735,6 +735,17 @@ class TestBillingResourceLimits:
|
||||
result = upload_document()
|
||||
assert result == "document_uploaded"
|
||||
|
||||
# Test 3: Form source must enforce the same quota as query source
|
||||
with app.test_request_context("/", method="POST", data={"source": "datasets"}):
|
||||
with patch(
|
||||
"controllers.console.wraps.current_account_with_tenant",
|
||||
return_value=(MockUser("test_user"), "tenant123"),
|
||||
):
|
||||
with patch("controllers.console.wraps.FeatureService.get_features", return_value=mock_features):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
upload_document()
|
||||
assert exc_info.value.code == 403
|
||||
|
||||
|
||||
class TestRateLimiting:
|
||||
"""Test rate limiting decorator"""
|
||||
|
||||
+43
-3
@@ -28,7 +28,14 @@ from sqlalchemy.orm import Session
|
||||
from werkzeug.datastructures import FileStorage
|
||||
from werkzeug.exceptions import Forbidden, NotFound
|
||||
|
||||
from controllers.common.errors import FilenameNotExistsError, NoFileUploadedError, TooManyFilesError
|
||||
from controllers.common.errors import (
|
||||
FilenameNotExistsError,
|
||||
NoFileUploadedError,
|
||||
TooManyFilesError,
|
||||
)
|
||||
from controllers.common.errors import (
|
||||
FileTooLargeError as FileTooLargeHTTPError,
|
||||
)
|
||||
from controllers.service_api.dataset.error import PipelineRunError
|
||||
from controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow import (
|
||||
DatasourceNodeRunApi,
|
||||
@@ -40,7 +47,8 @@ from controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow import (
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
from models.account import Account
|
||||
from models.dataset import Dataset
|
||||
from services.errors.file import FileTooLargeError, UnsupportedFileTypeError
|
||||
from services.errors.file import FileTooLargeError as FileTooLargeServiceError
|
||||
from services.errors.file import UnsupportedFileTypeError
|
||||
from services.rag_pipeline.entity.pipeline_service_api_entities import (
|
||||
DatasourceNodeRunApiEntity,
|
||||
PipelineRunApiEntity,
|
||||
@@ -143,7 +151,7 @@ class TestFileUploadErrors:
|
||||
|
||||
def test_file_too_large_error(self):
|
||||
"""Test FileTooLargeError can be raised."""
|
||||
error = FileTooLargeError("File exceeds size limit")
|
||||
error = FileTooLargeServiceError("File exceeds size limit")
|
||||
assert error is not None
|
||||
|
||||
def test_unsupported_file_type_error(self):
|
||||
@@ -684,6 +692,38 @@ class TestFileUploadApiPost:
|
||||
assert response["name"] == "doc.pdf"
|
||||
assert response["extension"] == "pdf"
|
||||
|
||||
@patch(
|
||||
"controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.FeatureService"
|
||||
".get_knowledge_file_size_limit",
|
||||
return_value=15,
|
||||
)
|
||||
@patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.FileService")
|
||||
@patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.current_user")
|
||||
@patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.db")
|
||||
def test_upload_file_too_large_returns_http_413(
|
||||
self, mock_db, mock_current_user, mock_file_svc_cls, mock_get_limit, app: Flask
|
||||
):
|
||||
mock_current_user.__bool__ = Mock(return_value=True)
|
||||
mock_file_svc_cls.return_value.upload_file.side_effect = FileTooLargeServiceError()
|
||||
file_data = FileStorage(
|
||||
stream=io.BytesIO(b"oversized content"),
|
||||
filename="doc.pdf",
|
||||
content_type="application/pdf",
|
||||
)
|
||||
|
||||
with app.test_request_context(
|
||||
"/datasets/pipeline/file-upload",
|
||||
method="POST",
|
||||
content_type="multipart/form-data",
|
||||
data={"file": file_data},
|
||||
):
|
||||
with pytest.raises(FileTooLargeHTTPError) as exc_info:
|
||||
KnowledgebasePipelineFileUploadApi().post(tenant_id="tenant-1")
|
||||
|
||||
assert exc_info.value.code == 413
|
||||
assert exc_info.value.error_code == "file_too_large"
|
||||
mock_get_limit.assert_called_once_with("tenant-1")
|
||||
|
||||
def test_upload_no_file(self, app: Flask):
|
||||
"""Test error when no file is uploaded."""
|
||||
with app.test_request_context(
|
||||
|
||||
@@ -26,6 +26,7 @@ import pytest
|
||||
from flask import Flask
|
||||
from werkzeug.exceptions import Forbidden, NotFound
|
||||
|
||||
from controllers.common.errors import FileTooLargeError as FileTooLargeHTTPError
|
||||
from controllers.service_api.dataset.document import (
|
||||
DeprecatedDocumentAddByTextApi,
|
||||
DeprecatedDocumentUpdateByFileApi,
|
||||
@@ -47,6 +48,7 @@ from models.dataset import Dataset, Document
|
||||
from models.enums import DataSourceType, DocumentCreatedFrom, DocumentDocType, IndexingStatus
|
||||
from services.dataset_service import DocumentService
|
||||
from services.entities.knowledge_entities.knowledge_entities import ProcessRule, RetrievalModel
|
||||
from services.errors.file import FileTooLargeError as FileTooLargeServiceError
|
||||
|
||||
|
||||
def _document_data_source_info() -> dict[str, str]:
|
||||
@@ -1155,6 +1157,9 @@ class TestDocumentIndexingStatusApi:
|
||||
"completed_at": 1609459204,
|
||||
"paused_at": None,
|
||||
"error": None,
|
||||
"error_code": None,
|
||||
"estimated_vector_space_mb": None,
|
||||
"vector_space_limit_mb": None,
|
||||
"stopped_at": None,
|
||||
"completed_segments": 5,
|
||||
"total_segments": 5,
|
||||
@@ -1593,6 +1598,52 @@ class TestDocumentAddByFileApiPost:
|
||||
200,
|
||||
)
|
||||
|
||||
@patch(
|
||||
"controllers.service_api.dataset.document.FeatureService.get_knowledge_file_size_limit",
|
||||
return_value=15,
|
||||
)
|
||||
@patch("controllers.service_api.dataset.document.FileService")
|
||||
@patch("controllers.service_api.dataset.document.current_user")
|
||||
@patch("controllers.service_api.dataset.document.db")
|
||||
def test_add_by_file_too_large_returns_http_413(
|
||||
self,
|
||||
mock_db,
|
||||
mock_current_user,
|
||||
mock_file_svc_cls,
|
||||
mock_get_limit,
|
||||
app: Flask,
|
||||
mock_tenant,
|
||||
mock_dataset,
|
||||
):
|
||||
mock_dataset.provider = "vendor"
|
||||
mock_dataset.indexing_technique = "economy"
|
||||
mock_dataset.chunk_structure = None
|
||||
mock_db.session.scalar.return_value = mock_dataset
|
||||
mock_current_user.__bool__ = Mock(return_value=True)
|
||||
mock_file_svc_cls.return_value.upload_file.side_effect = FileTooLargeServiceError()
|
||||
|
||||
from io import BytesIO
|
||||
|
||||
data = {
|
||||
"file": (BytesIO(b"oversized content"), "test.pdf", "application/pdf"),
|
||||
"data": json.dumps({"process_rule": {"mode": "automatic", "rules": None}}),
|
||||
}
|
||||
with app.test_request_context(
|
||||
f"/datasets/{mock_dataset.id}/document/create-by-file",
|
||||
method="POST",
|
||||
content_type="multipart/form-data",
|
||||
data=data,
|
||||
):
|
||||
api = DocumentAddByFileApi()
|
||||
with pytest.raises(FileTooLargeHTTPError) as exc_info:
|
||||
_unwrap_non_wrapped_controller(type(api).post)(
|
||||
api, mock_db.session, tenant_id=mock_tenant, dataset_id=mock_dataset.id
|
||||
)
|
||||
|
||||
assert exc_info.value.code == 413
|
||||
assert exc_info.value.error_code == "file_too_large"
|
||||
mock_get_limit.assert_called_once_with(mock_tenant)
|
||||
|
||||
@patch("controllers.service_api.dataset.document.db")
|
||||
@patch("controllers.service_api.wraps.FeatureService")
|
||||
@patch("controllers.service_api.wraps.validate_and_get_api_token")
|
||||
|
||||
@@ -10,7 +10,7 @@ import pytest
|
||||
from flask import Flask
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
from werkzeug.exceptions import Forbidden, NotFound, Unauthorized
|
||||
from werkzeug.exceptions import Forbidden, NotFound, ServiceUnavailable, Unauthorized
|
||||
|
||||
from controllers.service_api.wraps import (
|
||||
DatasetApiResource,
|
||||
@@ -338,6 +338,7 @@ class TestCloudEditionBillingResourceCheck:
|
||||
mock_vector_space = Mock()
|
||||
mock_vector_space.limit = 10
|
||||
mock_vector_space.size = 5
|
||||
mock_vector_space.usage_unknown = False
|
||||
mock_get_vector_space.return_value = mock_vector_space
|
||||
|
||||
@cloud_edition_billing_resource_check("vector_space", "dataset")
|
||||
@@ -356,6 +357,64 @@ class TestCloudEditionBillingResourceCheck:
|
||||
mock_get_vector_space.assert_called_once_with("tenant123")
|
||||
mock_get_features.assert_not_called()
|
||||
|
||||
@patch("controllers.service_api.wraps.validate_and_get_api_token")
|
||||
@patch("controllers.service_api.wraps.FeatureService.get_features")
|
||||
@patch("controllers.service_api.wraps.FeatureService.get_vector_space")
|
||||
def test_rejects_sandbox_when_vector_space_usage_is_unknown(
|
||||
self, mock_get_vector_space, mock_get_features, mock_validate_token, app: Flask
|
||||
):
|
||||
mock_validate_token.return_value = Mock(tenant_id="tenant123")
|
||||
mock_get_vector_space.return_value = Mock(size=0, limit=50, usage_unknown=True)
|
||||
mock_get_features.return_value = SimpleNamespace(
|
||||
billing=SimpleNamespace(
|
||||
enabled=True,
|
||||
subscription=SimpleNamespace(plan=CloudPlan.SANDBOX),
|
||||
)
|
||||
)
|
||||
|
||||
@cloud_edition_billing_resource_check("vector_space", "dataset")
|
||||
def upload_document():
|
||||
return "document_uploaded"
|
||||
|
||||
with (
|
||||
app.test_request_context("/", method="GET"),
|
||||
patch("controllers.service_api.wraps.dify_config.BILLING_ENABLED", True),
|
||||
pytest.raises(ServiceUnavailable) as exc_info,
|
||||
):
|
||||
upload_document()
|
||||
|
||||
assert "Please try again later" in str(exc_info.value)
|
||||
mock_get_features.assert_called_once_with("tenant123", exclude_vector_space=True)
|
||||
|
||||
@patch("controllers.service_api.wraps.validate_and_get_api_token")
|
||||
@patch("controllers.service_api.wraps.FeatureService.get_features")
|
||||
@patch("controllers.service_api.wraps.FeatureService.get_vector_space")
|
||||
@pytest.mark.parametrize("plan", [CloudPlan.PROFESSIONAL, CloudPlan.TEAM])
|
||||
def test_allows_paid_plan_when_vector_space_usage_is_unknown(
|
||||
self, mock_get_vector_space, mock_get_features, mock_validate_token, app: Flask, plan: CloudPlan
|
||||
):
|
||||
mock_validate_token.return_value = Mock(tenant_id="tenant123")
|
||||
mock_get_vector_space.return_value = Mock(size=0, limit=50, usage_unknown=True)
|
||||
mock_get_features.return_value = SimpleNamespace(
|
||||
billing=SimpleNamespace(
|
||||
enabled=True,
|
||||
subscription=SimpleNamespace(plan=plan),
|
||||
)
|
||||
)
|
||||
|
||||
@cloud_edition_billing_resource_check("vector_space", "dataset")
|
||||
def upload_document():
|
||||
return "document_uploaded"
|
||||
|
||||
with (
|
||||
app.test_request_context("/", method="GET"),
|
||||
patch("controllers.service_api.wraps.dify_config.BILLING_ENABLED", True),
|
||||
):
|
||||
result = upload_document()
|
||||
|
||||
assert result == "document_uploaded"
|
||||
mock_get_features.assert_called_once_with("tenant123", exclude_vector_space=True)
|
||||
|
||||
@patch("controllers.service_api.wraps.validate_and_get_api_token")
|
||||
@patch("controllers.service_api.wraps.FeatureService.get_features")
|
||||
def test_loads_features_when_checking_non_vector_space_limit(
|
||||
|
||||
@@ -179,7 +179,7 @@ def test_generate_published_pipeline_creates_documents_and_delay(generator, mock
|
||||
|
||||
mocker.patch("services.dataset_service.DocumentService.get_documents_position", return_value=1)
|
||||
features = SimpleNamespace()
|
||||
mocker.patch("services.feature_service.FeatureService.get_features", return_value=features)
|
||||
get_features = mocker.patch("services.feature_service.FeatureService.get_features", return_value=features)
|
||||
check_limits = mocker.patch("services.dataset_service.DocumentService.check_document_creation_limits")
|
||||
|
||||
document1 = SimpleNamespace(
|
||||
@@ -236,6 +236,7 @@ def test_generate_published_pipeline_creates_documents_and_delay(generator, mock
|
||||
session.flush.assert_called_once_with()
|
||||
session.commit.assert_called_once_with()
|
||||
task_proxy.delay.assert_called_once()
|
||||
get_features.assert_called_once_with("tenant")
|
||||
|
||||
|
||||
def test_generate_published_pipeline_rejects_when_document_creation_limits_exceeded(generator, mocker: MockerFixture):
|
||||
@@ -309,20 +310,26 @@ def test_generate_is_retry_calls_generate(generator, mocker: MockerFixture):
|
||||
return_value=MagicMock(),
|
||||
)
|
||||
|
||||
mocker.patch.object(generator, "_generate", return_value={"result": "ok"})
|
||||
generate = mocker.patch.object(generator, "_generate", return_value={"result": "ok"})
|
||||
|
||||
args = _build_args()
|
||||
args["original_document_id"] = "document-1"
|
||||
|
||||
result = generator.generate(
|
||||
session=session,
|
||||
pipeline=pipeline,
|
||||
workflow=workflow,
|
||||
user=_build_user(),
|
||||
args=_build_args(),
|
||||
args=args,
|
||||
invoke_from=InvokeFrom.PUBLISHED_PIPELINE,
|
||||
streaming=True,
|
||||
is_retry=True,
|
||||
)
|
||||
|
||||
assert result == {"result": "ok"}
|
||||
application_generate_entity = generate.call_args.kwargs["application_generate_entity"]
|
||||
assert application_generate_entity.document_id == "document-1"
|
||||
assert application_generate_entity.original_document_id is None
|
||||
|
||||
|
||||
def test_generate_worker_handles_errors(generator, mocker: MockerFixture):
|
||||
|
||||
@@ -47,19 +47,77 @@ class TestIndexProcessor:
|
||||
|
||||
index_processor = MagicMock()
|
||||
index_processor.index.side_effect = lambda *args: phase_events.append("index")
|
||||
processor = IndexProcessor()
|
||||
admission_service = MagicMock()
|
||||
chunks = {"general_chunks": ["content"]}
|
||||
|
||||
with patch("core.rag.index_processor.index_processor.IndexProcessorFactory") as index_processor_factory:
|
||||
with (
|
||||
patch(
|
||||
"core.rag.index_processor.index_processor.VectorSpaceAdmissionService",
|
||||
return_value=admission_service,
|
||||
),
|
||||
patch("core.rag.index_processor.index_processor.IndexProcessorFactory") as index_processor_factory,
|
||||
):
|
||||
index_processor_factory.return_value.init_index_processor.return_value = index_processor
|
||||
IndexProcessor().index_and_clean(
|
||||
processor.index_and_clean(
|
||||
dataset_id=dataset.id,
|
||||
document_id=document.id,
|
||||
original_document_id="",
|
||||
chunks={"general_chunks": ["content"]},
|
||||
chunks=chunks,
|
||||
batch="batch-1",
|
||||
session=session,
|
||||
)
|
||||
|
||||
assert phase_events == ["commit", "index", "commit"]
|
||||
admission_service.ensure_pipeline_can_be_indexed.assert_called_once_with(
|
||||
dataset=dataset,
|
||||
document_id=document.id,
|
||||
chunk_structure=dataset.chunk_structure,
|
||||
chunks=chunks,
|
||||
include_summaries=False,
|
||||
session=session,
|
||||
)
|
||||
|
||||
def test_index_and_clean_skips_admission_for_replacement_without_existing_vector_points(self) -> None:
|
||||
document = SimpleNamespace(
|
||||
id="document-1",
|
||||
name="Document",
|
||||
created_at=datetime.datetime(2026, 1, 1),
|
||||
indexing_latency=None,
|
||||
indexing_status=None,
|
||||
completed_at=None,
|
||||
word_count=0,
|
||||
need_summary=False,
|
||||
)
|
||||
dataset = SimpleNamespace(
|
||||
id="dataset-1",
|
||||
tenant_id="tenant-1",
|
||||
name="Dataset",
|
||||
chunk_structure="text_model",
|
||||
summary_index_setting=None,
|
||||
)
|
||||
session = MagicMock()
|
||||
session.scalar.side_effect = [dataset, document, 3]
|
||||
session.scalars.return_value.all.return_value = []
|
||||
index_processor = MagicMock()
|
||||
processor = IndexProcessor()
|
||||
chunks = {"general_chunks": ["content"]}
|
||||
|
||||
with (
|
||||
patch("core.rag.index_processor.index_processor.VectorSpaceAdmissionService") as admission_service_class,
|
||||
patch("core.rag.index_processor.index_processor.IndexProcessorFactory") as index_processor_factory,
|
||||
):
|
||||
index_processor_factory.return_value.init_index_processor.return_value = index_processor
|
||||
processor.index_and_clean(
|
||||
dataset_id=dataset.id,
|
||||
document_id=document.id,
|
||||
original_document_id=document.id,
|
||||
chunks=chunks,
|
||||
batch="batch-1",
|
||||
session=session,
|
||||
)
|
||||
|
||||
admission_service_class.assert_not_called()
|
||||
|
||||
def test_index_and_clean_scopes_replacement_queries_to_dataset_owner(self) -> None:
|
||||
dataset = SimpleNamespace(
|
||||
@@ -90,9 +148,13 @@ class TestIndexProcessor:
|
||||
session.scalar.side_effect = resolve_owner
|
||||
session.scalars.return_value.all.return_value = [segment]
|
||||
|
||||
with patch("core.rag.index_processor.index_processor.IndexProcessorFactory") as index_processor_factory:
|
||||
processor = IndexProcessor()
|
||||
with (
|
||||
patch("core.rag.index_processor.index_processor.VectorSpaceAdmissionService") as admission_service_class,
|
||||
patch("core.rag.index_processor.index_processor.IndexProcessorFactory") as index_processor_factory,
|
||||
):
|
||||
index_backend = index_processor_factory.return_value.init_index_processor.return_value
|
||||
IndexProcessor().index_and_clean(
|
||||
processor.index_and_clean(
|
||||
dataset_id="dataset-1",
|
||||
document_id="doc-1",
|
||||
original_document_id="original-doc",
|
||||
@@ -126,6 +188,7 @@ class TestIndexProcessor:
|
||||
session=session,
|
||||
)
|
||||
index_backend.index.assert_called_once_with(dataset, document, {}, session)
|
||||
admission_service_class.assert_not_called()
|
||||
|
||||
def test_get_preview_output_scopes_document_to_dataset_owner(self) -> None:
|
||||
dataset = SimpleNamespace(
|
||||
|
||||
@@ -71,6 +71,7 @@ from models.dataset import Dataset, DatasetProcessRule, DocumentSegment
|
||||
from models.dataset import Document as DatasetDocument
|
||||
from models.enums import SegmentStatus
|
||||
from models.model import Account
|
||||
from services.vector_space_admission_service import VectorSpaceAdmissionError
|
||||
|
||||
# ============================================================================
|
||||
# Helper Functions
|
||||
@@ -1084,6 +1085,65 @@ class TestIndexingRunnerRun:
|
||||
session=mock_dependencies["session"],
|
||||
)
|
||||
|
||||
@patch.object(Account, "set_tenant_id_with_session", autospec=True)
|
||||
def test_run_rejects_before_segment_or_vector_writes(
|
||||
self, set_tenant_id, mock_dependencies, sample_dataset_documents
|
||||
):
|
||||
runner = IndexingRunner(enforce_vector_space_admission=True)
|
||||
dataset_document = sample_dataset_documents[0]
|
||||
dataset_document.need_summary = False
|
||||
dataset = Dataset(
|
||||
id=dataset_document.dataset_id,
|
||||
tenant_id=dataset_document.tenant_id,
|
||||
indexing_technique=IndexTechniqueType.HIGH_QUALITY,
|
||||
)
|
||||
current_user = Account(name="Test Account", email="test@example.com")
|
||||
model_dispatch = {
|
||||
DatasetDocument: dataset_document,
|
||||
Dataset: dataset,
|
||||
Account: current_user,
|
||||
}
|
||||
mock_dependencies["session"].get.side_effect = lambda model, _: model_dispatch.get(model)
|
||||
process_rule = DatasetProcessRule(
|
||||
dataset_id="dataset-id", mode="automatic", rules="{}", created_by="account-id"
|
||||
)
|
||||
mock_dependencies["session"].scalar.return_value = process_rule
|
||||
transformed_documents = [Document(page_content="Chunk", metadata={"doc_id": "c1", "doc_hash": "h1"})]
|
||||
admission_error = VectorSpaceAdmissionError("estimated storage exceeds capacity")
|
||||
admission_service = Mock()
|
||||
admission_service.ensure_document_can_be_indexed.side_effect = admission_error
|
||||
|
||||
with (
|
||||
patch("core.indexing_runner.VectorSpaceAdmissionService", return_value=admission_service),
|
||||
patch.object(runner, "_extract", return_value=[Document(page_content="source", metadata={})]),
|
||||
patch.object(
|
||||
runner,
|
||||
"_transform",
|
||||
return_value=transformed_documents,
|
||||
),
|
||||
patch.object(runner, "_load_segments") as load_segments,
|
||||
patch.object(runner, "_load") as load,
|
||||
patch.object(runner, "_handle_indexing_error") as handle_error,
|
||||
):
|
||||
runner.run([dataset_document], mock_dependencies["session"])
|
||||
|
||||
load_segments.assert_not_called()
|
||||
load.assert_not_called()
|
||||
admission_service.ensure_document_can_be_indexed.assert_called_once_with(
|
||||
dataset=dataset,
|
||||
document_id=dataset_document.id,
|
||||
doc_form=dataset_document.doc_form,
|
||||
documents=transformed_documents,
|
||||
include_summaries=False,
|
||||
session=mock_dependencies["session"],
|
||||
)
|
||||
handle_error.assert_called_once_with(dataset_document.id, admission_error, mock_dependencies["session"])
|
||||
set_tenant_id.assert_called_once_with(
|
||||
current_user,
|
||||
dataset.tenant_id,
|
||||
session=mock_dependencies["session"],
|
||||
)
|
||||
|
||||
@patch.object(Account, "set_tenant_id_with_session", autospec=True)
|
||||
def test_run_in_splitting_status_counts_each_transformed_document_once(
|
||||
self, set_tenant_id, mock_dependencies, sample_dataset_documents
|
||||
|
||||
+7
-1
@@ -679,7 +679,13 @@ class TestInvokeKnowledgeIndex:
|
||||
dataset_id, document_id, False, summary_setting
|
||||
)
|
||||
mock_index_processor.index_and_clean.assert_called_once_with(
|
||||
dataset_id, document_id, original_document_id, chunks, batch, summary_setting, session=session
|
||||
dataset_id,
|
||||
document_id,
|
||||
original_document_id,
|
||||
chunks,
|
||||
batch,
|
||||
summary_setting,
|
||||
session=session,
|
||||
)
|
||||
session.commit.assert_called_once()
|
||||
assert result == {"status": "indexed"}
|
||||
|
||||
@@ -67,6 +67,7 @@ def test_remote_file_info_and_upload_config() -> None:
|
||||
|
||||
config = UploadConfig(
|
||||
file_size_limit=1,
|
||||
knowledge_file_size_limit=11,
|
||||
batch_count_limit=2,
|
||||
file_upload_limit=3,
|
||||
image_file_size_limit=4,
|
||||
@@ -81,6 +82,7 @@ def test_remote_file_info_and_upload_config() -> None:
|
||||
|
||||
dumped = config.model_dump(mode="json")
|
||||
assert dumped["file_upload_limit"] == 3
|
||||
assert dumped["knowledge_file_size_limit"] == 11
|
||||
assert dumped["skill_file_size_limit"] == 7
|
||||
assert dumped["attachment_image_file_size_limit"] == 11
|
||||
|
||||
|
||||
@@ -462,6 +462,37 @@ class TestBillingServiceSubscriptionInfo:
|
||||
params={"tenant_id": tenant_id},
|
||||
)
|
||||
|
||||
def test_get_vector_space_preserves_unknown_usage(self, mock_send_request):
|
||||
tenant_id = "tenant-123"
|
||||
expected_response = {"size": 0.0, "limit": 50, "usage_unknown": True}
|
||||
mock_send_request.return_value = expected_response
|
||||
|
||||
result = BillingService.get_vector_space(tenant_id)
|
||||
|
||||
assert result == expected_response
|
||||
|
||||
def test_get_info_preserves_unknown_vector_space_usage(self, mock_send_request):
|
||||
tenant_id = "tenant-123"
|
||||
expected_response = {
|
||||
"enabled": True,
|
||||
"subscription": {"plan": "sandbox", "interval": "", "education": False},
|
||||
"members": {"size": 1, "limit": 1},
|
||||
"apps": {"size": 1, "limit": 10},
|
||||
"vector_space": {"size": 0.0, "limit": 50, "usage_unknown": True},
|
||||
"knowledge_rate_limit": {"limit": 10},
|
||||
"documents_upload_quota": {"size": 1, "limit": 50},
|
||||
"annotation_quota_limit": {"size": 0, "limit": 10},
|
||||
"docs_processing": "standard",
|
||||
"can_replace_logo": False,
|
||||
"model_load_balancing_enabled": False,
|
||||
"knowledge_pipeline_publish_enabled": False,
|
||||
}
|
||||
mock_send_request.return_value = expected_response
|
||||
|
||||
result = BillingService.get_info(tenant_id)
|
||||
|
||||
assert result["vector_space"]["usage_unknown"] is True
|
||||
|
||||
def test_get_vector_space_bypasses_cache(self, mock_send_request):
|
||||
tenant_id = "tenant-123"
|
||||
mock_send_request.return_value = {"size": 4096, "limit": 20480}
|
||||
@@ -1989,6 +2020,8 @@ class TestBillingServiceSubscriptionInfoDataType:
|
||||
if "vector_space" in result:
|
||||
assert isinstance(result["vector_space"]["size"], float)
|
||||
assert isinstance(result["vector_space"]["limit"], int)
|
||||
if "usage_unknown" in result["vector_space"]:
|
||||
assert isinstance(result["vector_space"]["usage_unknown"], bool)
|
||||
|
||||
assert isinstance(result["knowledge_rate_limit"]["limit"], int)
|
||||
|
||||
|
||||
@@ -116,3 +116,19 @@ def test_get_vector_space_converts_billing_float_size(monkeypatch: pytest.Monkey
|
||||
|
||||
assert result.size == 5120
|
||||
assert result.limit == 20480
|
||||
assert result.usage_unknown is False
|
||||
|
||||
|
||||
def test_get_vector_space_preserves_unknown_usage(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(feature_service_module.dify_config, "BILLING_ENABLED", True)
|
||||
monkeypatch.setattr(
|
||||
feature_service_module.BillingService,
|
||||
"get_vector_space",
|
||||
lambda tenant_id: {"size": 0.0, "limit": 50, "usage_unknown": True},
|
||||
)
|
||||
|
||||
result = FeatureService.get_vector_space("tenant-1")
|
||||
|
||||
assert result.size == 0
|
||||
assert result.limit == 50
|
||||
assert result.usage_unknown is True
|
||||
|
||||
@@ -0,0 +1,69 @@
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
|
||||
from enums.cloud_plan import CloudPlan
|
||||
from services import feature_service as feature_service_module
|
||||
from services.feature_service import FeatureService
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("billing_enabled", "tenant_id", "billing_feature_enabled", "plan", "expected"),
|
||||
[
|
||||
(False, "tenant-1", True, CloudPlan.PROFESSIONAL, 15),
|
||||
(True, None, True, CloudPlan.PROFESSIONAL, 15),
|
||||
(True, "tenant-1", False, CloudPlan.PROFESSIONAL, 15),
|
||||
(True, "tenant-1", True, CloudPlan.SANDBOX, 15),
|
||||
(True, "tenant-1", True, CloudPlan.PROFESSIONAL, 50),
|
||||
(True, "tenant-1", True, CloudPlan.TEAM, 50),
|
||||
],
|
||||
)
|
||||
def test_get_knowledge_file_size_limit(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
billing_enabled: bool,
|
||||
tenant_id: str | None,
|
||||
billing_feature_enabled: bool,
|
||||
plan: CloudPlan,
|
||||
expected: int,
|
||||
) -> None:
|
||||
monkeypatch.setattr(feature_service_module.dify_config, "BILLING_ENABLED", billing_enabled)
|
||||
monkeypatch.setattr(feature_service_module.dify_config, "UPLOAD_FILE_SIZE_LIMIT", 15)
|
||||
monkeypatch.setattr(
|
||||
feature_service_module.dify_config,
|
||||
"KNOWLEDGE_UPLOAD_FILE_SIZE_LIMIT_FOR_PAID_PLAN",
|
||||
50,
|
||||
)
|
||||
get_info = Mock(
|
||||
return_value={
|
||||
"enabled": billing_feature_enabled,
|
||||
"subscription": {"plan": plan},
|
||||
}
|
||||
)
|
||||
monkeypatch.setattr(feature_service_module.BillingService, "get_info", get_info)
|
||||
|
||||
assert FeatureService.get_knowledge_file_size_limit(tenant_id) == expected
|
||||
|
||||
if billing_enabled and tenant_id:
|
||||
get_info.assert_called_once_with(tenant_id, exclude_vector_space=True)
|
||||
else:
|
||||
get_info.assert_not_called()
|
||||
|
||||
|
||||
def test_paid_knowledge_file_size_limit_never_reduces_default(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(feature_service_module.dify_config, "BILLING_ENABLED", True)
|
||||
monkeypatch.setattr(feature_service_module.dify_config, "UPLOAD_FILE_SIZE_LIMIT", 100)
|
||||
monkeypatch.setattr(
|
||||
feature_service_module.dify_config,
|
||||
"KNOWLEDGE_UPLOAD_FILE_SIZE_LIMIT_FOR_PAID_PLAN",
|
||||
50,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
feature_service_module.BillingService,
|
||||
"get_info",
|
||||
lambda *_args, **_kwargs: {
|
||||
"enabled": True,
|
||||
"subscription": {"plan": CloudPlan.PROFESSIONAL},
|
||||
},
|
||||
)
|
||||
|
||||
assert FeatureService.get_knowledge_file_size_limit("tenant-1") == 100
|
||||
@@ -1,6 +1,8 @@
|
||||
from typing import cast
|
||||
from unittest.mock import patch
|
||||
|
||||
from services.feature_service import FeatureService
|
||||
from services.billing_service import BillingInfo
|
||||
from services.feature_service import FeatureService, LimitationModel
|
||||
|
||||
|
||||
def test_get_features_exclude_vector_space_sets_vector_space_to_none():
|
||||
@@ -35,3 +37,15 @@ def test_get_features_exclude_vector_space_sets_vector_space_to_none():
|
||||
|
||||
assert features.vector_space is None
|
||||
get_info.assert_called_once_with(tenant_id, exclude_vector_space=True)
|
||||
|
||||
|
||||
def test_full_features_keep_treating_unknown_vector_usage_as_zero():
|
||||
vector_space = LimitationModel()
|
||||
|
||||
FeatureService._fulfill_vector_space_from_billing_info(
|
||||
vector_space,
|
||||
cast(BillingInfo, {"vector_space": {"size": 0.0, "limit": 50, "usage_unknown": True}}),
|
||||
)
|
||||
|
||||
assert vector_space.size == 0
|
||||
assert vector_space.limit == 50
|
||||
|
||||
@@ -224,6 +224,32 @@ class TestFileService:
|
||||
# Default
|
||||
assert FileService.is_file_size_within_limit(extension="txt", file_size=5 * 1024 * 1024) is True
|
||||
assert FileService.is_file_size_within_limit(extension="pdf", file_size=6 * 1024 * 1024) is False
|
||||
assert (
|
||||
FileService.is_file_size_within_limit(
|
||||
extension="pdf",
|
||||
file_size=6 * 1024 * 1024,
|
||||
default_file_size_limit=7,
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert (
|
||||
FileService.is_file_size_within_limit(
|
||||
extension="pdf",
|
||||
file_size=8 * 1024 * 1024,
|
||||
default_file_size_limit=7,
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
# Media-specific limits are not affected by the knowledge document override.
|
||||
assert (
|
||||
FileService.is_file_size_within_limit(
|
||||
extension="jpg",
|
||||
file_size=11 * 1024 * 1024,
|
||||
default_file_size_limit=100,
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
def test_get_file_base64_success(self, file_service: FileService, db_session: Session):
|
||||
self._persist_upload_file(db_session, key="test_key")
|
||||
|
||||
@@ -0,0 +1,570 @@
|
||||
import json
|
||||
import threading
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from types import SimpleNamespace, TracebackType
|
||||
from typing import cast
|
||||
from unittest.mock import PropertyMock, call, patch
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from configs import dify_config
|
||||
from core.rag.datasource.vdb.vector_type import VectorType
|
||||
from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType
|
||||
from core.rag.models.document import AttachmentDocument, ChildDocument, Document
|
||||
from enums.cloud_plan import CloudPlan
|
||||
from enums.deployment_edition import DeploymentEdition
|
||||
from models.dataset import Dataset
|
||||
from services.vector_space_admission_service import (
|
||||
VECTOR_SPACE_ADMISSION_ERROR_CODE,
|
||||
VectorSpaceAdmissionError,
|
||||
VectorSpaceAdmissionService,
|
||||
VectorStorageWorkload,
|
||||
build_document_workload,
|
||||
build_pipeline_workload,
|
||||
estimate_tidb_storage_bytes,
|
||||
format_vector_space_admission_error,
|
||||
get_vector_space_admission_error_fields,
|
||||
parse_vector_space_estimate_limits,
|
||||
)
|
||||
|
||||
_MEBIBYTE = 1024 * 1024
|
||||
_ESTIMATE_LIMITS = "sandbox:60,professional:6400,team:25600"
|
||||
|
||||
|
||||
class _FakeRedisLock:
|
||||
def __init__(self, lock: threading.Lock) -> None:
|
||||
self._lock = lock
|
||||
|
||||
def __enter__(self) -> "_FakeRedisLock":
|
||||
self._lock.acquire()
|
||||
return self
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exc_type: type[BaseException] | None,
|
||||
exc_value: BaseException | None,
|
||||
traceback: TracebackType | None,
|
||||
) -> None:
|
||||
self._lock.release()
|
||||
|
||||
|
||||
class _FakeRedis:
|
||||
def __init__(self) -> None:
|
||||
self.values: dict[str, str] = {}
|
||||
self.ttls: dict[str, int] = {}
|
||||
self._locks: dict[str, threading.Lock] = {}
|
||||
|
||||
def lock(self, key: str, **_kwargs: object) -> _FakeRedisLock:
|
||||
return _FakeRedisLock(self._locks.setdefault(key, threading.Lock()))
|
||||
|
||||
def get(self, key: str) -> str | None:
|
||||
return self.values.get(key)
|
||||
|
||||
def setex(self, key: str, ttl: int, value: str) -> None:
|
||||
self.values[key] = value
|
||||
self.ttls[key] = ttl
|
||||
|
||||
|
||||
def _dataset() -> Dataset:
|
||||
return cast(
|
||||
Dataset,
|
||||
SimpleNamespace(
|
||||
id="dataset-1",
|
||||
tenant_id="tenant-1",
|
||||
indexing_technique=IndexTechniqueType.HIGH_QUALITY,
|
||||
embedding_model_provider="provider",
|
||||
embedding_model="model",
|
||||
index_struct_dict={"type": VectorType.TIDB_ON_QDRANT},
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _workload() -> VectorStorageWorkload:
|
||||
return VectorStorageWorkload(text_points=1, summary_points=0, probe_text="probe")
|
||||
|
||||
|
||||
def _check_estimate(
|
||||
plan: CloudPlan,
|
||||
estimated_mb: float,
|
||||
*,
|
||||
usage_mb: float = 0,
|
||||
plan_limit_mb: int = 50,
|
||||
service: VectorSpaceAdmissionService | None = None,
|
||||
document_id: str = "document-1",
|
||||
redis: _FakeRedis | None = None,
|
||||
) -> VectorSpaceAdmissionService:
|
||||
service = service or VectorSpaceAdmissionService()
|
||||
redis = redis or _FakeRedis()
|
||||
with (
|
||||
patch.object(service, "_get_plan", return_value=plan),
|
||||
patch.object(service, "_get_embedding_dimension", return_value=3072),
|
||||
patch.object(
|
||||
type(dify_config),
|
||||
"DEPLOYMENT_EDITION",
|
||||
new_callable=PropertyMock,
|
||||
return_value=DeploymentEdition.CLOUD,
|
||||
),
|
||||
patch("services.vector_space_admission_service.dify_config.BILLING_ENABLED", True),
|
||||
patch(
|
||||
"services.vector_space_admission_service.dify_config.TIDB_ON_QDRANT_ESTIMATED_STORAGE_LIMITS_MB",
|
||||
_ESTIMATE_LIMITS,
|
||||
),
|
||||
patch(
|
||||
"services.vector_space_admission_service.Vector.resolve_vector_type",
|
||||
return_value=VectorType.TIDB_ON_QDRANT,
|
||||
),
|
||||
patch(
|
||||
"services.vector_space_admission_service.estimate_tidb_storage_bytes",
|
||||
return_value=estimated_mb * _MEBIBYTE,
|
||||
),
|
||||
patch(
|
||||
"services.vector_space_admission_service.BillingService.get_vector_space",
|
||||
return_value={"size": usage_mb, "limit": plan_limit_mb},
|
||||
),
|
||||
patch("services.vector_space_admission_service.redis_client", redis),
|
||||
):
|
||||
service._ensure_can_write(
|
||||
dataset=_dataset(),
|
||||
document_id=document_id,
|
||||
workload=_workload(),
|
||||
session=cast(Session, SimpleNamespace()),
|
||||
)
|
||||
return service
|
||||
|
||||
|
||||
def test_estimate_tidb_storage_bytes_counts_both_vector_copies_and_point_overhead() -> None:
|
||||
assert estimate_tidb_storage_bytes(point_count=10, dimension=1536) == 10 * (1536 * 4 * 2 + 3584)
|
||||
|
||||
|
||||
def test_parse_vector_space_estimate_limits_supports_all_plans() -> None:
|
||||
assert parse_vector_space_estimate_limits("sandbox:1,professional:2,team:3") == {
|
||||
CloudPlan.SANDBOX: 1,
|
||||
CloudPlan.PROFESSIONAL: 2,
|
||||
CloudPlan.TEAM: 3,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"value",
|
||||
[
|
||||
"",
|
||||
"sandbox",
|
||||
"sandbox:60",
|
||||
"unknown:60",
|
||||
"sandbox:not-a-number",
|
||||
"sandbox:0",
|
||||
"sandbox:-1",
|
||||
"sandbox:1,pro:2,team:3",
|
||||
"pro:6400,professional:6401",
|
||||
],
|
||||
)
|
||||
def test_parse_vector_space_estimate_limits_rejects_invalid_values(value: str) -> None:
|
||||
with pytest.raises(ValueError, match="Invalid vector-space estimate limit"):
|
||||
parse_vector_space_estimate_limits(value)
|
||||
|
||||
|
||||
def test_vector_space_admission_error_fields() -> None:
|
||||
message = format_vector_space_admission_error(61, 50)
|
||||
|
||||
assert get_vector_space_admission_error_fields(message) == {
|
||||
"error_code": VECTOR_SPACE_ADMISSION_ERROR_CODE,
|
||||
"estimated_vector_space_mb": 61,
|
||||
"vector_space_limit_mb": 50,
|
||||
}
|
||||
assert get_vector_space_admission_error_fields("another indexing error") == {
|
||||
"error_code": None,
|
||||
"estimated_vector_space_mb": None,
|
||||
"vector_space_limit_mb": None,
|
||||
}
|
||||
|
||||
|
||||
def test_workloads_ignore_images_and_attachments() -> None:
|
||||
document_workload = build_document_workload(
|
||||
IndexStructureType.PARAGRAPH_INDEX,
|
||||
[
|
||||
Document(
|
||||
page_content="text",
|
||||
attachments=[AttachmentDocument(page_content="image", metadata={"doc_id": "file-1"})],
|
||||
)
|
||||
],
|
||||
include_summaries=False,
|
||||
)
|
||||
pipeline_workload = build_pipeline_workload(
|
||||
IndexStructureType.PARAGRAPH_INDEX,
|
||||
{
|
||||
"general_chunks": [
|
||||
{
|
||||
"content": "text ",
|
||||
"files": [{"id": "file-1"}],
|
||||
}
|
||||
]
|
||||
},
|
||||
include_summaries=False,
|
||||
)
|
||||
|
||||
assert document_workload.total_points == 1
|
||||
assert pipeline_workload.total_points == 1
|
||||
|
||||
|
||||
def test_parent_child_workload_counts_child_and_summary_vectors() -> None:
|
||||
workload = build_document_workload(
|
||||
IndexStructureType.PARENT_CHILD_INDEX,
|
||||
[
|
||||
Document(
|
||||
page_content="parent-1",
|
||||
children=[ChildDocument(page_content="child-1"), ChildDocument(page_content="child-2")],
|
||||
),
|
||||
Document(page_content="parent-2", children=[ChildDocument(page_content="child-3")]),
|
||||
],
|
||||
include_summaries=True,
|
||||
)
|
||||
|
||||
assert workload.text_points == 3
|
||||
assert workload.summary_points == 2
|
||||
assert workload.total_points == 5
|
||||
|
||||
|
||||
def test_pipeline_qa_workload_counts_question_vectors_without_summaries() -> None:
|
||||
workload = build_pipeline_workload(
|
||||
IndexStructureType.QA_INDEX,
|
||||
{
|
||||
"qa_chunks": [
|
||||
{"question": "question-1", "answer": "answer-1"},
|
||||
{"question": "question-2", "answer": "answer-2"},
|
||||
]
|
||||
},
|
||||
include_summaries=True,
|
||||
)
|
||||
|
||||
assert workload.text_points == 2
|
||||
assert workload.summary_points == 0
|
||||
|
||||
|
||||
def test_admission_is_cloud_only() -> None:
|
||||
service = VectorSpaceAdmissionService()
|
||||
with (
|
||||
patch.object(
|
||||
type(dify_config),
|
||||
"DEPLOYMENT_EDITION",
|
||||
new_callable=PropertyMock,
|
||||
return_value=DeploymentEdition.COMMUNITY,
|
||||
),
|
||||
patch("services.vector_space_admission_service.dify_config.BILLING_ENABLED", True),
|
||||
patch("services.vector_space_admission_service.Vector.resolve_vector_type") as resolve_vector_type,
|
||||
patch("services.vector_space_admission_service.BillingService.get_info") as get_info,
|
||||
):
|
||||
service._ensure_can_write(
|
||||
dataset=_dataset(),
|
||||
document_id="document-1",
|
||||
workload=_workload(),
|
||||
session=cast(Session, SimpleNamespace()),
|
||||
)
|
||||
|
||||
resolve_vector_type.assert_not_called()
|
||||
get_info.assert_not_called()
|
||||
|
||||
|
||||
def test_admission_skips_non_tidb_vector_backends() -> None:
|
||||
service = VectorSpaceAdmissionService()
|
||||
with (
|
||||
patch.object(
|
||||
type(dify_config),
|
||||
"DEPLOYMENT_EDITION",
|
||||
new_callable=PropertyMock,
|
||||
return_value=DeploymentEdition.CLOUD,
|
||||
),
|
||||
patch("services.vector_space_admission_service.dify_config.BILLING_ENABLED", True),
|
||||
patch("services.vector_space_admission_service.Vector.resolve_vector_type", return_value=VectorType.QDRANT),
|
||||
patch("services.vector_space_admission_service.BillingService.get_info") as get_info,
|
||||
):
|
||||
service._ensure_can_write(
|
||||
dataset=_dataset(),
|
||||
document_id="document-1",
|
||||
workload=_workload(),
|
||||
session=cast(Session, SimpleNamespace()),
|
||||
)
|
||||
|
||||
get_info.assert_not_called()
|
||||
|
||||
|
||||
def test_sandbox_allows_60_mb_estimate() -> None:
|
||||
_check_estimate(CloudPlan.SANDBOX, 60)
|
||||
|
||||
|
||||
def test_sandbox_compares_current_usage_plus_document_estimate() -> None:
|
||||
_check_estimate(CloudPlan.SANDBOX, 20, usage_mb=40)
|
||||
|
||||
with pytest.raises(VectorSpaceAdmissionError):
|
||||
_check_estimate(CloudPlan.SANDBOX, 21, usage_mb=40)
|
||||
|
||||
|
||||
def test_admission_compares_fractional_usage_without_rounding_down() -> None:
|
||||
_check_estimate(CloudPlan.SANDBOX, 10.5, usage_mb=49.5)
|
||||
|
||||
with pytest.raises(VectorSpaceAdmissionError) as exc_info:
|
||||
_check_estimate(CloudPlan.SANDBOX, 10.6, usage_mb=49.5)
|
||||
|
||||
assert get_vector_space_admission_error_fields(str(exc_info.value)) == {
|
||||
"error_code": VECTOR_SPACE_ADMISSION_ERROR_CODE,
|
||||
"estimated_vector_space_mb": 61,
|
||||
"vector_space_limit_mb": 50,
|
||||
}
|
||||
|
||||
|
||||
def test_admission_uses_configured_threshold_above_nominal_limit() -> None:
|
||||
_check_estimate(CloudPlan.SANDBOX, 10, usage_mb=50)
|
||||
|
||||
with pytest.raises(VectorSpaceAdmissionError):
|
||||
_check_estimate(CloudPlan.SANDBOX, 10.1, usage_mb=50)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("plan", "usage_mb", "allowed_estimate_mb", "rejected_estimate_mb"),
|
||||
[
|
||||
(CloudPlan.PROFESSIONAL, 5000, 1400, 1401),
|
||||
(CloudPlan.TEAM, 20000, 5600, 5601),
|
||||
],
|
||||
)
|
||||
def test_paid_plan_projected_usage_boundaries(
|
||||
plan: CloudPlan,
|
||||
usage_mb: int,
|
||||
allowed_estimate_mb: int,
|
||||
rejected_estimate_mb: int,
|
||||
) -> None:
|
||||
_check_estimate(plan, allowed_estimate_mb, usage_mb=usage_mb)
|
||||
|
||||
with pytest.raises(VectorSpaceAdmissionError):
|
||||
_check_estimate(plan, rejected_estimate_mb, usage_mb=usage_mb)
|
||||
|
||||
|
||||
def test_same_batch_accumulates_projected_usage() -> None:
|
||||
service = VectorSpaceAdmissionService()
|
||||
redis = _FakeRedis()
|
||||
_check_estimate(
|
||||
CloudPlan.SANDBOX,
|
||||
10,
|
||||
usage_mb=40,
|
||||
service=service,
|
||||
document_id="document-1",
|
||||
redis=redis,
|
||||
)
|
||||
_check_estimate(
|
||||
CloudPlan.SANDBOX,
|
||||
10,
|
||||
usage_mb=40,
|
||||
service=service,
|
||||
document_id="document-2",
|
||||
redis=redis,
|
||||
)
|
||||
|
||||
with pytest.raises(VectorSpaceAdmissionError):
|
||||
_check_estimate(
|
||||
CloudPlan.SANDBOX,
|
||||
1,
|
||||
usage_mb=40,
|
||||
service=service,
|
||||
document_id="document-3",
|
||||
redis=redis,
|
||||
)
|
||||
|
||||
|
||||
def test_usage_lookup_is_refreshed_for_each_document() -> None:
|
||||
service = VectorSpaceAdmissionService()
|
||||
redis = _FakeRedis()
|
||||
with (
|
||||
patch.object(service, "_get_plan", return_value=CloudPlan.SANDBOX),
|
||||
patch.object(service, "_get_embedding_dimension", return_value=3072),
|
||||
patch.object(
|
||||
type(dify_config),
|
||||
"DEPLOYMENT_EDITION",
|
||||
new_callable=PropertyMock,
|
||||
return_value=DeploymentEdition.CLOUD,
|
||||
),
|
||||
patch("services.vector_space_admission_service.dify_config.BILLING_ENABLED", True),
|
||||
patch(
|
||||
"services.vector_space_admission_service.dify_config.TIDB_ON_QDRANT_ESTIMATED_STORAGE_LIMITS_MB",
|
||||
_ESTIMATE_LIMITS,
|
||||
),
|
||||
patch(
|
||||
"services.vector_space_admission_service.Vector.resolve_vector_type",
|
||||
return_value=VectorType.TIDB_ON_QDRANT,
|
||||
),
|
||||
patch(
|
||||
"services.vector_space_admission_service.estimate_tidb_storage_bytes",
|
||||
side_effect=[20 * _MEBIBYTE, 1 * _MEBIBYTE],
|
||||
),
|
||||
patch(
|
||||
"services.vector_space_admission_service.BillingService.get_vector_space",
|
||||
side_effect=[{"size": 40.0, "limit": 50}, {"size": 50.0, "limit": 50}],
|
||||
) as get_vector_space,
|
||||
patch("services.vector_space_admission_service.redis_client", redis),
|
||||
):
|
||||
service._ensure_can_write(
|
||||
dataset=_dataset(),
|
||||
document_id="document-1",
|
||||
workload=_workload(),
|
||||
session=cast(Session, SimpleNamespace()),
|
||||
)
|
||||
with pytest.raises(VectorSpaceAdmissionError):
|
||||
service._ensure_can_write(
|
||||
dataset=_dataset(),
|
||||
document_id="document-2",
|
||||
workload=_workload(),
|
||||
session=cast(Session, SimpleNamespace()),
|
||||
)
|
||||
|
||||
assert get_vector_space.call_args_list == [call("tenant-1"), call("tenant-1")]
|
||||
|
||||
|
||||
def test_independent_services_use_watermark_without_double_counting_fresh_usage() -> None:
|
||||
redis = _FakeRedis()
|
||||
_check_estimate(
|
||||
CloudPlan.SANDBOX,
|
||||
10,
|
||||
usage_mb=40,
|
||||
service=VectorSpaceAdmissionService(),
|
||||
document_id="document-1",
|
||||
redis=redis,
|
||||
)
|
||||
_check_estimate(
|
||||
CloudPlan.SANDBOX,
|
||||
10,
|
||||
usage_mb=50,
|
||||
service=VectorSpaceAdmissionService(),
|
||||
document_id="document-2",
|
||||
redis=redis,
|
||||
)
|
||||
|
||||
with pytest.raises(VectorSpaceAdmissionError):
|
||||
_check_estimate(
|
||||
CloudPlan.SANDBOX,
|
||||
1,
|
||||
usage_mb=50,
|
||||
service=VectorSpaceAdmissionService(),
|
||||
document_id="document-3",
|
||||
redis=redis,
|
||||
)
|
||||
|
||||
state = json.loads(redis.values["tenant:tenant-1:vector_space_estimate_watermark"])
|
||||
assert state["projected_usage_bytes"] == 60 * _MEBIBYTE
|
||||
assert state["document_ids"] == ["document-1", "document-2"]
|
||||
assert redis.ttls["tenant:tenant-1:vector_space_estimate_watermark"] == 1800
|
||||
|
||||
|
||||
def test_fresh_usage_above_watermark_becomes_next_projection_base() -> None:
|
||||
redis = _FakeRedis()
|
||||
_check_estimate(
|
||||
CloudPlan.SANDBOX,
|
||||
10,
|
||||
usage_mb=40,
|
||||
service=VectorSpaceAdmissionService(),
|
||||
document_id="document-1",
|
||||
redis=redis,
|
||||
)
|
||||
_check_estimate(
|
||||
CloudPlan.SANDBOX,
|
||||
5,
|
||||
usage_mb=55,
|
||||
service=VectorSpaceAdmissionService(),
|
||||
document_id="document-2",
|
||||
redis=redis,
|
||||
)
|
||||
|
||||
state = json.loads(redis.values["tenant:tenant-1:vector_space_estimate_watermark"])
|
||||
assert state["projected_usage_bytes"] == 60 * _MEBIBYTE
|
||||
|
||||
|
||||
def test_same_document_is_not_added_to_watermark_twice() -> None:
|
||||
redis = _FakeRedis()
|
||||
for _ in range(2):
|
||||
_check_estimate(
|
||||
CloudPlan.SANDBOX,
|
||||
10,
|
||||
usage_mb=40,
|
||||
service=VectorSpaceAdmissionService(),
|
||||
document_id="document-1",
|
||||
redis=redis,
|
||||
)
|
||||
|
||||
_check_estimate(
|
||||
CloudPlan.SANDBOX,
|
||||
10,
|
||||
usage_mb=40,
|
||||
service=VectorSpaceAdmissionService(),
|
||||
document_id="document-2",
|
||||
redis=redis,
|
||||
)
|
||||
|
||||
state = json.loads(redis.values["tenant:tenant-1:vector_space_estimate_watermark"])
|
||||
assert state["projected_usage_bytes"] == 60 * _MEBIBYTE
|
||||
assert state["document_ids"] == ["document-1", "document-2"]
|
||||
|
||||
|
||||
def test_concurrent_services_reserve_watermark_atomically() -> None:
|
||||
redis = _FakeRedis()
|
||||
barrier = threading.Barrier(2)
|
||||
|
||||
def reserve(document_id: str) -> bool:
|
||||
barrier.wait()
|
||||
_, projected_usage_bytes = VectorSpaceAdmissionService()._reserve_projected_usage(
|
||||
tenant_id="tenant-1",
|
||||
document_id=document_id,
|
||||
current_usage_bytes=40 * _MEBIBYTE,
|
||||
document_estimate_bytes=15 * _MEBIBYTE,
|
||||
estimate_limit_bytes=60 * _MEBIBYTE,
|
||||
)
|
||||
return projected_usage_bytes <= 60 * _MEBIBYTE
|
||||
|
||||
with (
|
||||
patch("services.vector_space_admission_service.redis_client", redis),
|
||||
ThreadPoolExecutor(max_workers=2) as executor,
|
||||
):
|
||||
results = list(executor.map(reserve, ["document-1", "document-2"]))
|
||||
|
||||
assert sorted(results) == [False, True]
|
||||
state = json.loads(redis.values["tenant:tenant-1:vector_space_estimate_watermark"])
|
||||
assert state["projected_usage_bytes"] == 55 * _MEBIBYTE
|
||||
assert len(state["document_ids"]) == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("plan", "estimated_mb", "plan_limit_mb"),
|
||||
[
|
||||
(CloudPlan.SANDBOX, 61, 55),
|
||||
(CloudPlan.PROFESSIONAL, 6401, 6000),
|
||||
(CloudPlan.TEAM, 25601, 24000),
|
||||
],
|
||||
)
|
||||
def test_plan_threshold_rejection_reports_billing_limit(
|
||||
plan: CloudPlan,
|
||||
estimated_mb: int,
|
||||
plan_limit_mb: int,
|
||||
) -> None:
|
||||
with pytest.raises(VectorSpaceAdmissionError) as exc_info:
|
||||
_check_estimate(plan, estimated_mb, plan_limit_mb=plan_limit_mb)
|
||||
|
||||
assert get_vector_space_admission_error_fields(str(exc_info.value)) == {
|
||||
"error_code": VECTOR_SPACE_ADMISSION_ERROR_CODE,
|
||||
"estimated_vector_space_mb": estimated_mb,
|
||||
"vector_space_limit_mb": plan_limit_mb,
|
||||
}
|
||||
|
||||
|
||||
def test_2060_mb_estimate_rejects_sandbox_but_allows_pro() -> None:
|
||||
with pytest.raises(VectorSpaceAdmissionError):
|
||||
_check_estimate(CloudPlan.SANDBOX, 2060)
|
||||
|
||||
_check_estimate(CloudPlan.PROFESSIONAL, 2060)
|
||||
|
||||
|
||||
def test_billing_plan_lookup_excludes_vector_space_and_is_cached() -> None:
|
||||
service = VectorSpaceAdmissionService()
|
||||
with patch(
|
||||
"services.vector_space_admission_service.BillingService.get_info",
|
||||
return_value={"enabled": True, "subscription": {"plan": "professional"}},
|
||||
) as get_info:
|
||||
assert service._get_plan("tenant-1") == CloudPlan.PROFESSIONAL
|
||||
assert service._get_plan("tenant-1") == CloudPlan.PROFESSIONAL
|
||||
|
||||
get_info.assert_called_once_with("tenant-1", exclude_vector_space=True)
|
||||
@@ -228,6 +228,7 @@ def mock_indexing_runner():
|
||||
with patch("tasks.document_indexing_task.IndexingRunner") as mock_runner_class:
|
||||
mock_runner = MagicMock()
|
||||
mock_runner_class.return_value = mock_runner
|
||||
mock_runner._constructor_mock = mock_runner_class
|
||||
yield mock_runner
|
||||
|
||||
|
||||
@@ -424,6 +425,7 @@ class TestBatchProcessing:
|
||||
assert doc.processing_started_at is not None
|
||||
|
||||
# IndexingRunner should be called with all documents
|
||||
mock_indexing_runner._constructor_mock.assert_called_once_with(enforce_vector_space_admission=True)
|
||||
mock_indexing_runner.run.assert_called_once()
|
||||
call_args = mock_indexing_runner.run.call_args[0][0]
|
||||
assert len(call_args) == len(document_ids)
|
||||
@@ -668,7 +670,12 @@ class TestErrorHandling:
|
||||
"""Test cases for error handling and retry mechanisms."""
|
||||
|
||||
def test_error_handling_sets_document_error_status(
|
||||
self, dataset_id, document_ids, mock_db_session, mock_dataset, mock_feature_service
|
||||
self,
|
||||
dataset_id,
|
||||
document_ids,
|
||||
mock_db_session,
|
||||
mock_dataset,
|
||||
mock_feature_service,
|
||||
):
|
||||
"""
|
||||
Test that errors during validation set document error status.
|
||||
@@ -694,8 +701,8 @@ class TestErrorHandling:
|
||||
# Set up to trigger vector space limit error
|
||||
mock_feature_service.get_features.return_value.billing.enabled = True
|
||||
mock_feature_service.get_features.return_value.billing.subscription.plan = CloudPlan.PROFESSIONAL
|
||||
mock_feature_service.get_features.return_value.vector_space.size = 100
|
||||
mock_feature_service.get_features.return_value.vector_space.limit = 100
|
||||
mock_feature_service.get_features.return_value.vector_space.size = 100 # At limit
|
||||
|
||||
# Act
|
||||
_document_indexing(dataset_id, document_ids)
|
||||
@@ -984,7 +991,12 @@ class TestAdvancedScenarios:
|
||||
assert mock_redis.setex.call_count >= concurrency_limit
|
||||
|
||||
def test_vector_space_limit_edge_case_at_exact_limit(
|
||||
self, dataset_id, document_ids, mock_db_session, mock_dataset, mock_feature_service
|
||||
self,
|
||||
dataset_id,
|
||||
document_ids,
|
||||
mock_db_session,
|
||||
mock_dataset,
|
||||
mock_feature_service,
|
||||
):
|
||||
"""
|
||||
Test vector space limit validation at exact boundary.
|
||||
@@ -1019,8 +1031,8 @@ class TestAdvancedScenarios:
|
||||
# Set vector space exactly at limit
|
||||
mock_feature_service.get_features.return_value.billing.enabled = True
|
||||
mock_feature_service.get_features.return_value.billing.subscription.plan = CloudPlan.PROFESSIONAL
|
||||
mock_feature_service.get_features.return_value.vector_space.size = 100
|
||||
mock_feature_service.get_features.return_value.vector_space.limit = 100
|
||||
mock_feature_service.get_features.return_value.vector_space.size = 100 # Exactly at limit
|
||||
|
||||
# Act
|
||||
_document_indexing(dataset_id, document_ids)
|
||||
@@ -1335,7 +1347,12 @@ class TestPerformanceScenarios:
|
||||
"""Test performance-related scenarios and optimizations."""
|
||||
|
||||
def test_large_document_batch_processing(
|
||||
self, dataset_id, mock_db_session, mock_dataset, mock_indexing_runner, mock_feature_service
|
||||
self,
|
||||
dataset_id,
|
||||
mock_db_session,
|
||||
mock_dataset,
|
||||
mock_indexing_runner,
|
||||
mock_feature_service,
|
||||
):
|
||||
"""
|
||||
Test processing a large batch of documents at batch limit.
|
||||
@@ -1373,8 +1390,8 @@ class TestPerformanceScenarios:
|
||||
# Configure billing with sufficient limits
|
||||
mock_feature_service.get_features.return_value.billing.enabled = True
|
||||
mock_feature_service.get_features.return_value.billing.subscription.plan = CloudPlan.PROFESSIONAL
|
||||
mock_feature_service.get_features.return_value.vector_space.size = 40.75
|
||||
mock_feature_service.get_features.return_value.vector_space.limit = 10000
|
||||
mock_feature_service.get_features.return_value.vector_space.size = 0
|
||||
|
||||
with patch("tasks.document_indexing_task.dify_config.BATCH_UPLOAD_LIMIT", str(batch_limit)):
|
||||
# Act
|
||||
@@ -1387,6 +1404,7 @@ class TestPerformanceScenarios:
|
||||
mock_indexing_runner.run.assert_called_once()
|
||||
call_args = mock_indexing_runner.run.call_args[0][0]
|
||||
assert len(call_args) == batch_limit
|
||||
mock_feature_service.get_features.assert_called_once_with(mock_dataset.tenant_id)
|
||||
|
||||
def test_tenant_queue_handles_burst_traffic(self, tenant_id, dataset_id, mock_redis, mock_db_session, mock_dataset):
|
||||
"""
|
||||
|
||||
@@ -0,0 +1,34 @@
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from tasks.retry_document_indexing_task import retry_document_indexing_task
|
||||
|
||||
|
||||
def test_retry_enforces_vector_space_admission() -> None:
|
||||
session = MagicMock()
|
||||
dataset = MagicMock(id="dataset-1", tenant_id="tenant-1", runtime_mode="general")
|
||||
user = MagicMock(id="user-1")
|
||||
tenant = MagicMock(id="tenant-1")
|
||||
document = MagicMock(id="document-1", dataset_id="dataset-1", doc_form="paragraph")
|
||||
session.scalar.side_effect = [dataset, user, tenant, document]
|
||||
empty_segments: list[MagicMock] = []
|
||||
session.scalars.return_value.all.return_value = empty_segments
|
||||
|
||||
session_context = MagicMock()
|
||||
session_context.__enter__.return_value = session
|
||||
features = MagicMock()
|
||||
features.billing.enabled = False
|
||||
|
||||
with (
|
||||
patch(
|
||||
"tasks.retry_document_indexing_task.session_factory.create_session",
|
||||
return_value=session_context,
|
||||
),
|
||||
patch("tasks.retry_document_indexing_task.FeatureService.get_features", return_value=features),
|
||||
patch("tasks.retry_document_indexing_task.IndexProcessorFactory"),
|
||||
patch("tasks.retry_document_indexing_task.IndexingRunner") as indexing_runner,
|
||||
patch("tasks.retry_document_indexing_task.redis_client"),
|
||||
):
|
||||
retry_document_indexing_task.run(dataset.id, [document.id], user.id)
|
||||
|
||||
indexing_runner.assert_called_once_with(enforce_vector_space_admission=True)
|
||||
indexing_runner.return_value.run.assert_called_once_with([document], session)
|
||||
Reference in New Issue
Block a user