diff --git a/api/extensions/ext_celery.py b/api/extensions/ext_celery.py index 2748c0736b0..2cf3505e918 100644 --- a/api/extensions/ext_celery.py +++ b/api/extensions/ext_celery.py @@ -157,6 +157,7 @@ def init_app(app: DifyApp) -> Celery: "tasks.regenerate_summary_index_task", # summary index regeneration "tasks.initialize_created_app_rbac_access_task", # app access initialization "tasks.install_default_plugins_task", # tenant default plugin installation + "tasks.refresh_billing_vector_space_task", # billing vector-space cache refresh "tasks.app_generate.resume_agent_app_task", # ENG-635: Agent v2 chat ask_human resume "tasks.workflow_run_archive_download_tasks", # workflow-run archive download preparation ] diff --git a/api/services/billing_service.py b/api/services/billing_service.py index ec00d5852fc..aef5ae2f02b 100644 --- a/api/services/billing_service.py +++ b/api/services/billing_service.py @@ -217,12 +217,18 @@ class BillingService: return _billing_info_adapter.validate_python(billing_info) @classmethod - def get_vector_space(cls, tenant_id: str) -> _VectorSpaceQuota: + def get_vector_space(cls, tenant_id: str, bypass_cache: bool = False) -> _VectorSpaceQuota: params = {"tenant_id": tenant_id} + if bypass_cache: + params["bypass_cache"] = "true" return _vector_space_quota_adapter.validate_python( cls._send_request("GET", "/subscription/vector-space", params=params) ) + @classmethod + def invalidate_vector_space_cache(cls, tenant_id: str) -> None: + cls.get_vector_space(tenant_id, bypass_cache=True) + @classmethod def get_tenant_feature_plan_usage_info(cls, tenant_id: str): """Deprecated: Use get_quota_info instead.""" diff --git a/api/services/dataset_service.py b/api/services/dataset_service.py index a8b83b78047..27650e68101 100644 --- a/api/services/dataset_service.py +++ b/api/services/dataset_service.py @@ -1974,7 +1974,10 @@ class DocumentService: if data_source_info and "upload_file_id" in data_source_info: file_id = data_source_info["upload_file_id"] document_was_deleted.send( - document.id, dataset_id=document.dataset_id, doc_form=document.doc_form, file_id=file_id + document.id, + dataset_id=document.dataset_id, + doc_form=document.doc_form, + file_id=file_id, ) session.delete(document) @@ -2013,7 +2016,12 @@ class DocumentService: # Dispatch cleanup task after commit to avoid lock contention # Task cleans up segments, files, and vector indexes if deleted_document_ids and doc_form is not None: - batch_clean_document_task.delay(deleted_document_ids, dataset_ref.dataset_id, doc_form, file_ids) + batch_clean_document_task.delay( + deleted_document_ids, + dataset_ref.dataset_id, + doc_form, + file_ids, + ) @staticmethod def rename_document(dataset_id: str, document_id: str, name: str, session: Session) -> Document: diff --git a/api/tasks/batch_clean_document_task.py b/api/tasks/batch_clean_document_task.py index d243663a428..11cf4b9835c 100644 --- a/api/tasks/batch_clean_document_task.py +++ b/api/tasks/batch_clean_document_task.py @@ -13,6 +13,7 @@ from core.tools.utils.web_reader_tool import get_image_upload_file_ids from extensions.ext_storage import storage from models.dataset import Dataset, DatasetMetadataBinding, DocumentSegment from models.model import UploadFile +from tasks.refresh_billing_vector_space_task import schedule_billing_vector_space_refresh logger = logging.getLogger(__name__) @@ -21,7 +22,12 @@ BATCH_SIZE = 1000 @shared_task(queue="dataset") -def batch_clean_document_task(document_ids: list[str], dataset_id: str, doc_form: str | None, file_ids: list[str]): +def batch_clean_document_task( + document_ids: list[str], + dataset_id: str, + doc_form: str | None, + file_ids: list[str], +) -> None: """ Clean document when document deleted. :param document_ids: document ids @@ -40,6 +46,7 @@ def batch_clean_document_task(document_ids: list[str], dataset_id: str, doc_form index_node_ids: list[str] = [] segment_ids: list[str] = [] total_image_upload_file_ids: list[str] = [] + dataset_tenant_id: str | None = None try: # ============ Step 1: Query segment and file data (short read-only transaction) ============ @@ -88,6 +95,7 @@ def batch_clean_document_task(document_ids: list[str], dataset_id: str, doc_form delete_summaries=True, session=session, ) + dataset_tenant_id = dataset.tenant_id except Exception: logger.exception( "Failed to clean vector index for dataset_id: %s, document_ids: %s, index_node_ids count: %d", @@ -203,6 +211,9 @@ def batch_clean_document_task(document_ids: list[str], dataset_id: str, doc_form dataset_id, ) + if dataset_tenant_id is not None: + schedule_billing_vector_space_refresh(dataset_tenant_id) + end_at = time.perf_counter() logger.info( click.style( diff --git a/api/tasks/clean_dataset_task.py b/api/tasks/clean_dataset_task.py index 195114499a0..5bf8784e3c2 100644 --- a/api/tasks/clean_dataset_task.py +++ b/api/tasks/clean_dataset_task.py @@ -24,6 +24,7 @@ from models.dataset import ( ) from models.model import UploadFile from models.workflow import Workflow +from tasks.refresh_billing_vector_space_task import schedule_billing_vector_space_refresh logger = logging.getLogger(__name__) @@ -52,6 +53,7 @@ def clean_dataset_task( """ logger.info(click.style(f"Start clean dataset when dataset deleted: {dataset_id}", fg="green")) start_at = time.perf_counter() + vector_cleanup_succeeded = False with session_factory.create_session() as session: try: @@ -93,6 +95,7 @@ def clean_dataset_task( try: index_processor = IndexProcessorFactory(doc_form).init_index_processor() index_processor.clean(dataset, None, with_keywords=True, delete_child_chunks=True, session=session) + vector_cleanup_succeeded = True logger.info(click.style(f"Successfully cleaned vector database for dataset: {dataset_id}", fg="green")) except Exception: logger.exception(click.style(f"Failed to clean vector database for dataset {dataset_id}", fg="red")) @@ -186,6 +189,8 @@ def clean_dataset_task( session.execute(file_delete_stmt) session.commit() + if vector_cleanup_succeeded: + schedule_billing_vector_space_refresh(dataset.tenant_id) end_at = time.perf_counter() logger.info( click.style( diff --git a/api/tasks/clean_document_task.py b/api/tasks/clean_document_task.py index 25887c9b704..e09743a0018 100644 --- a/api/tasks/clean_document_task.py +++ b/api/tasks/clean_document_task.py @@ -11,12 +11,18 @@ from core.tools.utils.web_reader_tool import get_image_upload_file_ids from extensions.ext_storage import storage from models.dataset import Dataset, DatasetMetadataBinding, DocumentSegment, SegmentAttachmentBinding from models.model import UploadFile +from tasks.refresh_billing_vector_space_task import schedule_billing_vector_space_refresh logger = logging.getLogger(__name__) @shared_task(queue="dataset") -def clean_document_task(document_id: str, dataset_id: str, doc_form: str, file_id: str | None): +def clean_document_task( + document_id: str, + dataset_id: str, + doc_form: str, + file_id: str | None, +) -> None: """ Clean document when document deleted. :param document_id: document id @@ -29,6 +35,7 @@ def clean_document_task(document_id: str, dataset_id: str, doc_form: str, file_i logger.info(click.style(f"Start clean document when document deleted: {document_id}", fg="green")) start_at = time.perf_counter() total_attachment_files = [] + vector_cleanup_succeeded = False with session_factory.create_session() as session: try: @@ -37,6 +44,7 @@ def clean_document_task(document_id: str, dataset_id: str, doc_form: str, file_i if not dataset: raise Exception("Document has no dataset") + dataset_tenant_id = dataset.tenant_id segments = session.scalars(select(DocumentSegment).where(DocumentSegment.document_id == document_id)).all() # Use JOIN to fetch attachments with bindings in a single query attachments_with_bindings = session.execute( @@ -82,6 +90,7 @@ def clean_document_task(document_id: str, dataset_id: str, doc_form: str, file_i delete_summaries=True, session=session, ) + vector_cleanup_succeeded = True except Exception: logger.exception( "Failed to clean vector / keyword index in clean_document_task, " @@ -154,6 +163,9 @@ def clean_document_task(document_id: str, dataset_id: str, doc_form: str, file_i ) ) + if vector_cleanup_succeeded: + schedule_billing_vector_space_refresh(dataset_tenant_id) + end_at = time.perf_counter() logger.info( click.style( diff --git a/api/tasks/refresh_billing_vector_space_task.py b/api/tasks/refresh_billing_vector_space_task.py new file mode 100644 index 00000000000..ff3da012e3e --- /dev/null +++ b/api/tasks/refresh_billing_vector_space_task.py @@ -0,0 +1,59 @@ +import logging + +from celery import shared_task +from opentelemetry import metrics + +from configs import dify_config +from services.billing_service import BillingService + +logger = logging.getLogger(__name__) + +_MAX_RETRIES = 3 +_RETRY_DELAY_SECONDS = 30 +_refresh_counter = metrics.get_meter(__name__).create_counter( + "billing.vector_space_cache_refresh.count", + description="Number of billing vector-space cache refresh outcomes", +) + + +@shared_task(queue="dataset", bind=True, max_retries=_MAX_RETRIES, default_retry_delay=_RETRY_DELAY_SECONDS) +def refresh_billing_vector_space_task(self, tenant_id: str) -> None: + """Refresh billing vector-space usage after vector cleanup has completed.""" + if not dify_config.BILLING_ENABLED: + return + + try: + BillingService.invalidate_vector_space_cache(tenant_id) + except Exception as exc: + if self.request.retries >= _MAX_RETRIES: + _refresh_counter.add(1, {"outcome": "exhausted"}) + logger.exception( + "Billing vector-space cache refresh retry budget exhausted, tenant_id=%s", + tenant_id, + ) + raise + + _refresh_counter.add(1, {"outcome": "retry"}) + logger.warning( + "Billing vector-space cache refresh failed, scheduling retry %d/%d, tenant_id=%s", + self.request.retries + 1, + _MAX_RETRIES, + tenant_id, + exc_info=True, + ) + raise self.retry(exc=exc, countdown=_RETRY_DELAY_SECONDS * (2**self.request.retries)) + + _refresh_counter.add(1, {"outcome": "success"}) + logger.info("Billing vector-space cache refreshed, tenant_id=%s", tenant_id) + + +def schedule_billing_vector_space_refresh(tenant_id: str) -> None: + """Dispatch a best-effort billing refresh without changing cleanup status.""" + if not dify_config.BILLING_ENABLED: + return + + try: + refresh_billing_vector_space_task.delay(tenant_id) + except Exception: + _refresh_counter.add(1, {"outcome": "dispatch_failure"}) + logger.exception("Failed to dispatch billing vector-space cache refresh, tenant_id=%s", tenant_id) diff --git a/api/tests/unit_tests/events/event_handlers/test_clean_when_document_deleted.py b/api/tests/unit_tests/events/event_handlers/test_clean_when_document_deleted.py new file mode 100644 index 00000000000..098a7b67880 --- /dev/null +++ b/api/tests/unit_tests/events/event_handlers/test_clean_when_document_deleted.py @@ -0,0 +1,15 @@ +from unittest.mock import patch + +from events.event_handlers.clean_when_document_deleted import handle + + +def test_handler_dispatches_cleanup_task(): + with patch("events.event_handlers.clean_when_document_deleted.clean_document_task.delay") as delay: + handle( + "document-1", + dataset_id="dataset-1", + doc_form="paragraph", + file_id="file-1", + ) + + delay.assert_called_once_with("document-1", "dataset-1", "paragraph", "file-1") diff --git a/api/tests/unit_tests/services/test_billing_service.py b/api/tests/unit_tests/services/test_billing_service.py index f771eabcf8c..a8d405ae8a3 100644 --- a/api/tests/unit_tests/services/test_billing_service.py +++ b/api/tests/unit_tests/services/test_billing_service.py @@ -462,6 +462,27 @@ class TestBillingServiceSubscriptionInfo: params={"tenant_id": tenant_id}, ) + def test_get_vector_space_bypasses_cache(self, mock_send_request): + tenant_id = "tenant-123" + mock_send_request.return_value = {"size": 4096, "limit": 20480} + + result = BillingService.get_vector_space(tenant_id, bypass_cache=True) + + assert result == {"size": 4096, "limit": 20480} + mock_send_request.assert_called_once_with( + "GET", + "/subscription/vector-space", + params={"tenant_id": tenant_id, "bypass_cache": "true"}, + ) + + def test_invalidate_vector_space_cache_bypasses_cache(self): + tenant_id = "tenant-123" + + with patch.object(BillingService, "get_vector_space") as get_vector_space: + BillingService.invalidate_vector_space_cache(tenant_id) + + get_vector_space.assert_called_once_with(tenant_id, bypass_cache=True) + def test_quota_get_balance_uses_quota_request(self): tenant_id = "tenant-123" with patch.object(BillingService, "_send_quota_request") as mock_send_quota_request: diff --git a/api/tests/unit_tests/tasks/test_batch_clean_document_task.py b/api/tests/unit_tests/tasks/test_batch_clean_document_task.py new file mode 100644 index 00000000000..6386c72188e --- /dev/null +++ b/api/tests/unit_tests/tasks/test_batch_clean_document_task.py @@ -0,0 +1,58 @@ +from unittest.mock import MagicMock, patch + +from tasks.batch_clean_document_task import batch_clean_document_task + + +def _setup_cleanup_dependencies(): + session = MagicMock() + segment = MagicMock(id="segment-1", index_node_id="node-1", content="content") + dataset = MagicMock(id="dataset-1", tenant_id="tenant-1") + session.scalars.return_value.all.return_value = [segment] + session.scalar.return_value = dataset + + context_manager = MagicMock() + context_manager.__enter__.return_value = session + context_manager.__exit__.return_value = None + return session, context_manager + + +def test_successful_vector_cleanup_schedules_billing_refresh(): + _, context_manager = _setup_cleanup_dependencies() + + with ( + patch("tasks.batch_clean_document_task.session_factory.create_session", return_value=context_manager), + patch("tasks.batch_clean_document_task.get_image_upload_file_ids", return_value=[]), + patch("tasks.batch_clean_document_task.IndexProcessorFactory") as processor_factory, + patch("tasks.batch_clean_document_task.schedule_billing_vector_space_refresh") as schedule_refresh, + ): + batch_clean_document_task( + document_ids=["document-1"], + dataset_id="dataset-1", + doc_form="paragraph", + file_ids=[], + ) + + processor_factory.return_value.init_index_processor.return_value.clean.assert_called_once() + schedule_refresh.assert_called_once_with("tenant-1") + + +def test_failed_vector_cleanup_does_not_schedule_billing_refresh(): + _, context_manager = _setup_cleanup_dependencies() + + with ( + patch("tasks.batch_clean_document_task.session_factory.create_session", return_value=context_manager), + patch("tasks.batch_clean_document_task.get_image_upload_file_ids", return_value=[]), + patch("tasks.batch_clean_document_task.IndexProcessorFactory") as processor_factory, + patch("tasks.batch_clean_document_task.schedule_billing_vector_space_refresh") as schedule_refresh, + ): + processor_factory.return_value.init_index_processor.return_value.clean.side_effect = RuntimeError( + "vector cleanup failed" + ) + batch_clean_document_task( + document_ids=["document-1"], + dataset_id="dataset-1", + doc_form="paragraph", + file_ids=[], + ) + + schedule_refresh.assert_not_called() diff --git a/api/tests/unit_tests/tasks/test_clean_dataset_task.py b/api/tests/unit_tests/tasks/test_clean_dataset_task.py index 826276086ba..4a2dcc44145 100644 --- a/api/tests/unit_tests/tasks/test_clean_dataset_task.py +++ b/api/tests/unit_tests/tasks/test_clean_dataset_task.py @@ -434,14 +434,15 @@ class TestIndexProcessorParameters: index_struct = '{"type": "paragraph"}' # Act - clean_dataset_task( - dataset_id=dataset_id, - tenant_id=tenant_id, - indexing_technique=indexing_technique, - index_struct=index_struct, - collection_binding_id=collection_binding_id, - doc_form=IndexStructureType.PARAGRAPH_INDEX, - ) + with patch("tasks.clean_dataset_task.schedule_billing_vector_space_refresh") as schedule_refresh: + clean_dataset_task( + dataset_id=dataset_id, + tenant_id=tenant_id, + indexing_technique=indexing_technique, + index_struct=index_struct, + collection_binding_id=collection_binding_id, + doc_form=IndexStructureType.PARAGRAPH_INDEX, + ) # Assert mock_index_processor_factory["processor"].clean.assert_called_once() @@ -462,3 +463,28 @@ class TestIndexProcessorParameters: assert call_args[1]["session"] is mock_db_session.session assert call_args[1]["with_keywords"] is True assert call_args[1]["delete_child_chunks"] is True + schedule_refresh.assert_called_once_with(tenant_id) + + def test_vector_cleanup_failure_does_not_schedule_billing_refresh( + self, + dataset_id: str, + tenant_id: str, + collection_binding_id: str, + mock_db_session, + mock_storage, + mock_index_processor_factory, + mock_get_image_upload_file_ids, + ): + mock_index_processor_factory["processor"].clean.side_effect = RuntimeError("vector cleanup failed") + + with patch("tasks.clean_dataset_task.schedule_billing_vector_space_refresh") as schedule_refresh: + clean_dataset_task( + dataset_id=dataset_id, + tenant_id=tenant_id, + indexing_technique=IndexTechniqueType.HIGH_QUALITY, + index_struct='{"type": "paragraph"}', + collection_binding_id=collection_binding_id, + doc_form=IndexStructureType.PARAGRAPH_INDEX, + ) + + schedule_refresh.assert_not_called() diff --git a/api/tests/unit_tests/tasks/test_clean_document_task.py b/api/tests/unit_tests/tasks/test_clean_document_task.py index 26d7b3e3b6b..2f517ce1ba4 100644 --- a/api/tests/unit_tests/tasks/test_clean_document_task.py +++ b/api/tests/unit_tests/tasks/test_clean_document_task.py @@ -169,12 +169,13 @@ class TestVectorCleanupResilience: ) # Act — must not raise out of the task even though clean() raises. - clean_document_task( - document_id=document_id, - dataset_id=dataset_id, - doc_form="paragraph", - file_id=None, - ) + with patch("tasks.clean_document_task.schedule_billing_vector_space_refresh") as schedule_refresh: + clean_document_task( + document_id=document_id, + dataset_id=dataset_id, + doc_form="paragraph", + file_id=None, + ) # Assert # 1. Vector cleanup was attempted. @@ -187,6 +188,7 @@ class TestVectorCleanupResilience: "Step 3+ DB cleanup did not run after vector cleanup failure; " "this regression would re-introduce the orphan-segment bug." ) + schedule_refresh.assert_not_called() def test_vector_cleanup_success_path_remains_unaffected( self, @@ -229,12 +231,13 @@ class TestVectorCleanupResilience: mock_sf.create_session.side_effect = [cm1, cm2] + [_default_cm() for _ in range(10)] - clean_document_task( - document_id=document_id, - dataset_id=dataset_id, - doc_form="paragraph", - file_id=None, - ) + with patch("tasks.clean_document_task.schedule_billing_vector_space_refresh") as schedule_refresh: + clean_document_task( + document_id=document_id, + dataset_id=dataset_id, + doc_form="paragraph", + file_id=None, + ) assert mock_index_processor_factory["processor"].clean.call_count == 1 # Index cleanup invoked with the expected delete_summaries / delete_child_chunks flags. @@ -242,6 +245,7 @@ class TestVectorCleanupResilience: assert kwargs.get("with_keywords") is True assert kwargs.get("delete_child_chunks") is True assert kwargs.get("delete_summaries") is True + schedule_refresh.assert_called_once_with(tenant_id) def test_no_segments_skips_vector_cleanup( self, @@ -279,13 +283,15 @@ class TestVectorCleanupResilience: mock_sf.create_session.side_effect = [cm1] + [_default_cm() for _ in range(10)] - clean_document_task( - document_id=document_id, - dataset_id=dataset_id, - doc_form="paragraph", - file_id=None, - ) + with patch("tasks.clean_document_task.schedule_billing_vector_space_refresh") as schedule_refresh: + clean_document_task( + document_id=document_id, + dataset_id=dataset_id, + doc_form="paragraph", + file_id=None, + ) # Vector cleanup is gated on ``index_node_ids``; when there are no # segments the IndexProcessorFactory path is never entered. mock_index_processor_factory["factory_cls"].assert_not_called() + schedule_refresh.assert_not_called() diff --git a/api/tests/unit_tests/tasks/test_refresh_billing_vector_space_task.py b/api/tests/unit_tests/tasks/test_refresh_billing_vector_space_task.py new file mode 100644 index 00000000000..8f4820be660 --- /dev/null +++ b/api/tests/unit_tests/tasks/test_refresh_billing_vector_space_task.py @@ -0,0 +1,45 @@ +from unittest.mock import patch + +import pytest + +from tasks.refresh_billing_vector_space_task import ( + refresh_billing_vector_space_task, + schedule_billing_vector_space_refresh, +) + + +def test_refresh_invalidates_vector_space_cache(): + with ( + patch("tasks.refresh_billing_vector_space_task.dify_config.BILLING_ENABLED", True), + patch( + "tasks.refresh_billing_vector_space_task.BillingService.invalidate_vector_space_cache" + ) as invalidate_cache, + ): + refresh_billing_vector_space_task.run("tenant-1") + + invalidate_cache.assert_called_once_with("tenant-1") + + +def test_refresh_failure_schedules_retry(): + error = RuntimeError("billing unavailable") + + with ( + patch("tasks.refresh_billing_vector_space_task.dify_config.BILLING_ENABLED", True), + patch( + "tasks.refresh_billing_vector_space_task.BillingService.invalidate_vector_space_cache", + side_effect=error, + ), + patch.object(refresh_billing_vector_space_task, "retry", side_effect=RuntimeError("retry scheduled")) as retry, + pytest.raises(RuntimeError, match="retry scheduled"), + ): + refresh_billing_vector_space_task.run("tenant-1") + + retry.assert_called_once_with(exc=error, countdown=30) + + +def test_dispatch_failure_does_not_propagate(): + with ( + patch("tasks.refresh_billing_vector_space_task.dify_config.BILLING_ENABLED", True), + patch.object(refresh_billing_vector_space_task, "delay", side_effect=RuntimeError("broker unavailable")), + ): + schedule_billing_vector_space_refresh("tenant-1")