From d9c99daf29edf37c8b89ac1a6ebd2dede82c6753 Mon Sep 17 00:00:00 2001 From: Pranav Agarwal Date: Mon, 6 Jul 2026 10:47:12 +0530 Subject: [PATCH 01/70] fix(api): isolate side-effect session writes in multimodal and RAG handlers (#38210) Co-authored-by: FFXN <31929997+FFXN@users.noreply.github.com> --- api/core/app/apps/base_app_runner.py | 11 +- .../index_tool_callback_handler.py | 76 +++-- .../chat/test_base_app_runner_multimodal.py | 298 ++++++++---------- .../test_index_tool_callback_handler.py | 132 ++++++-- 4 files changed, 286 insertions(+), 231 deletions(-) diff --git a/api/core/app/apps/base_app_runner.py b/api/core/app/apps/base_app_runner.py index 7b854fec34a..941ae6b330b 100644 --- a/api/core/app/apps/base_app_runner.py +++ b/api/core/app/apps/base_app_runner.py @@ -5,6 +5,8 @@ from collections.abc import Generator, Mapping, Sequence from mimetypes import guess_extension from typing import TYPE_CHECKING, Any, Union +from sqlalchemy.orm import sessionmaker + from core.app.app_config.entities import ExternalDataVariableEntity, PromptTemplateEntity from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom from core.app.apps.exc import GenerateTaskStoppedError @@ -423,7 +425,9 @@ class AppRunner: _logger.exception("Failed to save image file") return - # Create MessageFile record + # Create MessageFile record. + # Use an independent session so this side-effect write does not + # commit or close the caller's request-scoped session. message_file = MessageFile( message_id=message_id, type=FileType.IMAGE, @@ -437,9 +441,8 @@ class AppRunner: created_by=user_id, ) - db.session.add(message_file) - db.session.commit() - db.session.refresh(message_file) + with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as session: + session.add(message_file) # Publish QueueMessageFileEvent queue_manager.publish( diff --git a/api/core/callback_handler/index_tool_callback_handler.py b/api/core/callback_handler/index_tool_callback_handler.py index 5494769082e..26dc1a12a2c 100644 --- a/api/core/callback_handler/index_tool_callback_handler.py +++ b/api/core/callback_handler/index_tool_callback_handler.py @@ -2,7 +2,7 @@ import logging from collections.abc import Sequence from sqlalchemy import select, update -from sqlalchemy.orm import scoped_session +from sqlalchemy.orm import Session, scoped_session, sessionmaker from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom from core.app.entities.app_invoke_entities import InvokeFrom @@ -10,6 +10,7 @@ from core.app.entities.queue_entities import QueueRetrieverResourcesEvent from core.rag.entities import RetrievalSourceMetadata from core.rag.index_processor.constant.index_type import IndexStructureType from core.rag.models.document import Document +from extensions.ext_database import db from models.dataset import ChildChunk, DatasetQuery, DocumentSegment from models.dataset import Document as DatasetDocument from models.enums import CreatorUserRole, DatasetQuerySource @@ -46,47 +47,52 @@ class DatasetIndexToolCallbackHandler: created_by=self._user_id, ) - session.add(dataset_query) - session.commit() + # Use an independent session so this audit-log side effect does + # not commit or close the caller's request-scoped session. + with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as independent_session: + independent_session.add(dataset_query) def on_tool_end(self, documents: list[Document], session: scoped_session): """Handle tool end.""" - for document in documents: - if document.metadata is not None: - document_id = document.metadata["document_id"] - dataset_document_stmt = select(DatasetDocument).where(DatasetDocument.id == document_id) - dataset_document = session.scalar(dataset_document_stmt) - if not dataset_document: - _logger.warning( - "Expected DatasetDocument record to exist, but none was found, document_id=%s", - document_id, - ) - continue - if dataset_document.doc_form == IndexStructureType.PARENT_CHILD_INDEX: - child_chunk_stmt = select(ChildChunk).where( - ChildChunk.index_node_id == document.metadata["doc_id"], - ChildChunk.dataset_id == dataset_document.dataset_id, - ChildChunk.document_id == dataset_document.id, - ) - child_chunk = session.scalar(child_chunk_stmt) - if child_chunk: - session.execute( - update(DocumentSegment) - .where(DocumentSegment.id == child_chunk.segment_id) - .values(hit_count=DocumentSegment.hit_count + 1) + # Use an independent session so hit-count updates do not + # interfere with the caller's request-scoped session. + with Session(db.engine, expire_on_commit=False) as independent_session: + for document in documents: + if document.metadata is not None: + document_id = document.metadata["document_id"] + dataset_document_stmt = select(DatasetDocument).where(DatasetDocument.id == document_id) + dataset_document = independent_session.scalar(dataset_document_stmt) + if not dataset_document: + _logger.warning( + "Expected DatasetDocument record to exist, but none was found, document_id=%s", + document_id, ) - else: - conditions = [DocumentSegment.index_node_id == document.metadata["doc_id"]] + continue + if dataset_document.doc_form == IndexStructureType.PARENT_CHILD_INDEX: + child_chunk_stmt = select(ChildChunk).where( + ChildChunk.index_node_id == document.metadata["doc_id"], + ChildChunk.dataset_id == dataset_document.dataset_id, + ChildChunk.document_id == dataset_document.id, + ) + child_chunk = independent_session.scalar(child_chunk_stmt) + if child_chunk: + independent_session.execute( + update(DocumentSegment) + .where(DocumentSegment.id == child_chunk.segment_id) + .values(hit_count=DocumentSegment.hit_count + 1) + ) + else: + conditions = [DocumentSegment.index_node_id == document.metadata["doc_id"]] - if "dataset_id" in document.metadata: - conditions.append(DocumentSegment.dataset_id == document.metadata["dataset_id"]) + if "dataset_id" in document.metadata: + conditions.append(DocumentSegment.dataset_id == document.metadata["dataset_id"]) - # add hit count to document segment - session.execute( - update(DocumentSegment).where(*conditions).values(hit_count=DocumentSegment.hit_count + 1) - ) + # add hit count to document segment + independent_session.execute( + update(DocumentSegment).where(*conditions).values(hit_count=DocumentSegment.hit_count + 1) + ) - session.commit() + independent_session.commit() # TODO(-LAN-): Improve type check def return_retriever_resource_info(self, resource: Sequence[RetrievalSourceMetadata]): diff --git a/api/tests/unit_tests/core/app/apps/chat/test_base_app_runner_multimodal.py b/api/tests/unit_tests/core/app/apps/chat/test_base_app_runner_multimodal.py index b3ea1a464f8..130264972a3 100644 --- a/api/tests/unit_tests/core/app/apps/chat/test_base_app_runner_multimodal.py +++ b/api/tests/unit_tests/core/app/apps/chat/test_base_app_runner_multimodal.py @@ -5,7 +5,6 @@ from uuid import uuid4 import pytest -from core.app.apps.base_app_queue_manager import PublishFrom from core.app.apps.base_app_runner import AppRunner from core.app.entities.app_invoke_entities import InvokeFrom from core.app.entities.queue_entities import QueueMessageFileEvent @@ -81,59 +80,55 @@ class TestBaseAppRunnerMultimodal: # Setup mock message file mock_msg_file_class.return_value = mock_message_file - with patch("core.app.apps.base_app_runner.db.session", autospec=True) as mock_session: - mock_session.add = MagicMock() - mock_session.commit = MagicMock() - mock_session.refresh = MagicMock() + file_session = MagicMock() + mock_session_factory = MagicMock() + mock_session_factory.begin.return_value.__enter__ = MagicMock(return_value=file_session) + mock_session_factory.begin.return_value.__exit__ = MagicMock(return_value=False) - # Act - # Create a mock runner with the method bound - runner = MagicMock() + with patch("core.app.apps.base_app_runner.sessionmaker", return_value=mock_session_factory) as mock_sm: + with patch("core.app.apps.base_app_runner.db") as mock_db: + # Act + runner = MagicMock() + method = AppRunner._handle_multimodal_image_content + runner._handle_multimodal_image_content = lambda *args, **kwargs: method( + runner, *args, **kwargs + ) - method = AppRunner._handle_multimodal_image_content - runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) + runner._handle_multimodal_image_content( + content=content, + message_id=mock_message_id, + user_id=mock_user_id, + tenant_id=mock_tenant_id, + queue_manager=mock_queue_manager, + ) - runner._handle_multimodal_image_content( - content=content, - message_id=mock_message_id, - user_id=mock_user_id, - tenant_id=mock_tenant_id, - queue_manager=mock_queue_manager, - ) + # Assert + mock_mgr.create_file_by_url.assert_called_once_with( + user_id=mock_user_id, + tenant_id=mock_tenant_id, + file_url=image_url, + conversation_id=None, + ) - # Assert - # Verify tool file was created from URL - mock_mgr.create_file_by_url.assert_called_once_with( - user_id=mock_user_id, - tenant_id=mock_tenant_id, - file_url=image_url, - conversation_id=None, - ) + mock_msg_file_class.assert_called_once() + call_kwargs = mock_msg_file_class.call_args[1] + assert call_kwargs["message_id"] == mock_message_id + assert call_kwargs["type"] == FileType.IMAGE + assert call_kwargs["transfer_method"] == FileTransferMethod.TOOL_FILE + assert call_kwargs["belongs_to"] == "assistant" + assert call_kwargs["created_by"] == mock_user_id - # Verify message file was created with correct parameters - mock_msg_file_class.assert_called_once() - call_kwargs = mock_msg_file_class.call_args[1] - assert call_kwargs["message_id"] == mock_message_id - assert call_kwargs["type"] == FileType.IMAGE - assert call_kwargs["transfer_method"] == FileTransferMethod.TOOL_FILE - assert call_kwargs["belongs_to"] == "assistant" - assert call_kwargs["created_by"] == mock_user_id + # Verify independent session was used (not db.session) + mock_sm.assert_called_once_with(bind=mock_db.engine, expire_on_commit=False) + file_session.add.assert_called_once_with(mock_message_file) + mock_db.session.commit.assert_not_called() + mock_db.session.close.assert_not_called() - # Verify database operations - mock_session.add.assert_called_once_with(mock_message_file) - mock_session.commit.assert_called_once() - mock_session.refresh.assert_called_once_with(mock_message_file) - - # Verify event was published - mock_queue_manager.publish.assert_called_once() - publish_call = mock_queue_manager.publish.call_args - assert isinstance(publish_call[0][0], QueueMessageFileEvent) - assert publish_call[0][0].message_file_id == mock_message_file.id - # publish_from might be passed as positional or keyword argument - assert ( - publish_call[0][1] == PublishFrom.APPLICATION_MANAGER - or publish_call.kwargs.get("publish_from") == PublishFrom.APPLICATION_MANAGER - ) + # Verify event was published + mock_queue_manager.publish.assert_called_once() + publish_call = mock_queue_manager.publish.call_args + assert isinstance(publish_call[0][0], QueueMessageFileEvent) + assert publish_call[0][0].message_file_id == mock_message_file.id def test_handle_multimodal_image_content_with_base64( self, @@ -165,50 +160,44 @@ class TestBaseAppRunnerMultimodal: mock_mgr_class.return_value = mock_mgr with patch("core.app.apps.base_app_runner.MessageFile", autospec=True) as mock_msg_file_class: - # Setup mock message file mock_msg_file_class.return_value = mock_message_file - with patch("core.app.apps.base_app_runner.db.session", autospec=True) as mock_session: - mock_session.add = MagicMock() - mock_session.commit = MagicMock() - mock_session.refresh = MagicMock() + file_session = MagicMock() + mock_session_factory = MagicMock() + mock_session_factory.begin.return_value.__enter__ = MagicMock(return_value=file_session) + mock_session_factory.begin.return_value.__exit__ = MagicMock(return_value=False) - # Act - # Create a mock runner with the method bound - runner = MagicMock() - method = AppRunner._handle_multimodal_image_content - runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) + with patch("core.app.apps.base_app_runner.sessionmaker", return_value=mock_session_factory): + with patch("core.app.apps.base_app_runner.db") as mock_db: + runner = MagicMock() + method = AppRunner._handle_multimodal_image_content + runner._handle_multimodal_image_content = lambda *args, **kwargs: method( + runner, *args, **kwargs + ) - runner._handle_multimodal_image_content( - content=content, - message_id=mock_message_id, - user_id=mock_user_id, - tenant_id=mock_tenant_id, - queue_manager=mock_queue_manager, - ) + runner._handle_multimodal_image_content( + content=content, + message_id=mock_message_id, + user_id=mock_user_id, + tenant_id=mock_tenant_id, + queue_manager=mock_queue_manager, + ) - # Assert - # Verify tool file was created from base64 - mock_mgr.create_file_by_raw.assert_called_once() - call_kwargs = mock_mgr.create_file_by_raw.call_args[1] - assert call_kwargs["user_id"] == mock_user_id - assert call_kwargs["tenant_id"] == mock_tenant_id - assert call_kwargs["conversation_id"] is None - assert "file_binary" in call_kwargs - assert call_kwargs["mimetype"] == "image/png" - assert call_kwargs["filename"].startswith("generated_image") - assert call_kwargs["filename"].endswith(".png") + mock_mgr.create_file_by_raw.assert_called_once() + call_kwargs = mock_mgr.create_file_by_raw.call_args[1] + assert call_kwargs["user_id"] == mock_user_id + assert call_kwargs["tenant_id"] == mock_tenant_id + assert call_kwargs["conversation_id"] is None + assert "file_binary" in call_kwargs + assert call_kwargs["mimetype"] == "image/png" + assert call_kwargs["filename"].startswith("generated_image") + assert call_kwargs["filename"].endswith(".png") - # Verify message file was created - mock_msg_file_class.assert_called_once() + mock_msg_file_class.assert_called_once() + file_session.add.assert_called_once() + mock_db.session.commit.assert_not_called() - # Verify database operations - mock_session.add.assert_called_once() - mock_session.commit.assert_called_once() - mock_session.refresh.assert_called_once() - - # Verify event was published - mock_queue_manager.publish.assert_called_once() + mock_queue_manager.publish.assert_called_once() def test_handle_multimodal_image_content_with_base64_data_uri( self, @@ -238,33 +227,32 @@ class TestBaseAppRunnerMultimodal: mock_mgr_class.return_value = mock_mgr with patch("core.app.apps.base_app_runner.MessageFile", autospec=True) as mock_msg_file_class: - # Setup mock message file mock_msg_file_class.return_value = mock_message_file - with patch("core.app.apps.base_app_runner.db.session", autospec=True) as mock_session: - mock_session.add = MagicMock() - mock_session.commit = MagicMock() - mock_session.refresh = MagicMock() + file_session = MagicMock() + mock_session_factory = MagicMock() + mock_session_factory.begin.return_value.__enter__ = MagicMock(return_value=file_session) + mock_session_factory.begin.return_value.__exit__ = MagicMock(return_value=False) - # Act - # Create a mock runner with the method bound - runner = MagicMock() - method = AppRunner._handle_multimodal_image_content - runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) + with patch("core.app.apps.base_app_runner.sessionmaker", return_value=mock_session_factory): + with patch("core.app.apps.base_app_runner.db"): + runner = MagicMock() + method = AppRunner._handle_multimodal_image_content + runner._handle_multimodal_image_content = lambda *args, **kwargs: method( + runner, *args, **kwargs + ) - runner._handle_multimodal_image_content( - content=content, - message_id=mock_message_id, - user_id=mock_user_id, - tenant_id=mock_tenant_id, - queue_manager=mock_queue_manager, - ) + runner._handle_multimodal_image_content( + content=content, + message_id=mock_message_id, + user_id=mock_user_id, + tenant_id=mock_tenant_id, + queue_manager=mock_queue_manager, + ) - # Assert - verify that base64 data was extracted correctly (without prefix) - mock_mgr.create_file_by_raw.assert_called_once() - call_kwargs = mock_mgr.create_file_by_raw.call_args[1] - # The base64 data should be decoded, so we check the binary was passed - assert "file_binary" in call_kwargs + mock_mgr.create_file_by_raw.assert_called_once() + call_kwargs = mock_mgr.create_file_by_raw.call_args[1] + assert "file_binary" in call_kwargs def test_handle_multimodal_image_content_without_url_or_base64( self, @@ -284,9 +272,7 @@ class TestBaseAppRunnerMultimodal: with patch("core.app.apps.base_app_runner.ToolFileManager", autospec=True) as mock_mgr_class: with patch("core.app.apps.base_app_runner.MessageFile", autospec=True) as mock_msg_file_class: - with patch("core.app.apps.base_app_runner.db.session", autospec=True) as mock_session: - # Act - # Create a mock runner with the method bound + with patch("core.app.apps.base_app_runner.db"): runner = MagicMock() method = AppRunner._handle_multimodal_image_content runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) @@ -299,10 +285,8 @@ class TestBaseAppRunnerMultimodal: queue_manager=mock_queue_manager, ) - # Assert - should not create any files or publish events mock_mgr_class.assert_not_called() mock_msg_file_class.assert_not_called() - mock_session.add.assert_not_called() mock_queue_manager.publish.assert_not_called() def test_handle_multimodal_image_content_with_error( @@ -322,20 +306,16 @@ class TestBaseAppRunnerMultimodal: ) with patch("core.app.apps.base_app_runner.ToolFileManager", autospec=True) as mock_mgr_class: - # Setup mock to raise exception mock_mgr = MagicMock() mock_mgr.create_file_by_url.side_effect = Exception("Network error") mock_mgr_class.return_value = mock_mgr with patch("core.app.apps.base_app_runner.MessageFile", autospec=True) as mock_msg_file_class: - with patch("core.app.apps.base_app_runner.db.session", autospec=True) as mock_session: - # Act - # Create a mock runner with the method bound + with patch("core.app.apps.base_app_runner.db"): runner = MagicMock() method = AppRunner._handle_multimodal_image_content runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) - # Should not raise exception, just log it runner._handle_multimodal_image_content( content=content, message_id=mock_message_id, @@ -344,9 +324,7 @@ class TestBaseAppRunnerMultimodal: queue_manager=mock_queue_manager, ) - # Assert - should not create message file or publish event on error mock_msg_file_class.assert_not_called() - mock_session.add.assert_not_called() mock_queue_manager.publish.assert_not_called() def test_handle_multimodal_image_content_debugger_mode( @@ -369,37 +347,36 @@ class TestBaseAppRunnerMultimodal: mock_queue_manager.invoke_from = InvokeFrom.DEBUGGER with patch("core.app.apps.base_app_runner.ToolFileManager", autospec=True) as mock_mgr_class: - # Setup mock tool file manager mock_mgr = MagicMock() mock_mgr.create_file_by_url.return_value = mock_tool_file mock_mgr_class.return_value = mock_mgr with patch("core.app.apps.base_app_runner.MessageFile", autospec=True) as mock_msg_file_class: - # Setup mock message file mock_msg_file_class.return_value = mock_message_file - with patch("core.app.apps.base_app_runner.db.session", autospec=True) as mock_session: - mock_session.add = MagicMock() - mock_session.commit = MagicMock() - mock_session.refresh = MagicMock() + file_session = MagicMock() + mock_session_factory = MagicMock() + mock_session_factory.begin.return_value.__enter__ = MagicMock(return_value=file_session) + mock_session_factory.begin.return_value.__exit__ = MagicMock(return_value=False) - # Act - # Create a mock runner with the method bound - runner = MagicMock() - method = AppRunner._handle_multimodal_image_content - runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) + with patch("core.app.apps.base_app_runner.sessionmaker", return_value=mock_session_factory): + with patch("core.app.apps.base_app_runner.db"): + runner = MagicMock() + method = AppRunner._handle_multimodal_image_content + runner._handle_multimodal_image_content = lambda *args, **kwargs: method( + runner, *args, **kwargs + ) - runner._handle_multimodal_image_content( - content=content, - message_id=mock_message_id, - user_id=mock_user_id, - tenant_id=mock_tenant_id, - queue_manager=mock_queue_manager, - ) + runner._handle_multimodal_image_content( + content=content, + message_id=mock_message_id, + user_id=mock_user_id, + tenant_id=mock_tenant_id, + queue_manager=mock_queue_manager, + ) - # Assert - verify created_by_role is ACCOUNT for debugger mode - call_kwargs = mock_msg_file_class.call_args[1] - assert call_kwargs["created_by_role"] == CreatorUserRole.ACCOUNT + call_kwargs = mock_msg_file_class.call_args[1] + assert call_kwargs["created_by_role"] == CreatorUserRole.ACCOUNT def test_handle_multimodal_image_content_service_api_mode( self, @@ -421,34 +398,33 @@ class TestBaseAppRunnerMultimodal: mock_queue_manager.invoke_from = InvokeFrom.SERVICE_API with patch("core.app.apps.base_app_runner.ToolFileManager", autospec=True) as mock_mgr_class: - # Setup mock tool file manager mock_mgr = MagicMock() mock_mgr.create_file_by_url.return_value = mock_tool_file mock_mgr_class.return_value = mock_mgr with patch("core.app.apps.base_app_runner.MessageFile", autospec=True) as mock_msg_file_class: - # Setup mock message file mock_msg_file_class.return_value = mock_message_file - with patch("core.app.apps.base_app_runner.db.session", autospec=True) as mock_session: - mock_session.add = MagicMock() - mock_session.commit = MagicMock() - mock_session.refresh = MagicMock() + file_session = MagicMock() + mock_session_factory = MagicMock() + mock_session_factory.begin.return_value.__enter__ = MagicMock(return_value=file_session) + mock_session_factory.begin.return_value.__exit__ = MagicMock(return_value=False) - # Act - # Create a mock runner with the method bound - runner = MagicMock() - method = AppRunner._handle_multimodal_image_content - runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) + with patch("core.app.apps.base_app_runner.sessionmaker", return_value=mock_session_factory): + with patch("core.app.apps.base_app_runner.db"): + runner = MagicMock() + method = AppRunner._handle_multimodal_image_content + runner._handle_multimodal_image_content = lambda *args, **kwargs: method( + runner, *args, **kwargs + ) - runner._handle_multimodal_image_content( - content=content, - message_id=mock_message_id, - user_id=mock_user_id, - tenant_id=mock_tenant_id, - queue_manager=mock_queue_manager, - ) + runner._handle_multimodal_image_content( + content=content, + message_id=mock_message_id, + user_id=mock_user_id, + tenant_id=mock_tenant_id, + queue_manager=mock_queue_manager, + ) - # Assert - verify created_by_role is END_USER for service API - call_kwargs = mock_msg_file_class.call_args[1] - assert call_kwargs["created_by_role"] == CreatorUserRole.END_USER + call_kwargs = mock_msg_file_class.call_args[1] + assert call_kwargs["created_by_role"] == CreatorUserRole.END_USER diff --git a/api/tests/unit_tests/core/callback_handler/test_index_tool_callback_handler.py b/api/tests/unit_tests/core/callback_handler/test_index_tool_callback_handler.py index 4912badfc55..62c4ae9d411 100644 --- a/api/tests/unit_tests/core/callback_handler/test_index_tool_callback_handler.py +++ b/api/tests/unit_tests/core/callback_handler/test_index_tool_callback_handler.py @@ -14,6 +14,9 @@ def mock_queue_manager(mocker: MockerFixture): @pytest.fixture def handler(mock_queue_manager, mocker: MockerFixture): + mocker.patch( + "core.callback_handler.index_tool_callback_handler.db", + ) return DatasetIndexToolCallbackHandler( queue_manager=mock_queue_manager, app_id="app-1", @@ -33,8 +36,18 @@ class TestOnQuery: ], ) def test_on_query_success_roles(self, mocker: MockerFixture, mock_queue_manager, invoke_from, expected_role): - # Arrange - mock_session = mocker.Mock() + # Arrange — the caller passes a session, but our fix uses an independent one + caller_session = mocker.Mock() + + independent_session = mocker.MagicMock() + mock_session_factory = mocker.MagicMock() + mock_session_factory.begin.return_value.__enter__ = mocker.MagicMock(return_value=independent_session) + mock_session_factory.begin.return_value.__exit__ = mocker.MagicMock(return_value=False) + mocker.patch( + "core.callback_handler.index_tool_callback_handler.sessionmaker", + return_value=mock_session_factory, + ) + mocker.patch("core.callback_handler.index_tool_callback_handler.db") handler = DatasetIndexToolCallbackHandler( queue_manager=mock_queue_manager, @@ -46,17 +59,28 @@ class TestOnQuery: handler._invoke_from = invoke_from - # Act - handler.on_query("test query", "dataset-1", mock_session) + # Act — pass caller_session as required by signature + handler.on_query("test query", "dataset-1", caller_session) - # Assert - mock_session.add.assert_called_once() - dataset_query = mock_session.add.call_args.args[0] + # Assert — independent session used, not the caller's session + independent_session.add.assert_called_once() + dataset_query = independent_session.add.call_args.args[0] assert dataset_query.created_by_role == expected_role - mock_session.commit.assert_called_once() + caller_session.add.assert_not_called() + caller_session.commit.assert_not_called() def test_on_query_none_values(self, mocker: MockerFixture, mock_queue_manager): - mock_session = mocker.Mock() + caller_session = mocker.Mock() + + independent_session = mocker.MagicMock() + mock_session_factory = mocker.MagicMock() + mock_session_factory.begin.return_value.__enter__ = mocker.MagicMock(return_value=independent_session) + mock_session_factory.begin.return_value.__exit__ = mocker.MagicMock(return_value=False) + mocker.patch( + "core.callback_handler.index_tool_callback_handler.sessionmaker", + return_value=mock_session_factory, + ) + mocker.patch("core.callback_handler.index_tool_callback_handler.db") handler = DatasetIndexToolCallbackHandler( queue_manager=mock_queue_manager, @@ -66,40 +90,67 @@ class TestOnQuery: invoke_from=None, ) - handler.on_query(None, None, mock_session) + handler.on_query(None, None, caller_session) - mock_session.add.assert_called_once() - mock_session.commit.assert_called_once() + independent_session.add.assert_called_once() + caller_session.add.assert_not_called() class TestOnToolEnd: def test_on_tool_end_no_metadata(self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture): - mock_session = mocker.Mock() + caller_session = mocker.Mock() + + independent_session = mocker.MagicMock() + mocker.patch( + "core.callback_handler.index_tool_callback_handler.Session", + return_value=independent_session, + ) + independent_session.__enter__ = mocker.MagicMock(return_value=independent_session) + independent_session.__exit__ = mocker.MagicMock(return_value=False) document = mocker.Mock() document.metadata = None - handler.on_tool_end([document], mock_session) + handler.on_tool_end([document], caller_session) - mock_session.commit.assert_not_called() + independent_session.commit.assert_called_once() + independent_session.execute.assert_not_called() + caller_session.commit.assert_not_called() def test_on_tool_end_dataset_document_not_found( self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture ): - mock_session = mocker.Mock() - mock_session.scalar.return_value = None + caller_session = mocker.Mock() + + independent_session = mocker.MagicMock() + mocker.patch( + "core.callback_handler.index_tool_callback_handler.Session", + return_value=independent_session, + ) + independent_session.__enter__ = mocker.MagicMock(return_value=independent_session) + independent_session.__exit__ = mocker.MagicMock(return_value=False) + independent_session.scalar.return_value = None document = mocker.Mock() document.metadata = {"document_id": "doc-1", "doc_id": "node-1"} - handler.on_tool_end([document], mock_session) + handler.on_tool_end([document], caller_session) - mock_session.scalar.assert_called_once() + independent_session.scalar.assert_called_once() + caller_session.scalar.assert_not_called() def test_on_tool_end_parent_child_index_with_child( self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture ): - mock_session = mocker.Mock() + caller_session = mocker.Mock() + + independent_session = mocker.MagicMock() + mocker.patch( + "core.callback_handler.index_tool_callback_handler.Session", + return_value=independent_session, + ) + independent_session.__enter__ = mocker.MagicMock(return_value=independent_session) + independent_session.__exit__ = mocker.MagicMock(return_value=False) mock_dataset_doc = mocker.Mock() from core.callback_handler.index_tool_callback_handler import IndexStructureType @@ -111,23 +162,32 @@ class TestOnToolEnd: mock_child_chunk = mocker.Mock() mock_child_chunk.segment_id = "segment-1" - mock_session.scalar.side_effect = [mock_dataset_doc, mock_child_chunk] + independent_session.scalar.side_effect = [mock_dataset_doc, mock_child_chunk] document = mocker.Mock() document.metadata = {"document_id": "doc-1", "doc_id": "node-1"} - handler.on_tool_end([document], mock_session) + handler.on_tool_end([document], caller_session) - mock_session.execute.assert_called_once() - mock_session.commit.assert_called_once() + independent_session.execute.assert_called_once() + independent_session.commit.assert_called_once() + caller_session.execute.assert_not_called() def test_on_tool_end_non_parent_child_index(self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture): - mock_session = mocker.Mock() + caller_session = mocker.Mock() + + independent_session = mocker.MagicMock() + mocker.patch( + "core.callback_handler.index_tool_callback_handler.Session", + return_value=independent_session, + ) + independent_session.__enter__ = mocker.MagicMock(return_value=independent_session) + independent_session.__exit__ = mocker.MagicMock(return_value=False) mock_dataset_doc = mocker.Mock() mock_dataset_doc.doc_form = "OTHER" - mock_session.scalar.return_value = mock_dataset_doc + independent_session.scalar.return_value = mock_dataset_doc document = mocker.Mock() document.metadata = { @@ -136,14 +196,24 @@ class TestOnToolEnd: "dataset_id": "dataset-1", } - handler.on_tool_end([document], mock_session) + handler.on_tool_end([document], caller_session) - mock_session.execute.assert_called_once() - mock_session.commit.assert_called_once() + independent_session.execute.assert_called_once() + independent_session.commit.assert_called_once() + caller_session.execute.assert_not_called() def test_on_tool_end_empty_documents(self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture): - mock_session = mocker.Mock() - handler.on_tool_end([], mock_session) + caller_session = mocker.Mock() + + independent_session = mocker.MagicMock() + mocker.patch( + "core.callback_handler.index_tool_callback_handler.Session", + return_value=independent_session, + ) + independent_session.__enter__ = mocker.MagicMock(return_value=independent_session) + independent_session.__exit__ = mocker.MagicMock(return_value=False) + + handler.on_tool_end([], caller_session) class TestReturnRetrieverResourceInfo: From 8208b786ee1589522668a065470b280dec683dba Mon Sep 17 00:00:00 2001 From: yyh <92089059+lyzno1@users.noreply.github.com> Date: Mon, 6 Jul 2026 14:40:17 +0800 Subject: [PATCH 02/70] docs(dify-ui): clarify radio composition stories (#38456) --- .../dify-ui/src/radio-group/index.stories.tsx | 115 +++++++++++++++--- packages/dify-ui/src/radio/index.stories.tsx | 9 +- packages/dify-ui/src/radio/index.tsx | 10 +- 3 files changed, 106 insertions(+), 28 deletions(-) diff --git a/packages/dify-ui/src/radio-group/index.stories.tsx b/packages/dify-ui/src/radio-group/index.stories.tsx index 6c87f6583a5..be005f0491e 100644 --- a/packages/dify-ui/src/radio-group/index.stories.tsx +++ b/packages/dify-ui/src/radio-group/index.stories.tsx @@ -8,7 +8,7 @@ import { FieldRoot, } from '../field' import { FieldsetLegend, FieldsetRoot } from '../fieldset' -import { Radio, RadioControl, RadioRoot } from '../radio' +import { Radio, RadioControl, RadioIndicator, RadioRoot } from '../radio' const meta = { title: 'Base/Form/RadioGroup', @@ -17,7 +17,7 @@ const meta = { layout: 'centered', docs: { description: { - component: 'RadioGroup primitive built on Base UI. For normal form rows, compose FieldRoot, FieldsetRoot, FieldLabel, RadioGroup, and Radio. For option cards, wrap each option in FieldItem and make the card itself a RadioRoot with variant="unstyled".', + component: '`RadioGroup` owns single-selection state. Use `Radio` for plain form rows, `RadioRoot` when an entire row or card is the radio item, `RadioControl` for the standard visual dot inside custom roots, and `RadioIndicator` only when the design owns a custom control shell.', }, }, }, @@ -60,7 +60,7 @@ export const StandardFormRows: Story = { parameters: { docs: { description: { - story: 'Default form composition. Most product code should use this shape: RadioGroup owns value, FieldsetLegend names the group, and FieldLabel makes each row clickable.', + story: 'Plain form-row composition. `RadioGroup` owns value, `FieldsetLegend` names the group, `FieldLabel` makes each option label clickable, and `Radio` renders the default dot.', }, }, }, @@ -107,31 +107,39 @@ export const BooleanInline: Story = { }, } +type PromptMode = 'default' | 'custom' + function OptionCardsDemo() { - const [value, setValue] = React.useState('default') + const [value, setValue] = React.useState('default') + + const options = [ + { + value: 'default', + title: 'Default prompt', + description: 'Use the built-in prompt for consistent output.', + }, + { + value: 'custom', + title: 'Custom prompt', + description: 'Write a prompt for this app and keep full control.', + }, + ] satisfies Array<{ + value: PromptMode + title: string + description: string + }> return ( + value={value} onValueChange={setValue} className="flex-col items-stretch gap-3" /> )} > Prompt mode - {[ - { - value: 'default', - title: 'Default prompt', - description: 'Use the built-in prompt for consistent output.', - }, - { - value: 'custom', - title: 'Custom prompt', - description: 'Write a prompt for this app and keep full control.', - }, - ].map(option => ( + {options.map(option => ( - value={option.value} variant="unstyled" nativeButton @@ -162,7 +170,76 @@ export const OptionCards: Story = { parameters: { docs: { description: { - story: 'Wrap each option card in FieldItem, then use RadioRoot with variant="unstyled" when the entire card is the radio. RadioControl renders the visual dot inside the card.', + story: 'Product option cards should make the whole card the radio item with `RadioRoot variant="unstyled"`. `RadioControl` renders the standard visual dot inside the custom root.', + }, + }, + }, +} + +type ApprovalMode = 'automatic' | 'manual' + +function CustomIndicatorPartDemo() { + const [value, setValue] = React.useState('automatic') + + const options = [ + { + value: 'automatic', + title: 'Automatic approval', + description: 'Approve requests that match policy.', + }, + { + value: 'manual', + title: 'Manual review', + description: 'Ask an admin to review each request.', + }, + ] satisfies Array<{ + value: ApprovalMode + title: string + description: string + }> + + return ( + + value={value} onValueChange={setValue} className="flex-col items-stretch gap-2" /> + )} + > + Approval mode + {options.map(option => ( + + + value={option.value} + variant="unstyled" + nativeButton + render={ )} /> - + {footerNoticeTooltip} - - + + )} -
{footerNotice}
)} diff --git a/web/app/components/workflow/block-selector/__tests__/blocks.spec.tsx b/web/app/components/workflow/block-selector/__tests__/blocks.spec.tsx index 75c1b30d434..ac93a9ef407 100644 --- a/web/app/components/workflow/block-selector/__tests__/blocks.spec.tsx +++ b/web/app/components/workflow/block-selector/__tests__/blocks.spec.tsx @@ -1,3 +1,7 @@ +import type { + AgentInviteOptionResponse, + AgentInviteOptionsResponse, +} from '@dify/contracts/api/console/agent/types.gen' import type { NodeDefault } from '../../types' import { QueryClient, QueryClientProvider } from '@tanstack/react-query' import { render, screen, waitFor } from '@testing-library/react' @@ -15,7 +19,7 @@ const runtimeState = vi.hoisted(() => ({ })) const queryMocks = vi.hoisted(() => ({ - inviteOptionsQueryFn: vi.fn(), + request: vi.fn(), toastError: vi.fn(), })) @@ -35,19 +39,8 @@ vi.mock('@/app/components/app/store', () => ({ }), })) -vi.mock('@/service/client', () => ({ - consoleQuery: { - agent: { - inviteOptions: { - get: { - queryOptions: (options: unknown) => ({ - queryKey: ['agents', 'invite-options', options], - queryFn: () => queryMocks.inviteOptionsQueryFn(options), - }), - }, - }, - }, - }, +vi.mock('@/service/base', () => ({ + request: (...args: unknown[]) => queryMocks.request(...args), })) vi.mock('@langgenius/dify-ui/toast', () => ({ @@ -74,6 +67,58 @@ const createBlock = ( checkValid: () => ({ isValid: true }), }) +const createInviteOption = ( + overrides: Partial & Pick, +): AgentInviteOptionResponse => { + const { id, name, ...rest } = overrides + + return { + id, + name, + description: rest.description ?? 'Clarification Drafter', + active_config_snapshot_id: rest.active_config_snapshot_id ?? 'version-1', + role: rest.role ?? 'Researcher', + agent_kind: rest.agent_kind ?? 'dify_agent', + icon: rest.icon ?? 'A', + icon_background: rest.icon_background ?? '#E9D7FE', + icon_type: rest.icon_type ?? 'emoji', + scope: rest.scope ?? 'roster', + source: rest.source ?? 'workflow', + status: rest.status ?? 'active', + ...rest, + } +} + +const createInviteOptionsResponse = ( + agents: AgentInviteOptionResponse[], +): AgentInviteOptionsResponse => ({ + data: agents, + has_more: false, + limit: 8, + page: 1, + total: agents.length, +}) + +const createJsonResponse = (body: unknown) => + new Response(JSON.stringify(body), { + status: 200, + headers: { + 'Content-Type': 'application/json', + }, + }) + +const mockInviteOptionsResponse = (agents: AgentInviteOptionResponse[]) => { + queryMocks.request.mockImplementation(() => Promise.resolve(createJsonResponse(createInviteOptionsResponse(agents)))) +} + +const expectLastInviteOptionsRequest = () => { + const [url] = queryMocks.request.mock.calls.at(-1) ?? [] + const requestURL = new URL(String(url), window.location.origin) + + expect(requestURL.pathname).toBe('/console/api/agent/invite-options') + return requestURL +} + describe('Blocks', () => { beforeEach(() => { vi.clearAllMocks() @@ -125,13 +170,7 @@ describe('Blocks', () => { it('opens the agent selector on Agent block hover', async () => { const user = userEvent.setup() - queryMocks.inviteOptionsQueryFn.mockResolvedValue({ - data: [], - has_more: false, - limit: 8, - page: 1, - total: 0, - }) + mockInviteOptionsResponse([]) const queryClient = new QueryClient({ defaultOptions: { queries: { @@ -171,28 +210,12 @@ describe('Blocks', () => { it('opens the agent selector from the Agent block and selects an agent', async () => { const user = userEvent.setup() const onSelect = vi.fn() - queryMocks.inviteOptionsQueryFn.mockResolvedValue({ - data: [ - { - id: 'agent-1', - name: 'Nadia', - description: 'Clarification Drafter', - active_config_snapshot_id: 'version-1', - role: 'Researcher', - agent_kind: 'dify_agent', - icon: 'A', - icon_background: '#E9D7FE', - icon_type: 'emoji', - scope: 'roster', - source: 'workflow', - status: 'active', - }, - ], - has_more: false, - limit: 8, - page: 1, - total: 1, - }) + mockInviteOptionsResponse([ + createInviteOption({ + id: 'agent-1', + name: 'Nadia', + }), + ]) const queryClient = new QueryClient({ defaultOptions: { @@ -246,42 +269,84 @@ describe('Blocks', () => { agent_node_kind: 'dify_agent', version: '2', }) - expect(queryMocks.inviteOptionsQueryFn).toHaveBeenCalledWith({ - input: { - query: { - app_id: 'app-1', - limit: 8, - page: 1, + const requestURL = expectLastInviteOptionsRequest() + expect(requestURL.searchParams.get('app_id')).toBe('app-1') + expect(requestURL.searchParams.get('limit')).toBe('8') + expect(requestURL.searchParams.get('page')).toBe('1') + }) + + it('should refresh Agent v2 roster options when the selector is reopened', async () => { + const user = userEvent.setup() + queryMocks.request + .mockImplementationOnce(() => Promise.resolve(createJsonResponse(createInviteOptionsResponse([ + createInviteOption({ + id: 'agent-1', + name: 'Nadia', + }), + ])))) + .mockImplementation(() => Promise.resolve(createJsonResponse(createInviteOptionsResponse([ + createInviteOption({ + id: 'agent-2', + name: 'Bruno', + role: 'Planner', + }), + ])))) + const queryClient = new QueryClient({ + defaultOptions: { + queries: { + retry: false, + staleTime: 5 * 60 * 1000, }, }, }) + const hooksStore = createHooksStore({ + configsMap: { + flowId: 'app-1', + flowType: FlowType.appFlow, + fileSettings: {} as never, + }, + }) + + render( + + + + + , + ) + + await user.click(screen.getByRole('button', { name: /Agent/ })) + expect(await screen.findByText('Nadia')).toBeInTheDocument() + + await user.click(screen.getByRole('combobox', { name: 'agentV2.roster.searchLabel' })) + await user.keyboard('{Escape}') + await waitFor(() => { + expect(screen.queryByRole('dialog', { name: 'agentV2.roster.nodeSelector.dialogLabel' })).not.toBeInTheDocument() + }) + + await user.click(screen.getByRole('button', { name: /Agent/ })) + + expect(await screen.findByText('Bruno')).toBeInTheDocument() + expect(screen.getByText('Planner')).toBeInTheDocument() + await waitFor(() => expect(queryMocks.request).toHaveBeenCalledTimes(2)) + expect(screen.queryByText('Nadia')).not.toBeInTheDocument() }) it('does not select an Agent v2 roster agent without active config snapshot', async () => { const user = userEvent.setup() const onSelect = vi.fn() - queryMocks.inviteOptionsQueryFn.mockResolvedValue({ - data: [ - { - id: 'agent-1', - name: 'Nadia', - description: 'Clarification Drafter', - active_config_snapshot_id: null, - role: 'Researcher', - agent_kind: 'dify_agent', - icon: 'A', - icon_background: '#E9D7FE', - icon_type: 'emoji', - scope: 'roster', - source: 'workflow', - status: 'active', - }, - ], - has_more: false, - limit: 8, - page: 1, - total: 1, - }) + mockInviteOptionsResponse([ + createInviteOption({ + id: 'agent-1', + name: 'Nadia', + active_config_snapshot_id: null, + }), + ]) const queryClient = new QueryClient({ defaultOptions: { @@ -323,13 +388,7 @@ describe('Blocks', () => { it('inserts an inline Agent v2 node from the selector start action', async () => { const user = userEvent.setup() const onSelect = vi.fn() - queryMocks.inviteOptionsQueryFn.mockResolvedValue({ - data: [], - has_more: false, - limit: 8, - page: 1, - total: 0, - }) + mockInviteOptionsResponse([]) const queryClient = new QueryClient({ defaultOptions: { queries: { @@ -376,13 +435,7 @@ describe('Blocks', () => { it('closes the agent selector when Escape closes the combobox', async () => { const user = userEvent.setup() - queryMocks.inviteOptionsQueryFn.mockResolvedValue({ - data: [], - has_more: false, - limit: 8, - page: 1, - total: 0, - }) + mockInviteOptionsResponse([]) const queryClient = new QueryClient({ defaultOptions: { queries: { diff --git a/web/app/components/workflow/block-selector/agent-selector.tsx b/web/app/components/workflow/block-selector/agent-selector.tsx index f6d7c6db5d0..e3877355e69 100644 --- a/web/app/components/workflow/block-selector/agent-selector.tsx +++ b/web/app/components/workflow/block-selector/agent-selector.tsx @@ -66,6 +66,7 @@ export function AgentSelectorContent({ }, }, }), + staleTime: 0, }) const agents = agentsQuery.data?.data ?? [] const actionOptions: AgentSelectorActionOption[] = onStartFromScratch diff --git a/web/app/components/workflow/nodes/agent-v2/__tests__/default.spec.ts b/web/app/components/workflow/nodes/agent-v2/__tests__/default.spec.ts index 7460739778f..d040321886e 100644 --- a/web/app/components/workflow/nodes/agent-v2/__tests__/default.spec.ts +++ b/web/app/components/workflow/nodes/agent-v2/__tests__/default.spec.ts @@ -101,6 +101,10 @@ describe('agent/default', () => { }) }) + it('reuses the legacy agent node help document', () => { + expect(nodeDefault.metaData.helpLinkUri).toBe('agent') + }) + it('identifies version 2 agent data as Agent v2', () => { expect(isAgentV2NodeData(createPayload({ type: BlockEnum.Agent }))).toBe(true) expect(isAgentV2NodeData({ diff --git a/web/app/components/workflow/nodes/agent-v2/default.ts b/web/app/components/workflow/nodes/agent-v2/default.ts index b1273030112..c2b0c557839 100644 --- a/web/app/components/workflow/nodes/agent-v2/default.ts +++ b/web/app/components/workflow/nodes/agent-v2/default.ts @@ -7,6 +7,7 @@ import { hasValidAgentBinding } from './types' const metaData = genNodeMetaData({ sort: 3, type: BlockEnum.AgentV2, + helpLinkUri: 'agent', }) const nodeDefault: NodeDefault = { diff --git a/web/features/agent-v2/agent-composer/store-modules/__tests__/env.spec.ts b/web/features/agent-v2/agent-composer/store-modules/__tests__/env.spec.ts new file mode 100644 index 00000000000..a545055a0ba --- /dev/null +++ b/web/features/agent-v2/agent-composer/store-modules/__tests__/env.spec.ts @@ -0,0 +1,84 @@ +import { createStore } from 'jotai' +import { describe, expect, it } from 'vitest' +import { defaultAgentSoulConfigFormState } from '../../form-state' +import { agentComposerDraftAtom } from '../../store' +import { + addEnvVariableAtom, + importEnvVariablesAtom, + removeEnvVariableAtom, + setEnvVariableKeyAtom, + setEnvVariableValueAtom, +} from '../env' + +const starterVariable = { + id: 'starter', + key: '', + value: '', + scope: 'plain', +} as const + +describe('agent composer env store', () => { + it('should promote the starter variable when editing an empty env list', () => { + const store = createStore() + store.set(agentComposerDraftAtom, defaultAgentSoulConfigFormState) + + store.set(setEnvVariableKeyAtom, { + id: starterVariable.id, + key: 'API_KEY', + starterVariable, + }) + store.set(setEnvVariableValueAtom, { + id: starterVariable.id, + starterVariable, + value: 'secret-value', + }) + + expect(store.get(agentComposerDraftAtom).envVariables).toEqual([ + { + id: 'starter', + key: 'API_KEY', + value: 'secret-value', + scope: 'plain', + }, + ]) + }) + + it('should add, import, and remove variables from the latest draft state', () => { + const store = createStore() + store.set(agentComposerDraftAtom, defaultAgentSoulConfigFormState) + + store.set(addEnvVariableAtom, { + starterVariable, + variable: { + id: 'env-1', + key: 'FIRST_KEY', + value: '', + scope: 'plain', + }, + }) + store.set(importEnvVariablesAtom, [ + { + id: 'env-2', + key: 'SECOND_KEY', + value: 'enabled', + scope: 'plain', + }, + ]) + store.set(removeEnvVariableAtom, 'starter') + + expect(store.get(agentComposerDraftAtom).envVariables).toEqual([ + { + id: 'env-1', + key: 'FIRST_KEY', + value: '', + scope: 'plain', + }, + { + id: 'env-2', + key: 'SECOND_KEY', + value: 'enabled', + scope: 'plain', + }, + ]) + }) +}) diff --git a/web/features/agent-v2/agent-composer/store-modules/__tests__/files.spec.ts b/web/features/agent-v2/agent-composer/store-modules/__tests__/files.spec.ts new file mode 100644 index 00000000000..50f477b2a18 --- /dev/null +++ b/web/features/agent-v2/agent-composer/store-modules/__tests__/files.spec.ts @@ -0,0 +1,70 @@ +import { createStore } from 'jotai' +import { describe, expect, it } from 'vitest' +import { defaultAgentSoulConfigFormState } from '../../form-state' +import { agentComposerDraftAtom } from '../../store' +import { + clearAgentConfigNoteAtom, + removeAgentFileAtom, + upsertAgentFileAtom, +} from '../files' + +describe('agent composer files store', () => { + it('should upsert and remove files from the latest draft state', () => { + const store = createStore() + store.set(agentComposerDraftAtom, { + ...defaultAgentSoulConfigFormState, + files: [ + { + id: 'folder', + icon: 'folder', + name: 'Folder', + children: [ + { + id: 'brief.md', + icon: 'markdown', + name: 'brief.md', + }, + ], + }, + { + id: 'diagram.png', + icon: 'image', + name: 'diagram.png', + }, + ], + }) + + store.set(upsertAgentFileAtom, { + id: 'diagram.png', + icon: 'image', + name: 'updated-diagram.png', + }) + store.set(removeAgentFileAtom, 'brief.md') + + expect(store.get(agentComposerDraftAtom).files).toEqual([ + { + id: 'folder', + icon: 'folder', + name: 'Folder', + children: [], + }, + { + id: 'diagram.png', + icon: 'image', + name: 'updated-diagram.png', + }, + ]) + }) + + it('should clear config note through the file action surface', () => { + const store = createStore() + store.set(agentComposerDraftAtom, { + ...defaultAgentSoulConfigFormState, + configNote: 'Build note', + }) + + store.set(clearAgentConfigNoteAtom) + + expect(store.get(agentComposerDraftAtom).configNote).toBe('') + }) +}) diff --git a/web/features/agent-v2/agent-composer/store-modules/__tests__/knowledge.spec.ts b/web/features/agent-v2/agent-composer/store-modules/__tests__/knowledge.spec.ts new file mode 100644 index 00000000000..782e188e931 --- /dev/null +++ b/web/features/agent-v2/agent-composer/store-modules/__tests__/knowledge.spec.ts @@ -0,0 +1,41 @@ +import { createStore } from 'jotai' +import { describe, expect, it } from 'vitest' +import { defaultAgentSoulConfigFormState } from '../../form-state' +import { agentComposerDraftAtom } from '../../store' +import { + addKnowledgeRetrievalAtom, + removeKnowledgeRetrievalAtom, + updateKnowledgeRetrievalAtom, +} from '../knowledge' + +describe('agent composer knowledge store', () => { + it('should apply retrieval list actions against the latest draft state', () => { + const store = createStore() + store.set(agentComposerDraftAtom, { + ...defaultAgentSoulConfigFormState, + knowledgeRetrievals: [ + { + id: 'retrieval-1', + name: 'Docs Search', + }, + ], + }) + + store.set(addKnowledgeRetrievalAtom, { + id: 'retrieval-2', + name: 'Release Search', + }) + store.set(updateKnowledgeRetrievalAtom, { + id: 'retrieval-1', + name: 'Updated Docs Search', + }) + store.set(removeKnowledgeRetrievalAtom, 'retrieval-2') + + expect(store.get(agentComposerDraftAtom).knowledgeRetrievals).toEqual([ + { + id: 'retrieval-1', + name: 'Updated Docs Search', + }, + ]) + }) +}) diff --git a/web/features/agent-v2/agent-composer/store-modules/__tests__/skills.spec.ts b/web/features/agent-v2/agent-composer/store-modules/__tests__/skills.spec.ts new file mode 100644 index 00000000000..7127c944694 --- /dev/null +++ b/web/features/agent-v2/agent-composer/store-modules/__tests__/skills.spec.ts @@ -0,0 +1,46 @@ +import { createStore } from 'jotai' +import { describe, expect, it } from 'vitest' +import { defaultAgentSoulConfigFormState } from '../../form-state' +import { agentComposerDraftAtom } from '../../store' +import { + removeAgentSkillAtom, + upsertAgentSkillAtom, +} from '../skills' + +describe('agent composer skills store', () => { + it('should upsert and remove skills from the latest draft state', () => { + const store = createStore() + store.set(agentComposerDraftAtom, { + ...defaultAgentSoulConfigFormState, + skills: [ + { + id: 'Tender Analyzer', + name: 'Tender Analyzer', + description: 'Extracts tender requirements.', + fileId: 'tool-file-1', + }, + ], + }) + + store.set(upsertAgentSkillAtom, { + id: 'Tender Analyzer', + name: 'Tender Analyzer', + description: 'Updated skill.', + fileId: 'tool-file-1', + }) + store.set(upsertAgentSkillAtom, { + id: 'Invoice Helper', + name: 'Invoice Helper', + fileId: 'tool-file-2', + }) + store.set(removeAgentSkillAtom, 'Tender Analyzer') + + expect(store.get(agentComposerDraftAtom).skills).toEqual([ + { + id: 'Invoice Helper', + name: 'Invoice Helper', + fileId: 'tool-file-2', + }, + ]) + }) +}) diff --git a/web/features/agent-v2/agent-composer/store-modules/__tests__/tools.spec.ts b/web/features/agent-v2/agent-composer/store-modules/__tests__/tools.spec.ts new file mode 100644 index 00000000000..06723304a3f --- /dev/null +++ b/web/features/agent-v2/agent-composer/store-modules/__tests__/tools.spec.ts @@ -0,0 +1,181 @@ +import type { AgentProviderToolDefaultValue } from '../tools' +import { createStore } from 'jotai' +import { describe, expect, it } from 'vitest' +import { defaultAgentSoulConfigFormState } from '../../form-state' +import { agentComposerDraftAtom } from '../../store' +import { + addProviderTools, + addProviderToolsAtom, + removeProviderToolActionAtom, + saveCliToolAtom, +} from '../tools' + +const noCredentialTool = { + provider_id: 'duckduckgo', + provider_type: 'builtin', + provider_name: 'DuckDuckGo', + provider_show_name: 'DuckDuckGo', + tool_name: 'ddg_search', + tool_label: 'DuckDuckGo Search', + tool_description: 'Search the web.', + title: 'DuckDuckGo Search', + is_team_authorization: true, + params: {}, + paramSchemas: [], + allowDelete: false, + credentialRequired: false, +} satisfies AgentProviderToolDefaultValue + +const unauthorizedCredentialTool = { + ...noCredentialTool, + provider_id: 'google', + provider_name: 'google', + provider_show_name: 'Google', + tool_name: 'search', + tool_label: 'Google Search', + title: 'Google Search', + is_team_authorization: false, + credentialRequired: true, +} satisfies AgentProviderToolDefaultValue + +const unauthorizedOAuthTool = { + ...unauthorizedCredentialTool, + provider_id: 'slack', + provider_name: 'slack', + provider_show_name: 'Slack', + credentialType: 'oauth2', +} satisfies AgentProviderToolDefaultValue + +describe('agent composer tools store', () => { + describe('addProviderTools', () => { + it('should not mark tools that do not need credentials as unauthorized', () => { + const nextTools = addProviderTools([], [noCredentialTool]) + + expect(nextTools).toEqual([ + expect.objectContaining({ + credentialId: undefined, + credentialType: undefined, + credentialVariant: 'none', + }), + ]) + }) + + it('should mark credential-required tools without credentials as unauthorized', () => { + const nextTools = addProviderTools([], [unauthorizedCredentialTool]) + + expect(nextTools).toEqual([ + expect.objectContaining({ + credentialId: undefined, + credentialType: 'unauthorized', + credentialVariant: 'unauthorized', + }), + ]) + }) + + it('should preserve oauth credential type for credential-required OAuth tools', () => { + const nextTools = addProviderTools([], [unauthorizedOAuthTool]) + + expect(nextTools).toEqual([ + expect.objectContaining({ + credentialId: undefined, + credentialType: 'oauth2', + credentialVariant: 'unauthorized', + }), + ]) + }) + }) + + describe('write actions', () => { + it('should apply provider and CLI updates against the latest draft tools', () => { + const store = createStore() + store.set(agentComposerDraftAtom, defaultAgentSoulConfigFormState) + + store.set(addProviderToolsAtom, [noCredentialTool]) + store.set(saveCliToolAtom, { + id: 'cli-tool', + kind: 'cli', + name: 'CLI Tool', + installCommand: 'pnpm install', + }) + store.set(addProviderToolsAtom, [unauthorizedCredentialTool]) + + expect(store.get(agentComposerDraftAtom).tools).toEqual([ + expect.objectContaining({ + id: 'duckduckgo', + kind: 'provider', + }), + expect.objectContaining({ + id: 'cli-tool', + kind: 'cli', + }), + expect.objectContaining({ + id: 'google', + kind: 'provider', + }), + ]) + }) + + it('should update existing CLI tools instead of appending duplicates', () => { + const store = createStore() + store.set(agentComposerDraftAtom, defaultAgentSoulConfigFormState) + + store.set(saveCliToolAtom, { + id: 'cli-tool', + kind: 'cli', + name: 'CLI Tool', + }) + store.set(saveCliToolAtom, { + id: 'cli-tool', + kind: 'cli', + name: 'Updated CLI Tool', + installCommand: 'pnpm install', + }) + + expect(store.get(agentComposerDraftAtom).tools).toEqual([ + { + id: 'cli-tool', + kind: 'cli', + name: 'Updated CLI Tool', + installCommand: 'pnpm install', + }, + ]) + }) + + it('should remove provider action settings with the action', () => { + const store = createStore() + store.set(agentComposerDraftAtom, { + ...defaultAgentSoulConfigFormState, + tools: [ + { + id: 'duckduckgo', + kind: 'provider', + name: 'DuckDuckGo', + iconClassName: 'i-simple-icons-duckduckgo', + credentialVariant: 'none', + actions: [ + { + id: 'duckduckgo:ddg_search', + name: 'DuckDuckGo Search', + toolName: 'ddg_search', + description: 'Search the web.', + }, + ], + }, + ], + toolSettings: { + 'duckduckgo:ddg_search': { + query: 'docs', + }, + }, + }) + + store.set(removeProviderToolActionAtom, { + toolId: 'duckduckgo', + actionId: 'duckduckgo:ddg_search', + }) + + expect(store.get(agentComposerDraftAtom).tools).toEqual([]) + expect(store.get(agentComposerDraftAtom).toolSettings).toEqual({}) + }) + }) +}) diff --git a/web/features/agent-v2/agent-composer/store-modules/env.ts b/web/features/agent-v2/agent-composer/store-modules/env.ts index 8fcdd75b499..0945c01dd15 100644 --- a/web/features/agent-v2/agent-composer/store-modules/env.ts +++ b/web/features/agent-v2/agent-composer/store-modules/env.ts @@ -1,4 +1,4 @@ -import type { EnvVariable } from '../form-state' +import type { EnvScope, EnvVariable } from '../form-state' import type { DraftFieldUpdate } from './utils' import { atom } from 'jotai' import { agentComposerDraftAtom } from '../store' @@ -15,3 +15,95 @@ export const agentComposerEnvVariablesAtom = atom( }) }, ) + +const updateEnvVariable = ( + envVariables: EnvVariable[], + starterVariable: EnvVariable, + id: string, + updater: (variable: EnvVariable) => EnvVariable, +) => { + const existingVariable = envVariables.find(variable => variable.id === id) + + if (existingVariable) { + return envVariables.map(variable => ( + variable.id === id ? updater(variable) : variable + )) + } + + if (id === starterVariable.id) + return [updater(starterVariable)] + + return envVariables +} + +export const setEnvVariableKeyAtom = atom(null, (_get, set, { + id, + key, + starterVariable, +}: { + id: string + key: string + starterVariable: EnvVariable +}) => { + set(agentComposerEnvVariablesAtom, envVariables => updateEnvVariable( + envVariables, + starterVariable, + id, + variable => ({ ...variable, key }), + )) +}) + +export const setEnvVariableScopeAtom = atom(null, (_get, set, { + id, + scope, + starterVariable, +}: { + id: string + scope: EnvScope + starterVariable: EnvVariable +}) => { + set(agentComposerEnvVariablesAtom, envVariables => updateEnvVariable( + envVariables, + starterVariable, + id, + variable => ({ ...variable, scope }), + )) +}) + +export const setEnvVariableValueAtom = atom(null, (_get, set, { + id, + starterVariable, + value, +}: { + id: string + starterVariable: EnvVariable + value: string +}) => { + set(agentComposerEnvVariablesAtom, envVariables => updateEnvVariable( + envVariables, + starterVariable, + id, + variable => ({ ...variable, value }), + )) +}) + +export const addEnvVariableAtom = atom(null, (_get, set, { + starterVariable, + variable, +}: { + starterVariable: EnvVariable + variable: EnvVariable +}) => { + set(agentComposerEnvVariablesAtom, envVariables => [ + ...(envVariables.length > 0 ? envVariables : [starterVariable]), + variable, + ]) +}) + +export const importEnvVariablesAtom = atom(null, (_get, set, variables: EnvVariable[]) => { + set(agentComposerEnvVariablesAtom, envVariables => [...envVariables, ...variables]) +}) + +export const removeEnvVariableAtom = atom(null, (_get, set, id: string) => { + set(agentComposerEnvVariablesAtom, envVariables => envVariables.filter(variable => variable.id !== id)) +}) diff --git a/web/features/agent-v2/agent-composer/store-modules/files.ts b/web/features/agent-v2/agent-composer/store-modules/files.ts index e686116a8ff..1a931a4ac4b 100644 --- a/web/features/agent-v2/agent-composer/store-modules/files.ts +++ b/web/features/agent-v2/agent-composer/store-modules/files.ts @@ -22,3 +22,33 @@ export const agentComposerFilesAtom = atom files.flatMap((file) => { + if (file.id === fileId) + return [] + + if (file.children) + return [{ ...file, children: removeAgentFileNode(file.children, fileId) }] + + return [file] +}) + +export const upsertAgentFileAtom = atom(null, (_get, set, file: AgentFileNode) => { + set(agentComposerFilesAtom, files => [ + ...removeAgentFileNode(files, file.id), + file, + ]) +}) + +export const removeAgentFileAtom = atom(null, (_get, set, fileId: string) => { + set(agentComposerFilesAtom, files => removeAgentFileNode(files, fileId)) +}) + +export const clearAgentConfigNoteAtom = atom(null, (get, set) => { + const draft = get(agentComposerDraftAtom) + + set(agentComposerDraftAtom, { + ...draft, + configNote: '', + }) +}) diff --git a/web/features/agent-v2/agent-composer/store-modules/knowledge.ts b/web/features/agent-v2/agent-composer/store-modules/knowledge.ts index b33b4f6e951..53d25f74b78 100644 --- a/web/features/agent-v2/agent-composer/store-modules/knowledge.ts +++ b/web/features/agent-v2/agent-composer/store-modules/knowledge.ts @@ -22,3 +22,17 @@ export const agentComposerKnowledgeRetrievalsAtom = atom( }) }, ) + +export const addKnowledgeRetrievalAtom = atom(null, (_get, set, retrieval: AgentKnowledgeRetrievalItem) => { + set(agentComposerKnowledgeRetrievalsAtom, retrievals => [...retrievals, retrieval]) +}) + +export const updateKnowledgeRetrievalAtom = atom(null, (_get, set, retrieval: AgentKnowledgeRetrievalItem) => { + set(agentComposerKnowledgeRetrievalsAtom, retrievals => retrievals.map(currentRetrieval => ( + currentRetrieval.id === retrieval.id ? retrieval : currentRetrieval + ))) +}) + +export const removeKnowledgeRetrievalAtom = atom(null, (_get, set, retrievalId: string) => { + set(agentComposerKnowledgeRetrievalsAtom, retrievals => retrievals.filter(retrieval => retrieval.id !== retrievalId)) +}) diff --git a/web/features/agent-v2/agent-composer/store-modules/skills.ts b/web/features/agent-v2/agent-composer/store-modules/skills.ts index 2fa8bd47740..4fd342cf1d1 100644 --- a/web/features/agent-v2/agent-composer/store-modules/skills.ts +++ b/web/features/agent-v2/agent-composer/store-modules/skills.ts @@ -22,3 +22,14 @@ export const agentComposerSkillsAtom = atom { + set(agentComposerSkillsAtom, skills => [ + ...skills.filter(item => item.id !== skill.id), + skill, + ]) +}) + +export const removeAgentSkillAtom = atom(null, (_get, set, skillId: string) => { + set(agentComposerSkillsAtom, skills => skills.filter(item => item.id !== skillId)) +}) diff --git a/web/features/agent-v2/agent-composer/store-modules/tools.ts b/web/features/agent-v2/agent-composer/store-modules/tools.ts index 804f758b4a7..1f20efefdae 100644 --- a/web/features/agent-v2/agent-composer/store-modules/tools.ts +++ b/web/features/agent-v2/agent-composer/store-modules/tools.ts @@ -1,11 +1,17 @@ -import type { AgentProviderTool, AgentSoulConfigFormState, AgentTool } from '../form-state' +import type { AgentCliTool, AgentProviderTool, AgentSoulConfigFormState, AgentTool } from '../form-state' import type { DraftFieldUpdate } from './utils' -import { atom, useSetAtom } from 'jotai' -import { useCallback } from 'react' +import type { ToolDefaultValue } from '@/app/components/workflow/block-selector/types' +import { atom } from 'jotai' import { syncCliToolReferenceLabels } from '../reference-labels' import { agentComposerDraftAtom } from '../store' import { resolveDraftFieldUpdate } from './utils' +export type AgentProviderToolDefaultValue = ToolDefaultValue & { + allowDelete?: boolean + credentialType?: AgentProviderTool['credentialType'] + credentialRequired?: boolean +} + export const agentComposerToolsAtom = atom( get => get(agentComposerDraftAtom).tools, (get, set, toolsUpdate: DraftFieldUpdate) => { @@ -24,6 +30,104 @@ export const agentComposerToolsAtom = atom( }, ) +const toProviderToolAction = (tool: AgentProviderToolDefaultValue) => ({ + id: `${tool.provider_id}:${tool.tool_name}`, + name: tool.tool_label || tool.title || tool.tool_name, + toolName: tool.tool_name, + description: tool.tool_description || '', +}) + +const getCredentialVariant = (tool: AgentProviderToolDefaultValue) => { + if (!tool.credentialRequired) + return 'none' as const + + if (!tool.allowDelete) + return tool.credential_id ? 'authorized' as const : 'unauthorized' as const + + return tool.is_team_authorization ? 'authorized' as const : 'unauthorized' as const +} + +const getCredentialType = (tool: AgentProviderToolDefaultValue) => { + if (!tool.credentialRequired) + return undefined + + if (tool.credentialType === 'oauth2') + return 'oauth2' as const + + if (!tool.allowDelete) + return tool.credential_id ? 'api-key' as const : 'unauthorized' as const + + return tool.is_team_authorization ? 'api-key' as const : 'unauthorized' as const +} + +export const addProviderTools = ( + currentTools: AgentTool[], + selectedTools: AgentProviderToolDefaultValue[], +): AgentTool[] => { + if (selectedTools.length === 0) + return currentTools + + const nextTools = [...currentTools] + + selectedTools.forEach((selectedTool) => { + const action = toProviderToolAction(selectedTool) + const existingToolIndex = nextTools.findIndex(tool => tool.kind === 'provider' && tool.id === selectedTool.provider_id) + const existingTool = nextTools[existingToolIndex] + + if (existingTool?.kind === 'provider') { + if (existingTool.actions.some(existingAction => existingAction.toolName === action.toolName)) + return + + nextTools[existingToolIndex] = { + ...existingTool, + displayName: existingTool.displayName ?? selectedTool.provider_show_name, + icon: existingTool.icon ?? selectedTool.provider_icon, + iconDark: existingTool.iconDark ?? selectedTool.provider_icon_dark, + allowDelete: existingTool.allowDelete ?? selectedTool.allowDelete, + actions: [...existingTool.actions, action], + } + return + } + + nextTools.push({ + id: selectedTool.provider_id, + name: selectedTool.provider_name, + kind: 'provider', + displayName: selectedTool.provider_show_name, + iconClassName: 'i-custom-public-other-default-tool-icon text-text-tertiary', + icon: selectedTool.provider_icon, + iconDark: selectedTool.provider_icon_dark, + providerType: selectedTool.provider_type, + allowDelete: selectedTool.allowDelete, + credentialId: selectedTool.credential_id, + credentialKey: selectedTool.is_team_authorization + ? 'agentDetail.configure.tools.credential.authOne' + : undefined, + credentialType: getCredentialType(selectedTool), + credentialVariant: getCredentialVariant(selectedTool), + actions: [action], + }) + }) + + return nextTools +} + +export const addProviderToolsAtom = atom(null, (_get, set, selectedTools: AgentProviderToolDefaultValue[]) => { + set(agentComposerToolsAtom, tools => addProviderTools(tools, selectedTools)) +}) + +export const saveCliToolAtom = atom(null, (_get, set, cliTool: AgentCliTool) => { + set(agentComposerToolsAtom, tools => ( + tools.some(tool => tool.kind === 'cli' && tool.id === cliTool.id) + ? tools.map(tool => tool.id === cliTool.id ? cliTool : tool) + : [...tools, cliTool] + )) +}) + +export const removeCliToolAtom = atom(null, (_get, set, toolId: string) => { + set(agentComposerToolsAtom, tools => tools.filter(tool => tool.id !== toolId)) +}) + export const agentComposerToolSettingsAtom = atom( get => get(agentComposerDraftAtom).toolSettings, (get, set, toolSettingsUpdate: DraftFieldUpdate>>) => { @@ -49,63 +153,79 @@ const omitToolSettings = ( return nextToolSettings } -export function useRemoveProviderTool() { - const setDraft = useSetAtom(agentComposerDraftAtom) +export const removeProviderToolAtom = atom(null, (get, set, toolId: string) => { + const draft = get(agentComposerDraftAtom) + const toolToRemove = draft.tools.find(tool => tool.kind === 'provider' && tool.id === toolId) + const actionIds = toolToRemove?.kind === 'provider' + ? toolToRemove.actions.map(action => action.id) + : [] - return useCallback((toolId: string) => { - setDraft((draft) => { - const toolToRemove = draft.tools.find(tool => tool.kind === 'provider' && tool.id === toolId) - const actionIds = toolToRemove?.kind === 'provider' - ? toolToRemove.actions.map(action => action.id) - : [] + set(agentComposerDraftAtom, { + ...draft, + tools: draft.tools.filter(tool => tool.id !== toolId), + toolSettings: omitToolSettings(draft.toolSettings, actionIds), + }) +}) - return { - ...draft, - tools: draft.tools.filter(tool => tool.id !== toolId), - toolSettings: omitToolSettings(draft.toolSettings, actionIds), - } - }) - }, [setDraft]) -} +export const removeProviderToolActionAtom = atom(null, (get, set, { + toolId, + actionId, +}: { + toolId: string + actionId: string +}) => { + const draft = get(agentComposerDraftAtom) -export function useRemoveProviderToolAction() { - const setDraft = useSetAtom(agentComposerDraftAtom) - - return useCallback((toolId: string, actionId: string) => { - setDraft(draft => ({ - ...draft, - tools: draft.tools.flatMap((tool) => { - if (tool.kind !== 'provider' || tool.id !== toolId) - return [tool] - - const nextActions = tool.actions.filter(action => action.id !== actionId) - return nextActions.length > 0 - ? [{ ...tool, actions: nextActions }] - : [] - }), - toolSettings: omitToolSettings(draft.toolSettings, [actionId]), - })) - }, [setDraft]) -} - -export function useSetProviderToolCredential() { - const setTools = useSetAtom(agentComposerToolsAtom) - - return useCallback((toolId: string, credentialId?: string, credentialType?: AgentProviderTool['credentialType']) => { - setTools(tools => tools.map((tool) => { + set(agentComposerDraftAtom, { + ...draft, + tools: draft.tools.flatMap((tool) => { if (tool.kind !== 'provider' || tool.id !== toolId) - return tool + return [tool] - const nextCredentialType = credentialType === 'oauth2' || tool.credentialType === 'oauth2' - ? 'oauth2' - : 'api-key' + const nextActions = tool.actions.filter(action => action.id !== actionId) + return nextActions.length > 0 + ? [{ ...tool, actions: nextActions }] + : [] + }), + toolSettings: omitToolSettings(draft.toolSettings, [actionId]), + }) +}) - return { - ...tool, - credentialId, - credentialType: nextCredentialType, - credentialVariant: 'authorized', - } - })) - }, [setTools]) -} +export const setProviderToolCredentialAtom = atom(null, (_get, set, { + toolId, + credentialId, + credentialType, +}: { + toolId: string + credentialId?: string + credentialType?: AgentProviderTool['credentialType'] +}) => { + set(agentComposerToolsAtom, tools => tools.map((tool) => { + if (tool.kind !== 'provider' || tool.id !== toolId) + return tool + + const nextCredentialType = credentialType === 'oauth2' || tool.credentialType === 'oauth2' + ? 'oauth2' + : 'api-key' + + return { + ...tool, + credentialId, + credentialType: nextCredentialType, + credentialVariant: 'authorized', + } + })) +}) + +export const saveProviderToolActionSettingsAtom = atom(null, (_get, set, { + actionId, + value, +}: { + actionId: string + value: Record +}) => { + set(agentComposerToolSettingsAtom, toolSettings => ({ + ...toolSettings, + [actionId]: value, + })) +}) diff --git a/web/features/agent-v2/agent-detail/__tests__/layout.spec.tsx b/web/features/agent-v2/agent-detail/__tests__/layout.spec.tsx index 90d2b34166f..e1c5fe6258c 100644 --- a/web/features/agent-v2/agent-detail/__tests__/layout.spec.tsx +++ b/web/features/agent-v2/agent-detail/__tests__/layout.spec.tsx @@ -1,18 +1,33 @@ -import { render, screen } from '@testing-library/react' +import { render, screen, waitFor } from '@testing-library/react' import { AgentDetailLayout } from '../layout' +const mockReplace = vi.hoisted(() => vi.fn()) +const mockAgentQuery = vi.hoisted(() => ({ + data: { + name: 'Agent', + } as { name: string } | undefined, + error: null as unknown, +})) + vi.mock('@tanstack/react-query', async (importOriginal) => { const actual = await importOriginal() return { ...actual, - useQuery: vi.fn(() => ({ - data: { - name: 'Agent', - }, - })), + useQuery: vi.fn(() => mockAgentQuery), } }) +vi.mock('@/next/navigation', () => ({ + useRouter: () => ({ + back: vi.fn(), + forward: vi.fn(), + refresh: vi.fn(), + push: vi.fn(), + replace: mockReplace, + prefetch: vi.fn(), + }), +})) + vi.mock('@/hooks/use-document-title', () => ({ default: vi.fn(), })) @@ -32,6 +47,10 @@ vi.mock('@/service/client', () => ({ describe('AgentDetailLayout', () => { beforeEach(() => { vi.clearAllMocks() + mockAgentQuery.data = { + name: 'Agent', + } + mockAgentQuery.error = null }) it('should render detail content without owning navigation landmarks', () => { @@ -45,4 +64,20 @@ describe('AgentDetailLayout', () => { expect(screen.queryByRole('main')).not.toBeInTheDocument() expect(screen.queryByRole('complementary', { name: 'Detail sidebar' })).not.toBeInTheDocument() }) + + it('should redirect to roster when agent detail returns 404', async () => { + mockAgentQuery.data = undefined + mockAgentQuery.error = new Response(null, { status: 404 }) + + render( + +
Agent detail content
+
, + ) + + await waitFor(() => { + expect(mockReplace).toHaveBeenCalledWith('/roster') + }) + expect(screen.queryByText('Agent detail content')).not.toBeInTheDocument() + }) }) diff --git a/web/features/agent-v2/agent-detail/access/components/__tests__/workflow-references-table.spec.tsx b/web/features/agent-v2/agent-detail/access/components/__tests__/workflow-references-table.spec.tsx index e8f012397fb..eda95721bbf 100644 --- a/web/features/agent-v2/agent-detail/access/components/__tests__/workflow-references-table.spec.tsx +++ b/web/features/agent-v2/agent-detail/access/components/__tests__/workflow-references-table.spec.tsx @@ -5,9 +5,10 @@ import { WorkflowReferencesTable } from '../workflow-references-table' const mocks = vi.hoisted(() => ({ queryFn: vi.fn(), - queryOptions: vi.fn((input: unknown) => ({ + queryOptions: vi.fn(({ enabled = true, input }: { enabled?: boolean, input: unknown }) => ({ queryKey: ['agent-referencing-workflows', input], queryFn: () => mocks.queryFn(input), + enabled, })), })) @@ -31,7 +32,7 @@ vi.mock('@/hooks/use-timestamp', () => ({ }), })) -const renderTable = () => { +const renderTable = ({ enabled }: { enabled?: boolean } = {}) => { const queryClient = new QueryClient({ defaultOptions: { queries: { @@ -42,7 +43,7 @@ const renderTable = () => { render( - + , ) @@ -67,9 +68,27 @@ describe('WorkflowReferencesTable', () => { agent_id: 'agent-1', }, }, + enabled: true, }) }) }) + + it('should not fetch workflow references when disabled', async () => { + renderTable({ enabled: false }) + + await waitFor(() => { + expect(mocks.queryOptions).toHaveBeenCalledWith({ + input: { + params: { + agent_id: 'agent-1', + }, + }, + enabled: false, + }) + }) + expect(mocks.queryFn).not.toHaveBeenCalled() + expect(screen.queryByText('agentV2.agentDetail.access.workflow.loading')).not.toBeInTheDocument() + }) }) describe('Rendering', () => { diff --git a/web/features/agent-v2/agent-detail/access/components/workflow-references-table.tsx b/web/features/agent-v2/agent-detail/access/components/workflow-references-table.tsx index 41029edc93a..9cfa89a6167 100644 --- a/web/features/agent-v2/agent-detail/access/components/workflow-references-table.tsx +++ b/web/features/agent-v2/agent-detail/access/components/workflow-references-table.tsx @@ -12,6 +12,7 @@ import { consoleQuery } from '@/service/client' type WorkflowReferencesTableProps = { agentId: string + enabled?: boolean } const workflowTableColSpan = 5 @@ -20,6 +21,7 @@ const getWorkflowReferenceHref = (reference: AgentReferencingWorkflowResponse) = export function WorkflowReferencesTable({ agentId, + enabled = true, }: WorkflowReferencesTableProps) { const { t } = useTranslation('agentV2') const { t: tCommon } = useTranslation('common') @@ -29,6 +31,7 @@ export function WorkflowReferencesTable({ agent_id: agentId, }, }, + enabled, })) const workflowReferences = workflowReferencesQuery.data?.data ?? [] @@ -62,12 +65,12 @@ export function WorkflowReferencesTable({ - {workflowReferencesQuery.isPending && ( + {enabled && workflowReferencesQuery.isPending && ( {t('agentDetail.access.workflow.loading')} )} - {workflowReferencesQuery.isError && ( + {enabled && workflowReferencesQuery.isError && (
{t('agentDetail.access.workflow.loadFailed')} @@ -83,12 +86,12 @@ export function WorkflowReferencesTable({
)} - {workflowReferencesQuery.isSuccess && workflowReferences.length === 0 && ( + {enabled && workflowReferencesQuery.isSuccess && workflowReferences.length === 0 && ( {t('agentDetail.access.workflow.empty')} )} - {workflowReferencesQuery.isSuccess && workflowReferences.map(reference => ( + {enabled && workflowReferencesQuery.isSuccess && workflowReferences.map(reference => ( - + diff --git a/web/features/agent-v2/agent-detail/configure/__tests__/page.spec.tsx b/web/features/agent-v2/agent-detail/configure/__tests__/page.spec.tsx index cdce3ee690a..0dbe6f0354e 100644 --- a/web/features/agent-v2/agent-detail/configure/__tests__/page.spec.tsx +++ b/web/features/agent-v2/agent-detail/configure/__tests__/page.spec.tsx @@ -54,6 +54,19 @@ const mocks = vi.hoisted(() => ({ }, })) +const toastMock = vi.hoisted(() => ({ + error: vi.fn(), +})) + +const modelHooksState = vi.hoisted(() => ({ + defaultTextGenerationModel: { + provider: { + provider: 'langgenius/openai/openai', + }, + model: 'gpt-4o-mini', + } as { provider: { provider: string }, model: string } | undefined, +})) + function createDeferredPromise() { let resolve!: (value: T) => void const promise = new Promise((promiseResolve) => { @@ -101,6 +114,10 @@ vi.mock('@tanstack/react-query', async (importOriginal) => { } }) +vi.mock('@langgenius/dify-ui/toast', () => ({ + toast: toastMock, +})) + vi.mock('@/service/client', () => ({ consoleQuery: { agent: { @@ -194,7 +211,7 @@ vi.mock('@/service/client', () => ({ })) vi.mock('@/app/components/header/account-setting/model-provider-page/hooks', () => ({ - useDefaultModel: () => ({ data: undefined }), + useDefaultModel: () => ({ data: modelHooksState.defaultTextGenerationModel }), useTextGenerationCurrentProviderAndModelAndModelList: () => ({ textGenerationModelList: [], }), @@ -271,7 +288,7 @@ vi.mock('../components/preview/build-chat', async () => { void props.onSaveDraftBeforeRun?.().then(() => { setMessageSent(true) props.onConversationIdChange?.('build-conversation-new') - }) + }).catch(() => undefined) }} > send build message @@ -359,6 +376,12 @@ vi.mock('../components/preview/versions-panel', () => ({ describe('AgentConfigurePage', () => { beforeEach(() => { vi.clearAllMocks() + modelHooksState.defaultTextGenerationModel = { + provider: { + provider: 'langgenius/openai/openai', + }, + model: 'gpt-4o-mini', + } mocks.refreshDebugConversation.mockResolvedValue({ debug_conversation_has_messages: false, debug_conversation_id: 'debug-conversation-new', @@ -1036,6 +1059,50 @@ describe('AgentConfigurePage', () => { expect(screen.getByRole('button', { name: 'discard build draft' })).toBeDisabled() }) + it('should block build chat checkout when no model is configured', async () => { + const queryClient = new QueryClient() + modelHooksState.defaultTextGenerationModel = undefined + mocks.queryState.composer = { + data: { + agent_soul: { + prompt: { + system_prompt: 'draft prompt', + }, + }, + }, + isFetching: false, + isError: false, + isPending: false, + isSuccess: true, + refetch: vi.fn(), + } + mocks.queryState.buildDraft = { + data: undefined as unknown, + dataUpdatedAt: 0, + error: new Response(null, { status: 404 }), + isFetching: false, + isError: true, + isPending: false, + isSuccess: false, + refetch: vi.fn(), + } + + render( + + + , + ) + + fireEvent.click(screen.getByRole('button', { name: 'send build message' })) + + await waitFor(() => { + expect(toastMock.error).toHaveBeenCalledWith('common.modelProvider.selectModel') + }) + expect(mocks.checkoutBuildDraft).not.toHaveBeenCalled() + expect(screen.getByRole('region', { name: 'build-chat' })).toHaveTextContent('sent:no') + expect(screen.getByRole('region', { name: 'orchestrate-panel' })).toHaveTextContent('buildDraft:no') + }) + it('should keep the build draft bar disabled while a build conversation is responding', async () => { vi.useFakeTimers() const queryClient = new QueryClient() diff --git a/web/features/agent-v2/agent-detail/configure/__tests__/use-agent-configure-sync.spec.tsx b/web/features/agent-v2/agent-detail/configure/__tests__/use-agent-configure-sync.spec.tsx index ad402fb2e38..55148b316b3 100644 --- a/web/features/agent-v2/agent-detail/configure/__tests__/use-agent-configure-sync.spec.tsx +++ b/web/features/agent-v2/agent-detail/configure/__tests__/use-agent-configure-sync.spec.tsx @@ -91,6 +91,11 @@ function setDocumentVisibilityState(visibilityState: DocumentVisibilityState) { }) } +const configuredModel = { + provider: 'langgenius/openai/openai', + model: 'gpt-4o-mini', +} + vi.mock('@langgenius/dify-ui/toast', () => ({ toast: toastMock, })) @@ -607,7 +612,9 @@ describe('useAgentConfigureSync', () => { }) it('should publish only when publishDraft is called explicitly', async () => { - const { queryClient, result, store } = renderUseAgentConfigureSync() + const { queryClient, result, store } = renderUseAgentConfigureSync({ + currentModel: configuredModel, + }) const invalidateQueries = vi.spyOn(queryClient, 'invalidateQueries') queryClient.setQueryData(['agent-detail', 'agent-1'], { active_config_is_published: false, @@ -654,12 +661,28 @@ describe('useAgentConfigureSync', () => { expect(toastMock.success).toHaveBeenCalledWith('common.api.actionSuccess') }) + it('should toast and skip publish when no model is configured', async () => { + const { result, store } = renderUseAgentConfigureSync() + + act(() => { + store.set(agentComposerDraftAtom, { + ...defaultAgentSoulConfigFormState, + prompt: 'Published prompt', + }) + }) + + await act(async () => { + await result.current.publishDraft() + }) + + expect(composerPutMutationFn).not.toHaveBeenCalled() + expect(publishAgentMutationFn).not.toHaveBeenCalled() + expect(toastMock.error).toHaveBeenCalledWith('common.modelProvider.selectModel') + }) + it('should keep default model fallback from creating unpublished changes after publish', async () => { const { result, store } = renderUseAgentConfigureSync({ - currentModel: { - provider: 'langgenius/openai/openai', - model: 'gpt-4o-mini', - }, + currentModel: configuredModel, }) act(() => { store.set(agentComposerDraftAtom, { @@ -681,6 +704,7 @@ describe('useAgentConfigureSync', () => { it('should keep base config fallback fields from creating unpublished changes after publish', async () => { const { result, store } = renderUseAgentConfigureSync({ + currentModel: configuredModel, baseConfig: { app_features: { file_upload: { @@ -708,7 +732,9 @@ describe('useAgentConfigureSync', () => { }) it('should publish the current draft snapshot instead of a stale caller payload', async () => { - const { result, store } = renderUseAgentConfigureSync() + const { result, store } = renderUseAgentConfigureSync({ + currentModel: configuredModel, + }) act(() => { store.set(agentComposerDraftAtom, { @@ -736,7 +762,9 @@ describe('useAgentConfigureSync', () => { it('should reject publish and keep the publish mutation untouched when saving the draft fails', async () => { composerPutMutationFn.mockRejectedValueOnce(new Error('save failed')) - const { queryClient, result, store } = renderUseAgentConfigureSync() + const { queryClient, result, store } = renderUseAgentConfigureSync({ + currentModel: configuredModel, + }) queryClient.setQueryData(['agent-detail', 'agent-1'], { active_config_is_published: false, name: 'Agent', @@ -760,7 +788,9 @@ describe('useAgentConfigureSync', () => { }) it('should toast and skip publish when knowledge retrieval validation fails', async () => { - const { result, store } = renderUseAgentConfigureSync() + const { result, store } = renderUseAgentConfigureSync({ + currentModel: configuredModel, + }) act(() => { store.set(agentComposerDraftAtom, { @@ -785,7 +815,9 @@ describe('useAgentConfigureSync', () => { }) it('should toast metadata filtering model error when publishing with automatic metadata filtering and no model', async () => { - const { result, store } = renderUseAgentConfigureSync() + const { result, store } = renderUseAgentConfigureSync({ + currentModel: configuredModel, + }) act(() => { store.set(agentComposerDraftAtom, { @@ -813,7 +845,9 @@ describe('useAgentConfigureSync', () => { it('should expose publishing status from the publish mutation while publish is pending', async () => { const publishDeferred = createDeferredPromise() publishAgentMutationFn.mockReturnValueOnce(publishDeferred.promise) - const { result } = renderUseAgentConfigureSync() + const { result } = renderUseAgentConfigureSync({ + currentModel: configuredModel, + }) let publishPromise!: Promise act(() => { publishPromise = result.current.publishDraft() diff --git a/web/features/agent-v2/agent-detail/configure/components/__tests__/agent-prompt-editor.spec.tsx b/web/features/agent-v2/agent-detail/configure/components/__tests__/agent-prompt-editor.spec.tsx index e2165c288c5..5447ab28e4e 100644 --- a/web/features/agent-v2/agent-detail/configure/components/__tests__/agent-prompt-editor.spec.tsx +++ b/web/features/agent-v2/agent-detail/configure/components/__tests__/agent-prompt-editor.spec.tsx @@ -16,6 +16,42 @@ const mockPromptEditor = vi.hoisted(() => vi.fn()) const mockCopy = vi.hoisted(() => vi.fn()) const mockReset = vi.hoisted(() => vi.fn()) const mockUseClipboard = vi.hoisted(() => vi.fn()) +const mockConfigFiles = vi.hoisted(() => ({ + current: [] as Array<{ + id: string + name: string + driveKey?: string + children?: Array<{ + id: string + name: string + driveKey?: string + }> + }>, +})) +const mockLexical = vi.hoisted(() => ({ + selection: null as null | { + __range: true + isCollapsed: () => boolean + anchor: { + getNode: () => { + __text: true + getKey: () => string + getTextContent: () => string + getTextContentSize: () => number + select: (anchorOffset: number, focusOffset: number) => void + } + offset: number + } + }, + rootChildren: [] as Array<{ + __text: true + getKey: () => string + getTextContent: () => string + getTextContentSize: () => number + select: (anchorOffset: number, focusOffset: number) => void + }>, + rootSelectEnd: vi.fn(), +})) const mockBuiltInTools = vi.hoisted(() => [ { id: 'duckduckgo', @@ -62,11 +98,37 @@ vi.mock('@/app/components/base/prompt-editor', () => ({ return (
+ {props.children}
) }, })) +vi.mock('@lexical/react/LexicalComposerContext', () => ({ + useLexicalComposerContext: () => [{ + focus: (callback: () => void) => callback(), + getEditorState: () => ({ + read: (callback: () => void) => callback(), + }), + registerCommand: () => vi.fn(), + registerUpdateListener: () => vi.fn(), + update: (callback: () => void) => callback(), + }], +})) + +vi.mock('lexical', () => ({ + $getRoot: () => ({ + getChildren: () => mockLexical.rootChildren, + selectEnd: mockLexical.rootSelectEnd, + }), + $getSelection: () => mockLexical.selection, + $isElementNode: (node: { __element?: boolean } | null | undefined) => !!node?.__element, + $isRangeSelection: (selection: { __range?: boolean } | null | undefined) => !!selection?.__range, + $isTextNode: (node: { __text?: boolean } | null | undefined) => !!node?.__text, + COMMAND_PRIORITY_LOW: 1, + SELECTION_CHANGE_COMMAND: Symbol('selection-change-command'), +})) + vi.mock('@/app/components/base/infotip', () => ({ Infotip: ({ children }: { children: ReactNode }) => {children}, })) @@ -103,10 +165,11 @@ vi.mock('../orchestrate/config-context', () => ({ { id: 'playwright', name: 'Playwright', + skillMdKey: 'skills/playwright/SKILL.md', }, ], }), - useAgentConfigFiles: () => ({ files: [] }), + useAgentConfigFiles: () => ({ files: mockConfigFiles.current }), })) const duckDuckGoSearchAction = { @@ -167,6 +230,10 @@ const renderAgentPromptEditor = ( describe('AgentPromptEditor', () => { beforeEach(() => { vi.clearAllMocks() + mockConfigFiles.current = [] + mockLexical.selection = null + mockLexical.rootChildren = [] + mockLexical.rootSelectEnd.mockClear() mockUseClipboard.mockReturnValue({ copied: false, copy: mockCopy, @@ -290,6 +357,40 @@ describe('AgentPromptEditor', () => { expect(container.querySelector('.i-ri-terminal-box-line')).not.toBeInTheDocument() }) + + it('should warn only for prompt references missing from the current configuration', () => { + mockConfigFiles.current = [{ + id: 'folder', + name: 'Folder', + children: [{ + id: 'file-1', + name: 'Spec.md', + driveKey: 'drive/spec.md', + }], + }] + renderAgentPromptEditor('Review these tenders', { + knowledgeRetrievals: [{ id: 'retrieval-1', name: 'Release Notes' }], + tools: [ + duckDuckGoProviderTool, + { id: 'cli-1', kind: 'cli', name: 'Lark CLI' }, + ], + }) + + const promptEditorProps = mockPromptEditor.mock.calls.at(-1)?.[0] as PromptEditorProps + const getWarning = promptEditorProps.rosterReferenceBlock?.getWarning + expect(getWarning).toBeDefined() + + expect(getWarning?.({ kind: 'skill', id: 'skills%2Fplaywright%2FSKILL.md', label: 'Playwright' })).toBeUndefined() + expect(getWarning?.({ kind: 'file', id: 'drive%2Fspec.md', label: 'Spec.md' })).toBeUndefined() + expect(getWarning?.({ kind: 'knowledge', id: 'retrieval-1', label: 'Release Notes' })).toBeUndefined() + expect(getWarning?.({ kind: 'tool', id: 'duckduckgo/ddg_search', label: 'DuckDuckGo Search' })).toBeUndefined() + expect(getWarning?.({ kind: 'tool-all', id: 'duckduckgo/*', label: 'DuckDuckGo' })).toBeUndefined() + + expect(getWarning?.({ kind: 'skill', id: 'missing-skill', label: 'Missing Skill' })).toContain('agentDetail.configure.prompt.referenceMissing') + expect(getWarning?.({ kind: 'file', id: 'missing-file', label: 'Missing File' })).toContain('agentDetail.configure.prompt.referenceMissing') + expect(getWarning?.({ kind: 'knowledge', id: 'missing-retrieval', label: 'Missing Retrieval' })).toContain('agentDetail.configure.prompt.referenceMissing') + expect(getWarning?.({ kind: 'tool', id: 'missing/action', label: 'Missing Tool' })).toContain('agentDetail.configure.prompt.referenceMissing') + }) }) // Prompt slash commands should use the Agent Roster category menu and replace it with submenus. @@ -320,6 +421,35 @@ describe('AgentPromptEditor', () => { }) }) + it('should replace the slash at the current lexical selection instead of appending', async () => { + const textNode = { + __text: true as const, + getKey: () => 'text-node', + getTextContent: () => 'Review / now', + getTextContentSize: () => 'Review / now'.length, + select: vi.fn(), + } + mockLexical.rootChildren = [textNode] + mockLexical.selection = { + __range: true, + isCollapsed: () => true, + anchor: { + getNode: () => textNode, + offset: 'Review /'.length, + }, + } + const { store } = renderAgentPromptEditor('Review / now') + + fireEvent.keyDown(screen.getByRole('textbox'), { key: '/' }) + fireEvent.click(screen.getByRole('button', { name: /agentDetail\.configure\.skills\.label/i })) + fireEvent.click(screen.getByRole('button', { name: /Playwright/i })) + + expect(store.get(agentComposerPromptAtom)).toBe('Review [§skill:playwright:Playwright§] now') + await waitFor(() => { + expect(mockLexical.rootSelectEnd).toHaveBeenCalled() + }) + }) + it('should insert slash from the focused footer insert action', () => { const { store } = renderAgentPromptEditor('Review these tenders') @@ -350,7 +480,7 @@ describe('AgentPromptEditor', () => { skills={[]} files={[]} tools={[]} - onToolsChange={vi.fn()} + onAddProviderTools={vi.fn()} onAddSkill={options => options?.onAdded?.({ id: 'skill-1', name: 'Skill One' })} retrievals={[]} onBack={vi.fn()} @@ -368,7 +498,7 @@ describe('AgentPromptEditor', () => { skills={[]} files={[]} tools={[]} - onToolsChange={vi.fn()} + onAddProviderTools={vi.fn()} onAddFile={options => options?.onAdded?.({ id: 'file-1', name: 'Guide.md', icon: 'markdown', configName: 'Guide.md' })} retrievals={[]} onBack={vi.fn()} @@ -386,7 +516,7 @@ describe('AgentPromptEditor', () => { skills={[]} files={[]} tools={[]} - onToolsChange={vi.fn()} + onAddProviderTools={vi.fn()} onAddKnowledge={options => options?.onAdded?.({ id: 'retrieval-1', name: 'Retrieval One', queryMode: 'agent' })} retrievals={[]} onBack={vi.fn()} @@ -404,7 +534,7 @@ describe('AgentPromptEditor', () => { skills={[]} files={[]} tools={[]} - onToolsChange={vi.fn()} + onAddProviderTools={vi.fn()} onAddCliTool={options => options?.onAdded?.({ id: 'cli-1', kind: 'cli', name: 'Lark CLI' })} retrievals={[]} onBack={vi.fn()} diff --git a/web/features/agent-v2/agent-detail/configure/components/composer-session.tsx b/web/features/agent-v2/agent-detail/configure/components/composer-session.tsx index ac5e8b886ea..9ddf1986d7c 100644 --- a/web/features/agent-v2/agent-detail/configure/components/composer-session.tsx +++ b/web/features/agent-v2/agent-detail/configure/components/composer-session.tsx @@ -2,6 +2,7 @@ import type { AgentAppDetailWithSite, AgentIconType, AgentSoulConfig } from '@dify/contracts/api/console/agent/types.gen' import type { useAgentConfigureData } from '../hooks' +import { toast } from '@langgenius/dify-ui/toast' import { useMutation, useQueryClient } from '@tanstack/react-query' import { useAtomValue, useSetAtom } from 'jotai' import { ScopeProvider } from 'jotai-scope' @@ -214,6 +215,7 @@ function AgentConfigurePageComposerContent({ activeConfigSnapshot, agentSoulConfig, } = configureData + const { t: tCommon } = useTranslation('common') const [buildDraftActionsDisabled, setBuildDraftActionsDisabled] = useState(false) const [clearPreviewChat, setClearPreviewChat] = useState(false) const [completedBuildConversationId, setCompletedBuildConversationId] = useState(null) @@ -333,6 +335,7 @@ function AgentConfigurePageComposerContent({ isBuildDraftActive={buildDraft.isActive} buildDraftChangedKeys={buildDraft.changedKeys} showPublishBar={!buildDraft.isActive} + workflowReferencesEnabled={agentQuery.isSuccess} bottomAction={showBuildDraftBar ? ( { + if (!currentModel?.provider || !currentModel.model) { + toast.error(tCommon('modelProvider.selectModel')) + throw new Error('Agent model is required.') + } + setBuildDraftActionsDisabled(true) try { return await buildDraftActions.prepareBuildDraftBeforeRun() diff --git a/web/features/agent-v2/agent-detail/configure/components/orchestrate/__tests__/publish-bar.spec.tsx b/web/features/agent-v2/agent-detail/configure/components/orchestrate/__tests__/publish-bar.spec.tsx index a2400bc8e72..b5782ad795f 100644 --- a/web/features/agent-v2/agent-detail/configure/components/orchestrate/__tests__/publish-bar.spec.tsx +++ b/web/features/agent-v2/agent-detail/configure/components/orchestrate/__tests__/publish-bar.spec.tsx @@ -74,8 +74,9 @@ vi.mock('@/service/client', () => ({ }, referencingWorkflows: { get: { - queryOptions: ({ input }: { input: { params: { agent_id: string } } }) => ({ + queryOptions: ({ enabled = true, input }: { enabled?: boolean, input: { params: { agent_id: string } } }) => ({ queryKey: ['agent-referencing-workflows', input], + enabled, queryFn: async () => ({ data: (workflowReferences.fetchCount++, workflowReferences.data), }), @@ -162,6 +163,7 @@ function renderPublishBar({ selectedVersionSnapshot, setupStore, usedByAppReferences = [], + workflowReferencesEnabled, }: { activeConfigIsPublished?: boolean activeConfigSnapshot?: AgentConfigSnapshotSummaryResponse | null @@ -173,6 +175,7 @@ function renderPublishBar({ selectedVersionSnapshot?: AgentConfigSnapshotSummaryResponse | null setupStore?: (store: ReturnType) => void usedByAppReferences?: AgentReferencingWorkflowResponse[] + workflowReferencesEnabled?: boolean } = {}) { workflowReferences.data = usedByAppReferences const queryClient = new QueryClient({ @@ -200,6 +203,7 @@ function renderPublishBar({ agentName="Iris" isPublishing={nextProps?.isPublishing ?? isPublishing} selectedVersionSnapshot={selectedVersionSnapshot} + workflowReferencesEnabled={workflowReferencesEnabled} onPublish={onPublish} onExitVersions={onExitVersions} onOpenVersions={vi.fn()} @@ -407,6 +411,28 @@ describe('AgentConfigurePublishBar', () => { }) }) + it('should publish without loading workflow references when references are disabled', async () => { + const { onPublish } = renderPublishBar({ + activeConfigSnapshot, + prompt: 'Updated system prompt', + usedByAppReferences: publishedReferences, + workflowReferencesEnabled: false, + }) + + await waitFor(() => { + expect(workflowReferences.fetchCount).toBe(0) + }) + fireEvent.click(screen.getByRole('button', { name: /agentV2\.agentDetail\.configure\.publishBar\.publishUpdate/ })) + + await waitFor(() => { + expect(onPublish).toHaveBeenCalledTimes(1) + }) + expect(workflowReferences.fetchCount).toBe(0) + expect(screen.queryByRole('region', { + name: /agentV2\.agentDetail\.configure\.publishImpact\.title/, + })).not.toBeInTheDocument() + }) + it('should mark non-prompt draft changes as unpublished', () => { renderPublishBar({ activeConfigSnapshot, diff --git a/web/features/agent-v2/agent-detail/configure/components/orchestrate/advanced/env.tsx b/web/features/agent-v2/agent-detail/configure/components/orchestrate/advanced/env.tsx index 555c8e5185d..ddd7da61d09 100644 --- a/web/features/agent-v2/agent-detail/configure/components/orchestrate/advanced/env.tsx +++ b/web/features/agent-v2/agent-detail/configure/components/orchestrate/advanced/env.tsx @@ -7,10 +7,18 @@ import { Input } from '@langgenius/dify-ui/input' import { Select, SelectContent, SelectItem, SelectItemIndicator, SelectItemText, SelectTrigger } from '@langgenius/dify-ui/select' import { toast } from '@langgenius/dify-ui/toast' import { Tooltip, TooltipContent, TooltipTrigger } from '@langgenius/dify-ui/tooltip' -import { useAtom } from 'jotai' +import { useAtomValue, useSetAtom } from 'jotai' import { useEffect, useRef, useState } from 'react' import { useTranslation } from 'react-i18next' -import { agentComposerEnvVariablesAtom } from '@/features/agent-v2/agent-composer/store-modules/env' +import { + addEnvVariableAtom, + agentComposerEnvVariablesAtom, + importEnvVariablesAtom, + removeEnvVariableAtom, + setEnvVariableKeyAtom, + setEnvVariableScopeAtom, + setEnvVariableValueAtom, +} from '@/features/agent-v2/agent-composer/store-modules/env' import { checkKeys } from '@/utils/var' import { ConfigureSection } from '../common/section' import { AgentConfigureTipContent } from '../common/tip-content' @@ -410,7 +418,13 @@ export function EnvVariablesTable({ export function AgentEnvEditor() { const { t } = useTranslation('agentV2') const readOnly = useAgentOrchestrateReadOnly() - const [envVariables, setEnvVariables] = useAtom(agentComposerEnvVariablesAtom) + const envVariables = useAtomValue(agentComposerEnvVariablesAtom) + const addEnvVariable = useSetAtom(addEnvVariableAtom) + const importEnvVariables = useSetAtom(importEnvVariablesAtom) + const removeEnvVariable = useSetAtom(removeEnvVariableAtom) + const setEnvVariableKey = useSetAtom(setEnvVariableKeyAtom) + const setEnvVariableScope = useSetAtom(setEnvVariableScopeAtom) + const setEnvVariableValue = useSetAtom(setEnvVariableValueAtom) const starterVariableRef = useRef(undefined) if (!starterVariableRef.current) starterVariableRef.current = createEnvVariable() @@ -422,20 +436,6 @@ export function AgentEnvEditor() { const envEditorTableId = 'agent-configure-env-editor-table' const visibleEnvVariables = envVariables.length > 0 ? envVariables : [starterVariable] - const updateVariable = (id: string, updater: (variable: EnvVariable) => EnvVariable) => { - const existingVariable = envVariables.find(variable => variable.id === id) - - if (existingVariable) { - setEnvVariables(envVariables.map(variable => ( - variable.id === id ? updater(variable) : variable - ))) - return - } - - if (id === starterVariable.id) - setEnvVariables([updater(starterVariable)]) - } - const addVariable = ({ focusField = 'key', scope, @@ -448,13 +448,13 @@ export function AgentEnvEditor() { ...(scope ? { scope } : {}), } - setEnvVariables([ - ...(envVariables.length > 0 ? envVariables : [starterVariable]), + addEnvVariable({ + starterVariable, variable, - ]) + }) setFocusedVariable({ id: variable.id, field: focusField }) } - const importEnvVariables = async (file: File) => { + const handleImportEnvVariables = async (file: File) => { const { invalidLineCount, variables, @@ -470,19 +470,19 @@ export function AgentEnvEditor() { if (importedVariables.length === 0) return - setEnvVariables([...envVariables, ...importedVariables]) + importEnvVariables(importedVariables) } const updateVariableKey = (id: string, key: string) => { - updateVariable(id, variable => ({ ...variable, key })) + setEnvVariableKey({ id, key, starterVariable }) } const updateVariableScope = (id: string, scope: EnvScope) => { - updateVariable(id, variable => ({ ...variable, scope })) + setEnvVariableScope({ id, scope, starterVariable }) } const updateVariableValue = (id: string, value: string) => { - updateVariable(id, variable => ({ ...variable, value })) + setEnvVariableValue({ id, starterVariable, value }) } const deleteVariable = (id: string) => { - setEnvVariables(envVariables.filter(variable => variable.id !== id)) + removeEnvVariable(id) } return ( @@ -508,7 +508,7 @@ export function AgentEnvEditor() { event.target.value = '' if (file) - void importEnvVariables(file) + void handleImportEnvVariables(file) }} /> diff --git a/web/features/agent-v2/agent-detail/configure/components/orchestrate/files/__tests__/index.spec.tsx b/web/features/agent-v2/agent-detail/configure/components/orchestrate/files/__tests__/index.spec.tsx index 028bd17f304..3919244a6bf 100644 --- a/web/features/agent-v2/agent-detail/configure/components/orchestrate/files/__tests__/index.spec.tsx +++ b/web/features/agent-v2/agent-detail/configure/components/orchestrate/files/__tests__/index.spec.tsx @@ -3,7 +3,7 @@ import type { AgentConfigApiContext } from '../../config-context' import type { AgentSoulConfigFormState } from '@/features/agent-v2/agent-composer/form-state' import { toast } from '@langgenius/dify-ui/toast' import { QueryClient, QueryClientProvider } from '@tanstack/react-query' -import { fireEvent, render, screen, waitFor } from '@testing-library/react' +import { fireEvent, render, screen, waitFor, within } from '@testing-library/react' import userEvent from '@testing-library/user-event' import { useAtomValue } from 'jotai' import { beforeEach, describe, expect, it, vi } from 'vitest' @@ -39,6 +39,8 @@ const mocks = vi.hoisted(() => ({ deleteFileMutationFn: vi.fn(async (_input: unknown) => ({ removed_names: ['brief.md'], result: 'success' })), previewQueryOptions: vi.fn((_options: ConfigFileQueryOptionsInput) => ({})), downloadQueryOptions: vi.fn((_options: ConfigFileQueryOptionsInput) => ({})), + downloadBlob: vi.fn(), + downloadUrl: vi.fn(), })) vi.mock('@langgenius/dify-ui/toast', () => ({ @@ -48,6 +50,11 @@ vi.mock('@langgenius/dify-ui/toast', () => ({ }, })) +vi.mock('@/utils/download', () => ({ + downloadBlob: mocks.downloadBlob, + downloadUrl: mocks.downloadUrl, +})) + vi.mock('@/service/client', () => ({ consoleQuery: { agent: { @@ -368,6 +375,57 @@ describe('AgentFiles', () => { }) }) + it('should download configured files from the row action by config name', async () => { + const user = userEvent.setup() + renderAgentFiles() + + await user.click(screen.getByRole('button', { + name: /agentV2\.agentDetail\.configure\.files\.download.*diagram\.png/, + })) + + await waitFor(() => { + expect(mocks.downloadQueryOptions).toHaveBeenCalledWith(expect.objectContaining({ + input: expect.objectContaining({ + params: { + agent_id: 'agent-1', + name: 'diagram.png', + }, + }), + })) + }) + expect(mocks.downloadUrl).toHaveBeenCalledWith({ + url: 'https://example.com/diagram.png', + fileName: 'diagram.png', + }) + }) + + it('should download the selected file from the preview header action', async () => { + const user = userEvent.setup() + renderAgentFiles() + + await user.click(screen.getByText('diagram.png').closest('button')!) + const dialog = await screen.findByRole('dialog') + + await user.click(within(dialog).getByRole('button', { + name: /common\.operation\.download.*diagram\.png/, + })) + + await waitFor(() => { + expect(mocks.downloadQueryOptions).toHaveBeenCalledWith(expect.objectContaining({ + input: expect.objectContaining({ + params: { + agent_id: 'agent-1', + name: 'diagram.png', + }, + }), + })) + }) + expect(mocks.downloadUrl).toHaveBeenCalledWith({ + url: 'https://example.com/diagram.png', + fileName: 'diagram.png', + }) + }) + it('should show config note as a virtual build note file and preview its content locally', async () => { const user = userEvent.setup() renderAgentFiles({ @@ -399,15 +457,65 @@ describe('AgentFiles', () => { })) }) + it('should download the virtual build note file as markdown content', async () => { + const user = userEvent.setup() + renderAgentFiles({ + initialDraft: createInitialDraft({ configNote: 'Build context from the latest build chat.' }), + }) + + await user.click(screen.getByRole('button', { + name: /agentV2\.agentDetail\.configure\.files\.download.*build_note\.md/, + })) + + expect(mocks.downloadBlob).toHaveBeenCalledWith({ + data: expect.any(Blob), + fileName: 'build_note.md', + }) + const blob = mocks.downloadBlob.mock.calls[0]?.[0].data as Blob + await expect(blob.text()).resolves.toBe('Build context from the latest build chat.') + expect(mocks.downloadQueryOptions).not.toHaveBeenCalledWith(expect.objectContaining({ + input: expect.objectContaining({ + params: expect.objectContaining({ + name: 'build_note.md', + }), + }), + })) + }) + + it('should download the virtual build note from the preview header action', async () => { + const user = userEvent.setup() + renderAgentFiles({ + initialDraft: createInitialDraft({ configNote: 'Build context from the latest build chat.' }), + }) + + await user.click(screen.getByText('build_note.md').closest('button')!) + const dialog = await screen.findByRole('dialog') + + await user.click(within(dialog).getByRole('button', { + name: /common\.operation\.download.*build_note\.md/, + })) + + expect(mocks.downloadBlob).toHaveBeenCalledWith({ + data: expect.any(Blob), + fileName: 'build_note.md', + }) + const blob = mocks.downloadBlob.mock.calls[0]?.[0].data as Blob + await expect(blob.text()).resolves.toBe('Build context from the latest build chat.') + }) + it('should show generated build note metadata with an explanatory infotip', async () => { const user = userEvent.setup() renderAgentFiles({ initialDraft: createInitialDraft({ configNote: 'Build context from the latest build chat.' }), }) - expect(screen.getByText('agentV2.agentDetail.configure.files.buildNote.generated')).toBeInTheDocument() + const generatedBadge = screen.getByText('agentV2.agentDetail.configure.files.buildNote.generated') + const buildNoteRow = generatedBadge.closest('li') - await user.click(screen.getByRole('button', { name: 'agentV2.agentDetail.configure.files.buildNote.tooltip' })) + expect(generatedBadge).toBeInTheDocument() + expect(buildNoteRow).not.toBeNull() + + await user.click(within(buildNoteRow!).getByRole('button', { name: 'agentV2.agentDetail.configure.files.buildNote.tooltip' })) expect(await screen.findByText('agentDetail.configure.files.buildNote.richTooltip')).toBeInTheDocument() }) diff --git a/web/features/agent-v2/agent-detail/configure/components/orchestrate/files/index.tsx b/web/features/agent-v2/agent-detail/configure/components/orchestrate/files/index.tsx index da9fa0bbc35..1c13730d78b 100644 --- a/web/features/agent-v2/agent-detail/configure/components/orchestrate/files/index.tsx +++ b/web/features/agent-v2/agent-detail/configure/components/orchestrate/files/index.tsx @@ -1,10 +1,9 @@ 'use client' -import type { ReactNode } from 'react' +import type { MouseEvent, ReactNode } from 'react' import type { AgentOrchestrateAddActionOptions } from '../add-actions-context' import type { AgentConfigApiContext } from '../config-context' import type { AgentFileNode } from '@/features/agent-v2/agent-composer/form-state' -import { cn } from '@langgenius/dify-ui/cn' import { Dialog, DialogTrigger, @@ -15,15 +14,21 @@ import { FileTreeIcon, FileTreeLabel, } from '@langgenius/dify-ui/file-tree' -import { useMutation, useQuery } from '@tanstack/react-query' +import { useMutation, useQuery, useQueryClient } from '@tanstack/react-query' import { useAtomValue, useSetAtom } from 'jotai' import { useCallback, useRef, useState } from 'react' import { Trans, useTranslation } from 'react-i18next' import { Infotip } from '@/app/components/base/infotip' import { useDocLink } from '@/context/i18n' import { agentComposerDraftAtom } from '@/features/agent-v2/agent-composer/store' -import { agentComposerFilesAtom } from '@/features/agent-v2/agent-composer/store-modules/files' +import { + agentComposerFilesAtom, + clearAgentConfigNoteAtom, + removeAgentFileAtom, + upsertAgentFileAtom, +} from '@/features/agent-v2/agent-composer/store-modules/files' import { consoleQuery } from '@/service/client' +import { downloadBlob, downloadUrl } from '@/utils/download' import { useRegisterAgentOrchestrateAddAction } from '../add-actions-context' import { ConfigureSectionAddButton } from '../common/add-button' import { DocsLink } from '../common/docs-link' @@ -64,16 +69,6 @@ const findAgentFileNode = (files: AgentFileNode[], fileId: string): AgentFileNod } } -const removeAgentFileNode = (files: AgentFileNode[], fileId: string): AgentFileNode[] => files.flatMap((file) => { - if (file.id === fileId) - return [] - - if (file.children) - return [{ ...file, children: removeAgentFileNode(file.children, fileId) }] - - return [file] -}) - function AgentFileItem({ children, depth, @@ -93,6 +88,7 @@ function AgentFileItem({ }) { const { t } = useTranslation('agentV2') const readOnly = useAgentOrchestrateReadOnly() + const queryClient = useQueryClient() const [isPreviewOpen, setIsPreviewOpen] = useState(false) const [selectedFileId, setSelectedFileId] = useState() const selectedFile = selectedFileId ? findAgentFileNode(files, selectedFileId) : undefined @@ -169,25 +165,71 @@ function AgentFileItem({ const handleRemove = useCallback(() => { onRemove(file.id) }, [file.id, onRemove]) + const downloadFile = useCallback(async (targetFile: AgentFileNode) => { + if (targetFile.virtualContent !== undefined) { + downloadBlob({ + data: new Blob([targetFile.virtualContent], { type: 'text/markdown;charset=utf-8' }), + fileName: targetFile.name, + }) + return + } + + const fileName = getAgentFilePreviewKey(targetFile) + if (apiContext.workflow) { + const result = await queryClient.fetchQuery(consoleQuery.apps.byAppId.agent.config.files.byName.download.get.queryOptions({ + input: { + params: { + app_id: apiContext.workflow.appId, + name: fileName, + }, + query: { + node_id: apiContext.workflow.nodeId, + draft_type: apiContext.draftType, + version_id: apiContext.versionId, + }, + }, + })) + downloadUrl({ url: result.url, fileName: targetFile.name }) + return + } + + const result = await queryClient.fetchQuery(consoleQuery.agent.byAgentId.config.files.byName.download.get.queryOptions({ + input: { + params: { + agent_id: apiContext.agentId, + name: fileName, + }, + query: { + draft_type: apiContext.draftType, + version_id: apiContext.versionId, + }, + }, + })) + downloadUrl({ url: result.url, fileName: targetFile.name }) + }, [apiContext, queryClient]) + const handleDownload = useCallback(async (event: MouseEvent) => { + event.stopPropagation() + await downloadFile(file) + }, [downloadFile, file]) const handlePreviewOpenChange = useCallback((open: boolean) => { if (open) setSelectedFileId(file.id) setIsPreviewOpen(open) }, [file.id]) + const canRemoveFile = !readOnly && (!file.virtualContent || isBuildNoteFile) return ( -
  • +
  • )} > @@ -214,62 +256,73 @@ function AgentFileItem({ isImage: isImagePreviewFile, isLoading: !isVirtualPreviewFile && previewQuery.isPending, }, + onDownloadFile: () => downloadFile(selectedPreviewFile), onSelectFile: selectedFile => setSelectedFileId(selectedFile.id), selectedFileId: selectedFileId ?? file.id, sections: [], }} /> - {isBuildNoteFile && ( - - )} - {!readOnly && (!file.virtualContent || isBuildNoteFile) && ( +
    - )} + {canRemoveFile && ( + + )} +
  • ) } function AgentBuildNoteFileRow() { - const { t } = useTranslation('agentV2') - return ( <> - + {BUILD_NOTE_FILE_NAME} - - - {t('agentDetail.configure.files.buildNote.generated')} - +
    + + +
    ) } -function AgentBuildNoteInfotip({ - className, -}: { - className?: string -}) { +function AgentBuildNoteBadge() { + const { t } = useTranslation('agentV2') + + return ( + + + {t('agentDetail.configure.files.buildNote.generated')} + + ) +} + +function AgentBuildNoteInfotip() { const { t } = useTranslation('agentV2') const docLink = useDocLink() return (

    @@ -293,19 +346,17 @@ export function AgentFiles() { const promptAddCallbackRef = useRef(undefined) const apiContext = useAgentConfigApiContext() const draft = useAtomValue(agentComposerDraftAtom) - const setDraft = useSetAtom(agentComposerDraftAtom) const files = useAtomValue(agentComposerFilesAtom) - const setFiles = useSetAtom(agentComposerFilesAtom) + const clearAgentConfigNote = useSetAtom(clearAgentConfigNoteAtom) + const removeAgentFile = useSetAtom(removeAgentFileAtom) + const upsertAgentFile = useSetAtom(upsertAgentFileAtom) const buildNoteFile = getBuildNoteFile(draft.configNote) const visibleFiles = buildNoteFile ? [buildNoteFile, ...files] : files const { mutate: deleteAgentFile } = useMutation(consoleQuery.agent.byAgentId.config.files.byName.delete.mutationOptions()) const { mutate: deleteWorkflowAgentFile } = useMutation(consoleQuery.apps.byAppId.agent.config.files.byName.delete.mutationOptions()) const removeFile = useCallback((fileId: string) => { if (fileId === BUILD_NOTE_FILE_ID) { - setDraft(draft => ({ - ...draft, - configNote: '', - })) + clearAgentConfigNote() return } @@ -316,7 +367,7 @@ export function AgentFiles() { return const onSuccess = () => { - setFiles(files => removeAgentFileNode(files, fileId)) + removeAgentFile(fileId) } if (apiContext.workflow) { deleteWorkflowAgentFile({ @@ -343,20 +394,17 @@ export function AgentFiles() { version_id: apiContext.versionId, }, }, { onSuccess }) - }, [apiContext, deleteAgentFile, deleteWorkflowAgentFile, files, setDraft, setFiles]) + }, [apiContext, clearAgentConfigNote, deleteAgentFile, deleteWorkflowAgentFile, files, removeAgentFile]) const handleOpenUpload = useCallback((options?: AgentOrchestrateAddActionOptions) => { promptAddCallbackRef.current = options?.onAdded setIsUploadOpen(true) }, []) useRegisterAgentOrchestrateAddAction('files', handleOpenUpload) const handleUploaded = useCallback((file: AgentFileNode) => { - setFiles(files => [ - ...removeAgentFileNode(files, file.id), - file, - ]) + upsertAgentFile(file) promptAddCallbackRef.current?.(file) promptAddCallbackRef.current = undefined - }, [setFiles]) + }, [upsertAgentFile]) const handleUploadOpenChange = useCallback((open: boolean) => { if (!open) promptAddCallbackRef.current = undefined diff --git a/web/features/agent-v2/agent-detail/configure/components/orchestrate/files/tree.tsx b/web/features/agent-v2/agent-detail/configure/components/orchestrate/files/tree.tsx index c1fb78524ea..8f0ddb4421c 100644 --- a/web/features/agent-v2/agent-detail/configure/components/orchestrate/files/tree.tsx +++ b/web/features/agent-v2/agent-detail/configure/components/orchestrate/files/tree.tsx @@ -182,7 +182,7 @@ export function AgentFileTree({ label={label} labelledBy={labelledBy} slotClassNames={{ - viewport: 'max-h-[inherit] overscroll-contain', + viewport: 'max-h-[inherit]', content: 'w-full max-w-full min-w-0!', scrollbar: 'hidden', }} diff --git a/web/features/agent-v2/agent-detail/configure/components/orchestrate/header.tsx b/web/features/agent-v2/agent-detail/configure/components/orchestrate/header.tsx index 269a453028e..85607335da8 100644 --- a/web/features/agent-v2/agent-detail/configure/components/orchestrate/header.tsx +++ b/web/features/agent-v2/agent-detail/configure/components/orchestrate/header.tsx @@ -1,6 +1,7 @@ 'use client' import type { ReactNode } from 'react' +import { Popover, PopoverContent, PopoverTrigger } from '@langgenius/dify-ui/popover' import { useTranslation } from 'react-i18next' type AgentOrchestrateHeaderProps = { @@ -15,6 +16,7 @@ export function AgentOrchestrateHeader({ isBuildDraftActive = false, }: AgentOrchestrateHeaderProps) { const { t } = useTranslation('agentV2') + const communityEditionIsolationTip = t('agentDetail.configure.communityEditionIsolationTip') return (

    @@ -23,6 +25,28 @@ export function AgentOrchestrateHeader({

    {t('agentDetail.configure.title')}

    + + + + + )} + /> + + {communityEditionIsolationTip} + + {isBuildDraftActive && ( {t('agentDetail.configure.buildDraft.modeBadge')} diff --git a/web/features/agent-v2/agent-detail/configure/components/orchestrate/index.tsx b/web/features/agent-v2/agent-detail/configure/components/orchestrate/index.tsx index 894265d0ad6..18869d2b762 100644 --- a/web/features/agent-v2/agent-detail/configure/components/orchestrate/index.tsx +++ b/web/features/agent-v2/agent-detail/configure/components/orchestrate/index.tsx @@ -41,6 +41,7 @@ type AgentOrchestratePanelProps = { className?: string readOnly?: boolean selectedVersionSnapshot?: AgentConfigSnapshotSummaryResponse | null + workflowReferencesEnabled?: boolean isBuildDraftActive?: boolean buildDraftChangedKeys?: readonly AgentBuildDraftChangedKey[] showHeader?: boolean @@ -68,6 +69,7 @@ export function AgentOrchestratePanel({ className, readOnly = false, selectedVersionSnapshot, + workflowReferencesEnabled, isBuildDraftActive = false, buildDraftChangedKeys = [], showHeader = true, @@ -92,6 +94,7 @@ export function AgentOrchestratePanel({ draftSavedAt={draftSavedAt} isPublishing={isPublishing} selectedVersionSnapshot={selectedVersionSnapshot} + workflowReferencesEnabled={workflowReferencesEnabled} onPublish={onPublish} onExitVersions={onExitVersions} onOpenVersions={onOpenVersions} diff --git a/web/features/agent-v2/agent-detail/configure/components/orchestrate/knowledge/index.tsx b/web/features/agent-v2/agent-detail/configure/components/orchestrate/knowledge/index.tsx index c0feb0ed44f..fd95479099c 100644 --- a/web/features/agent-v2/agent-detail/configure/components/orchestrate/knowledge/index.tsx +++ b/web/features/agent-v2/agent-detail/configure/components/orchestrate/knowledge/index.tsx @@ -2,10 +2,15 @@ import type { AgentOrchestrateAddActionOptions } from '../add-actions-context' import type { AgentKnowledgeRetrievalItem } from '@/features/agent-v2/agent-composer/form-state' -import { useAtom } from 'jotai' +import { useAtomValue, useSetAtom } from 'jotai' import { useRef, useState } from 'react' import { useTranslation } from 'react-i18next' -import { agentComposerKnowledgeRetrievalsAtom } from '@/features/agent-v2/agent-composer/store-modules/knowledge' +import { + addKnowledgeRetrievalAtom, + agentComposerKnowledgeRetrievalsAtom, + removeKnowledgeRetrievalAtom, + updateKnowledgeRetrievalAtom, +} from '@/features/agent-v2/agent-composer/store-modules/knowledge' import { useRegisterAgentOrchestrateAddAction } from '../add-actions-context' import { ConfigureSectionAddButton } from '../common/add-button' import { ConfigureSectionConfigurableItem } from '../common/configurable-item' @@ -48,7 +53,10 @@ function AgentKnowledgeRetrievalRow({ export function AgentKnowledgeRetrieval() { const { t } = useTranslation('agentV2') - const [retrievals, setRetrievals] = useAtom(agentComposerKnowledgeRetrievalsAtom) + const retrievals = useAtomValue(agentComposerKnowledgeRetrievalsAtom) + const addKnowledgeRetrieval = useSetAtom(addKnowledgeRetrievalAtom) + const updateKnowledgeRetrieval = useSetAtom(updateKnowledgeRetrievalAtom) + const removeKnowledgeRetrieval = useSetAtom(removeKnowledgeRetrievalAtom) const [isAddDialogOpen, setIsAddDialogOpen] = useState(false) const [addDialogName, setAddDialogName] = useState() const [editingRetrieval, setEditingRetrieval] = useState(null) @@ -57,7 +65,7 @@ export function AgentKnowledgeRetrieval() { const retrievalListId = 'agent-configure-knowledge-retrieval-list' const isDialogOpen = isAddDialogOpen || !!editingRetrieval const updateRetrieval = (nextRetrieval: AgentKnowledgeRetrievalItem) => { - setRetrievals(retrievals.map(retrieval => retrieval.id === nextRetrieval.id ? nextRetrieval : retrieval)) + updateKnowledgeRetrieval(nextRetrieval) setEditingRetrieval(nextRetrieval) } const getDefaultRetrievalName = (index: number) => { @@ -74,7 +82,7 @@ export function AgentKnowledgeRetrieval() { setIsAddDialogOpen(true) } const createRetrieval = (nextRetrieval: AgentKnowledgeRetrievalItem) => { - setRetrievals(current => [...current, nextRetrieval]) + addKnowledgeRetrieval(nextRetrieval) setEditingRetrieval(nextRetrieval) setIsAddDialogOpen(false) addOptionsRef.current?.onAdded?.(nextRetrieval) @@ -110,7 +118,7 @@ export function AgentKnowledgeRetrieval() { setRetrievals(retrievals.filter(retrieval => retrieval.id !== item.id))} + onDelete={() => removeKnowledgeRetrieval(item.id)} onEdit={() => setEditingRetrieval(item)} /> ))} diff --git a/web/features/agent-v2/agent-detail/configure/components/orchestrate/prompt-editor/__tests__/options.spec.ts b/web/features/agent-v2/agent-detail/configure/components/orchestrate/prompt-editor/__tests__/options.spec.ts new file mode 100644 index 00000000000..12af5c3af62 --- /dev/null +++ b/web/features/agent-v2/agent-detail/configure/components/orchestrate/prompt-editor/__tests__/options.spec.ts @@ -0,0 +1,57 @@ +import { describe, expect, it } from 'vitest' +import { insertTokenAtTextRange, replaceTrailingSlashWithToken } from '../options' + +describe('prompt editor token replacement', () => { + // Replacing the tracked slash range keeps insertion at the user's caret instead of appending. + describe('insertTokenAtTextRange', () => { + it('should replace a slash in the middle of the prompt and place the cursor after the token', () => { + expect(insertTokenAtTextRange( + 'Review / before replying', + { start: 7, end: 8 }, + '[§file:file-1:Spec§]', + )).toEqual({ + value: 'Review [§file:file-1:Spec§] before replying', + cursorOffset: 'Review [§file:file-1:Spec§]'.length, + }) + }) + + it('should add spacing when the slash is adjacent to text and place the cursor after the spacer', () => { + expect(insertTokenAtTextRange( + 'Review/now', + { start: 6, end: 7 }, + '[§skill:analysis:Analysis§]', + )).toEqual({ + value: 'Review [§skill:analysis:Analysis§] now', + cursorOffset: 'Review [§skill:analysis:Analysis§] '.length, + }) + }) + + it('should clamp out-of-bound ranges before replacing', () => { + expect(insertTokenAtTextRange( + 'Review/', + { start: 6, end: 99 }, + '[§knowledge:kb-1:KB§]', + )).toEqual({ + value: 'Review [§knowledge:kb-1:KB§]', + cursorOffset: 'Review [§knowledge:kb-1:KB§]'.length, + }) + }) + }) + + // Existing fallback behavior is retained for callers that only know about a trailing slash. + describe('replaceTrailingSlashWithToken', () => { + it('should replace a trailing slash', () => { + expect(replaceTrailingSlashWithToken( + 'Review /', + '[§file:file-1:Spec§]', + )).toBe('Review [§file:file-1:Spec§]') + }) + + it('should append when no trailing slash exists', () => { + expect(replaceTrailingSlashWithToken( + 'Review', + '[§file:file-1:Spec§]', + )).toBe('Review [§file:file-1:Spec§]') + }) + }) +}) diff --git a/web/features/agent-v2/agent-detail/configure/components/orchestrate/prompt-editor/index.tsx b/web/features/agent-v2/agent-detail/configure/components/orchestrate/prompt-editor/index.tsx index aef94ad0593..639ebed2b63 100644 --- a/web/features/agent-v2/agent-detail/configure/components/orchestrate/prompt-editor/index.tsx +++ b/web/features/agent-v2/agent-detail/configure/components/orchestrate/prompt-editor/index.tsx @@ -1,6 +1,8 @@ 'use client' +import type { LexicalNode } from 'lexical' import type { KeyboardEvent, MouseEvent, PointerEvent as ReactPointerEvent } from 'react' +import type { TextRange } from './options' import type { SlashMenuCategory, SlashMenuView } from './slash' import type { RosterReferenceToken } from '@/app/components/base/prompt-editor/plugins/roster-reference-block/utils' import type { AgentFileNode, AgentProviderTool, AgentTool } from '@/features/agent-v2/agent-composer/form-state' @@ -8,9 +10,21 @@ import { cn } from '@langgenius/dify-ui/cn' import { Kbd } from '@langgenius/dify-ui/kbd' import { toast } from '@langgenius/dify-ui/toast' import { Tooltip, TooltipContent, TooltipTrigger } from '@langgenius/dify-ui/tooltip' +import { useLexicalComposerContext } from '@lexical/react/LexicalComposerContext' +import { mergeRegister } from '@lexical/utils' import { useClipboard } from 'foxact/use-clipboard' -import { useAtom, useAtomValue } from 'jotai' +import { useAtom, useAtomValue, useSetAtom } from 'jotai' +import { + $getRoot, + $getSelection, + $isElementNode, + $isRangeSelection, + $isTextNode, + COMMAND_PRIORITY_LOW, + SELECTION_CHANGE_COMMAND, +} from 'lexical' import { useCallback, useEffect, useMemo, useRef, useState, useSyncExternalStore } from 'react' +import { createPortal } from 'react-dom' import { useTranslation } from 'react-i18next' import { Infotip } from '@/app/components/base/infotip' import PromptEditor from '@/app/components/base/prompt-editor' @@ -18,14 +32,17 @@ import BlockIcon from '@/app/components/workflow/block-icon' import { BlockEnum } from '@/app/components/workflow/types' import { agentComposerKnowledgeRetrievalsAtom } from '@/features/agent-v2/agent-composer/store-modules/knowledge' import { agentComposerPromptAtom } from '@/features/agent-v2/agent-composer/store-modules/prompt' -import { agentComposerToolsAtom } from '@/features/agent-v2/agent-composer/store-modules/tools' +import { + addProviderToolsAtom, + agentComposerToolsAtom, +} from '@/features/agent-v2/agent-composer/store-modules/tools' import { ENABLE_AGENT_CLI_TOOLS } from '@/features/agent-v2/agent-detail/configure/feature-flags' import { useAgentOrchestrateAddActions } from '../add-actions-context' import { AgentConfigureTipContent } from '../common/tip-content' import { useAgentConfigFiles, useAgentConfigSkills } from '../config-context' import { useAgentOrchestrateReadOnly } from '../read-only-context' import { useAgentPromptToolIconResolver } from './hooks' -import { replaceTrailingSlashWithToken } from './options' +import { insertTokenAtTextRange, replaceTrailingSlashWithToken } from './options' import { AgentPromptSlashMenu } from './slash' const subscribeHydrationState = () => () => {} @@ -152,13 +169,232 @@ const isSelectionAfterSlash = (rootElement: HTMLElement | null, fallbackValue: s return previousChild ? getLastTextContent(previousChild).endsWith('/') : false } +/* v8 ignore start -- Lexical selection offsets and DOM range geometry are browser-editor integration glue; user-visible slash insertion behavior is covered by AgentPromptEditor tests. @preserve */ +const getNodeOffset = ( + node: LexicalNode, + anchorNode: LexicalNode, + anchorOffset: number, +): { found: boolean, offset: number } => { + if (node.getKey() === anchorNode.getKey()) + return { found: true, offset: anchorOffset } + + if (!$isElementNode(node)) + return { found: false, offset: node.getTextContent().length } + + let offset = 0 + for (const child of node.getChildren()) { + const childOffset = getNodeOffset(child, anchorNode, anchorOffset) + if (childOffset.found) + return { found: true, offset: offset + childOffset.offset } + + offset += childOffset.offset + } + + return { found: false, offset } +} + +const getSelectionTextOffset = () => { + const selection = $getSelection() + if (!$isRangeSelection(selection) || !selection.isCollapsed()) + return null + + const anchor = selection.anchor + const anchorNode = anchor.getNode() + const root = $getRoot() + let offset = 0 + + for (const child of root.getChildren()) { + const childOffset = getNodeOffset(child, anchorNode, anchor.offset) + if (childOffset.found) + return offset + childOffset.offset + + offset += childOffset.offset + 1 + } + + return null +} + +const readSlashInsertRange = (): TextRange | null => { + const offset = getSelectionTextOffset() + if (!offset) + return null + + const value = $getRoot().getChildren().map(node => node.getTextContent()).join('\n') + if (value[offset - 1] !== '/') + return null + + return { + start: offset - 1, + end: offset, + } +} + +const selectNodeTextOffset = (node: LexicalNode, textOffset: number): boolean => { + if ($isTextNode(node)) { + const offset = Math.max(0, Math.min(textOffset, node.getTextContentSize())) + node.select(offset, offset) + return true + } + + if (!$isElementNode(node)) + return false + + const children = node.getChildren() + let currentOffset = 0 + + for (let index = 0; index < children.length; index++) { + const child = children[index]! + const childLength = child.getTextContent().length + if (textOffset > currentOffset + childLength) { + currentOffset += childLength + continue + } + + if ($isElementNode(child) || $isTextNode(child)) + return selectNodeTextOffset(child, textOffset - currentOffset) + + const childSelectionOffset = textOffset <= currentOffset ? index : index + 1 + node.select(childSelectionOffset, childSelectionOffset) + return true + } + + node.select(children.length, children.length) + return true +} + +const selectTextOffset = (textOffset: number) => { + const root = $getRoot() + let currentOffset = 0 + + for (const child of root.getChildren()) { + const childLength = child.getTextContent().length + if (textOffset <= currentOffset + childLength) { + selectNodeTextOffset(child, textOffset - currentOffset) + return + } + + currentOffset += childLength + 1 + } + + root.selectEnd() +} + +type SelectionRestoreRequest = { + id: number + offset: number +} + +type SlashMenuPosition = { + left: number + top: number +} + +const slashMenuViewportPadding = 8 +const slashMenuMainWidth = 200 +const slashMenuSubmenuWidth = 360 + +const getSlashMenuPosition = (editorElement: HTMLElement): SlashMenuPosition | null => { + const selection = window.getSelection() + if (!selection || !selection.isCollapsed || selection.rangeCount === 0) + return null + + const anchorNode = selection.anchorNode + if (!anchorNode || !editorElement.contains(anchorNode)) + return null + + const range = selection.getRangeAt(0).cloneRange() + let rect: DOMRect | null = null + const rects = range.getClientRects() + if (rects.length) + rect = rects[rects.length - 1]! + else + rect = range.getBoundingClientRect() + + if (!rect || (rect.top === 0 && rect.left === 0 && rect.width === 0 && rect.height === 0)) { + const node = anchorNode.nodeType === Node.ELEMENT_NODE + ? anchorNode as Element + : anchorNode.parentElement + + rect = node?.getBoundingClientRect() ?? editorElement.getBoundingClientRect() + } + + const editorRect = editorElement.getBoundingClientRect() + if (!rect || rect.bottom < editorRect.top || rect.top > editorRect.bottom) + return null + + return { + left: rect.right, + top: rect.bottom + 4, + } +} + +const getSlashMenuLeft = (position: SlashMenuPosition, width: number) => { + if (typeof window === 'undefined') + return position.left + + return Math.max( + slashMenuViewportPadding, + Math.min(position.left, window.innerWidth - width - slashMenuViewportPadding), + ) +} +/* v8 ignore stop */ + +function AgentPromptSelectionBridge({ + restoreRequest, + onSlashRangeChange, +}: { + restoreRequest: SelectionRestoreRequest | null + onSlashRangeChange: (range: TextRange | null) => void +}) { + const [editor] = useLexicalComposerContext() + + useEffect(() => { + const updateSlashRange = () => { + editor.getEditorState().read(() => { + onSlashRangeChange(readSlashInsertRange()) + }) + + return false + } + + updateSlashRange() + + return mergeRegister( + editor.registerCommand( + SELECTION_CHANGE_COMMAND, + updateSlashRange, + COMMAND_PRIORITY_LOW, + ), + editor.registerUpdateListener(({ editorState }) => { + editorState.read(() => { + onSlashRangeChange(readSlashInsertRange()) + }) + }), + ) + }, [editor, onSlashRangeChange]) + + useEffect(() => { + if (!restoreRequest) + return + + editor.focus(() => { + editor.update(() => { + selectTextOffset(restoreRequest.offset) + }) + }) + }, [editor, restoreRequest]) + + return null +} + export function AgentPromptEditor() { const { t } = useTranslation('agentV2') const readOnly = useAgentOrchestrateReadOnly() const [value, setValue] = useAtom(agentComposerPromptAtom) const { skills } = useAgentConfigSkills() const { files } = useAgentConfigFiles() - const [tools, setTools] = useAtom(agentComposerToolsAtom) + const tools = useAtomValue(agentComposerToolsAtom) + const addProviderTools = useSetAtom(addProviderToolsAtom) const { getConfiguredToolIcon } = useAgentPromptToolIconResolver() const retrievals = useAtomValue(agentComposerKnowledgeRetrievalsAtom) const addActions = useAgentOrchestrateAddActions() @@ -178,8 +414,12 @@ export function AgentPromptEditor() { }) const [slashMenuView, setSlashMenuView] = useState('main') const [isSlashMenuOpen, setIsSlashMenuOpen] = useState(false) + const [slashMenuPosition, setSlashMenuPosition] = useState(null) + const [selectionRestoreRequest, setSelectionRestoreRequest] = useState(null) const rootRef = useRef(null) const editorRef = useRef(null) + const slashInsertRangeRef = useRef(null) + const selectionRestoreRequestIdRef = useRef(0) const configuredReferenceIds = useMemo(() => { const skillIds = new Set() skills.forEach((skill) => { @@ -215,23 +455,46 @@ export function AgentPromptEditor() { const closeSlashMenu = () => { setIsSlashMenuOpen(false) + setSlashMenuPosition(null) setSlashMenuView('main') } - const openSlashMenu = () => { + const updateSlashMenuPosition = useCallback(() => { + const editorElement = editorRef.current + if (!editorElement) + return + + const position = getSlashMenuPosition(editorElement) + if (!position) + return + + setSlashMenuPosition(position) + }, []) + + const openSlashMenu = useCallback(() => { setSlashMenuView('main') + updateSlashMenuPosition() setIsSlashMenuOpen(true) - } + }, [updateSlashMenuPosition]) const syncSlashMenuWithSelection = useCallback(() => { if (!isHydrated || readOnly) return - if (isSelectionAfterSlash(editorRef.current, value)) + if (isSelectionAfterSlash(editorRef.current, value)) { + updateSlashMenuPosition() openSlashMenu() - else + } + else { + slashInsertRangeRef.current = null closeSlashMenu() - }, [isHydrated, readOnly, value]) + } + }, [isHydrated, openSlashMenu, readOnly, updateSlashMenuPosition, value]) + + const handleSlashRangeChange = useCallback((range: TextRange | null) => { + if (range) + slashInsertRangeRef.current = range + }, []) const handleEditorKeyDown = (event: KeyboardEvent) => { if (!isHydrated || readOnly) @@ -287,7 +550,25 @@ export function AgentPromptEditor() { } const handleSlashSelect = (token: string) => { - setValue(replaceTrailingSlashWithToken(value, token)) + const slashRange = slashInsertRangeRef.current + let insertionResult + if (slashRange) { + insertionResult = insertTokenAtTextRange(value, slashRange, token) + } + else { + const nextValue = replaceTrailingSlashWithToken(value, token) + insertionResult = { + value: nextValue, + cursorOffset: nextValue.length, + } + } + setValue(insertionResult.value) + slashInsertRangeRef.current = null + selectionRestoreRequestIdRef.current += 1 + setSelectionRestoreRequest({ + id: selectionRestoreRequestIdRef.current, + offset: insertionResult.cursorOffset, + }) closeSlashMenu() } @@ -343,6 +624,13 @@ export function AgentPromptEditor() { if (!(target instanceof Node)) return + if ( + target instanceof Element + && target.closest('[data-agent-prompt-slash-menu]') + ) { + return + } + if (!rootRef.current?.contains(target)) closeSlashMenu() } @@ -375,6 +663,37 @@ export function AgentPromptEditor() { icon: 'i-ri-book-open-line', }, ] + const slashMenuWidth = slashMenuView === 'main' ? slashMenuMainWidth : slashMenuSubmenuWidth + const slashMenu = isHydrated && !readOnly && isSlashMenuOpen + ? createPortal( +
    + setSlashMenuView('main')} + onOpenCategory={setSlashMenuView} + onSelect={handleSlashSelect} + /> +
    , + document.body, + ) + : null return (
    @@ -442,7 +761,12 @@ export function AgentPromptEditor() { }} disableSlashPicker disableBracePicker - /> + > + +
    {!readOnly && (
    - {isHydrated && !readOnly && isSlashMenuOpen && ( -
    - setSlashMenuView('main')} - onOpenCategory={setSlashMenuView} - onSelect={handleSlashSelect} - /> -
    - )} + {slashMenu}
    ) diff --git a/web/features/agent-v2/agent-detail/configure/components/orchestrate/prompt-editor/options.ts b/web/features/agent-v2/agent-detail/configure/components/orchestrate/prompt-editor/options.ts index bc4e0631f2f..79b29430504 100644 --- a/web/features/agent-v2/agent-detail/configure/components/orchestrate/prompt-editor/options.ts +++ b/web/features/agent-v2/agent-detail/configure/components/orchestrate/prompt-editor/options.ts @@ -71,6 +71,20 @@ const appendToken = (value: string, token: string) => { return `${value}${value.endsWith(' ') || value.endsWith('\n') ? '' : ' '}${token}` } +export type TextRange = { + start: number + end: number +} + +export type TokenInsertionResult = { + value: string + cursorOffset: number +} + +const hasTrailingSpace = (value: string) => value.endsWith(' ') || value.endsWith('\n') + +const hasLeadingSpace = (value: string) => value.startsWith(' ') || value.startsWith('\n') + export const replaceTrailingSlashWithToken = (value: string, token: string) => { if (!value.endsWith('/')) return appendToken(value, token) @@ -79,5 +93,19 @@ export const replaceTrailingSlashWithToken = (value: string, token: string) => { if (!valueWithoutSlash) return token - return `${valueWithoutSlash}${valueWithoutSlash.endsWith(' ') || valueWithoutSlash.endsWith('\n') ? '' : ' '}${token}` + return `${valueWithoutSlash}${hasTrailingSpace(valueWithoutSlash) ? '' : ' '}${token}` +} + +export const insertTokenAtTextRange = (value: string, range: TextRange, token: string): TokenInsertionResult => { + const start = Math.max(0, Math.min(range.start, value.length)) + const end = Math.max(start, Math.min(range.end, value.length)) + const prefix = value.slice(0, start) + const suffix = value.slice(end) + const beforeToken = prefix && !hasTrailingSpace(prefix) ? ' ' : '' + const afterToken = suffix && !hasLeadingSpace(suffix) ? ' ' : '' + + return { + value: `${prefix}${beforeToken}${token}${afterToken}${suffix}`, + cursorOffset: prefix.length + beforeToken.length + token.length + afterToken.length, + } } diff --git a/web/features/agent-v2/agent-detail/configure/components/orchestrate/prompt-editor/slash.tsx b/web/features/agent-v2/agent-detail/configure/components/orchestrate/prompt-editor/slash.tsx index 41bd5254582..40e2df9babb 100644 --- a/web/features/agent-v2/agent-detail/configure/components/orchestrate/prompt-editor/slash.tsx +++ b/web/features/agent-v2/agent-detail/configure/components/orchestrate/prompt-editor/slash.tsx @@ -2,11 +2,11 @@ import type { ReactNode } from 'react' import type { AgentOrchestrateAddAction, AgentOrchestrateAddedItem } from '../add-actions-context' -import type { AgentProviderToolDefaultValue } from '../tools/types' import type { Tool } from '@/app/components/tools/types' import type { ToolTypeEnum, ToolValue } from '@/app/components/workflow/block-selector/types' import type { ToolWithProvider } from '@/app/components/workflow/types' import type { AgentFileNode, AgentKnowledgeRetrievalItem, AgentSkill, AgentTool } from '@/features/agent-v2/agent-composer/form-state' +import type { AgentProviderToolDefaultValue } from '@/features/agent-v2/agent-composer/store-modules/tools' import { cn } from '@langgenius/dify-ui/cn' import { FileTreeIcon } from '@langgenius/dify-ui/file-tree' import { useMemo, useState } from 'react' @@ -25,7 +25,6 @@ import { useAllMCPTools, useAllWorkflowTools, } from '@/service/use-tools' -import { addProviderTools } from '../tools/hooks' import { useAgentPromptToolIconResolver } from './hooks' export type SlashMenuView = 'main' | 'skills' | 'files' | 'tools' | 'knowledge' @@ -42,7 +41,7 @@ type AgentPromptSlashMenuProps = { skills: AgentSkill[] files: AgentFileNode[] tools: AgentTool[] - onToolsChange: (tools: AgentTool[]) => void + onAddProviderTools: (tools: AgentProviderToolDefaultValue[]) => void onAddCliTool?: AgentOrchestrateAddAction onAddFile?: AgentOrchestrateAddAction onAddKnowledge?: AgentOrchestrateAddAction @@ -83,7 +82,7 @@ export function AgentPromptSlashMenu({ skills, files, tools, - onToolsChange, + onAddProviderTools, onAddCliTool, onAddFile, onAddKnowledge, @@ -167,7 +166,7 @@ export function AgentPromptSlashMenu({ {view === 'tools' && ( )} @@ -279,11 +278,11 @@ function AgentPromptFileRows({ function AgentPromptToolRows({ configuredTools, - onConfiguredToolsChange, + onAddProviderTools, onSelect, }: { configuredTools: AgentTool[] - onConfiguredToolsChange: (tools: AgentTool[]) => void + onAddProviderTools: (tools: AgentProviderToolDefaultValue[]) => void onSelect: (token: string) => void }) { const { t } = useTranslation('agentV2') @@ -329,7 +328,7 @@ function AgentPromptToolRows({ ] const selectTools = (tools: AgentProviderToolDefaultValue[]) => { - onConfiguredToolsChange(addProviderTools(configuredTools, tools)) + onAddProviderTools(tools) } const toggleProvider = (providerId: string) => { diff --git a/web/features/agent-v2/agent-detail/configure/components/orchestrate/publish-bar/index.tsx b/web/features/agent-v2/agent-detail/configure/components/orchestrate/publish-bar/index.tsx index 9ce15bed43c..2d24fcbf119 100644 --- a/web/features/agent-v2/agent-detail/configure/components/orchestrate/publish-bar/index.tsx +++ b/web/features/agent-v2/agent-detail/configure/components/orchestrate/publish-bar/index.tsx @@ -33,6 +33,7 @@ type AgentConfigurePublishBarProps = { draftSavedAt?: number isPublishing?: boolean selectedVersionSnapshot?: AgentConfigSnapshotSummaryResponse | null + workflowReferencesEnabled?: boolean onPublish?: () => void | Promise onExitVersions?: () => void onOpenVersions?: () => void @@ -90,6 +91,7 @@ export function AgentConfigurePublishBar({ draftSavedAt, isPublishing = false, selectedVersionSnapshot, + workflowReferencesEnabled = true, onPublish, onExitVersions, onOpenVersions, @@ -128,11 +130,9 @@ export function AgentConfigurePublishBar({ agent_id: agentId, }, }, + enabled: workflowReferencesEnabled && publishIsAvailable && !selectedVersionSnapshot, }) - const workflowReferencesQuery = useQuery({ - ...workflowReferencesQueryOptions, - enabled: publishIsAvailable && !selectedVersionSnapshot, - }) + const workflowReferencesQuery = useQuery(workflowReferencesQueryOptions) const restoreVersionMutation = useMutation(consoleQuery.agent.byAgentId.versions.byVersionId.restore.post.mutationOptions()) const canPublish = publishIsAvailable @@ -195,7 +195,9 @@ export function AgentConfigurePublishBar({ } const cachedReferences = queryClient.getQueryData(workflowReferencesQueryOptions.queryKey) - const references = (cachedReferences ?? workflowReferencesQuery.data ?? await queryClient.ensureQueryData(workflowReferencesQueryOptions))?.data ?? [] + const references = workflowReferencesEnabled + ? (cachedReferences ?? workflowReferencesQuery.data ?? await queryClient.ensureQueryData(workflowReferencesQueryOptions))?.data ?? [] + : [] if (references.length > 0) { setPublishBarMode({ status: 'confirmingImpact', references }) diff --git a/web/features/agent-v2/agent-detail/configure/components/orchestrate/skills/__tests__/index.spec.tsx b/web/features/agent-v2/agent-detail/configure/components/orchestrate/skills/__tests__/index.spec.tsx index 3c3815e625d..a9951b441cd 100644 --- a/web/features/agent-v2/agent-detail/configure/components/orchestrate/skills/__tests__/index.spec.tsx +++ b/web/features/agent-v2/agent-detail/configure/components/orchestrate/skills/__tests__/index.spec.tsx @@ -30,6 +30,14 @@ type ConfigSkillFileQueryOptionsInput = { } } +type ConfigSkillDownloadQueryOptionsInput = { + input: { + params: { + name: string + } + } +} + const mocks = vi.hoisted(() => ({ deleteSkillMutationFn: vi.fn(async (_input: unknown) => ({ removed_names: ['Tender Analyzer'], result: 'success' })), uploadSkillMutationFn: vi.fn(async (_input: unknown) => ({ @@ -44,9 +52,12 @@ const mocks = vi.hoisted(() => ({ size: 128, }, })), + skillDownloadQueryOptions: vi.fn((_options: ConfigSkillDownloadQueryOptionsInput) => ({})), inspectQueryOptions: vi.fn((_options: ConfigSkillInspectQueryOptionsInput) => ({})), previewQueryOptions: vi.fn((_options: ConfigSkillFileQueryOptionsInput) => ({})), downloadQueryOptions: vi.fn((_options: ConfigSkillFileQueryOptionsInput) => ({})), + downloadBlob: vi.fn(), + downloadUrl: vi.fn(), })) vi.mock('@langgenius/dify-ui/toast', () => ({ @@ -56,6 +67,11 @@ vi.mock('@langgenius/dify-ui/toast', () => ({ }, })) +vi.mock('@/utils/download', () => ({ + downloadBlob: mocks.downloadBlob, + downloadUrl: mocks.downloadUrl, +})) + vi.mock('@/service/client', () => ({ consoleQuery: { agent: { @@ -71,6 +87,11 @@ vi.mock('@/service/client', () => ({ delete: { mutationOptions: () => ({ mutationFn: mocks.deleteSkillMutationFn }), }, + download: { + get: { + queryOptions: mocks.skillDownloadQueryOptions, + }, + }, inspect: { get: { queryOptions: mocks.inspectQueryOptions, @@ -107,6 +128,11 @@ vi.mock('@/service/client', () => ({ delete: { mutationOptions: () => ({ mutationFn: mocks.deleteSkillMutationFn }), }, + download: { + get: { + queryOptions: mocks.skillDownloadQueryOptions, + }, + }, inspect: { get: { queryOptions: mocks.inspectQueryOptions, @@ -235,6 +261,12 @@ describe('AgentSkills', () => { url: `https://example.com/${input.query.path}`, }), })) + mocks.skillDownloadQueryOptions.mockImplementation(({ input }) => ({ + queryKey: ['download-skill', input], + queryFn: async () => ({ + url: `https://example.com/${input.params.name}.skill`, + }), + })) }) it('should delete a configured skill by config name', async () => { @@ -390,6 +422,69 @@ describe('AgentSkills', () => { }) }) + it('should download a whole skill package from the row action', async () => { + const user = userEvent.setup() + renderAgentSkills() + + await user.click(screen.getByRole('button', { + name: /common\.operation\.download.*Tender Analyzer/, + })) + + await waitFor(() => { + expect(mocks.skillDownloadQueryOptions).toHaveBeenCalledWith(expect.objectContaining({ + input: expect.objectContaining({ + params: { + agent_id: 'agent-1', + name: 'Tender Analyzer', + }, + query: { + draft_type: 'draft', + version_id: undefined, + }, + }), + })) + }) + expect(mocks.downloadUrl).toHaveBeenCalledWith({ + url: 'https://example.com/Tender Analyzer.skill', + fileName: 'Tender Analyzer', + }) + }) + + it('should download a whole workflow skill package with node_id', async () => { + const user = userEvent.setup() + renderAgentSkills({ + apiContext: { + agentId: 'agent-1', + draftType: 'draft', + versionId: 'draft-1', + workflow: { + appId: 'app-1', + nodeId: 'node-1', + }, + }, + }) + + await user.click(screen.getByRole('button', { + name: /common\.operation\.download.*Tender Analyzer/, + })) + + await waitFor(() => { + expect(mocks.skillDownloadQueryOptions).toHaveBeenCalledWith(expect.objectContaining({ + input: expect.objectContaining({ + params: { + app_id: 'app-1', + name: 'Tender Analyzer', + }, + query: { + draft_type: 'draft', + node_id: 'node-1', + version_id: 'draft-1', + }, + }), + })) + }) + }) + it('should inspect skills by config name and preview package members by member path', async () => { const user = userEvent.setup() renderAgentSkills() @@ -425,6 +520,75 @@ describe('AgentSkills', () => { }) }) + it('should wrap long preview lines instead of forcing a horizontal code block', async () => { + const user = userEvent.setup() + renderAgentSkills() + + await user.click(screen.getByText('Tender Analyzer').closest('button')!) + + const skillMdCode = await screen.findByText('# Skill') + expect(skillMdCode.tagName).toBe('CODE') + expect(skillMdCode).toHaveClass('[overflow-wrap:anywhere]') + expect(skillMdCode).toHaveClass('break-words') + expect(skillMdCode).toHaveClass('whitespace-pre-wrap') + expect(skillMdCode).not.toHaveClass('whitespace-pre') + expect(skillMdCode).not.toHaveClass('min-w-max') + }) + + it('should download skill package members from the detail file tree', async () => { + const user = userEvent.setup() + renderAgentSkills() + + await user.click(screen.getByText('Tender Analyzer').closest('button')!) + await user.click(await screen.findByText('references')) + await user.click(screen.getByText('guide.md').closest('button')!) + await user.click(screen.getByRole('button', { + name: /common\.operation\.download.*guide\.md/, + })) + + await waitFor(() => { + expect(mocks.downloadQueryOptions).toHaveBeenCalledWith(expect.objectContaining({ + input: expect.objectContaining({ + params: { + agent_id: 'agent-1', + name: 'Tender Analyzer', + }, + query: expect.objectContaining({ + path: 'references/guide.md', + }), + }), + })) + }) + expect(mocks.downloadUrl).toHaveBeenCalledWith({ + url: 'https://example.com/references/guide.md', + fileName: 'guide.md', + }) + }) + + it('should download inspected SKILL.md content as markdown', async () => { + const user = userEvent.setup() + renderAgentSkills() + + await user.click(screen.getByText('Tender Analyzer').closest('button')!) + await user.click(await screen.findByRole('button', { + name: /common\.operation\.download.*SKILL\.md/, + })) + + expect(mocks.downloadBlob).toHaveBeenCalledWith({ + data: expect.any(Blob), + fileName: 'SKILL.md', + }) + const blob = mocks.downloadBlob.mock.calls[0]?.[0].data as Blob + await expect(blob.text()).resolves.toBe('# Skill\n') + expect(mocks.downloadQueryOptions).not.toHaveBeenCalledWith(expect.objectContaining({ + input: expect.objectContaining({ + query: expect.objectContaining({ + path: 'SKILL.md', + }), + }), + })) + }) + it('should disable add and remove actions when the section is read only', () => { const { container } = renderAgentSkills({ readOnly: true }) diff --git a/web/features/agent-v2/agent-detail/configure/components/orchestrate/skills/detail-dialog.tsx b/web/features/agent-v2/agent-detail/configure/components/orchestrate/skills/detail-dialog.tsx index 36be64434e9..4275689c69d 100644 --- a/web/features/agent-v2/agent-detail/configure/components/orchestrate/skills/detail-dialog.tsx +++ b/web/features/agent-v2/agent-detail/configure/components/orchestrate/skills/detail-dialog.tsx @@ -48,6 +48,7 @@ export type AgentSkillDetail = { } onFolderOpenChange?: (context: { file: AgentSkillFileNode, depth: number, open: boolean }) => void onFolderDoubleClick?: (context: { file: AgentSkillFileNode, depth: number }) => void + onDownloadFile?: () => void onSelectFile?: (file: AgentSkillFileNode) => void renderFolderSuffix?: (context: { file: AgentSkillFileNode, depth: number }) => ReactNode selectedFileId?: string @@ -210,29 +211,25 @@ function AgentFilePreviewContent({ } if (binary) { - if (downloadUrl) { - return ( -
    - - {t('agentDetail.configure.files.preview.unsupported')} - - - - {tCommon('operation.download')} - -
    - ) - } - return ( -

    - {t('agentDetail.configure.files.preview.empty')} -

    + ) } @@ -244,19 +241,27 @@ function AgentFilePreviewContent({ ) } - const lines = content.split('\n') + const lines = content.split('\n').map((line, index) => ({ + content: line, + key: `${index}:${line}`, + lineNumber: String(index + 1).padStart(2, '0'), + })) return ( -
    - -
    -        {content}
    -      
    +
    + {lines.map(line => ( +
    + + + {line.content} + +
    + ))}
    ) } @@ -269,6 +274,7 @@ export function AgentSkillDetailDialog({ detail: AgentSkillDetail }) { const { t } = useTranslation('agentV2') + const { t: tCommon } = useTranslation('common') const previewTitle = detail.filePreview?.fileName return ( @@ -306,7 +312,19 @@ export function AgentSkillDetailDialog({ )}
    - +
    + {detail.onDownloadFile && previewTitle && ( + + )} + +
    (undefined) const apiContext = useAgentConfigApiContext() const skills = useAtomValue(agentComposerSkillsAtom) - const setSkills = useSetAtom(agentComposerSkillsAtom) + const upsertAgentSkill = useSetAtom(upsertAgentSkillAtom) + const removeAgentSkill = useSetAtom(removeAgentSkillAtom) const { mutate: deleteAgentSkill } = useMutation(consoleQuery.agent.byAgentId.config.skills.byName.delete.mutationOptions()) const { mutate: deleteAppSkill } = useMutation(consoleQuery.apps.byAppId.agent.config.skills.byName.delete.mutationOptions()) @@ -36,13 +41,10 @@ export function AgentSkills() { useRegisterAgentOrchestrateAddAction('skills', handleOpenUpload) const handleUploaded = useCallback((skill: AgentSkill) => { - setSkills(skills => [ - ...skills.filter(item => item.id !== skill.id), - skill, - ]) + upsertAgentSkill(skill) promptAddCallbackRef.current?.(skill) promptAddCallbackRef.current = undefined - }, [setSkills]) + }, [upsertAgentSkill]) const handleUploadOpenChange = useCallback((open: boolean) => { if (!open) @@ -56,7 +58,7 @@ export function AgentSkills() { return const onSuccess = () => { - setSkills(skills => skills.filter(item => item.id !== skillId)) + removeAgentSkill(skillId) } if (apiContext.workflow) { deleteAppSkill({ @@ -83,7 +85,7 @@ export function AgentSkills() { version_id: apiContext.versionId, }, }, { onSuccess }) - }, [apiContext, deleteAgentSkill, deleteAppSkill, setSkills, skills]) + }, [apiContext, deleteAgentSkill, deleteAppSkill, removeAgentSkill, skills]) return ( <> diff --git a/web/features/agent-v2/agent-detail/configure/components/orchestrate/skills/item.tsx b/web/features/agent-v2/agent-detail/configure/components/orchestrate/skills/item.tsx index f81e9d5d111..8209782b160 100644 --- a/web/features/agent-v2/agent-detail/configure/components/orchestrate/skills/item.tsx +++ b/web/features/agent-v2/agent-detail/configure/components/orchestrate/skills/item.tsx @@ -6,8 +6,11 @@ import { cn } from '@langgenius/dify-ui/cn' import { Dialog, } from '@langgenius/dify-ui/dialog' +import { useQueryClient } from '@tanstack/react-query' import { useCallback, useState } from 'react' import { useTranslation } from 'react-i18next' +import { consoleQuery } from '@/service/client' +import { downloadUrl } from '@/utils/download' import { useAgentOrchestrateReadOnly } from '../read-only-context' import { AgentSkillDetailDialog } from './detail-dialog' import { useAgentSkillDetail } from './use-skill-detail' @@ -22,11 +25,46 @@ export function AgentSkillItem({ onRemove: (skillId: string) => void }) { const { t } = useTranslation('agentV2') + const { t: tCommon } = useTranslation('common') + const queryClient = useQueryClient() const readOnly = useAgentOrchestrateReadOnly() const [isPreviewOpen, setIsPreviewOpen] = useState(false) const handleRemove = useCallback(() => { onRemove(skill.id) }, [onRemove, skill.id]) + const handleDownload = useCallback(async () => { + if (apiContext.workflow) { + const result = await queryClient.fetchQuery(consoleQuery.apps.byAppId.agent.config.skills.byName.download.get.queryOptions({ + input: { + params: { + app_id: apiContext.workflow.appId, + name: skill.name, + }, + query: { + node_id: apiContext.workflow.nodeId, + draft_type: apiContext.draftType, + version_id: apiContext.versionId, + }, + }, + })) + downloadUrl({ url: result.url, fileName: skill.name }) + return + } + + const result = await queryClient.fetchQuery(consoleQuery.agent.byAgentId.config.skills.byName.download.get.queryOptions({ + input: { + params: { + agent_id: apiContext.agentId, + name: skill.name, + }, + query: { + draft_type: apiContext.draftType, + version_id: apiContext.versionId, + }, + }, + })) + downloadUrl({ url: result.url, fileName: skill.name }) + }, [apiContext, queryClient, skill.name]) const handleOpenPreview = useCallback(() => { setIsPreviewOpen(true) }, []) @@ -53,12 +91,23 @@ export function AgentSkillItem({ {t('agentDetail.configure.skills.itemType')} + {!readOnly && ( + )} + /> + + {communityEditionBuildModeTip} + +

    {t('agentDetail.configure.build.empty.description')} diff --git a/web/features/agent-v2/agent-detail/configure/components/preview/chat-features-panel.tsx b/web/features/agent-v2/agent-detail/configure/components/preview/chat-features-panel.tsx index 61d7259d665..92d8c9025ae 100644 --- a/web/features/agent-v2/agent-detail/configure/components/preview/chat-features-panel.tsx +++ b/web/features/agent-v2/agent-detail/configure/components/preview/chat-features-panel.tsx @@ -1,6 +1,6 @@ 'use client' -import type { AgentSoulAppFeaturesConfig } from '@dify/contracts/api/console/agent/types.gen' +import type { AgentSoulAppFeaturesConfig, FileTransferMethod, FileType } from '@dify/contracts/api/console/agent/types.gen' import type { Features } from '@/app/components/base/features/types' import { useCallback, useMemo } from 'react' import { useTranslation } from 'react-i18next' @@ -37,6 +37,41 @@ const defaultFeatureState: Features = { annotationReply: { enabled: false }, } +const agentFileTypes = new Set(['audio', 'custom', 'document', 'image', 'video']) +const agentFileTransferMethods = new Set(['datasource_file', 'local_file', 'remote_url', 'tool_file']) + +function isAgentFileType(value: string): value is FileType { + return agentFileTypes.has(value) +} + +function isAgentFileTransferMethod(value: string): value is FileTransferMethod { + return agentFileTransferMethods.has(value) +} + +function toAgentFileTransferMethods(values?: readonly string[]): FileTransferMethod[] | undefined { + return values?.filter(isAgentFileTransferMethod) +} + +function toAgentFileUploadFeatureConfig(file: Features['file']): AgentSoulAppFeaturesConfig['file_upload'] { + if (!file) + return undefined + + const { allowed_file_types, allowed_file_upload_methods } = file + const fileUpload: Record = { ...file } + delete fileUpload.allowed_file_types + delete fileUpload.allowed_file_upload_methods + + return { + ...fileUpload, + ...(allowed_file_types + ? { allowed_file_types: allowed_file_types.filter(isAgentFileType) } + : {}), + ...(allowed_file_upload_methods + ? { allowed_file_upload_methods: toAgentFileTransferMethods(allowed_file_upload_methods) } + : {}), + } +} + function toPanelFeatures(appFeatures?: AgentSoulAppFeaturesConfig): Features { return { ...defaultFeatureState, @@ -65,7 +100,7 @@ function toAppFeatures(features: Features, appFeatures?: AgentSoulAppFeaturesCon speech_to_text: features.speech2text, retriever_resource: features.citation, sensitive_word_avoidance: features.moderation as AgentSoulAppFeaturesConfig['sensitive_word_avoidance'], - file_upload: features.file, + file_upload: toAgentFileUploadFeatureConfig(features.file), annotation_reply: features.annotationReply, } } diff --git a/web/features/agent-v2/agent-detail/configure/components/preview/working-directory-panel.tsx b/web/features/agent-v2/agent-detail/configure/components/preview/working-directory-panel.tsx index 7e327e9c276..3599a6958b5 100644 --- a/web/features/agent-v2/agent-detail/configure/components/preview/working-directory-panel.tsx +++ b/web/features/agent-v2/agent-detail/configure/components/preview/working-directory-panel.tsx @@ -6,9 +6,10 @@ import type { AgentFileNode } from '@/features/agent-v2/agent-composer/form-stat import { Dialog } from '@langgenius/dify-ui/dialog' import { Tooltip, TooltipContent, TooltipTrigger } from '@langgenius/dify-ui/tooltip' import { skipToken, useQueries, useQuery } from '@tanstack/react-query' -import { useState } from 'react' +import { useCallback, useState } from 'react' import { useTranslation } from 'react-i18next' import { consoleQuery } from '@/service/client' +import { downloadBlob } from '@/utils/download' import { getFileIconType } from '../orchestrate/files/file-icon' import { AgentSkillDetailDialog } from '../orchestrate/skills/detail-dialog' import { AgentWorkingDirectoryBreadcrumb } from './working-directory-breadcrumb' @@ -431,6 +432,20 @@ export function AgentWorkingDirectoryPanel({ retry: false, }) const isFileReadLoading = !!selectedWorkingDirectoryFile && fileReadQuery.isPending + const { data: fileReadData, refetch: refetchFileRead } = fileReadQuery + const handleDownloadFile = useCallback(async () => { + if (!selectedWorkingDirectoryFile) + return + + const readResult = fileReadData ?? (await refetchFileRead()).data + if (readResult?.binary || readResult?.text === undefined || readResult.text === null) + return + + downloadBlob({ + data: new Blob([readResult.text], { type: 'text/plain;charset=utf-8' }), + fileName: selectedWorkingDirectoryFile.name, + }) + }, [fileReadData, refetchFileRead, selectedWorkingDirectoryFile]) return (

    @@ -485,6 +500,9 @@ export function AgentWorkingDirectoryPanel({ isError: fileListQuery.isError || fileReadQuery.isError, isLoading: isFileListLoading || isFileReadLoading, }, + onDownloadFile: selectedWorkingDirectoryFile && !fileReadQuery.data?.binary + ? handleDownloadFile + : undefined, folderOpenState: ({ file }) => { const queryIndex = loadedFolderPathIndexes.get(file.id) const folderLoaded = queryIndex !== undefined && expandedFolderQueries[queryIndex]?.isSuccess diff --git a/web/features/agent-v2/agent-detail/configure/use-agent-configure-sync.ts b/web/features/agent-v2/agent-detail/configure/use-agent-configure-sync.ts index a3493702cf2..96ee788e0f9 100644 --- a/web/features/agent-v2/agent-detail/configure/use-agent-configure-sync.ts +++ b/web/features/agent-v2/agent-detail/configure/use-agent-configure-sync.ts @@ -237,6 +237,16 @@ export function useAgentConfigureSync({ return const draft = store.get(agentComposerDraftAtom) + const configSnapshot = formStateToAgentSoulConfig({ + baseConfig: baseConfigRef.current, + formState: draft, + currentModel: currentModelRef.current, + }) + if (!configSnapshot.model?.model_provider || !configSnapshot.model.model) { + toast.error(tCommon('modelProvider.selectModel')) + return + } + const knowledgeValidation = validateKnowledgeRetrievals(draft.knowledgeRetrievals) if (!knowledgeValidation.isValid) { toast.error(getKnowledgeValidationMessage(knowledgeValidation.firstIssue?.code) ?? tCommon('api.actionFailed')) @@ -247,11 +257,6 @@ export function useAgentConfigureSync({ setIsPublishInFlight(true) try { debouncedSaveDraft.cancel?.() - const configSnapshot = formStateToAgentSoulConfig({ - baseConfig: baseConfigRef.current, - formState: draft, - currentModel: currentModelRef.current, - }) const saved = await saveComposer({ configSnapshot, draftBaseline: draft, diff --git a/web/features/agent-v2/agent-detail/layout.tsx b/web/features/agent-v2/agent-detail/layout.tsx index b7247064516..1dad110f8af 100644 --- a/web/features/agent-v2/agent-detail/layout.tsx +++ b/web/features/agent-v2/agent-detail/layout.tsx @@ -2,8 +2,10 @@ import type { ReactNode } from 'react' import { useQuery } from '@tanstack/react-query' +import { useEffect } from 'react' import { useTranslation } from 'react-i18next' import useDocumentTitle from '@/hooks/use-document-title' +import { useRouter } from '@/next/navigation' import { consoleQuery } from '@/service/client' type AgentDetailLayoutProps = { @@ -11,11 +13,14 @@ type AgentDetailLayoutProps = { children: ReactNode } +const isNotFoundResponse = (error: unknown) => error instanceof Response && error.status === 404 + export function AgentDetailLayout({ agentId, children, }: AgentDetailLayoutProps) { const { t } = useTranslation('agentV2') + const router = useRouter() const agentQuery = useQuery(consoleQuery.agent.byAgentId.get.queryOptions({ input: { params: { @@ -23,9 +28,18 @@ export function AgentDetailLayout({ }, }, })) + const shouldRedirectToRoster = isNotFoundResponse(agentQuery.error) useDocumentTitle(agentQuery.data?.name ?? t('agentDetail.documentTitle')) + useEffect(() => { + if (shouldRedirectToRoster) + router.replace('/roster') + }, [router, shouldRedirectToRoster]) + + if (shouldRedirectToRoster) + return null + return (
    diff --git a/web/i18n/ar-TN/agent-v-2.json b/web/i18n/ar-TN/agent-v-2.json index c32f0b0dac2..56c75195270 100644 --- a/web/i18n/ar-TN/agent-v-2.json +++ b/web/i18n/ar-TN/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "الإعدادات المتقدمة", "agentDetail.configure.advancedSettings.toggle": "تبديل الإعدادات المتقدمة", + "agentDetail.configure.build.empty.communityEditionTip": "توفر Community Edition إدارة إصدارات للتكوينات المستخرجة، لكنها لا تدعم إدارة إصدارات نظام الملفات نفسه. كن حذراً في إجراءاتك ضمن وضع Build، لأن التغييرات التي تُجرى على نظام الملفات تحدث في الوقت الفعلي ولا يمكن دائماً التراجع عنها بشكل منظم. Community Edition ليست الخيار المثالي لخدمة جماهير خارجية واسعة.", "agentDetail.configure.build.empty.description": "صِف ما تريده وسيتم ملء النموذج على اليسار أثناء المحادثة.", "agentDetail.configure.build.empty.title": "ابنِ وكيلك عبر الدردشة", "agentDetail.configure.build.inputPlaceholder": "صِف ما يجب أن يفعله وكيلك", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "{{count}} تغييرات للتطبيق", "agentDetail.configure.buildDraft.discard": "تجاهل", "agentDetail.configure.buildDraft.modeBadge": "وضع البناء", - "agentDetail.configure.buildDraft.modeDescription": "أنت في وضع البناء. شكّل هذا الإعداد عبر الدردشة على اليمين، ثم طبّق.", + "agentDetail.configure.buildDraft.modeDescription": "أنت في وضع البناء. لا يمكن تحديث Configure في هذا الوضع إلا بواسطة الوكيل. شكّل هذا الإعداد عبر الدردشة على اليمين، ثم طبّق.", "agentDetail.configure.buildDraft.rewritten": "أُعيدت صياغته", "agentDetail.configure.buildDraft.title": "مسودة البناء", "agentDetail.configure.buildDraft.updated": "تم التحديث", "agentDetail.configure.chatFeatures.description": "شكّل تجربة الدردشة للمستخدم النهائي على Web app وأسطح الدردشة.", "agentDetail.configure.chatFeatures.title": "ميزات الدردشة", + "agentDetail.configure.communityEditionIsolationTip": "لا توفر Community Edition عزلاً صارماً لنظام الملفات بين المستخدمين النهائيين أو بين عمليات التشغيل. لا تعرض وكيل CE نفسه لعدة مستخدمين نهائيين مستقلين عندما يكون عزل البيانات أو الامتثال الصارم مطلوباً.", "agentDetail.configure.files.add": "إضافة ملف", "agentDetail.configure.files.buildNote.generated": "تم إنشاؤه", "agentDetail.configure.files.buildNote.richTooltip": "سجل الوكيل لما أعدّه في وضع Build. يقرأه في بداية كل محادثة إلى جانب Prompt الخاص بك. معرفة المزيد", "agentDetail.configure.files.buildNote.tooltip": "سجل الوكيل لما أعدّه في وضع Build. يقرأه في بداية كل محادثة إلى جانب Prompt الخاص بك. معرفة المزيد", + "agentDetail.configure.files.download": "تنزيل {{name}}", "agentDetail.configure.files.empty.description": "قم بتحميل المستندات التي يمكن للوكيل قراءتها، مثل المواصفات أو القوالب أو الإرشادات", "agentDetail.configure.files.empty.title": "لا توجد ملفات بعد", "agentDetail.configure.files.label": "الملفات", diff --git a/web/i18n/de-DE/agent-v-2.json b/web/i18n/de-DE/agent-v-2.json index ce433129037..b9c777680ce 100644 --- a/web/i18n/de-DE/agent-v-2.json +++ b/web/i18n/de-DE/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "Erweiterte Einstellungen", "agentDetail.configure.advancedSettings.toggle": "Erweiterte Einstellungen umschalten", + "agentDetail.configure.build.empty.communityEditionTip": "Community Edition bietet zwar Versionierung für extrahierte Konfigurationen, unterstützt aber keine Versionierung des Dateisystems selbst. Seien Sie bei Aktionen im Build-Modus vorsichtig, da Änderungen am Dateisystem in Echtzeit erfolgen und nicht immer sauber rückgängig gemacht werden können. Community Edition ist nicht die ideale Wahl, um große externe Zielgruppen zu bedienen.", "agentDetail.configure.build.empty.description": "Beschreiben Sie, was Sie möchten, und das Formular links wird während des Chats ausgefüllt.", "agentDetail.configure.build.empty.title": "Agent per Chat erstellen", "agentDetail.configure.build.inputPlaceholder": "Beschreiben Sie, was Ihr Agent tun soll", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "{{count}} Änderungen anzuwenden", "agentDetail.configure.buildDraft.discard": "Verwerfen", "agentDetail.configure.buildDraft.modeBadge": "Build-Modus", - "agentDetail.configure.buildDraft.modeDescription": "Sie sind im Build-Modus. Formen Sie diese Einrichtung über den Chat rechts und wenden Sie sie dann an.", + "agentDetail.configure.buildDraft.modeDescription": "Sie sind im Build-Modus. Configure kann in diesem Modus nur vom Agenten aktualisiert werden. Formen Sie diese Einrichtung über den Chat rechts und wenden Sie sie dann an.", "agentDetail.configure.buildDraft.rewritten": "Neu geschrieben", "agentDetail.configure.buildDraft.title": "Build-Entwurf", "agentDetail.configure.buildDraft.updated": "Aktualisiert", "agentDetail.configure.chatFeatures.description": "Gestalten Sie das Chat-Erlebnis für Endnutzer in Ihrer Webapp und in Chat-Oberflächen.", "agentDetail.configure.chatFeatures.title": "Chat-Funktionen", + "agentDetail.configure.communityEditionIsolationTip": "Community Edition bietet keine harte Dateisystemisolierung zwischen Endbenutzern oder Ausführungen. Stellen Sie denselben CE-Agenten nicht mehreren voneinander unabhängigen Endbenutzern bereit, wenn Datenisolierung oder strenge Compliance erforderlich ist.", "agentDetail.configure.files.add": "Datei hinzufügen", "agentDetail.configure.files.buildNote.generated": "Generiert", "agentDetail.configure.files.buildNote.richTooltip": "Die Aufzeichnung des Agenten darüber, was er im Build-Modus eingerichtet hat. Er liest sie zu Beginn jeder Unterhaltung zusammen mit Ihrem Prompt. Mehr erfahren", "agentDetail.configure.files.buildNote.tooltip": "Die Aufzeichnung des Agenten darüber, was er im Build-Modus eingerichtet hat. Er liest sie zu Beginn jeder Unterhaltung zusammen mit Ihrem Prompt. Mehr erfahren", + "agentDetail.configure.files.download": "{{name}} herunterladen", "agentDetail.configure.files.empty.description": "Laden Sie Dokumente hoch, die der Agent lesen kann, z. B. Spezifikationen, Vorlagen oder Richtlinien", "agentDetail.configure.files.empty.title": "Noch keine Dateien", "agentDetail.configure.files.label": "Dateien", diff --git a/web/i18n/en-US/agent-v-2.json b/web/i18n/en-US/agent-v-2.json index dd489624891..59c2e6f2172 100644 --- a/web/i18n/en-US/agent-v-2.json +++ b/web/i18n/en-US/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "Advanced Settings", "agentDetail.configure.advancedSettings.toggle": "Toggle advanced settings", + "agentDetail.configure.build.empty.communityEditionTip": "Community Edition, while offering versioning to the extracted configs, does not support versioning on the file system itself. Be careful with your actions in Build mode, as changes made to the file system happen in real time and cannot always be neatly reverted. Community Edition is not your ideal choice for serving mass external audiences.", "agentDetail.configure.build.empty.description": "Describe what you want and it fills in the form on the left as you go.", "agentDetail.configure.build.empty.title": "Build your agent by chatting", "agentDetail.configure.build.inputPlaceholder": "Describe what your agent should do", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "{{count}} changes to apply", "agentDetail.configure.buildDraft.discard": "Discard", "agentDetail.configure.buildDraft.modeBadge": "Build mode", - "agentDetail.configure.buildDraft.modeDescription": "You're in build mode. Shape this setup through the chat on the right, then Apply.", + "agentDetail.configure.buildDraft.modeDescription": "You’re in build mode. Configure can only be updated by the agent in this mode. Shape this setup through the chat on the right, then Apply.", "agentDetail.configure.buildDraft.rewritten": "Rewritten", "agentDetail.configure.buildDraft.title": "Build draft", "agentDetail.configure.buildDraft.updated": "Updated", "agentDetail.configure.chatFeatures.description": "Shape the end-user chat experience on your web app and chat surfaces.", "agentDetail.configure.chatFeatures.title": "Chat Features", + "agentDetail.configure.communityEditionIsolationTip": "Community Edition does not provide hard file system isolation between end users or runs. Do not expose the same CE agent to multiple independent end users where data isolation or strict compliance is required.", "agentDetail.configure.files.add": "Add file", "agentDetail.configure.files.buildNote.generated": "Generated", "agentDetail.configure.files.buildNote.richTooltip": "The agent's record of what it set up in Build mode. It reads this at the start of every conversation, alongside your Prompt. Learn more", "agentDetail.configure.files.buildNote.tooltip": "The agent's record of what it set up in Build mode. It reads this at the start of every conversation, alongside your Prompt. Learn more", + "agentDetail.configure.files.download": "Download {{name}}", "agentDetail.configure.files.empty.description": "Upload docs the agent can read, like specs, templates, or guidelines", "agentDetail.configure.files.empty.title": "No files yet", "agentDetail.configure.files.label": "Files", diff --git a/web/i18n/es-ES/agent-v-2.json b/web/i18n/es-ES/agent-v-2.json index 67394756bc7..9ec0bc588a4 100644 --- a/web/i18n/es-ES/agent-v-2.json +++ b/web/i18n/es-ES/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "Configuración avanzada", "agentDetail.configure.advancedSettings.toggle": "Alternar configuración avanzada", + "agentDetail.configure.build.empty.communityEditionTip": "Community Edition, aunque ofrece versionado para las configuraciones extraídas, no admite versionado del propio sistema de archivos. Ten cuidado con tus acciones en el modo Build, ya que los cambios realizados en el sistema de archivos ocurren en tiempo real y no siempre se pueden revertir limpiamente. Community Edition no es la opción ideal para atender audiencias externas masivas.", "agentDetail.configure.build.empty.description": "Describe lo que quieres y se irá completando el formulario de la izquierda mientras avanzas.", "agentDetail.configure.build.empty.title": "Crea tu agente chateando", "agentDetail.configure.build.inputPlaceholder": "Describe qué debe hacer tu agente", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "{{count}} cambios por aplicar", "agentDetail.configure.buildDraft.discard": "Descartar", "agentDetail.configure.buildDraft.modeBadge": "Modo build", - "agentDetail.configure.buildDraft.modeDescription": "Estás en modo build. Ajusta esta configuración con el chat de la derecha y luego aplica los cambios.", + "agentDetail.configure.buildDraft.modeDescription": "Estás en modo build. Configure solo puede ser actualizado por el agente en este modo. Ajusta esta configuración con el chat de la derecha y luego aplica los cambios.", "agentDetail.configure.buildDraft.rewritten": "Reescrito", "agentDetail.configure.buildDraft.title": "Borrador de compilación", "agentDetail.configure.buildDraft.updated": "Actualizado", "agentDetail.configure.chatFeatures.description": "Da forma a la experiencia de chat del usuario final en tu webapp y superficies de chat.", "agentDetail.configure.chatFeatures.title": "Funciones de chat", + "agentDetail.configure.communityEditionIsolationTip": "Community Edition no proporciona aislamiento estricto del sistema de archivos entre usuarios finales ni entre ejecuciones. No expongas el mismo agente CE a varios usuarios finales independientes cuando se requiera aislamiento de datos o cumplimiento estricto.", "agentDetail.configure.files.add": "Agregar archivo", "agentDetail.configure.files.buildNote.generated": "Generado", "agentDetail.configure.files.buildNote.richTooltip": "El registro del agente sobre lo que configuró en modo Build. Lo lee al inicio de cada conversación, junto con tu Prompt. Más información", "agentDetail.configure.files.buildNote.tooltip": "El registro del agente sobre lo que configuró en modo Build. Lo lee al inicio de cada conversación, junto con tu Prompt. Más información", + "agentDetail.configure.files.download": "Descargar {{name}}", "agentDetail.configure.files.empty.description": "Sube documentos que el agente pueda leer, como especificaciones, plantillas o guías", "agentDetail.configure.files.empty.title": "Aún no hay archivos", "agentDetail.configure.files.label": "Archivos", diff --git a/web/i18n/fa-IR/agent-v-2.json b/web/i18n/fa-IR/agent-v-2.json index 1dcab1bff12..d8565933495 100644 --- a/web/i18n/fa-IR/agent-v-2.json +++ b/web/i18n/fa-IR/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "تنظیمات پیشرفته", "agentDetail.configure.advancedSettings.toggle": "تغییر تنظیمات پیشرفته", + "agentDetail.configure.build.empty.communityEditionTip": "Community Edition با اینکه برای پیکربندی‌های استخراج‌شده نسخه‌بندی ارائه می‌کند، از نسخه‌بندی خود سیستم فایل پشتیبانی نمی‌کند. در حالت Build با اقدامات خود محتاط باشید، زیرا تغییرات روی سیستم فایل به‌صورت بلادرنگ انجام می‌شوند و همیشه نمی‌توان آن‌ها را به‌صورت تمیز بازگرداند. Community Edition انتخاب ایده‌آلی برای خدمت‌رسانی به مخاطبان خارجی گسترده نیست.", "agentDetail.configure.build.empty.description": "آنچه می‌خواهید را توضیح دهید تا فرم سمت چپ در طول گفتگو تکمیل شود.", "agentDetail.configure.build.empty.title": "عامل خود را با چت بسازید", "agentDetail.configure.build.inputPlaceholder": "توضیح دهید عامل شما باید چه کاری انجام دهد", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "{{count}} تغییر برای اعمال", "agentDetail.configure.buildDraft.discard": "رد کردن", "agentDetail.configure.buildDraft.modeBadge": "حالت ساخت", - "agentDetail.configure.buildDraft.modeDescription": "شما در حالت ساخت هستید. این تنظیمات را از طریق چت سمت راست شکل دهید، سپس اعمال کنید.", + "agentDetail.configure.buildDraft.modeDescription": "شما در حالت ساخت هستید. در این حالت Configure فقط می‌تواند توسط عامل به‌روزرسانی شود. این تنظیمات را از طریق چت سمت راست شکل دهید، سپس اعمال کنید.", "agentDetail.configure.buildDraft.rewritten": "بازنویسی شد", "agentDetail.configure.buildDraft.title": "پیش نویس ساخت", "agentDetail.configure.buildDraft.updated": "به‌روزرسانی شد", "agentDetail.configure.chatFeatures.description": "تجربه چت کاربر نهایی را در Web app و سطوح چت خود شکل دهید.", "agentDetail.configure.chatFeatures.title": "ویژگی‌های چت", + "agentDetail.configure.communityEditionIsolationTip": "Community Edition جداسازی سخت‌گیرانهٔ سیستم فایل را بین کاربران نهایی یا اجراها فراهم نمی‌کند. در جایی که جداسازی داده یا رعایت الزامات سخت‌گیرانه لازم است، همان عامل CE را در اختیار چند کاربر نهایی مستقل قرار ندهید.", "agentDetail.configure.files.add": "افزودن فایل", "agentDetail.configure.files.buildNote.generated": "تولید شده", "agentDetail.configure.files.buildNote.richTooltip": "رکورد عامل از چیزهایی که در حالت Build تنظیم کرده است. در آغاز هر گفتگو، آن را همراه با Prompt شما می‌خواند. بیشتر بدانید", "agentDetail.configure.files.buildNote.tooltip": "رکورد عامل از چیزهایی که در حالت Build تنظیم کرده است. در آغاز هر گفتگو، آن را همراه با Prompt شما می‌خواند. بیشتر بدانید", + "agentDetail.configure.files.download": "دانلود {{name}}", "agentDetail.configure.files.empty.description": "اسنادی را که عامل می‌تواند بخواند بارگذاری کنید، مانند مشخصات، قالب‌ها یا دستورالعمل‌ها", "agentDetail.configure.files.empty.title": "هنوز فایلی وجود ندارد", "agentDetail.configure.files.label": "فایل‌ها", diff --git a/web/i18n/fr-FR/agent-v-2.json b/web/i18n/fr-FR/agent-v-2.json index 697e50cca8c..1831035a2d0 100644 --- a/web/i18n/fr-FR/agent-v-2.json +++ b/web/i18n/fr-FR/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "Paramètres avancés", "agentDetail.configure.advancedSettings.toggle": "Basculer les paramètres avancés", + "agentDetail.configure.build.empty.communityEditionTip": "Community Edition, bien qu’elle offre la gestion des versions pour les configurations extraites, ne prend pas en charge la gestion des versions du système de fichiers lui-même. Soyez prudent avec vos actions en mode Build, car les modifications apportées au système de fichiers se produisent en temps réel et ne peuvent pas toujours être annulées proprement. Community Edition n’est pas le choix idéal pour servir un public externe massif.", "agentDetail.configure.build.empty.description": "Décrivez ce que vous voulez et le formulaire de gauche se remplit au fil de la conversation.", "agentDetail.configure.build.empty.title": "Créez votre agent par chat", "agentDetail.configure.build.inputPlaceholder": "Décrivez ce que votre agent doit faire", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "{{count}} modifications à appliquer", "agentDetail.configure.buildDraft.discard": "Ignorer", "agentDetail.configure.buildDraft.modeBadge": "Mode build", - "agentDetail.configure.buildDraft.modeDescription": "Vous êtes en mode build. Ajustez cette configuration avec le chat à droite, puis appliquez.", + "agentDetail.configure.buildDraft.modeDescription": "Vous êtes en mode build. Configure ne peut être mis à jour que par l’agent dans ce mode. Ajustez cette configuration avec le chat à droite, puis appliquez.", "agentDetail.configure.buildDraft.rewritten": "Réécrit", "agentDetail.configure.buildDraft.title": "Brouillon de build", "agentDetail.configure.buildDraft.updated": "Mis à jour", "agentDetail.configure.chatFeatures.description": "Façonnez l’expérience de chat de l’utilisateur final sur votre webapp et vos surfaces de chat.", "agentDetail.configure.chatFeatures.title": "Fonctionnalités de chat", + "agentDetail.configure.communityEditionIsolationTip": "Community Edition ne fournit pas d’isolation stricte du système de fichiers entre les utilisateurs finaux ni entre les exécutions. N’exposez pas le même agent CE à plusieurs utilisateurs finaux indépendants lorsque l’isolation des données ou une conformité stricte est requise.", "agentDetail.configure.files.add": "Ajouter un fichier", "agentDetail.configure.files.buildNote.generated": "Généré", "agentDetail.configure.files.buildNote.richTooltip": "Le registre de l'agent sur ce qu'il a configuré en mode Build. Il le lit au début de chaque conversation, avec votre Prompt. En savoir plus", "agentDetail.configure.files.buildNote.tooltip": "Le registre de l'agent sur ce qu'il a configuré en mode Build. Il le lit au début de chaque conversation, avec votre Prompt. En savoir plus", + "agentDetail.configure.files.download": "Télécharger {{name}}", "agentDetail.configure.files.empty.description": "Téléchargez des documents que l’agent peut lire, comme des spécifications, des modèles ou des directives", "agentDetail.configure.files.empty.title": "Pas encore de fichiers", "agentDetail.configure.files.label": "Fichiers", diff --git a/web/i18n/hi-IN/agent-v-2.json b/web/i18n/hi-IN/agent-v-2.json index 55acc170321..e605fab7e8b 100644 --- a/web/i18n/hi-IN/agent-v-2.json +++ b/web/i18n/hi-IN/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "उन्नत सेटिंग्स", "agentDetail.configure.advancedSettings.toggle": "उन्नत सेटिंग्स टॉगल करें", + "agentDetail.configure.build.empty.communityEditionTip": "Community Edition निकाले गए कॉन्फ़िगरेशन के लिए वर्शनिंग प्रदान करता है, लेकिन फ़ाइल सिस्टम पर स्वयं वर्शनिंग का समर्थन नहीं करता। Build मोड में अपनी कार्रवाइयों के साथ सावधान रहें, क्योंकि फ़ाइल सिस्टम में किए गए बदलाव वास्तविक समय में होते हैं और हमेशा साफ़-सुथरे ढंग से वापस नहीं किए जा सकते। बड़े पैमाने पर बाहरी दर्शकों को सेवा देने के लिए Community Edition आदर्श विकल्प नहीं है।", "agentDetail.configure.build.empty.description": "आप जो चाहते हैं उसका वर्णन करें और बाईं ओर का फ़ॉर्म बातचीत के साथ भरता जाएगा।", "agentDetail.configure.build.empty.title": "चैट करके अपना एजेंट बनाएं", "agentDetail.configure.build.inputPlaceholder": "बताएं कि आपका एजेंट क्या करे", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "लागू करने के लिए {{count}} बदलाव", "agentDetail.configure.buildDraft.discard": "छोड़ें", "agentDetail.configure.buildDraft.modeBadge": "बिल्ड मोड", - "agentDetail.configure.buildDraft.modeDescription": "आप बिल्ड मोड में हैं। दाईं ओर की चैट से इस सेटअप को आकार दें, फिर लागू करें।", + "agentDetail.configure.buildDraft.modeDescription": "आप बिल्ड मोड में हैं। इस मोड में Configure को केवल एजेंट ही अपडेट कर सकता है। दाईं ओर की चैट से इस सेटअप को आकार दें, फिर लागू करें।", "agentDetail.configure.buildDraft.rewritten": "फिर से लिखा गया", "agentDetail.configure.buildDraft.title": "बिल्ड ड्राफ्ट", "agentDetail.configure.buildDraft.updated": "अपडेट किया गया", "agentDetail.configure.chatFeatures.description": "अपने Web app और चैट सतहों पर अंतिम-उपयोगकर्ता चैट अनुभव को आकार दें।", "agentDetail.configure.chatFeatures.title": "चैट सुविधाएँ", + "agentDetail.configure.communityEditionIsolationTip": "Community Edition अंतिम उपयोगकर्ताओं या रन के बीच सख्त फ़ाइल सिस्टम आइसोलेशन प्रदान नहीं करता है। जहाँ डेटा आइसोलेशन या कड़े अनुपालन की आवश्यकता हो, वहाँ एक ही CE एजेंट को कई स्वतंत्र अंतिम उपयोगकर्ताओं के लिए उजागर न करें।", "agentDetail.configure.files.add": "फ़ाइल जोड़ें", "agentDetail.configure.files.buildNote.generated": "जनरेट किया गया", "agentDetail.configure.files.buildNote.richTooltip": "Build mode में एजेंट ने जो सेट अप किया उसका रिकॉर्ड। हर बातचीत की शुरुआत में यह इसे आपके Prompt के साथ पढ़ता है। और जानें", "agentDetail.configure.files.buildNote.tooltip": "Build mode में एजेंट ने जो सेट अप किया उसका रिकॉर्ड। हर बातचीत की शुरुआत में यह इसे आपके Prompt के साथ पढ़ता है। और जानें", + "agentDetail.configure.files.download": "{{name}} डाउनलोड करें", "agentDetail.configure.files.empty.description": "ऐसे दस्तावेज़ अपलोड करें जिन्हें एजेंट पढ़ सके, जैसे विनिर्देश, टेम्पलेट या दिशानिर्देश", "agentDetail.configure.files.empty.title": "अभी तक कोई फ़ाइल नहीं", "agentDetail.configure.files.label": "फ़ाइलें", diff --git a/web/i18n/id-ID/agent-v-2.json b/web/i18n/id-ID/agent-v-2.json index 983bea96f3e..33e473ead14 100644 --- a/web/i18n/id-ID/agent-v-2.json +++ b/web/i18n/id-ID/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "Pengaturan Lanjutan", "agentDetail.configure.advancedSettings.toggle": "Alihkan pengaturan lanjutan", + "agentDetail.configure.build.empty.communityEditionTip": "Community Edition, meskipun menawarkan versioning untuk konfigurasi yang diekstrak, tidak mendukung versioning pada sistem file itu sendiri. Berhati-hatilah dengan tindakan Anda dalam mode Build, karena perubahan pada sistem file terjadi secara real time dan tidak selalu dapat dikembalikan dengan rapi. Community Edition bukan pilihan ideal untuk melayani audiens eksternal dalam skala besar.", "agentDetail.configure.build.empty.description": "Jelaskan yang Anda inginkan dan formulir di kiri akan terisi seiring percakapan.", "agentDetail.configure.build.empty.title": "Bangun agen Anda lewat chat", "agentDetail.configure.build.inputPlaceholder": "Jelaskan apa yang harus dilakukan agen Anda", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "{{count}} perubahan untuk diterapkan", "agentDetail.configure.buildDraft.discard": "Buang", "agentDetail.configure.buildDraft.modeBadge": "Mode build", - "agentDetail.configure.buildDraft.modeDescription": "Anda berada dalam mode build. Bentuk pengaturan ini lewat chat di kanan, lalu terapkan.", + "agentDetail.configure.buildDraft.modeDescription": "Anda berada dalam mode build. Configure hanya dapat diperbarui oleh agen dalam mode ini. Bentuk pengaturan ini lewat chat di kanan, lalu terapkan.", "agentDetail.configure.buildDraft.rewritten": "Ditulis ulang", "agentDetail.configure.buildDraft.title": "Draf build", "agentDetail.configure.buildDraft.updated": "Diperbarui", "agentDetail.configure.chatFeatures.description": "Bentuk pengalaman chat pengguna akhir di Web app dan permukaan chat Anda.", "agentDetail.configure.chatFeatures.title": "Fitur Chat", + "agentDetail.configure.communityEditionIsolationTip": "Community Edition tidak menyediakan isolasi sistem file yang ketat antar pengguna akhir atau antar eksekusi. Jangan mengekspos agen CE yang sama kepada beberapa pengguna akhir independen ketika isolasi data atau kepatuhan ketat diperlukan.", "agentDetail.configure.files.add": "Tambahkan file", "agentDetail.configure.files.buildNote.generated": "Dihasilkan", "agentDetail.configure.files.buildNote.richTooltip": "Catatan agen tentang apa yang disiapkannya dalam mode Build. Agen membaca ini di awal setiap percakapan, bersama Prompt Anda. Pelajari selengkapnya", "agentDetail.configure.files.buildNote.tooltip": "Catatan agen tentang apa yang disiapkannya dalam mode Build. Agen membaca ini di awal setiap percakapan, bersama Prompt Anda. Pelajari selengkapnya", + "agentDetail.configure.files.download": "Unduh {{name}}", "agentDetail.configure.files.empty.description": "Unggah dokumen yang dapat dibaca agen, seperti spesifikasi, templat, atau pedoman", "agentDetail.configure.files.empty.title": "Belum ada file", "agentDetail.configure.files.label": "File", diff --git a/web/i18n/it-IT/agent-v-2.json b/web/i18n/it-IT/agent-v-2.json index 4f0bb90010a..d626ac1356e 100644 --- a/web/i18n/it-IT/agent-v-2.json +++ b/web/i18n/it-IT/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "Impostazioni avanzate", "agentDetail.configure.advancedSettings.toggle": "Attiva/disattiva impostazioni avanzate", + "agentDetail.configure.build.empty.communityEditionTip": "Community Edition, pur offrendo il versionamento delle configurazioni estratte, non supporta il versionamento del file system stesso. Fai attenzione alle azioni in modalità Build, perché le modifiche al file system avvengono in tempo reale e non sempre possono essere annullate in modo pulito. Community Edition non è la scelta ideale per servire un pubblico esterno di massa.", "agentDetail.configure.build.empty.description": "Descrivi ciò che vuoi e il modulo a sinistra verrà compilato man mano.", "agentDetail.configure.build.empty.title": "Crea il tuo agente con la chat", "agentDetail.configure.build.inputPlaceholder": "Descrivi cosa dovrebbe fare il tuo agente", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "{{count}} modifiche da applicare", "agentDetail.configure.buildDraft.discard": "Scarta", "agentDetail.configure.buildDraft.modeBadge": "Modalità build", - "agentDetail.configure.buildDraft.modeDescription": "Sei in modalità build. Modella questa configurazione tramite la chat a destra, poi applica.", + "agentDetail.configure.buildDraft.modeDescription": "Sei in modalità build. Configure può essere aggiornato solo dall’agente in questa modalità. Modella questa configurazione tramite la chat a destra, poi applica.", "agentDetail.configure.buildDraft.rewritten": "Riscritto", "agentDetail.configure.buildDraft.title": "Bozza di build", "agentDetail.configure.buildDraft.updated": "Aggiornato", "agentDetail.configure.chatFeatures.description": "Definisci l’esperienza di chat per l’utente finale sulla tua webapp e sulle superfici di chat.", "agentDetail.configure.chatFeatures.title": "Funzionalità chat", + "agentDetail.configure.communityEditionIsolationTip": "Community Edition non fornisce un isolamento rigido del file system tra utenti finali o esecuzioni. Non esporre lo stesso agente CE a più utenti finali indipendenti quando sono richiesti isolamento dei dati o conformità rigorosa.", "agentDetail.configure.files.add": "Aggiungi file", "agentDetail.configure.files.buildNote.generated": "Generato", "agentDetail.configure.files.buildNote.richTooltip": "Il registro dell'agente di ciò che ha configurato in modalità Build. Lo legge all'inizio di ogni conversazione, insieme al tuo Prompt. Scopri di più", "agentDetail.configure.files.buildNote.tooltip": "Il registro dell'agente di ciò che ha configurato in modalità Build. Lo legge all'inizio di ogni conversazione, insieme al tuo Prompt. Scopri di più", + "agentDetail.configure.files.download": "Scarica {{name}}", "agentDetail.configure.files.empty.description": "Carica documenti che l’agente possa leggere, come specifiche, modelli o linee guida", "agentDetail.configure.files.empty.title": "Nessun file al momento", "agentDetail.configure.files.label": "File", diff --git a/web/i18n/ja-JP/agent-v-2.json b/web/i18n/ja-JP/agent-v-2.json index a45519caf11..30279b3121c 100644 --- a/web/i18n/ja-JP/agent-v-2.json +++ b/web/i18n/ja-JP/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "詳細設定", "agentDetail.configure.advancedSettings.toggle": "詳細設定の表示を切り替え", + "agentDetail.configure.build.empty.communityEditionTip": "Community Edition は抽出された設定のバージョン管理を提供しますが、ファイルシステム自体のバージョン管理には対応していません。Build モードでの操作には注意してください。ファイルシステムへの変更はリアルタイムで発生し、常にきれいに元に戻せるとは限りません。Community Edition は大規模な外部利用者向けサービスには理想的な選択ではありません。", "agentDetail.configure.build.empty.description": "やりたいことを説明すると、左側のフォームが会話に合わせて入力されます。", "agentDetail.configure.build.empty.title": "チャットでエージェントを作成", "agentDetail.configure.build.inputPlaceholder": "エージェントに実行させたいことを説明", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "適用する変更 {{count}} 件", "agentDetail.configure.buildDraft.discard": "破棄", "agentDetail.configure.buildDraft.modeBadge": "ビルドモード", - "agentDetail.configure.buildDraft.modeDescription": "ビルドモードです。右側のチャットでこの設定を調整してから適用してください。", + "agentDetail.configure.buildDraft.modeDescription": "ビルドモードです。このモードでは、Configure はエージェントのみが更新できます。右側のチャットでこの設定を調整してから適用してください。", "agentDetail.configure.buildDraft.rewritten": "書き換え済み", "agentDetail.configure.buildDraft.title": "ビルドドラフト", "agentDetail.configure.buildDraft.updated": "更新済み", "agentDetail.configure.chatFeatures.description": "Web app やチャット画面でのエンドユーザー向けチャット体験を設定します。", "agentDetail.configure.chatFeatures.title": "チャット機能", + "agentDetail.configure.communityEditionIsolationTip": "Community Edition では、エンドユーザー間または実行間で厳密なファイルシステム分離は提供されません。データ分離や厳格なコンプライアンスが必要な場合は、同じ CE エージェントを複数の独立したエンドユーザーに公開しないでください。", "agentDetail.configure.files.add": "ファイルを追加", "agentDetail.configure.files.buildNote.generated": "生成済み", "agentDetail.configure.files.buildNote.richTooltip": "Build モードでエージェントが設定した内容の記録です。各会話の開始時に、Prompt と一緒にこれを読み取ります。詳しく見る", "agentDetail.configure.files.buildNote.tooltip": "Build モードでエージェントが設定した内容の記録です。各会話の開始時に、Prompt と一緒にこれを読み取ります。詳しく見る", + "agentDetail.configure.files.download": "{{name}} をダウンロード", "agentDetail.configure.files.empty.description": "仕様、テンプレート、ガイドラインなど、エージェントが読めるドキュメントをアップロード", "agentDetail.configure.files.empty.title": "ファイルはまだありません", "agentDetail.configure.files.label": "ファイル", diff --git a/web/i18n/ko-KR/agent-v-2.json b/web/i18n/ko-KR/agent-v-2.json index 08cf45d7c70..4f0ba8d8151 100644 --- a/web/i18n/ko-KR/agent-v-2.json +++ b/web/i18n/ko-KR/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "고급 설정", "agentDetail.configure.advancedSettings.toggle": "고급 설정 전환", + "agentDetail.configure.build.empty.communityEditionTip": "Community Edition은 추출된 설정에 대한 버전 관리는 제공하지만 파일 시스템 자체의 버전 관리는 지원하지 않습니다. Build 모드에서의 작업에 주의하세요. 파일 시스템 변경은 실시간으로 발생하며 항상 깔끔하게 되돌릴 수 있는 것은 아닙니다. Community Edition은 대규모 외부 사용자에게 서비스를 제공하기에 이상적인 선택이 아닙니다.", "agentDetail.configure.build.empty.description": "원하는 내용을 설명하면 대화에 맞춰 왼쪽 양식이 채워집니다.", "agentDetail.configure.build.empty.title": "채팅으로 에이전트 만들기", "agentDetail.configure.build.inputPlaceholder": "에이전트가 해야 할 일을 설명하세요", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "적용할 변경 사항 {{count}}개", "agentDetail.configure.buildDraft.discard": "폐기", "agentDetail.configure.buildDraft.modeBadge": "빌드 모드", - "agentDetail.configure.buildDraft.modeDescription": "빌드 모드입니다. 오른쪽 채팅으로 이 설정을 다듬은 뒤 적용하세요.", + "agentDetail.configure.buildDraft.modeDescription": "빌드 모드입니다. 이 모드에서는 에이전트만 Configure를 업데이트할 수 있습니다. 오른쪽 채팅으로 이 설정을 다듬은 뒤 적용하세요.", "agentDetail.configure.buildDraft.rewritten": "다시 작성됨", "agentDetail.configure.buildDraft.title": "빌드 초안", "agentDetail.configure.buildDraft.updated": "업데이트됨", "agentDetail.configure.chatFeatures.description": "Web app 및 채팅 화면에서의 최종 사용자 채팅 경험을 구성합니다.", "agentDetail.configure.chatFeatures.title": "채팅 기능", + "agentDetail.configure.communityEditionIsolationTip": "Community Edition은 최종 사용자 간 또는 실행 간에 강력한 파일 시스템 격리를 제공하지 않습니다. 데이터 격리나 엄격한 규정 준수가 필요한 경우 동일한 CE 에이전트를 여러 독립 최종 사용자에게 노출하지 마세요.", "agentDetail.configure.files.add": "파일 추가", "agentDetail.configure.files.buildNote.generated": "생성됨", "agentDetail.configure.files.buildNote.richTooltip": "에이전트가 Build mode에서 설정한 내용의 기록입니다. 모든 대화 시작 시 Prompt와 함께 이 기록을 읽습니다. 자세히 알아보기", "agentDetail.configure.files.buildNote.tooltip": "에이전트가 Build mode에서 설정한 내용의 기록입니다. 모든 대화 시작 시 Prompt와 함께 이 기록을 읽습니다. 자세히 알아보기", + "agentDetail.configure.files.download": "{{name}} 다운로드", "agentDetail.configure.files.empty.description": "사양, 템플릿, 가이드라인 등 에이전트가 읽을 수 있는 문서를 업로드하세요", "agentDetail.configure.files.empty.title": "아직 파일이 없습니다", "agentDetail.configure.files.label": "파일", diff --git a/web/i18n/nl-NL/agent-v-2.json b/web/i18n/nl-NL/agent-v-2.json index d4d48eb77e6..5fe85b681df 100644 --- a/web/i18n/nl-NL/agent-v-2.json +++ b/web/i18n/nl-NL/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "Geavanceerde instellingen", "agentDetail.configure.advancedSettings.toggle": "Geavanceerde instellingen in/uit", + "agentDetail.configure.build.empty.communityEditionTip": "Community Edition biedt wel versiebeheer voor geëxtraheerde configuraties, maar ondersteunt geen versiebeheer van het bestandssysteem zelf. Wees voorzichtig met je acties in de Build-modus, omdat wijzigingen aan het bestandssysteem in realtime plaatsvinden en niet altijd netjes kunnen worden teruggedraaid. Community Edition is niet de ideale keuze voor het bedienen van grote externe doelgroepen.", "agentDetail.configure.build.empty.description": "Beschrijf wat je wilt en het formulier links wordt gaandeweg ingevuld.", "agentDetail.configure.build.empty.title": "Bouw je agent via chat", "agentDetail.configure.build.inputPlaceholder": "Beschrijf wat je agent moet doen", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "{{count}} wijzigingen om toe te passen", "agentDetail.configure.buildDraft.discard": "Negeren", "agentDetail.configure.buildDraft.modeBadge": "Buildmodus", - "agentDetail.configure.buildDraft.modeDescription": "Je bent in buildmodus. Werk deze configuratie bij via de chat rechts en pas daarna toe.", + "agentDetail.configure.buildDraft.modeDescription": "Je bent in buildmodus. Configure kan in deze modus alleen door de agent worden bijgewerkt. Werk deze configuratie bij via de chat rechts en pas daarna toe.", "agentDetail.configure.buildDraft.rewritten": "Herschreven", "agentDetail.configure.buildDraft.title": "Buildconcept", "agentDetail.configure.buildDraft.updated": "Bijgewerkt", "agentDetail.configure.chatFeatures.description": "Geef vorm aan de chatervaring voor eindgebruikers in je webapp en chatoppervlakken.", "agentDetail.configure.chatFeatures.title": "Chatfuncties", + "agentDetail.configure.communityEditionIsolationTip": "Community Edition biedt geen harde bestandssysteemisolatie tussen eindgebruikers of runs. Stel dezelfde CE-agent niet beschikbaar aan meerdere onafhankelijke eindgebruikers wanneer gegevensisolatie of strikte compliance vereist is.", "agentDetail.configure.files.add": "Bestand toevoegen", "agentDetail.configure.files.buildNote.generated": "Gegenereerd", "agentDetail.configure.files.buildNote.richTooltip": "Het verslag van de agent van wat hij in de Build-modus heeft ingesteld. Hij leest dit aan het begin van elk gesprek, samen met uw Prompt. Meer informatie", "agentDetail.configure.files.buildNote.tooltip": "Het verslag van de agent van wat hij in de Build-modus heeft ingesteld. Hij leest dit aan het begin van elk gesprek, samen met uw Prompt. Meer informatie", + "agentDetail.configure.files.download": "{{name}} downloaden", "agentDetail.configure.files.empty.description": "Upload documenten die de agent kan lezen, zoals specificaties, sjablonen of richtlijnen", "agentDetail.configure.files.empty.title": "Nog geen bestanden", "agentDetail.configure.files.label": "Bestanden", diff --git a/web/i18n/pl-PL/agent-v-2.json b/web/i18n/pl-PL/agent-v-2.json index fe4c96745d1..5492b29fcf9 100644 --- a/web/i18n/pl-PL/agent-v-2.json +++ b/web/i18n/pl-PL/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "Ustawienia zaawansowane", "agentDetail.configure.advancedSettings.toggle": "Przełącz ustawienia zaawansowane", + "agentDetail.configure.build.empty.communityEditionTip": "Community Edition oferuje wersjonowanie wyodrębnionych konfiguracji, ale nie obsługuje wersjonowania samego systemu plików. Zachowaj ostrożność podczas działań w trybie Build, ponieważ zmiany w systemie plików zachodzą w czasie rzeczywistym i nie zawsze można je czysto cofnąć. Community Edition nie jest idealnym wyborem do obsługi masowej publiczności zewnętrznej.", "agentDetail.configure.build.empty.description": "Opisz, czego chcesz, a formularz po lewej będzie wypełniany w trakcie rozmowy.", "agentDetail.configure.build.empty.title": "Buduj agenta przez czat", "agentDetail.configure.build.inputPlaceholder": "Opisz, co agent ma robić", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "{{count}} zmiany do zastosowania", "agentDetail.configure.buildDraft.discard": "Odrzuć", "agentDetail.configure.buildDraft.modeBadge": "Tryb budowania", - "agentDetail.configure.buildDraft.modeDescription": "Jesteś w trybie budowania. Dostosuj tę konfigurację przez czat po prawej, a następnie zastosuj.", + "agentDetail.configure.buildDraft.modeDescription": "Jesteś w trybie budowania. W tym trybie Configure może być aktualizowane tylko przez agenta. Dostosuj tę konfigurację przez czat po prawej, a następnie zastosuj.", "agentDetail.configure.buildDraft.rewritten": "Przepisano", "agentDetail.configure.buildDraft.title": "Szkic budowania", "agentDetail.configure.buildDraft.updated": "Zaktualizowano", "agentDetail.configure.chatFeatures.description": "Ukształtuj doświadczenie czatu użytkownika końcowego w aplikacji webowej i powierzchniach czatu.", "agentDetail.configure.chatFeatures.title": "Funkcje czatu", + "agentDetail.configure.communityEditionIsolationTip": "Community Edition nie zapewnia twardej izolacji systemu plików między użytkownikami końcowymi ani uruchomieniami. Nie udostępniaj tego samego agenta CE wielu niezależnym użytkownikom końcowym, gdy wymagana jest izolacja danych lub ścisła zgodność.", "agentDetail.configure.files.add": "Dodaj plik", "agentDetail.configure.files.buildNote.generated": "Wygenerowano", "agentDetail.configure.files.buildNote.richTooltip": "Zapis agenta tego, co skonfigurował w trybie Build. Odczytuje go na początku każdej rozmowy razem z Twoim Promptem. Dowiedz się więcej", "agentDetail.configure.files.buildNote.tooltip": "Zapis agenta tego, co skonfigurował w trybie Build. Odczytuje go na początku każdej rozmowy razem z Twoim Promptem. Dowiedz się więcej", + "agentDetail.configure.files.download": "Pobierz {{name}}", "agentDetail.configure.files.empty.description": "Prześlij dokumenty, które agent może czytać, np. specyfikacje, szablony lub wytyczne", "agentDetail.configure.files.empty.title": "Brak plików", "agentDetail.configure.files.label": "Pliki", diff --git a/web/i18n/pt-BR/agent-v-2.json b/web/i18n/pt-BR/agent-v-2.json index 39671d80b88..867ef4b8e6b 100644 --- a/web/i18n/pt-BR/agent-v-2.json +++ b/web/i18n/pt-BR/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "Configurações avançadas", "agentDetail.configure.advancedSettings.toggle": "Alternar configurações avançadas", + "agentDetail.configure.build.empty.communityEditionTip": "A Community Edition, embora ofereça versionamento para as configurações extraídas, não oferece suporte a versionamento do próprio sistema de arquivos. Tenha cuidado com suas ações no modo Build, pois as alterações feitas no sistema de arquivos acontecem em tempo real e nem sempre podem ser revertidas de forma limpa. A Community Edition não é a escolha ideal para atender grandes públicos externos.", "agentDetail.configure.build.empty.description": "Descreva o que você quer e o formulário à esquerda será preenchido conforme a conversa avança.", "agentDetail.configure.build.empty.title": "Crie seu agente conversando", "agentDetail.configure.build.inputPlaceholder": "Descreva o que seu agente deve fazer", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "{{count}} alterações para aplicar", "agentDetail.configure.buildDraft.discard": "Descartar", "agentDetail.configure.buildDraft.modeBadge": "Modo build", - "agentDetail.configure.buildDraft.modeDescription": "Você está no modo build. Ajuste esta configuração pelo chat à direita e depois aplique.", + "agentDetail.configure.buildDraft.modeDescription": "Você está no modo build. O Configure só pode ser atualizado pelo agente neste modo. Ajuste esta configuração pelo chat à direita e depois aplique.", "agentDetail.configure.buildDraft.rewritten": "Reescrito", "agentDetail.configure.buildDraft.title": "Rascunho de build", "agentDetail.configure.buildDraft.updated": "Atualizado", "agentDetail.configure.chatFeatures.description": "Modele a experiência de chat do usuário final no seu webapp e superfícies de chat.", "agentDetail.configure.chatFeatures.title": "Recursos de chat", + "agentDetail.configure.communityEditionIsolationTip": "A Community Edition não fornece isolamento rígido do sistema de arquivos entre usuários finais ou execuções. Não exponha o mesmo agente CE a vários usuários finais independentes quando isolamento de dados ou conformidade rigorosa forem necessários.", "agentDetail.configure.files.add": "Adicionar arquivo", "agentDetail.configure.files.buildNote.generated": "Gerado", "agentDetail.configure.files.buildNote.richTooltip": "O registro do agente do que ele configurou no modo Build. Ele lê isso no início de cada conversa, junto com seu Prompt. Saiba mais", "agentDetail.configure.files.buildNote.tooltip": "O registro do agente do que ele configurou no modo Build. Ele lê isso no início de cada conversa, junto com seu Prompt. Saiba mais", + "agentDetail.configure.files.download": "Baixar {{name}}", "agentDetail.configure.files.empty.description": "Envie documentos que o agente possa ler, como especificações, modelos ou diretrizes", "agentDetail.configure.files.empty.title": "Ainda não há arquivos", "agentDetail.configure.files.label": "Arquivos", diff --git a/web/i18n/ro-RO/agent-v-2.json b/web/i18n/ro-RO/agent-v-2.json index 7390fe70517..d56d9c17b33 100644 --- a/web/i18n/ro-RO/agent-v-2.json +++ b/web/i18n/ro-RO/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "Setări avansate", "agentDetail.configure.advancedSettings.toggle": "Comută setările avansate", + "agentDetail.configure.build.empty.communityEditionTip": "Community Edition, deși oferă versionare pentru configurațiile extrase, nu acceptă versionarea sistemului de fișiere în sine. Fii atent la acțiunile din modul Build, deoarece modificările făcute în sistemul de fișiere au loc în timp real și nu pot fi întotdeauna anulate curat. Community Edition nu este alegerea ideală pentru servirea unor audiențe externe numeroase.", "agentDetail.configure.build.empty.description": "Descrie ce dorești, iar formularul din stânga se completează pe măsură ce avansezi.", "agentDetail.configure.build.empty.title": "Construiește agentul prin chat", "agentDetail.configure.build.inputPlaceholder": "Descrie ce ar trebui să facă agentul tău", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "{{count}} modificări de aplicat", "agentDetail.configure.buildDraft.discard": "Renunță", "agentDetail.configure.buildDraft.modeBadge": "Mod build", - "agentDetail.configure.buildDraft.modeDescription": "Ești în modul build. Ajustează această configurare prin chatul din dreapta, apoi aplică.", + "agentDetail.configure.buildDraft.modeDescription": "Ești în modul build. Configure poate fi actualizat doar de agent în acest mod. Ajustează această configurare prin chatul din dreapta, apoi aplică.", "agentDetail.configure.buildDraft.rewritten": "Rescris", "agentDetail.configure.buildDraft.title": "Schiță de build", "agentDetail.configure.buildDraft.updated": "Actualizat", "agentDetail.configure.chatFeatures.description": "Modelează experiența de chat a utilizatorului final pe webapp-ul tău și pe suprafețele de chat.", "agentDetail.configure.chatFeatures.title": "Funcții de chat", + "agentDetail.configure.communityEditionIsolationTip": "Community Edition nu oferă izolare strictă a sistemului de fișiere între utilizatorii finali sau între rulări. Nu expune același agent CE către mai mulți utilizatori finali independenți atunci când este necesară izolarea datelor sau conformitatea strictă.", "agentDetail.configure.files.add": "Adaugă fișier", "agentDetail.configure.files.buildNote.generated": "Generat", "agentDetail.configure.files.buildNote.richTooltip": "Înregistrarea agentului despre ce a configurat în modul Build. O citește la începutul fiecărei conversații, împreună cu Promptul dvs. Aflați mai multe", "agentDetail.configure.files.buildNote.tooltip": "Înregistrarea agentului despre ce a configurat în modul Build. O citește la începutul fiecărei conversații, împreună cu Promptul dvs. Aflați mai multe", + "agentDetail.configure.files.download": "Descarcă {{name}}", "agentDetail.configure.files.empty.description": "Încarcă documente pe care agentul le poate citi, precum specificații, șabloane sau ghiduri", "agentDetail.configure.files.empty.title": "Niciun fișier încă", "agentDetail.configure.files.label": "Fișiere", diff --git a/web/i18n/ru-RU/agent-v-2.json b/web/i18n/ru-RU/agent-v-2.json index c3378ae51fd..62ff29d2d05 100644 --- a/web/i18n/ru-RU/agent-v-2.json +++ b/web/i18n/ru-RU/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "Расширенные настройки", "agentDetail.configure.advancedSettings.toggle": "Переключить расширенные настройки", + "agentDetail.configure.build.empty.communityEditionTip": "Community Edition предоставляет версионирование извлеченных конфигураций, но не поддерживает версионирование самой файловой системы. Будьте осторожны с действиями в режиме Build: изменения файловой системы происходят в реальном времени и не всегда могут быть аккуратно отменены. Community Edition не является идеальным выбором для обслуживания массовой внешней аудитории.", "agentDetail.configure.build.empty.description": "Опишите, что вам нужно, и форма слева будет заполняться по ходу диалога.", "agentDetail.configure.build.empty.title": "Создайте агента в чате", "agentDetail.configure.build.inputPlaceholder": "Опишите, что должен делать агент", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "{{count}} изменений для применения", "agentDetail.configure.buildDraft.discard": "Отменить", "agentDetail.configure.buildDraft.modeBadge": "Режим сборки", - "agentDetail.configure.buildDraft.modeDescription": "Вы в режиме сборки. Настройте эту конфигурацию через чат справа, затем примените изменения.", + "agentDetail.configure.buildDraft.modeDescription": "Вы в режиме сборки. В этом режиме Configure может обновлять только агент. Настройте эту конфигурацию через чат справа, затем примените изменения.", "agentDetail.configure.buildDraft.rewritten": "Переписано", "agentDetail.configure.buildDraft.title": "Черновик сборки", "agentDetail.configure.buildDraft.updated": "Обновлено", "agentDetail.configure.chatFeatures.description": "Настройте чат-опыт конечного пользователя в вашем веб-приложении и чат-поверхностях.", "agentDetail.configure.chatFeatures.title": "Функции чата", + "agentDetail.configure.communityEditionIsolationTip": "Community Edition не обеспечивает жесткую изоляцию файловой системы между конечными пользователями или запусками. Не предоставляйте один и тот же CE-агент нескольким независимым конечным пользователям, если требуется изоляция данных или строгое соответствие требованиям.", "agentDetail.configure.files.add": "Добавить файл", "agentDetail.configure.files.buildNote.generated": "Сгенерировано", "agentDetail.configure.files.buildNote.richTooltip": "Запись агента о том, что он настроил в режиме Build. Он читает ее в начале каждого разговора вместе с вашим Prompt. Подробнее", "agentDetail.configure.files.buildNote.tooltip": "Запись агента о том, что он настроил в режиме Build. Он читает ее в начале каждого разговора вместе с вашим Prompt. Подробнее", + "agentDetail.configure.files.download": "Скачать {{name}}", "agentDetail.configure.files.empty.description": "Загрузите документы, которые может прочитать агент, например спецификации, шаблоны или руководства", "agentDetail.configure.files.empty.title": "Пока нет файлов", "agentDetail.configure.files.label": "Файлы", diff --git a/web/i18n/sl-SI/agent-v-2.json b/web/i18n/sl-SI/agent-v-2.json index 533f2813f30..900a6ed6075 100644 --- a/web/i18n/sl-SI/agent-v-2.json +++ b/web/i18n/sl-SI/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "Napredne nastavitve", "agentDetail.configure.advancedSettings.toggle": "Preklopi napredne nastavitve", + "agentDetail.configure.build.empty.communityEditionTip": "Community Edition sicer ponuja različice za izvlečene konfiguracije, vendar ne podpira različic samega datotečnega sistema. Pri dejanjih v načinu Build bodite previdni, saj se spremembe datotečnega sistema zgodijo v realnem času in jih ni vedno mogoče lepo razveljaviti. Community Edition ni idealna izbira za množično zunanjo publiko.", "agentDetail.configure.build.empty.description": "Opišite, kaj želite, in obrazec na levi se bo sproti izpolnjeval.", "agentDetail.configure.build.empty.title": "Izdelajte agenta s klepetom", "agentDetail.configure.build.inputPlaceholder": "Opišite, kaj naj vaš agent počne", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "{{count}} spremembe za uveljavitev", "agentDetail.configure.buildDraft.discard": "Zavrzi", "agentDetail.configure.buildDraft.modeBadge": "Način gradnje", - "agentDetail.configure.buildDraft.modeDescription": "Ste v načinu gradnje. Nastavitev oblikujte s klepetom na desni, nato jo uporabite.", + "agentDetail.configure.buildDraft.modeDescription": "Ste v načinu gradnje. Configure lahko v tem načinu posodobi samo agent. Nastavitev oblikujte s klepetom na desni, nato jo uporabite.", "agentDetail.configure.buildDraft.rewritten": "Prepisano", "agentDetail.configure.buildDraft.title": "Osnutek gradnje", "agentDetail.configure.buildDraft.updated": "Posodobljeno", "agentDetail.configure.chatFeatures.description": "Oblikujte uporabniško izkušnjo klepeta v vaši spletni aplikaciji in klepetalnih površinah.", "agentDetail.configure.chatFeatures.title": "Funkcije klepeta", + "agentDetail.configure.communityEditionIsolationTip": "Community Edition ne zagotavlja stroge izolacije datotečnega sistema med končnimi uporabniki ali zagoni. Istega agenta CE ne izpostavljajte več neodvisnim končnim uporabnikom, kadar sta potrebni izolacija podatkov ali stroga skladnost.", "agentDetail.configure.files.add": "Dodaj datoteko", "agentDetail.configure.files.buildNote.generated": "Ustvarjeno", "agentDetail.configure.files.buildNote.richTooltip": "Agentov zapis tega, kar je nastavil v načinu Build. Prebere ga na začetku vsakega pogovora skupaj z vašim Promptom. Več informacij", "agentDetail.configure.files.buildNote.tooltip": "Agentov zapis tega, kar je nastavil v načinu Build. Prebere ga na začetku vsakega pogovora skupaj z vašim Promptom. Več informacij", + "agentDetail.configure.files.download": "Prenesi {{name}}", "agentDetail.configure.files.empty.description": "Naložite dokumente, ki jih lahko agent bere, npr. specifikacije, predloge ali smernice", "agentDetail.configure.files.empty.title": "Še ni datotek", "agentDetail.configure.files.label": "Datoteke", diff --git a/web/i18n/th-TH/agent-v-2.json b/web/i18n/th-TH/agent-v-2.json index 1683bb55939..1042946c10a 100644 --- a/web/i18n/th-TH/agent-v-2.json +++ b/web/i18n/th-TH/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "การตั้งค่าขั้นสูง", "agentDetail.configure.advancedSettings.toggle": "สลับการตั้งค่าขั้นสูง", + "agentDetail.configure.build.empty.communityEditionTip": "Community Edition แม้จะมีการจัดการเวอร์ชันสำหรับการกำหนดค่าที่แยกออกมา แต่ไม่รองรับการจัดการเวอร์ชันของระบบไฟล์เอง โปรดระมัดระวังการดำเนินการในโหมด Build เนื่องจากการเปลี่ยนแปลงระบบไฟล์เกิดขึ้นแบบเรียลไทม์และอาจไม่สามารถย้อนกลับได้อย่างเรียบร้อยเสมอไป Community Edition ไม่ใช่ตัวเลือกที่เหมาะสมที่สุดสำหรับการให้บริการผู้ชมภายนอกจำนวนมาก", "agentDetail.configure.build.empty.description": "อธิบายสิ่งที่คุณต้องการ แล้วแบบฟอร์มด้านซ้ายจะถูกกรอกไปพร้อมกับการสนทนา", "agentDetail.configure.build.empty.title": "สร้างเอเจนต์ของคุณด้วยการแชท", "agentDetail.configure.build.inputPlaceholder": "อธิบายว่าเอเจนต์ของคุณควรทำอะไร", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "{{count}} การเปลี่ยนแปลงที่ต้องนำไปใช้", "agentDetail.configure.buildDraft.discard": "ทิ้ง", "agentDetail.configure.buildDraft.modeBadge": "โหมด Build", - "agentDetail.configure.buildDraft.modeDescription": "คุณอยู่ในโหมด Build ปรับแต่งการตั้งค่านี้ผ่านแชททางขวา แล้วกดนำไปใช้", + "agentDetail.configure.buildDraft.modeDescription": "คุณอยู่ในโหมด Build ในโหมดนี้ Configure จะอัปเดตได้โดยเอเจนต์เท่านั้น ปรับแต่งการตั้งค่านี้ผ่านแชททางขวา แล้วกดนำไปใช้", "agentDetail.configure.buildDraft.rewritten": "เขียนใหม่แล้ว", "agentDetail.configure.buildDraft.title": "ฉบับร่าง Build", "agentDetail.configure.buildDraft.updated": "อัปเดตแล้ว", "agentDetail.configure.chatFeatures.description": "กำหนดประสบการณ์การแชทของผู้ใช้ปลายทางบน Web app และหน้าจอแชท", "agentDetail.configure.chatFeatures.title": "ฟีเจอร์แชท", + "agentDetail.configure.communityEditionIsolationTip": "Community Edition ไม่มีการแยกระบบไฟล์อย่างเข้มงวดระหว่างผู้ใช้ปลายทางหรือระหว่างการรัน อย่าเปิดเผยเอเจนต์ CE ตัวเดียวกันให้กับผู้ใช้ปลายทางอิสระหลายรายเมื่อจำเป็นต้องมีการแยกข้อมูลหรือการปฏิบัติตามข้อกำหนดอย่างเข้มงวด", "agentDetail.configure.files.add": "เพิ่มไฟล์", "agentDetail.configure.files.buildNote.generated": "สร้างแล้ว", "agentDetail.configure.files.buildNote.richTooltip": "บันทึกของ agent เกี่ยวกับสิ่งที่ตั้งค่าไว้ในโหมด Build โดยจะอ่านสิ่งนี้ตอนเริ่มทุกบทสนทนา พร้อมกับ Prompt ของคุณ เรียนรู้เพิ่มเติม", "agentDetail.configure.files.buildNote.tooltip": "บันทึกของ agent เกี่ยวกับสิ่งที่ตั้งค่าไว้ในโหมด Build โดยจะอ่านสิ่งนี้ตอนเริ่มทุกบทสนทนา พร้อมกับ Prompt ของคุณ เรียนรู้เพิ่มเติม", + "agentDetail.configure.files.download": "ดาวน์โหลด {{name}}", "agentDetail.configure.files.empty.description": "อัปโหลดเอกสารที่ตัวแทนสามารถอ่านได้ เช่น ข้อกำหนด เทมเพลต หรือแนวทาง", "agentDetail.configure.files.empty.title": "ยังไม่มีไฟล์", "agentDetail.configure.files.label": "ไฟล์", diff --git a/web/i18n/tr-TR/agent-v-2.json b/web/i18n/tr-TR/agent-v-2.json index d15a40df8fc..ea90af5e5d9 100644 --- a/web/i18n/tr-TR/agent-v-2.json +++ b/web/i18n/tr-TR/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "Gelişmiş Ayarlar", "agentDetail.configure.advancedSettings.toggle": "Gelişmiş ayarları değiştir", + "agentDetail.configure.build.empty.communityEditionTip": "Community Edition, çıkarılan yapılandırmalar için sürümleme sunsa da dosya sisteminin kendisi için sürümleme desteklemez. Build modundaki işlemlerinizde dikkatli olun; dosya sisteminde yapılan değişiklikler gerçek zamanlı gerçekleşir ve her zaman temiz bir şekilde geri alınamayabilir. Community Edition, kitlesel dış hedef kitlelere hizmet vermek için ideal seçiminiz değildir.", "agentDetail.configure.build.empty.description": "Ne istediğinizi açıklayın; soldaki form ilerledikçe doldurulur.", "agentDetail.configure.build.empty.title": "Aracınızı sohbet ederek oluşturun", "agentDetail.configure.build.inputPlaceholder": "Aracınızın ne yapması gerektiğini açıklayın", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "Uygulanacak {{count}} değişiklik", "agentDetail.configure.buildDraft.discard": "Vazgeç", "agentDetail.configure.buildDraft.modeBadge": "Build modu", - "agentDetail.configure.buildDraft.modeDescription": "Build modundasınız. Bu kurulumu sağdaki sohbetle şekillendirin, ardından uygulayın.", + "agentDetail.configure.buildDraft.modeDescription": "Build modundasınız. Bu modda Configure yalnızca ajan tarafından güncellenebilir. Bu kurulumu sağdaki sohbetle şekillendirin, ardından uygulayın.", "agentDetail.configure.buildDraft.rewritten": "Yeniden yazıldı", "agentDetail.configure.buildDraft.title": "Build taslağı", "agentDetail.configure.buildDraft.updated": "Güncellendi", "agentDetail.configure.chatFeatures.description": "Web app ve sohbet yüzeylerinizde son kullanıcı sohbet deneyimini şekillendirin.", "agentDetail.configure.chatFeatures.title": "Sohbet Özellikleri", + "agentDetail.configure.communityEditionIsolationTip": "Community Edition, son kullanıcılar veya çalıştırmalar arasında katı dosya sistemi yalıtımı sağlamaz. Veri yalıtımı veya sıkı uyumluluk gerektiğinde aynı CE ajanını birden fazla bağımsız son kullanıcıya açmayın.", "agentDetail.configure.files.add": "Dosya ekle", "agentDetail.configure.files.buildNote.generated": "Oluşturuldu", "agentDetail.configure.files.buildNote.richTooltip": "Agent'ın Build mode'da kurduğu şeylerin kaydı. Her konuşmanın başında bunu Prompt'unuzla birlikte okur. Daha fazla bilgi", "agentDetail.configure.files.buildNote.tooltip": "Agent'ın Build mode'da kurduğu şeylerin kaydı. Her konuşmanın başında bunu Prompt'unuzla birlikte okur. Daha fazla bilgi", + "agentDetail.configure.files.download": "{{name}} indir", "agentDetail.configure.files.empty.description": "Ajanın okuyabileceği belgeleri yükleyin, örneğin spesifikasyonlar, şablonlar veya yönergeler", "agentDetail.configure.files.empty.title": "Henüz dosya yok", "agentDetail.configure.files.label": "Dosyalar", diff --git a/web/i18n/uk-UA/agent-v-2.json b/web/i18n/uk-UA/agent-v-2.json index 717923d30a2..6da2f4e3dde 100644 --- a/web/i18n/uk-UA/agent-v-2.json +++ b/web/i18n/uk-UA/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "Розширені налаштування", "agentDetail.configure.advancedSettings.toggle": "Перемкнути розширені налаштування", + "agentDetail.configure.build.empty.communityEditionTip": "Community Edition, хоча й пропонує версіонування витягнутих конфігурацій, не підтримує версіонування самої файлової системи. Будьте обережні з діями в режимі Build, адже зміни файлової системи відбуваються в реальному часі й не завжди можуть бути чисто скасовані. Community Edition не є ідеальним вибором для обслуговування масової зовнішньої аудиторії.", "agentDetail.configure.build.empty.description": "Опишіть, що вам потрібно, і форма ліворуч заповнюватиметься під час розмови.", "agentDetail.configure.build.empty.title": "Створіть агента через чат", "agentDetail.configure.build.inputPlaceholder": "Опишіть, що має робити ваш агент", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "{{count}} змін для застосування", "agentDetail.configure.buildDraft.discard": "Відхилити", "agentDetail.configure.buildDraft.modeBadge": "Режим збірки", - "agentDetail.configure.buildDraft.modeDescription": "Ви в режимі збірки. Налаштуйте цю конфігурацію через чат праворуч, а потім застосуйте.", + "agentDetail.configure.buildDraft.modeDescription": "Ви в режимі збірки. У цьому режимі Configure може оновлювати лише агент. Налаштуйте цю конфігурацію через чат праворуч, а потім застосуйте.", "agentDetail.configure.buildDraft.rewritten": "Переписано", "agentDetail.configure.buildDraft.title": "Чернетка збірки", "agentDetail.configure.buildDraft.updated": "Оновлено", "agentDetail.configure.chatFeatures.description": "Налаштуйте чат-досвід кінцевого користувача у вашому веб-застосунку та чат-поверхнях.", "agentDetail.configure.chatFeatures.title": "Функції чату", + "agentDetail.configure.communityEditionIsolationTip": "Community Edition не забезпечує жорсткої ізоляції файлової системи між кінцевими користувачами або запусками. Не надавайте один і той самий CE-агент кільком незалежним кінцевим користувачам, якщо потрібна ізоляція даних або сувора відповідність вимогам.", "agentDetail.configure.files.add": "Додати файл", "agentDetail.configure.files.buildNote.generated": "Згенеровано", "agentDetail.configure.files.buildNote.richTooltip": "Запис агента про те, що він налаштував у режимі Build. Він читає його на початку кожної розмови разом із вашим Prompt. Докладніше", "agentDetail.configure.files.buildNote.tooltip": "Запис агента про те, що він налаштував у режимі Build. Він читає його на початку кожної розмови разом із вашим Prompt. Докладніше", + "agentDetail.configure.files.download": "Завантажити {{name}}", "agentDetail.configure.files.empty.description": "Завантажте документи, які може читати агент, наприклад специфікації, шаблони чи інструкції", "agentDetail.configure.files.empty.title": "Файлів ще немає", "agentDetail.configure.files.label": "Файли", diff --git a/web/i18n/vi-VN/agent-v-2.json b/web/i18n/vi-VN/agent-v-2.json index 8a15e581f67..f3a16a876ad 100644 --- a/web/i18n/vi-VN/agent-v-2.json +++ b/web/i18n/vi-VN/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "Cài đặt nâng cao", "agentDetail.configure.advancedSettings.toggle": "Bật/tắt cài đặt nâng cao", + "agentDetail.configure.build.empty.communityEditionTip": "Community Edition, mặc dù cung cấp phiên bản cho các cấu hình được trích xuất, không hỗ trợ phiên bản cho chính hệ thống tệp. Hãy cẩn thận với các thao tác trong chế độ Build, vì các thay đổi đối với hệ thống tệp diễn ra theo thời gian thực và không phải lúc nào cũng có thể hoàn tác gọn gàng. Community Edition không phải là lựa chọn lý tưởng để phục vụ lượng lớn khán giả bên ngoài.", "agentDetail.configure.build.empty.description": "Mô tả điều bạn muốn và biểu mẫu bên trái sẽ được điền dần khi trò chuyện.", "agentDetail.configure.build.empty.title": "Xây dựng tác nhân bằng trò chuyện", "agentDetail.configure.build.inputPlaceholder": "Mô tả tác nhân của bạn nên làm gì", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "{{count}} thay đổi để áp dụng", "agentDetail.configure.buildDraft.discard": "Hủy bỏ", "agentDetail.configure.buildDraft.modeBadge": "Chế độ build", - "agentDetail.configure.buildDraft.modeDescription": "Bạn đang ở chế độ build. Điều chỉnh thiết lập này qua khung chat bên phải, rồi Áp dụng.", + "agentDetail.configure.buildDraft.modeDescription": "Bạn đang ở chế độ build. Configure chỉ có thể được cập nhật bởi tác nhân trong chế độ này. Điều chỉnh thiết lập này qua khung chat bên phải, rồi Áp dụng.", "agentDetail.configure.buildDraft.rewritten": "Đã viết lại", "agentDetail.configure.buildDraft.title": "Bản nháp build", "agentDetail.configure.buildDraft.updated": "Đã cập nhật", "agentDetail.configure.chatFeatures.description": "Định hình trải nghiệm trò chuyện cho người dùng cuối trên Web app và các bề mặt trò chuyện.", "agentDetail.configure.chatFeatures.title": "Tính năng trò chuyện", + "agentDetail.configure.communityEditionIsolationTip": "Community Edition không cung cấp cách ly hệ thống tệp cứng giữa người dùng cuối hoặc giữa các lần chạy. Không cung cấp cùng một tác nhân CE cho nhiều người dùng cuối độc lập khi cần cách ly dữ liệu hoặc tuân thủ nghiêm ngặt.", "agentDetail.configure.files.add": "Thêm tệp", "agentDetail.configure.files.buildNote.generated": "Đã tạo", "agentDetail.configure.files.buildNote.richTooltip": "Bản ghi của tác nhân về những gì nó đã thiết lập trong chế độ Build. Nó đọc bản ghi này ở đầu mỗi cuộc trò chuyện, cùng với Prompt của bạn. Tìm hiểu thêm", "agentDetail.configure.files.buildNote.tooltip": "Bản ghi của tác nhân về những gì nó đã thiết lập trong chế độ Build. Nó đọc bản ghi này ở đầu mỗi cuộc trò chuyện, cùng với Prompt của bạn. Tìm hiểu thêm", + "agentDetail.configure.files.download": "Tải xuống {{name}}", "agentDetail.configure.files.empty.description": "Tải lên tài liệu mà tác nhân có thể đọc, như đặc tả, mẫu hoặc hướng dẫn", "agentDetail.configure.files.empty.title": "Chưa có tệp nào", "agentDetail.configure.files.label": "Tệp", diff --git a/web/i18n/zh-Hans/agent-v-2.json b/web/i18n/zh-Hans/agent-v-2.json index bff704ebff3..613d125939a 100644 --- a/web/i18n/zh-Hans/agent-v-2.json +++ b/web/i18n/zh-Hans/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "高级设置", "agentDetail.configure.advancedSettings.toggle": "展开或收起高级设置", + "agentDetail.configure.build.empty.communityEditionTip": "Community Edition 虽然会对提取出的配置提供版本管理,但不支持文件系统本身的版本管理。请谨慎使用 Build 模式,因为对文件系统的更改会实时发生,并且不一定总能干净地回退。Community Edition 并不是面向大量外部受众提供服务的理想选择。", "agentDetail.configure.build.empty.description": "描述你的需求,它会随着对话填写左侧表单。", "agentDetail.configure.build.empty.title": "通过对话构建 Agent", "agentDetail.configure.build.inputPlaceholder": "描述你的 Agent 应该做什么", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "{{count}} 项变更待应用", "agentDetail.configure.buildDraft.discard": "放弃", "agentDetail.configure.buildDraft.modeBadge": "构建模式", - "agentDetail.configure.buildDraft.modeDescription": "你正在使用构建模式。通过右侧聊天调整此配置,然后应用。", + "agentDetail.configure.buildDraft.modeDescription": "你正在使用构建模式。在此模式下,Configure 只能由 Agent 更新。通过右侧聊天调整此配置,然后应用。", "agentDetail.configure.buildDraft.rewritten": "已重写", "agentDetail.configure.buildDraft.title": "Build 草稿", "agentDetail.configure.buildDraft.updated": "已更新", "agentDetail.configure.chatFeatures.description": "配置 Web app 和聊天界面的终端用户聊天体验。", "agentDetail.configure.chatFeatures.title": "Chat 功能", + "agentDetail.configure.communityEditionIsolationTip": "Community Edition 不在最终用户之间或不同运行之间提供严格的文件系统隔离。如果需要数据隔离或严格合规,请勿将同一个 CE Agent 暴露给多个相互独立的最终用户。", "agentDetail.configure.files.add": "添加文件", "agentDetail.configure.files.buildNote.generated": "已生成", "agentDetail.configure.files.buildNote.richTooltip": "Agent 在构建模式中完成的设置都记录在这里。每次对话开始时,它会连同提示词一起读取这份记录。了解更多", "agentDetail.configure.files.buildNote.tooltip": "Agent 在构建模式中完成的设置都记录在这里。每次对话开始时,它会连同提示词一起读取这份记录。了解更多", + "agentDetail.configure.files.download": "下载 {{name}}", "agentDetail.configure.files.empty.description": "上传 Agent 可读取的文档,例如规格、模板或指南", "agentDetail.configure.files.empty.title": "暂无文件", "agentDetail.configure.files.label": "文件", diff --git a/web/i18n/zh-Hans/common.json b/web/i18n/zh-Hans/common.json index 83e309caa57..3f81738d058 100644 --- a/web/i18n/zh-Hans/common.json +++ b/web/i18n/zh-Hans/common.json @@ -288,7 +288,7 @@ "menus.explore": "探索", "menus.exploreMarketplace": "探索 Marketplace", "menus.plugins": "集成", - "menus.roster": "Agent 名册", + "menus.roster": "名册", "menus.status": "beta", "menus.tools": "工具", "model.capabilities": "多模态能力", diff --git a/web/i18n/zh-Hant/agent-v-2.json b/web/i18n/zh-Hant/agent-v-2.json index 4c0daade372..2e3e2007894 100644 --- a/web/i18n/zh-Hant/agent-v-2.json +++ b/web/i18n/zh-Hant/agent-v-2.json @@ -63,6 +63,7 @@ "agentDetail.configure.advancedSettings.envEditor.valuePlaceholder": "Value", "agentDetail.configure.advancedSettings.label": "進階設定", "agentDetail.configure.advancedSettings.toggle": "展開或收合進階設定", + "agentDetail.configure.build.empty.communityEditionTip": "Community Edition 雖然會對提取出的設定提供版本管理,但不支援檔案系統本身的版本管理。請謹慎使用 Build 模式,因為對檔案系統的變更會即時發生,且不一定總能乾淨地回復。Community Edition 並不是面向大量外部受眾提供服務的理想選擇。", "agentDetail.configure.build.empty.description": "描述你的需求,它會隨著對話填寫左側表單。", "agentDetail.configure.build.empty.title": "透過對話建置 Agent", "agentDetail.configure.build.inputPlaceholder": "描述你的 Agent 應該做什麼", @@ -72,16 +73,18 @@ "agentDetail.configure.buildDraft.changesToApply_other": "{{count}} 項變更待套用", "agentDetail.configure.buildDraft.discard": "放棄", "agentDetail.configure.buildDraft.modeBadge": "建置模式", - "agentDetail.configure.buildDraft.modeDescription": "你正在使用建置模式。透過右側聊天調整此設定,然後套用。", + "agentDetail.configure.buildDraft.modeDescription": "你正在使用建置模式。在此模式下,Configure 只能由 Agent 更新。透過右側聊天調整此設定,然後套用。", "agentDetail.configure.buildDraft.rewritten": "已重寫", "agentDetail.configure.buildDraft.title": "Build 草稿", "agentDetail.configure.buildDraft.updated": "已更新", "agentDetail.configure.chatFeatures.description": "配置 Web app 和聊天介面的終端使用者聊天體驗。", "agentDetail.configure.chatFeatures.title": "Chat 功能", + "agentDetail.configure.communityEditionIsolationTip": "Community Edition 不會在最終使用者之間或不同執行之間提供嚴格的檔案系統隔離。如果需要資料隔離或嚴格合規,請勿將同一個 CE Agent 暴露給多個相互獨立的最終使用者。", "agentDetail.configure.files.add": "新增檔案", "agentDetail.configure.files.buildNote.generated": "已生成", "agentDetail.configure.files.buildNote.richTooltip": "Agent 在建置模式中完成的設定都記錄在這裡。每次對話開始時,它會連同提示詞一起讀取這份記錄。了解更多", "agentDetail.configure.files.buildNote.tooltip": "Agent 在建置模式中完成的設定都記錄在這裡。每次對話開始時,它會連同提示詞一起讀取這份記錄。了解更多", + "agentDetail.configure.files.download": "下載 {{name}}", "agentDetail.configure.files.empty.description": "上傳 Agent 可讀取的文件,例如規格、範本或指南", "agentDetail.configure.files.empty.title": "暫無檔案", "agentDetail.configure.files.label": "檔案", diff --git a/web/i18n/zh-Hant/common.json b/web/i18n/zh-Hant/common.json index f84b7fde1d0..bd5bee143b5 100644 --- a/web/i18n/zh-Hant/common.json +++ b/web/i18n/zh-Hant/common.json @@ -288,7 +288,7 @@ "menus.explore": "探索", "menus.exploreMarketplace": "探索 Marketplace", "menus.plugins": "集成", - "menus.roster": "Agent 名冊", + "menus.roster": "名冊", "menus.status": "beta", "menus.tools": "工具", "model.capabilities": "多模式功能", diff --git a/web/service/client.spec.ts b/web/service/client.spec.ts index d596c7da478..4e94673f6d5 100644 --- a/web/service/client.spec.ts +++ b/web/service/client.spec.ts @@ -59,6 +59,16 @@ type AgentMutationResponse = Parameters['onSuccess']>>[0] type AgentPublishMutationResponse = Parameters['onSuccess']>>[0] type WorkflowAgentComposerMutationResponse = Parameters['onSuccess']>>[0] +type RetryFn = (failureCount: number, error: unknown) => boolean + +const getRetryFn = (queryOptions: object): RetryFn => { + const retry = (queryOptions as { retry?: unknown }).retry + expect(typeof retry).toBe('function') + if (typeof retry !== 'function') + throw new TypeError('Expected query retry to be a function.') + + return retry as RetryFn +} const createAgent = (overrides: Partial = {}): AgentMutationResponse => ({ ...overrides, @@ -335,6 +345,42 @@ describe('normalizeConsoleOpenAPIURL', () => { }) }) +// Scenario: oRPC query defaults own shared Agent detail fetch behavior. +describe('consoleQuery agent query defaults', () => { + afterEach(() => { + vi.restoreAllMocks() + }) + + it('should not retry missing agent detail errors', async () => { + const consoleQuery = await loadConsoleQuery() + const queryOptions = consoleQuery.agent.byAgentId.get.queryOptions({ + input: { + params: { + agent_id: 'agent-1', + }, + }, + }) + const retry = getRetryFn(queryOptions) + + expect(retry(0, new Response(null, { status: 404 }))).toBe(false) + }) + + it('should retry other agent detail errors fewer than three times', async () => { + const consoleQuery = await loadConsoleQuery() + const queryOptions = consoleQuery.agent.byAgentId.get.queryOptions({ + input: { + params: { + agent_id: 'agent-1', + }, + }, + }) + const retry = getRetryFn(queryOptions) + + expect(retry(2, new Error('temporary failure'))).toBe(true) + expect(retry(3, new Error('temporary failure'))).toBe(false) + }) +}) + // Scenario: oRPC mutation defaults own shared Agent roster cache behavior. describe('consoleQuery agent mutation defaults', () => { beforeEach(() => { diff --git a/web/service/client.ts b/web/service/client.ts index e0527885df6..3adfc82b0ea 100644 --- a/web/service/client.ts +++ b/web/service/client.ts @@ -458,6 +458,16 @@ export const consoleQuery: RouterUtils = createTanstackQue }, }, byAgentId: { + get: { + queryOptions: { + retry: (failureCount, error) => { + if (error instanceof Response && error.status === 404) + return false + + return failureCount < 3 + }, + }, + }, copy: { post: { mutationOptions: { From fc01d112a06cc286f5b47d34112387c46503ddb2 Mon Sep 17 00:00:00 2001 From: QuantumGhost Date: Tue, 7 Jul 2026 02:14:54 +0800 Subject: [PATCH 14/70] refactor(api): Stop masking refresh-token service errors as 401 (#38463) Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- api/controllers/console/auth/login.py | 26 ++--- api/services/account_service.py | 6 +- api/services/errors/account.py | 8 ++ .../console/auth/test_token_refresh.py | 96 ++++++++++++++++--- 4 files changed, 108 insertions(+), 28 deletions(-) diff --git a/api/controllers/console/auth/login.py b/api/controllers/console/auth/login.py index 74d1ecc38f7..5165fc3003a 100644 --- a/api/controllers/console/auth/login.py +++ b/api/controllers/console/auth/login.py @@ -56,7 +56,7 @@ from models.account import Account from services.account_service import AccountService, InvitationDetailDict, RegisterService, TenantService from services.billing_service import BillingService from services.entities.auth_entities import LoginFailureReason, LoginPayloadBase -from services.errors.account import AccountRegisterError +from services.errors.account import AccountRegisterError, RefreshTokenAccountNotFoundError, RefreshTokenNotFoundError from services.errors.workspace import WorkSpaceNotAllowedCreateError, WorkspacesLimitExceededError from services.feature_service import FeatureService @@ -359,18 +359,22 @@ class RefreshTokenApi(Resource): try: new_token_pair = AccountService.refresh_token(refresh_token, session=db.session) + except Unauthorized as exc: + return SimpleResultMessageResponse(result="fail", message=exc.description or "Unauthorized.").model_dump( + mode="json" + ), 401 + except (RefreshTokenNotFoundError, RefreshTokenAccountNotFoundError) as exc: + return SimpleResultMessageResponse(result="fail", message=str(exc)).model_dump(mode="json"), 401 - # Create response with new cookies - # response-contract:ignore cookie-bearing Flask response - response = make_response(SimpleResultResponse(result="success").model_dump(mode="json")) + # Create response with new cookies + # response-contract:ignore cookie-bearing Flask response + response = make_response(SimpleResultResponse(result="success").model_dump(mode="json")) - # Update cookies with new tokens - set_csrf_token_to_cookie(request, response, new_token_pair.csrf_token) - set_access_token_to_cookie(request, response, new_token_pair.access_token) - set_refresh_token_to_cookie(request, response, new_token_pair.refresh_token) - return response - except Exception as e: - return SimpleResultMessageResponse(result="fail", message=str(e)).model_dump(mode="json"), 401 + # Update cookies with new tokens + set_csrf_token_to_cookie(request, response, new_token_pair.csrf_token) + set_access_token_to_cookie(request, response, new_token_pair.access_token) + set_refresh_token_to_cookie(request, response, new_token_pair.refresh_token) + return response def _get_account_with_case_fallback(email: str): diff --git a/api/services/account_service.py b/api/services/account_service.py index 9bfa586c457..1b9fd724a71 100644 --- a/api/services/account_service.py +++ b/api/services/account_service.py @@ -65,6 +65,8 @@ from services.errors.account import ( LinkAccountIntegrateError, MemberNotInTenantError, NoPermissionError, + RefreshTokenAccountNotFoundError, + RefreshTokenNotFoundError, RoleAlreadyAssignedError, TenantNotFoundError, ) @@ -654,11 +656,11 @@ class AccountService: # Verify the refresh token account_id = redis_client.get(AccountService._get_refresh_token_key(refresh_token)) if not account_id: - raise ValueError("Invalid refresh token") + raise RefreshTokenNotFoundError("Invalid refresh token") account = AccountService.load_user(account_id.decode("utf-8"), session) if not account: - raise ValueError("Invalid account") + raise RefreshTokenAccountNotFoundError("Invalid account") # Generate new access token and refresh token new_access_token = AccountService.get_account_jwt_token(account) diff --git a/api/services/errors/account.py b/api/services/errors/account.py index 4d3d150e072..700c1dd4aaf 100644 --- a/api/services/errors/account.py +++ b/api/services/errors/account.py @@ -17,6 +17,14 @@ class AccountPasswordError(BaseServiceError): pass +class RefreshTokenNotFoundError(BaseServiceError): + pass + + +class RefreshTokenAccountNotFoundError(BaseServiceError): + pass + + class AccountNotLinkTenantError(BaseServiceError): pass diff --git a/api/tests/unit_tests/controllers/console/auth/test_token_refresh.py b/api/tests/unit_tests/controllers/console/auth/test_token_refresh.py index 34fff57b0ad..8effbb96887 100644 --- a/api/tests/unit_tests/controllers/console/auth/test_token_refresh.py +++ b/api/tests/unit_tests/controllers/console/auth/test_token_refresh.py @@ -13,8 +13,10 @@ from unittest.mock import ANY, MagicMock, patch import pytest from flask import Flask from flask_restx import Api +from werkzeug.exceptions import Unauthorized from controllers.console.auth.login import RefreshTokenApi +from services.errors.account import RefreshTokenAccountNotFoundError, RefreshTokenNotFoundError class TestRefreshTokenApi: @@ -98,18 +100,19 @@ class TestRefreshTokenApi: @patch("controllers.console.auth.login.extract_refresh_token", autospec=True) @patch("controllers.console.auth.login.AccountService.refresh_token", autospec=True) - def test_refresh_fails_with_invalid_token(self, mock_refresh_token, mock_extract_token, app: Flask): + def test_refresh_returns_unauthorized_for_invalid_refresh_token( + self, mock_refresh_token, mock_extract_token, app: Flask + ): """ - Test token refresh failure with invalid refresh token. + Test token refresh maps invalid refresh tokens to unauthorized responses. Verifies that: - - Exception is caught when token is invalid - - 401 status code is returned - - Error message is included in response + - Invalid refresh token validation failures return 401 + - The failure response preserves the validation message """ # Arrange mock_extract_token.return_value = "invalid_refresh_token" - mock_refresh_token.side_effect = Exception("Invalid refresh token") + mock_refresh_token.side_effect = RefreshTokenNotFoundError("Invalid refresh token") # Act with app.test_request_context("/refresh-token", method="POST"): @@ -119,22 +122,21 @@ class TestRefreshTokenApi: # Assert assert status_code == 401 assert response["result"] == "fail" - assert "Invalid refresh token" in response["message"] + assert response["message"] == "Invalid refresh token" @patch("controllers.console.auth.login.extract_refresh_token", autospec=True) @patch("controllers.console.auth.login.AccountService.refresh_token", autospec=True) - def test_refresh_fails_with_expired_token(self, mock_refresh_token, mock_extract_token, app: Flask): + def test_refresh_returns_unauthorized_for_invalid_account(self, mock_refresh_token, mock_extract_token, app: Flask): """ - Test token refresh failure with expired refresh token. + Test token refresh maps missing accounts to unauthorized responses. Verifies that: - - Expired tokens are rejected - - 401 status code is returned - - Appropriate error handling + - Invalid account validation failures return 401 + - The failure response preserves the validation message """ # Arrange - mock_extract_token.return_value = "expired_refresh_token" - mock_refresh_token.side_effect = Exception("Refresh token expired") + mock_extract_token.return_value = "refresh_token_for_missing_account" + mock_refresh_token.side_effect = RefreshTokenAccountNotFoundError("Invalid account") # Act with app.test_request_context("/refresh-token", method="POST"): @@ -144,7 +146,71 @@ class TestRefreshTokenApi: # Assert assert status_code == 401 assert response["result"] == "fail" - assert "expired" in response["message"].lower() + assert response["message"] == "Invalid account" + + @patch("controllers.console.auth.login.extract_refresh_token", autospec=True) + @patch("controllers.console.auth.login.AccountService.refresh_token", autospec=True) + def test_refresh_returns_unauthorized_for_banned_account(self, mock_refresh_token, mock_extract_token, app: Flask): + """ + Test token refresh maps banned accounts to unauthorized responses. + + Verifies that: + - Authorization failures raised during account loading return 401 + - The failure response preserves the authorization message + """ + # Arrange + mock_extract_token.return_value = "refresh_token_for_banned_account" + mock_refresh_token.side_effect = Unauthorized("Account is banned.") + + # Act + with app.test_request_context("/refresh-token", method="POST"): + refresh_api = RefreshTokenApi() + response, status_code = refresh_api.post() + + # Assert + assert status_code == 401 + assert response["result"] == "fail" + assert response["message"] == "Account is banned." + + @patch("controllers.console.auth.login.extract_refresh_token", autospec=True) + @patch("controllers.console.auth.login.AccountService.refresh_token", autospec=True) + def test_refresh_propagates_non_whitelisted_value_error(self, mock_refresh_token, mock_extract_token, app: Flask): + """ + Test token refresh preserves non-whitelisted ValueError failures. + + Verifies that: + - Only known refresh-token validation errors are mapped to 401 + - Unexpected ValueError instances continue to propagate + """ + # Arrange + mock_extract_token.return_value = "valid_refresh_token" + mock_refresh_token.side_effect = ValueError("unexpected parse failure") + + # Act & Assert + with app.test_request_context("/refresh-token", method="POST"): + refresh_api = RefreshTokenApi() + with pytest.raises(ValueError, match="unexpected parse failure"): + refresh_api.post() + + @patch("controllers.console.auth.login.extract_refresh_token", autospec=True) + @patch("controllers.console.auth.login.AccountService.refresh_token", autospec=True) + def test_refresh_propagates_unexpected_service_errors(self, mock_refresh_token, mock_extract_token, app: Flask): + """ + Test token refresh preserves unexpected service failures. + + Verifies that: + - Operational errors are not misreported as authentication failures + - The original exception is preserved for higher-level error handling + """ + # Arrange + mock_extract_token.return_value = "valid_refresh_token" + mock_refresh_token.side_effect = RuntimeError("redis unavailable") + + # Act & Assert + with app.test_request_context("/refresh-token", method="POST"): + refresh_api = RefreshTokenApi() + with pytest.raises(RuntimeError, match="redis unavailable"): + refresh_api.post() @patch("controllers.console.auth.login.extract_refresh_token", autospec=True) @patch("controllers.console.auth.login.AccountService.refresh_token", autospec=True) From b9c7199d3458618455275abec108616ce6588aad Mon Sep 17 00:00:00 2001 From: Jingyi Date: Mon, 6 Jul 2026 18:44:14 -0700 Subject: [PATCH 15/70] fix(web): unify detail sidebar home control (#38487) --- .../__tests__/dataset-detail-top.spec.tsx | 15 ++-------- .../app-sidebar/dataset-detail-top.tsx | 27 ++++++------------ .../__tests__/navigation.spec.tsx | 16 ++++++++++- .../agent-v2/agent-detail/navigation.tsx | 28 ++++++------------- .../detail/__tests__/index.spec.tsx | 19 +++++++++++++ .../deployments/detail/deployment-sidebar.tsx | 28 ++++++------------- 6 files changed, 63 insertions(+), 70 deletions(-) diff --git a/web/app/components/app-sidebar/__tests__/dataset-detail-top.spec.tsx b/web/app/components/app-sidebar/__tests__/dataset-detail-top.spec.tsx index 18f1c959e3d..1f23bf5c334 100644 --- a/web/app/components/app-sidebar/__tests__/dataset-detail-top.spec.tsx +++ b/web/app/components/app-sidebar/__tests__/dataset-detail-top.spec.tsx @@ -4,14 +4,6 @@ import { createStore, Provider as JotaiProvider } from 'jotai' import { useGotoAnythingOpen } from '@/app/components/goto-anything/atoms' import DatasetDetailTop from '../dataset-detail-top' -const mockBack = vi.fn() - -vi.mock('@/next/navigation', () => ({ - useRouter: () => ({ - back: mockBack, - }), -})) - vi.mock('../toggle-button', () => ({ default: ({ expand, handleToggle, icon }: { expand: boolean, handleToggle: () => void, icon?: ReactNode }) => ( - - - -
    + + + + {expand && ( <> diff --git a/web/features/agent-v2/agent-detail/__tests__/navigation.spec.tsx b/web/features/agent-v2/agent-detail/__tests__/navigation.spec.tsx index b5f3f48b6e1..acc67ea7389 100644 --- a/web/features/agent-v2/agent-detail/__tests__/navigation.spec.tsx +++ b/web/features/agent-v2/agent-detail/__tests__/navigation.spec.tsx @@ -1,6 +1,6 @@ import type { AgentAppDetailWithSite } from '@dify/contracts/api/console/agent/types.gen' import { render, screen } from '@testing-library/react' -import { AgentDetailSection } from '../navigation' +import { AgentDetailSection, AgentDetailTop } from '../navigation' const mocks = vi.hoisted(() => ({ pathname: '/roster/agent/agent-1/configure', @@ -71,3 +71,17 @@ describe('AgentDetailSection', () => { expect(agentName.parentElement?.parentElement).toHaveClass('h-13', 'py-1.5', 'pl-1.5', 'pr-2') }) }) + +describe('AgentDetailTop', () => { + beforeEach(() => { + vi.clearAllMocks() + }) + + it('links the combined home control to home', () => { + render() + + expect(screen.getByRole('link', { name: 'common.mainNav.home' })).toHaveAttribute('href', '/') + expect(screen.getByRole('link', { name: 'common.menus.roster' })).toHaveAttribute('href', '/roster') + expect(screen.queryByRole('button', { name: 'common.operation.back' })).not.toBeInTheDocument() + }) +}) diff --git a/web/features/agent-v2/agent-detail/navigation.tsx b/web/features/agent-v2/agent-detail/navigation.tsx index 3f68a409702..714bbc57e94 100644 --- a/web/features/agent-v2/agent-detail/navigation.tsx +++ b/web/features/agent-v2/agent-detail/navigation.tsx @@ -17,7 +17,7 @@ import Divider from '@/app/components/base/divider' import SidebarLeftArrowIcon from '@/app/components/base/icons/src/vender/SidebarLeftArrowIcon' import { useSetGotoAnythingOpen } from '@/app/components/goto-anything/atoms' import Link from '@/next/link' -import { usePathname, useRouter } from '@/next/navigation' +import { usePathname } from '@/next/navigation' import { consoleQuery } from '@/service/client' import { getAgentDetailPath, getAgentIdFromPathname } from './routes' @@ -88,7 +88,6 @@ export function AgentDetailTop({ }: AgentDetailTopProps) { const { t: tApp } = useTranslation('app') const { t: tCommon } = useTranslation('common') - const router = useRouter() const setGotoAnythingOpen = useSetGotoAnythingOpen() if (!expand) { @@ -109,23 +108,14 @@ export function AgentDetailTop({ return (
    -
    - - - - -
    + + + + / diff --git a/web/features/deployments/detail/__tests__/index.spec.tsx b/web/features/deployments/detail/__tests__/index.spec.tsx index 67486675dcd..81e18849fec 100644 --- a/web/features/deployments/detail/__tests__/index.spec.tsx +++ b/web/features/deployments/detail/__tests__/index.spec.tsx @@ -4,6 +4,7 @@ import { Provider as JotaiProvider } from 'jotai' import { NextRouteStateBridge } from '@/app/components/next-route-state' import { useParams, usePathname } from '@/next/navigation' import { InstanceDetail } from '..' +import { DeploymentDetailTop } from '../deployment-sidebar' vi.mock('@/next/navigation', async (importOriginal) => { const actual = await importOriginal() @@ -60,3 +61,21 @@ describe('InstanceDetail', () => { expect(screen.queryByRole('complementary', { name: 'Detail sidebar' })).not.toBeInTheDocument() }) }) + +describe('DeploymentDetailTop', () => { + beforeEach(() => { + vi.clearAllMocks() + }) + + it('links the combined home control to home', () => { + render( + + + , + ) + + expect(screen.getByRole('link', { name: 'common.mainNav.home' })).toHaveAttribute('href', '/') + expect(screen.getByRole('link', { name: 'common.menus.deployments' })).toHaveAttribute('href', '/deployments') + expect(screen.queryByRole('button', { name: 'common.operation.back' })).not.toBeInTheDocument() + }) +}) diff --git a/web/features/deployments/detail/deployment-sidebar.tsx b/web/features/deployments/detail/deployment-sidebar.tsx index afbfd18b440..95f50c768d6 100644 --- a/web/features/deployments/detail/deployment-sidebar.tsx +++ b/web/features/deployments/detail/deployment-sidebar.tsx @@ -16,7 +16,7 @@ import SidebarLeftArrowIcon from '@/app/components/base/icons/src/vender/Sidebar import { SkeletonContainer, SkeletonRectangle } from '@/app/components/base/skeleton' import { useSetGotoAnythingOpen } from '@/app/components/goto-anything/atoms' import Link from '@/next/link' -import { usePathname, useRouter } from '@/next/navigation' +import { usePathname } from '@/next/navigation' import { DeploymentActionsMenu } from '../deployment-actions' import { deploymentRouteAppInstanceIdAtom } from '../route-state' import { TitleTooltip } from '../shared/components/title-tooltip' @@ -187,7 +187,6 @@ export function DeploymentDetailTop({ onToggle?: () => void }) { const { t } = useTranslation() - const router = useRouter() const setGotoAnythingOpen = useSetGotoAnythingOpen() if (!expand) { @@ -208,23 +207,14 @@ export function DeploymentDetailTop({ return (
    -
    - - - - -
    + + + + / From 0a3426ea38dce322321dc962a5e8cd4abd34a373 Mon Sep 17 00:00:00 2001 From: WH-2099 Date: Tue, 7 Jul 2026 10:06:49 +0800 Subject: [PATCH 16/70] refactor(api): clarify DSL import and plugin migration boundaries (#38483) Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- api/commands/plugin.py | 26 ++-- api/services/app_dsl_service.py | 13 +- api/services/dsl_content.py | 9 ++ api/services/plugin/plugin_migration.py | 116 ++++++++++-------- .../rag_pipeline/rag_pipeline_dsl_service.py | 13 +- .../rag_pipeline_transform_service.py | 20 +-- api/services/snippet_dsl_service.py | 54 ++------ .../services/plugin/test_plugin_migration.py | 58 ++++++++- .../test_rag_pipeline_dsl_service.py | 53 ++++++++ .../test_rag_pipeline_transform_service.py | 4 +- .../services/test_app_dsl_service.py | 53 ++++++++ .../services/test_snippet_dsl_service.py | 43 ++++++- 12 files changed, 329 insertions(+), 133 deletions(-) create mode 100644 api/services/dsl_content.py create mode 100644 api/tests/unit_tests/services/test_app_dsl_service.py diff --git a/api/commands/plugin.py b/api/commands/plugin.py index 445ecb31415..718fa60761c 100644 --- a/api/commands/plugin.py +++ b/api/commands/plugin.py @@ -188,13 +188,13 @@ def transform_datasource_credentials(environment: str): firecrawl_plugin_id = "langgenius/firecrawl_datasource" jina_plugin_id = "langgenius/jina_datasource" if environment == "online": - notion_plugin_unique_identifier = plugin_migration._fetch_plugin_unique_identifier(notion_plugin_id) - firecrawl_plugin_unique_identifier = plugin_migration._fetch_plugin_unique_identifier(firecrawl_plugin_id) - jina_plugin_unique_identifier = plugin_migration._fetch_plugin_unique_identifier(jina_plugin_id) + notion_package_identifier = plugin_migration._fetch_latest_package_identifier(notion_plugin_id) + firecrawl_package_identifier = plugin_migration._fetch_latest_package_identifier(firecrawl_plugin_id) + jina_package_identifier = plugin_migration._fetch_latest_package_identifier(jina_plugin_id) else: - notion_plugin_unique_identifier = None - firecrawl_plugin_unique_identifier = None - jina_plugin_unique_identifier = None + notion_package_identifier = None + firecrawl_package_identifier = None + jina_package_identifier = None oauth_credential_type = CredentialType.OAUTH2 api_key_credential_type = CredentialType.API_KEY @@ -219,9 +219,9 @@ def transform_datasource_credentials(environment: str): installed_plugins = installer_manager.list_plugins(tenant_id) installed_plugins_ids = [plugin.plugin_id for plugin in installed_plugins] if notion_plugin_id not in installed_plugins_ids: - if notion_plugin_unique_identifier: + if notion_package_identifier: # install notion plugin - PluginService.install_from_marketplace_pkg(tenant_id, [notion_plugin_unique_identifier]) + PluginService.install_from_marketplace_pkg(tenant_id, [notion_package_identifier]) auth_count = 0 for notion_tenant_credential in notion_tenant_credentials: auth_count += 1 @@ -279,9 +279,9 @@ def transform_datasource_credentials(environment: str): installed_plugins = installer_manager.list_plugins(tenant_id) installed_plugins_ids = [plugin.plugin_id for plugin in installed_plugins] if firecrawl_plugin_id not in installed_plugins_ids: - if firecrawl_plugin_unique_identifier: + if firecrawl_package_identifier: # install firecrawl plugin - PluginService.install_from_marketplace_pkg(tenant_id, [firecrawl_plugin_unique_identifier]) + PluginService.install_from_marketplace_pkg(tenant_id, [firecrawl_package_identifier]) auth_count = 0 for firecrawl_tenant_credential in firecrawl_tenant_credentials: @@ -343,10 +343,10 @@ def transform_datasource_credentials(environment: str): installed_plugins = installer_manager.list_plugins(tenant_id) installed_plugins_ids = [plugin.plugin_id for plugin in installed_plugins] if jina_plugin_id not in installed_plugins_ids: - if jina_plugin_unique_identifier: + if jina_package_identifier: # install jina plugin - logger.debug("Installing Jina plugin %s", jina_plugin_unique_identifier) - PluginService.install_from_marketplace_pkg(tenant_id, [jina_plugin_unique_identifier]) + logger.debug("Installing Jina plugin %s", jina_package_identifier) + PluginService.install_from_marketplace_pkg(tenant_id, [jina_package_identifier]) auth_count = 0 for jina_tenant_credential in jina_tenant_credentials: diff --git a/api/services/app_dsl_service.py b/api/services/app_dsl_service.py index 52e936bf1ee..d042ad69f88 100644 --- a/api/services/app_dsl_service.py +++ b/api/services/app_dsl_service.py @@ -39,6 +39,7 @@ from libs.datetime_utils import naive_utc_now from models import Account, App, AppMode from models.model import AppModelConfig, AppModelConfigDict, IconType from models.workflow import Workflow +from services.dsl_content import DSL_MAX_SIZE, dsl_content_size from services.dsl_version import check_version_compatibility from services.entities.dsl_entities import CheckDependenciesResult, ImportMode, ImportStatus from services.errors.app import WorkflowNotFoundError @@ -51,7 +52,6 @@ logger = logging.getLogger(__name__) IMPORT_INFO_REDIS_KEY_PREFIX = "app_import_info:" CHECK_DEPENDENCIES_REDIS_KEY_PREFIX = "app_check_dependencies:" IMPORT_INFO_REDIS_EXPIRY = 10 * 60 # 10 minutes -DSL_MAX_SIZE = 10 * 1024 * 1024 # 10MB CURRENT_DSL_VERSION = CURRENT_APP_DSL_VERSION @@ -131,15 +131,16 @@ class AppDslService: yaml_url = yaml_url.replace("/blob/", "/") response = remote_fetcher.make_request("GET", yaml_url.strip(), follow_redirects=True, timeout=(10, 10)) response.raise_for_status() - content = response.content.decode() + raw_content = response.content - if len(content) > DSL_MAX_SIZE: + if dsl_content_size(raw_content) > DSL_MAX_SIZE: return Import( id=import_id, status=ImportStatus.FAILED, error="File size exceeds the limit of 10MB", ) + content = raw_content.decode("utf-8") if not content: return Import( id=import_id, @@ -160,6 +161,12 @@ class AppDslService: error="yaml_content is required when import_mode is yaml-content", ) content = yaml_content + if dsl_content_size(content) > DSL_MAX_SIZE: + return Import( + id=import_id, + status=ImportStatus.FAILED, + error="File size exceeds the limit of 10MB", + ) # Process YAML content try: diff --git a/api/services/dsl_content.py b/api/services/dsl_content.py new file mode 100644 index 00000000000..3874d0546db --- /dev/null +++ b/api/services/dsl_content.py @@ -0,0 +1,9 @@ +"""Shared DSL content size and decoding rules.""" + +DSL_MAX_SIZE = 10 * 1024 * 1024 # 10MB + + +def dsl_content_size(content: str | bytes) -> int: + if isinstance(content, bytes): + return len(content) + return len(content.encode("utf-8")) diff --git a/api/services/plugin/plugin_migration.py b/api/services/plugin/plugin_migration.py index 82eeb5a7261..d6f154df812 100644 --- a/api/services/plugin/plugin_migration.py +++ b/api/services/plugin/plugin_migration.py @@ -307,9 +307,9 @@ class PluginMigration: return result @classmethod - def _fetch_plugin_unique_identifier(cls, plugin_id: str) -> str | None: + def _fetch_latest_package_identifier(cls, plugin_id: str) -> str | None: """ - Fetch plugin unique identifier using plugin id. + Fetch the latest marketplace package identifier using a plugin id. """ if not dify_config.MARKETPLACE_ENABLED: return None @@ -328,7 +328,7 @@ class PluginMigration: @classmethod def extract_unique_plugins(cls, extracted_plugins: str) -> ExtractedPluginsDict: - plugins: dict[str, str] = {} + package_identifier_by_plugin_id: dict[str, str] = {} plugin_ids = [] plugin_not_exist = [] logger.info("Extracting unique plugins from %s", extracted_plugins) @@ -341,19 +341,19 @@ class PluginMigration: def fetch_plugin(plugin_id): try: - unique_identifier = cls._fetch_plugin_unique_identifier(plugin_id) - if unique_identifier: - plugins[plugin_id] = unique_identifier + latest_package_identifier = cls._fetch_latest_package_identifier(plugin_id) + if latest_package_identifier: + package_identifier_by_plugin_id[plugin_id] = latest_package_identifier else: plugin_not_exist.append(plugin_id) except Exception: - logger.exception("Failed to fetch plugin unique identifier for %s", plugin_id) + logger.exception("Failed to fetch latest package identifier for %s", plugin_id) plugin_not_exist.append(plugin_id) with ThreadPoolExecutor(max_workers=10) as executor: list(tqdm.tqdm(executor.map(fetch_plugin, plugin_ids), total=len(plugin_ids))) - return {"plugins": plugins, "plugin_not_exist": plugin_not_exist} + return {"plugins": package_identifier_by_plugin_id, "plugin_not_exist": plugin_not_exist} @classmethod def install_plugins(cls, extracted_plugins: str, output_file: str, workers: int = 100): @@ -362,17 +362,22 @@ class PluginMigration: """ manager = PluginInstaller() - plugins = cls.extract_unique_plugins(extracted_plugins) + extracted = cls.extract_unique_plugins(extracted_plugins) + package_identifier_by_plugin_id = extracted["plugins"] not_installed = [] plugin_install_failed = [] # use a fake tenant id to install all the plugins fake_tenant_id = uuid4().hex - logger.info("Installing %s plugin instances for fake tenant %s", len(plugins["plugins"]), fake_tenant_id) + logger.info( + "Installing %s plugin instances for fake tenant %s", + len(package_identifier_by_plugin_id), + fake_tenant_id, + ) thread_pool = ThreadPoolExecutor(max_workers=workers) - response = cls.handle_plugin_instance_install(fake_tenant_id, plugins["plugins"]) + response = cls.handle_plugin_instance_install(fake_tenant_id, package_identifier_by_plugin_id) if response.get("failed"): plugin_install_failed.extend(response.get("failed", [])) @@ -384,21 +389,21 @@ class PluginMigration: # at most 64 plugins one batch for i in range(0, len(plugin_ids), 64): batch_plugin_ids = plugin_ids[i : i + 64] - batch_plugin_identifiers = [ - plugins["plugins"][plugin_id] + batch_package_identifiers = [ + package_identifier_by_plugin_id[plugin_id] for plugin_id in batch_plugin_ids - if plugin_id not in installed_plugins_ids and plugin_id in plugins["plugins"] + if plugin_id not in installed_plugins_ids and plugin_id in package_identifier_by_plugin_id ] - if batch_plugin_identifiers: + if batch_package_identifiers: manager.install_from_identifiers( tenant_id, - batch_plugin_identifiers, + batch_package_identifiers, PluginInstallationSource.Marketplace, metas=[ { - "plugin_unique_identifier": identifier, + "plugin_unique_identifier": package_identifier, } - for identifier in batch_plugin_identifiers + for package_identifier in batch_package_identifiers ], ) PluginService.invalidate_plugin_model_providers_cache(tenant_id) @@ -412,10 +417,8 @@ class PluginMigration: tenant_id = data["tenant_id"] plugin_ids = data["plugins"] plugin_not_exist: list[str] = [] - # get plugin unique identifier for plugin_id in plugin_ids: - unique_identifier = plugins.get(plugin_id) - if unique_identifier: + if plugin_id not in package_identifier_by_plugin_id: plugin_not_exist.append(plugin_id) if plugin_not_exist: @@ -459,36 +462,44 @@ class PluginMigration: """ manager = PluginInstaller() - plugins = cls.extract_unique_plugins(extracted_plugins) + extracted = cls.extract_unique_plugins(extracted_plugins) + package_identifier_by_plugin_id = extracted["plugins"] plugin_install_failed = [] # use a fake tenant id to install all the plugins fake_tenant_id = uuid4().hex - logger.info("Installing %s plugin instances for fake tenant %s", len(plugins["plugins"]), fake_tenant_id) + logger.info( + "Installing %s plugin instances for fake tenant %s", + len(package_identifier_by_plugin_id), + fake_tenant_id, + ) thread_pool = ThreadPoolExecutor(max_workers=workers) - response = cls.handle_plugin_instance_install(fake_tenant_id, plugins["plugins"]) + response = cls.handle_plugin_instance_install(fake_tenant_id, package_identifier_by_plugin_id) if response.get("failed"): plugin_install_failed.extend(response.get("failed", [])) def install( - tenant_id: str, plugin_ids: dict[str, str], total_success_tenant: int, total_failed_tenant: int + tenant_id: str, + package_identifier_by_plugin_id: dict[str, str], + total_success_tenant: int, + total_failed_tenant: int, ) -> None: - logger.info("Installing %s plugins for tenant %s", len(plugin_ids), tenant_id) + logger.info("Installing %s plugins for tenant %s", len(package_identifier_by_plugin_id), tenant_id) try: # fetch plugin already installed installed_plugins = manager.list_plugins(tenant_id) installed_plugins_ids = [plugin.plugin_id for plugin in installed_plugins] # at most 64 plugins one batch - for i in range(0, len(plugin_ids), 64): - batch_plugin_ids = list(plugin_ids.keys())[i : i + 64] - batch_plugin_identifiers = [ - plugin_ids[plugin_id] + for i in range(0, len(package_identifier_by_plugin_id), 64): + batch_plugin_ids = list(package_identifier_by_plugin_id.keys())[i : i + 64] + batch_package_identifiers = [ + package_identifier_by_plugin_id[plugin_id] for plugin_id in batch_plugin_ids - if plugin_id not in installed_plugins_ids and plugin_id in plugin_ids + if plugin_id not in installed_plugins_ids and plugin_id in package_identifier_by_plugin_id ] - PluginService.install_from_marketplace_pkg(tenant_id, batch_plugin_identifiers) + PluginService.install_from_marketplace_pkg(tenant_id, batch_package_identifiers) total_success_tenant += 1 except Exception: @@ -510,7 +521,7 @@ class PluginMigration: thread_pool.submit( install, tenant_id, - plugins.get("plugins", {}), + package_identifier_by_plugin_id, total_success_tenant, total_failed_tenant, ) @@ -542,12 +553,12 @@ class PluginMigration: @classmethod def handle_plugin_instance_install( - cls, tenant_id: str, plugin_identifiers_map: Mapping[str, str] + cls, tenant_id: str, package_identifier_by_plugin_id: Mapping[str, str] ) -> PluginInstallResultDict: """ Install plugins for a tenant. """ - if plugin_identifiers_map and not dify_config.MARKETPLACE_ENABLED: + if package_identifier_by_plugin_id and not dify_config.MARKETPLACE_ENABLED: raise ValueError( "Marketplace disabled in offline mode; cannot bulk-install plugins. " "Pre-upload plugin packages via Console first." @@ -558,17 +569,17 @@ class PluginMigration: thread_pool = ThreadPoolExecutor(max_workers=10) futures = [] - for plugin_id, plugin_identifier in plugin_identifiers_map.items(): + for plugin_id, package_identifier in package_identifier_by_plugin_id.items(): - def download_and_upload(tenant_id, plugin_id, plugin_identifier): - plugin_package = marketplace.download_plugin_pkg(plugin_identifier) + def download_and_upload(tenant_id, plugin_id, package_identifier): + plugin_package = marketplace.download_plugin_pkg(package_identifier) if not plugin_package: - raise Exception(f"Failed to download plugin {plugin_identifier}") + raise Exception(f"Failed to download plugin {package_identifier}") # upload manager.upload_pkg(tenant_id, plugin_package, verify_signature=True) - futures.append(thread_pool.submit(download_and_upload, tenant_id, plugin_id, plugin_identifier)) + futures.append(thread_pool.submit(download_and_upload, tenant_id, plugin_id, package_identifier)) # Wait for all downloads to complete for future in futures: @@ -578,33 +589,33 @@ class PluginMigration: success = [] failed = [] - reverse_map = {v: k for k, v in plugin_identifiers_map.items()} + plugin_id_by_package_identifier = {v: k for k, v in package_identifier_by_plugin_id.items()} # at most 8 plugins one batch - for i in range(0, len(plugin_identifiers_map), 8): - batch_plugin_ids = list(plugin_identifiers_map.keys())[i : i + 8] - batch_plugin_identifiers = [plugin_identifiers_map[plugin_id] for plugin_id in batch_plugin_ids] + for i in range(0, len(package_identifier_by_plugin_id), 8): + batch_plugin_ids = list(package_identifier_by_plugin_id.keys())[i : i + 8] + batch_package_identifiers = [package_identifier_by_plugin_id[plugin_id] for plugin_id in batch_plugin_ids] try: response = manager.install_from_identifiers( tenant_id=tenant_id, - identifiers=batch_plugin_identifiers, + identifiers=batch_package_identifiers, source=PluginInstallationSource.Marketplace, metas=[ { - "plugin_unique_identifier": identifier, + "plugin_unique_identifier": package_identifier, } - for identifier in batch_plugin_identifiers + for package_identifier in batch_package_identifiers ], ) PluginService.invalidate_plugin_model_providers_cache(tenant_id) except Exception: # add to failed - failed.extend(batch_plugin_identifiers) + failed.extend(batch_plugin_ids) continue if response.all_installed: - success.extend(batch_plugin_identifiers) + success.extend(batch_plugin_ids) continue task_id = response.task_id @@ -614,10 +625,13 @@ class PluginMigration: if status.status in [PluginInstallTaskStatus.Failed, PluginInstallTaskStatus.Success]: PluginService.invalidate_plugin_model_providers_cache(tenant_id) for plugin in status.plugins: + plugin_id = plugin_id_by_package_identifier.get( + plugin.plugin_unique_identifier, plugin.plugin_unique_identifier.split(":", 1)[0] + ) if plugin.status == PluginInstallTaskStatus.Success: - success.append(reverse_map[plugin.plugin_unique_identifier]) + success.append(plugin_id) else: - failed.append(reverse_map[plugin.plugin_unique_identifier]) + failed.append(plugin_id) logger.error( "Failed to install plugin %s, error: %s", plugin.plugin_unique_identifier, diff --git a/api/services/rag_pipeline/rag_pipeline_dsl_service.py b/api/services/rag_pipeline/rag_pipeline_dsl_service.py index 4f9cde37a7b..5459c3e5f1f 100644 --- a/api/services/rag_pipeline/rag_pipeline_dsl_service.py +++ b/api/services/rag_pipeline/rag_pipeline_dsl_service.py @@ -36,6 +36,7 @@ from models import Account from models.dataset import Dataset, DatasetCollectionBinding, Pipeline from models.enums import CollectionBindingType, DatasetRuntimeMode from models.workflow import Workflow, WorkflowType +from services.dsl_content import DSL_MAX_SIZE, dsl_content_size from services.dsl_version import check_version_compatibility from services.entities.dsl_entities import CheckDependenciesResult, ImportMode, ImportStatus from services.entities.knowledge_entities.rag_pipeline_entities import ( @@ -50,7 +51,6 @@ logger = logging.getLogger(__name__) IMPORT_INFO_REDIS_KEY_PREFIX = "app_import_info:" CHECK_DEPENDENCIES_REDIS_KEY_PREFIX = "app_check_dependencies:" IMPORT_INFO_REDIS_EXPIRY = 10 * 60 # 10 minutes -DSL_MAX_SIZE = 10 * 1024 * 1024 # 10MB CURRENT_DSL_VERSION = "0.1.0" @@ -127,15 +127,16 @@ class RagPipelineDslService: yaml_url = yaml_url.replace("/blob/", "/") response = remote_fetcher.make_request("GET", yaml_url.strip(), follow_redirects=True, timeout=(10, 10)) response.raise_for_status() - content = response.content.decode() + raw_content = response.content - if len(content) > DSL_MAX_SIZE: + if dsl_content_size(raw_content) > DSL_MAX_SIZE: return RagPipelineImportInfo( id=import_id, status=ImportStatus.FAILED, error="File size exceeds the limit of 10MB", ) + content = raw_content.decode("utf-8") if not content: return RagPipelineImportInfo( id=import_id, @@ -156,6 +157,12 @@ class RagPipelineDslService: error="yaml_content is required when import_mode is yaml-content", ) content = yaml_content + if dsl_content_size(content) > DSL_MAX_SIZE: + return RagPipelineImportInfo( + id=import_id, + status=ImportStatus.FAILED, + error="File size exceeds the limit of 10MB", + ) # Process YAML content try: diff --git a/api/services/rag_pipeline/rag_pipeline_transform_service.py b/api/services/rag_pipeline/rag_pipeline_transform_service.py index daefaa9e30e..1b922b3f7b9 100644 --- a/api/services/rag_pipeline/rag_pipeline_transform_service.py +++ b/api/services/rag_pipeline/rag_pipeline_transform_service.py @@ -269,11 +269,13 @@ class RagPipelineTransformService: installed_plugins_ids = [plugin.plugin_id for plugin in installed_plugins] dependencies = pipeline_yaml.get("dependencies", []) - need_install_plugin_unique_identifiers = [] + package_identifiers_to_install = [] for dependency in dependencies: if dependency.get("type") == "marketplace": - plugin_unique_identifier = dependency.get("value", {}).get("plugin_unique_identifier") - plugin_id = plugin_unique_identifier.split(":")[0] + package_identifier = dependency.get("value", {}).get("plugin_unique_identifier") + if not package_identifier: + continue + plugin_id = package_identifier.split(":", 1)[0] if plugin_id not in installed_plugins_ids: if not dify_config.MARKETPLACE_ENABLED: logger.warning( @@ -282,12 +284,12 @@ class RagPipelineTransformService: plugin_id, ) continue - plugin_unique_identifier = plugin_migration._fetch_plugin_unique_identifier(plugin_id) # type: ignore - if plugin_unique_identifier: - need_install_plugin_unique_identifiers.append(plugin_unique_identifier) - if need_install_plugin_unique_identifiers: - logger.debug("Installing missing pipeline plugins %s", need_install_plugin_unique_identifiers) - PluginService.install_from_marketplace_pkg(tenant_id, need_install_plugin_unique_identifiers) + latest_package_identifier = plugin_migration._fetch_latest_package_identifier(plugin_id) # type: ignore + if latest_package_identifier: + package_identifiers_to_install.append(latest_package_identifier) + if package_identifiers_to_install: + logger.debug("Installing missing pipeline plugins %s", package_identifiers_to_install) + PluginService.install_from_marketplace_pkg(tenant_id, package_identifiers_to_install) def _transform_to_empty_pipeline(self, dataset: Dataset, session: Session): pipeline = Pipeline( diff --git a/api/services/snippet_dsl_service.py b/api/services/snippet_dsl_service.py index 19e0dabab95..ae1cd7f7e31 100644 --- a/api/services/snippet_dsl_service.py +++ b/api/services/snippet_dsl_service.py @@ -3,12 +3,10 @@ import logging import uuid from collections.abc import Mapping from datetime import UTC, datetime -from enum import StrEnum from urllib.parse import urlparse import yaml -from packaging import version -from pydantic import BaseModel, Field +from pydantic import BaseModel from sqlalchemy import select from sqlalchemy.orm import Session @@ -20,6 +18,9 @@ from graphon.model_runtime.utils.encoders import jsonable_encoder from models import Account from models.snippet import CustomizedSnippet, SnippetType from models.workflow import Workflow +from services.dsl_content import DSL_MAX_SIZE, dsl_content_size +from services.dsl_version import check_version_compatibility +from services.entities.dsl_entities import CheckDependenciesResult, ImportMode, ImportStatus from services.plugin.dependencies_analysis import DependenciesAnalysisService from services.snippet_service import SNIPPET_FORBIDDEN_NODE_TYPES, SnippetService @@ -28,22 +29,9 @@ logger = logging.getLogger(__name__) IMPORT_INFO_REDIS_KEY_PREFIX = "snippet_import_info:" CHECK_DEPENDENCIES_REDIS_KEY_PREFIX = "snippet_check_dependencies:" IMPORT_INFO_REDIS_EXPIRY = 10 * 60 # 10 minutes -DSL_MAX_SIZE = 10 * 1024 * 1024 # 10MB CURRENT_DSL_VERSION = "0.1.0" -class ImportMode(StrEnum): - YAML_CONTENT = "yaml-content" - YAML_URL = "yaml-url" - - -class ImportStatus(StrEnum): - COMPLETED = "completed" - COMPLETED_WITH_WARNINGS = "completed-with-warnings" - PENDING = "pending" - FAILED = "failed" - - class SnippetImportInfo(BaseModel): id: str status: ImportStatus @@ -53,32 +41,9 @@ class SnippetImportInfo(BaseModel): error: str = "" -class CheckDependenciesResult(BaseModel): - leaked_dependencies: list[PluginDependency] = Field(default_factory=list) - - def _check_version_compatibility(imported_version: str) -> ImportStatus: - """Determine import status based on version comparison""" - try: - current_ver = version.parse(CURRENT_DSL_VERSION) - imported_ver = version.parse(imported_version) - except version.InvalidVersion: - return ImportStatus.FAILED - - # If imported version is newer than current, always return PENDING - if imported_ver > current_ver: - return ImportStatus.PENDING - - # If imported version is older than current's major, return PENDING - if imported_ver.major < current_ver.major: - return ImportStatus.PENDING - - # If imported version is older than current's minor, return COMPLETED_WITH_WARNINGS - if imported_ver.minor < current_ver.minor: - return ImportStatus.COMPLETED_WITH_WARNINGS - - # If imported version equals or is older than current's micro, return COMPLETED - return ImportStatus.COMPLETED + """Determine import status based on version comparison.""" + return check_version_compatibility(imported_version, CURRENT_DSL_VERSION) class SnippetPendingData(BaseModel): @@ -145,13 +110,14 @@ class SnippetDslService: status=ImportStatus.FAILED, error=f"Failed to fetch YAML from URL: {response.status_code}", ) - content = response.text - if len(content) > DSL_MAX_SIZE: + raw_content = response.content + if dsl_content_size(raw_content) > DSL_MAX_SIZE: return SnippetImportInfo( id=import_id, status=ImportStatus.FAILED, error=f"YAML content size exceeds maximum limit of {DSL_MAX_SIZE} bytes", ) + content = raw_content.decode("utf-8") except Exception as e: logger.exception("Failed to fetch YAML from URL") return SnippetImportInfo( @@ -167,7 +133,7 @@ class SnippetDslService: error="yaml_content is required when import_mode is yaml-content", ) content = yaml_content - if len(content) > DSL_MAX_SIZE: + if dsl_content_size(content) > DSL_MAX_SIZE: return SnippetImportInfo( id=import_id, status=ImportStatus.FAILED, diff --git a/api/tests/unit_tests/services/plugin/test_plugin_migration.py b/api/tests/unit_tests/services/plugin/test_plugin_migration.py index 8f730d4ed31..fa10d775aa6 100644 --- a/api/tests/unit_tests/services/plugin/test_plugin_migration.py +++ b/api/tests/unit_tests/services/plugin/test_plugin_migration.py @@ -1,3 +1,4 @@ +import json from unittest.mock import MagicMock, patch import pytest @@ -8,17 +9,17 @@ from services.plugin.plugin_migration import PluginMigration MIGRATION_MODULE = "services.plugin.plugin_migration" -def test_fetch_plugin_unique_identifier_returns_none_when_disabled(mocker: MockerFixture) -> None: +def test_fetch_latest_package_identifier_returns_none_when_disabled(mocker: MockerFixture) -> None: mocker.patch("services.plugin.plugin_migration.dify_config.MARKETPLACE_ENABLED", False) batch_fetch = mocker.patch("services.plugin.plugin_migration.marketplace.batch_fetch_plugin_manifests") - result = PluginMigration._fetch_plugin_unique_identifier("langgenius/openai") + result = PluginMigration._fetch_latest_package_identifier("langgenius/openai") assert result is None batch_fetch.assert_not_called() -def test_fetch_plugin_unique_identifier_calls_marketplace_when_enabled(mocker: MockerFixture) -> None: +def test_fetch_latest_package_identifier_calls_marketplace_when_enabled(mocker: MockerFixture) -> None: mocker.patch("services.plugin.plugin_migration.dify_config.MARKETPLACE_ENABLED", True) manifest = mocker.MagicMock() manifest.latest_package_identifier = "langgenius/openai:1.0.0@abc" @@ -27,7 +28,7 @@ def test_fetch_plugin_unique_identifier_calls_marketplace_when_enabled(mocker: M return_value=[manifest], ) - result = PluginMigration._fetch_plugin_unique_identifier("langgenius/openai") + result = PluginMigration._fetch_latest_package_identifier("langgenius/openai") assert result == "langgenius/openai:1.0.0@abc" @@ -75,7 +76,27 @@ class TestHandlePluginInstanceInstall: mock_marketplace.download_plugin_pkg.assert_called_once() invalidate_cache.assert_called_once_with("tenant1") - assert "success" in result or "failed" in result + assert result["success"] == ["langgenius/openai"] + assert result["failed"] == [] + + def test_reports_failed_plugin_ids_when_install_batch_raises(self) -> None: + with ( + patch(f"{MIGRATION_MODULE}.dify_config") as mock_cfg, + patch(f"{MIGRATION_MODULE}.marketplace") as mock_marketplace, + patch(f"{MIGRATION_MODULE}.PluginInstaller") as mock_installer_cls, + ): + mock_cfg.MARKETPLACE_ENABLED = True + mock_marketplace.download_plugin_pkg.return_value = b"pkg_data" + mock_installer = MagicMock() + mock_installer_cls.return_value = mock_installer + mock_installer.install_from_identifiers.side_effect = RuntimeError("install failed") + + result = PluginMigration.handle_plugin_instance_install( + "tenant1", {"langgenius/openai": "langgenius/openai:1.0.0@abc"} + ) + + assert result["success"] == [] + assert result["failed"] == ["langgenius/openai"] def test_install_plugins_invalidates_cache_after_direct_tenant_install(self, tmp_path) -> None: extracted_plugins = tmp_path / "plugins.jsonl" @@ -102,3 +123,30 @@ class TestHandlePluginInstanceInstall: mock_installer.install_from_identifiers.assert_called_once() invalidate_cache.assert_called_once_with("tenant1") + + def test_install_plugins_reports_missing_plugin_ids(self, tmp_path) -> None: + extracted_plugins = tmp_path / "plugins.jsonl" + output_file = tmp_path / "output.json" + extracted_plugins.write_text('{"tenant_id":"tenant1","plugins":["langgenius/openai","langgenius/missing"]}\n') + + with ( + patch( + f"{MIGRATION_MODULE}.PluginMigration.extract_unique_plugins", + return_value={ + "plugins": {"langgenius/openai": "langgenius/openai:1.0.0@abc"}, + "plugin_not_exist": ["langgenius/missing"], + }, + ), + patch(f"{MIGRATION_MODULE}.PluginMigration.handle_plugin_instance_install", return_value={}), + patch(f"{MIGRATION_MODULE}.PluginInstaller") as mock_installer_cls, + patch(f"{MIGRATION_MODULE}.PluginService.invalidate_plugin_model_providers_cache"), + ): + mock_installer = MagicMock() + mock_installer.list_plugins.return_value = [] + mock_installer_cls.return_value = mock_installer + + PluginMigration.install_plugins(str(extracted_plugins), str(output_file), workers=1) + + assert json.loads(output_file.read_text())["not_installed"] == [ + {"tenant_id": "tenant1", "plugin_not_exist": ["langgenius/missing"]} + ] diff --git a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_dsl_service.py b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_dsl_service.py index 5bb577ba3dc..93884d07c5c 100644 --- a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_dsl_service.py +++ b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_dsl_service.py @@ -643,6 +643,19 @@ def test_import_rag_pipeline_yaml_content_requires_mapping() -> None: assert "content must be a mapping" in result.error +def test_import_rag_pipeline_rejects_oversized_yaml_content_by_bytes( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr("services.rag_pipeline.rag_pipeline_dsl_service.DSL_MAX_SIZE", 1) + service = RagPipelineDslService(session=Mock()) + account = Mock(current_tenant_id="t1") + + result = service.import_rag_pipeline(account=account, import_mode="yaml-content", yaml_content="é") + + assert result.status == ImportStatus.FAILED + assert "10MB" in result.error + + def test_confirm_import_returns_failed_when_pending_data_is_invalid_type(mocker: MockerFixture) -> None: mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.redis_client.get", return_value=object()) service = RagPipelineDslService(session=Mock()) @@ -901,6 +914,46 @@ def test_import_rag_pipeline_url_size_exceeds_limit(mocker: MockerFixture) -> No assert "10MB" in result.error +def test_import_rag_pipeline_rejects_oversized_yaml_url_bytes_before_decode( + mocker: MockerFixture, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr("services.rag_pipeline.rag_pipeline_dsl_service.DSL_MAX_SIZE", 1) + response = Mock() + response.raise_for_status.return_value = None + response.content = b"\xff\xff" + mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.remote_fetcher.make_request", return_value=response) + service = RagPipelineDslService(session=Mock()) + account = Mock(current_tenant_id="t1") + + result = service.import_rag_pipeline( + account=account, + import_mode="yaml-url", + yaml_url="https://example.com/pipeline.yaml", + ) + + assert result.status == ImportStatus.FAILED + assert "10MB" in result.error + + +def test_import_rag_pipeline_returns_decode_error_for_invalid_yaml_url_bytes(mocker: MockerFixture) -> None: + response = Mock() + response.raise_for_status.return_value = None + response.content = b"\xff" + mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.remote_fetcher.make_request", return_value=response) + service = RagPipelineDslService(session=Mock()) + account = Mock(current_tenant_id="t1") + + result = service.import_rag_pipeline( + account=account, + import_mode="yaml-url", + yaml_url="https://example.com/pipeline.yaml", + ) + + assert result.status == ImportStatus.FAILED + assert "utf-8" in result.error + + def test_import_rag_pipeline_fails_when_rag_pipeline_data_missing() -> None: service = RagPipelineDslService(session=Mock()) account = Mock(current_tenant_id="t1") diff --git a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_transform_service.py b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_transform_service.py index 21df81d3ea8..cee8e55f8cc 100644 --- a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_transform_service.py +++ b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_transform_service.py @@ -89,7 +89,7 @@ def test_deal_dependencies_installs_missing_marketplace_plugins(mocker: MockerFi installer_cls.return_value.list_plugins.return_value = [SimpleNamespace(plugin_id="installed-plugin")] migration_cls = mocker.patch("services.rag_pipeline.rag_pipeline_transform_service.PluginMigration") - migration_cls.return_value._fetch_plugin_unique_identifier.return_value = "missing-plugin:1.0.0" + migration_cls.return_value._fetch_latest_package_identifier.return_value = "missing-plugin:1.0.0" install_mock = mocker.patch( "services.rag_pipeline.rag_pipeline_transform_service.PluginService.install_from_marketplace_pkg" @@ -518,7 +518,7 @@ def test_deal_dependencies_installs_when_enabled(mocker: MockerFixture) -> None: installer = mocker.patch("services.rag_pipeline.rag_pipeline_transform_service.PluginInstaller").return_value installer.list_plugins.return_value = [] migration = mocker.patch("services.rag_pipeline.rag_pipeline_transform_service.PluginMigration").return_value - migration._fetch_plugin_unique_identifier.return_value = "langgenius/openai:1.0.0@abc" + migration._fetch_latest_package_identifier.return_value = "langgenius/openai:1.0.0@abc" install_call = mocker.patch( "services.rag_pipeline.rag_pipeline_transform_service.PluginService.install_from_marketplace_pkg" ) diff --git a/api/tests/unit_tests/services/test_app_dsl_service.py b/api/tests/unit_tests/services/test_app_dsl_service.py new file mode 100644 index 00000000000..64236ea5a90 --- /dev/null +++ b/api/tests/unit_tests/services/test_app_dsl_service.py @@ -0,0 +1,53 @@ +from types import SimpleNamespace +from unittest.mock import Mock + +from services.app_dsl_service import AppDslService, ImportStatus + + +def test_import_app_rejects_oversized_yaml_content_by_bytes(monkeypatch) -> None: + monkeypatch.setattr("services.app_dsl_service.DSL_MAX_SIZE", 1) + service = AppDslService(session=SimpleNamespace()) + + result = service.import_app( + account=SimpleNamespace(current_tenant_id="tenant-1"), + import_mode="yaml-content", + yaml_content="é", + ) + + assert result.status == ImportStatus.FAILED + assert "10MB" in result.error + + +def test_import_app_rejects_oversized_yaml_url_bytes_before_decode(monkeypatch) -> None: + monkeypatch.setattr("services.app_dsl_service.DSL_MAX_SIZE", 1) + response = Mock() + response.raise_for_status.return_value = None + response.content = b"\xff\xff" + monkeypatch.setattr("services.app_dsl_service.remote_fetcher.make_request", Mock(return_value=response)) + service = AppDslService(session=SimpleNamespace()) + + result = service.import_app( + account=SimpleNamespace(current_tenant_id="tenant-1"), + import_mode="yaml-url", + yaml_url="https://example.com/app.yaml", + ) + + assert result.status == ImportStatus.FAILED + assert "10MB" in result.error + + +def test_import_app_returns_decode_error_for_invalid_yaml_url_bytes(monkeypatch) -> None: + response = Mock() + response.raise_for_status.return_value = None + response.content = b"\xff" + monkeypatch.setattr("services.app_dsl_service.remote_fetcher.make_request", Mock(return_value=response)) + service = AppDslService(session=SimpleNamespace()) + + result = service.import_app( + account=SimpleNamespace(current_tenant_id="tenant-1"), + import_mode="yaml-url", + yaml_url="https://example.com/app.yaml", + ) + + assert result.status == ImportStatus.FAILED + assert "utf-8" in result.error diff --git a/api/tests/unit_tests/services/test_snippet_dsl_service.py b/api/tests/unit_tests/services/test_snippet_dsl_service.py index c155d3f8330..b234d9f91fe 100644 --- a/api/tests/unit_tests/services/test_snippet_dsl_service.py +++ b/api/tests/unit_tests/services/test_snippet_dsl_service.py @@ -95,7 +95,7 @@ def test_import_snippet_rejects_oversized_yaml_url_content(monkeypatch: pytest.M monkeypatch.setattr("services.snippet_dsl_service.DSL_MAX_SIZE", 3) monkeypatch.setattr( "services.snippet_dsl_service.ssrf_proxy.get", - Mock(return_value=SimpleNamespace(status_code=200, text="too large")), + Mock(return_value=SimpleNamespace(status_code=200, content=b"too large")), ) result = service.import_snippet( @@ -108,6 +108,43 @@ def test_import_snippet_rejects_oversized_yaml_url_content(monkeypatch: pytest.M assert "YAML content size exceeds maximum limit" in result.error +def test_import_snippet_rejects_oversized_yaml_url_bytes_before_decode(monkeypatch: pytest.MonkeyPatch) -> None: + service = SnippetDslService(session=SimpleNamespace()) + monkeypatch.setattr("services.snippet_dsl_service.DSL_MAX_SIZE", 1) + monkeypatch.setattr( + "services.snippet_dsl_service.ssrf_proxy.get", + Mock(return_value=SimpleNamespace(status_code=200, content=b"\xff\xff")), + ) + + result = service.import_snippet( + account=SimpleNamespace(current_tenant_id="tenant-1"), + import_mode=ImportMode.YAML_URL.value, + yaml_url="https://example.com/snippet.yaml", + ) + + assert result.status == ImportStatus.FAILED + assert "YAML content size exceeds maximum limit" in result.error + + +def test_import_snippet_returns_decode_error_for_invalid_yaml_url_bytes( + monkeypatch: pytest.MonkeyPatch, +) -> None: + service = SnippetDslService(session=SimpleNamespace()) + monkeypatch.setattr( + "services.snippet_dsl_service.ssrf_proxy.get", + Mock(return_value=SimpleNamespace(status_code=200, content=b"\xff")), + ) + + result = service.import_snippet( + account=SimpleNamespace(current_tenant_id="tenant-1"), + import_mode=ImportMode.YAML_URL.value, + yaml_url="https://example.com/snippet.yaml", + ) + + assert result.status == ImportStatus.FAILED + assert "utf-8" in result.error + + def test_import_snippet_returns_failed_when_yaml_url_fetch_raises(monkeypatch: pytest.MonkeyPatch) -> None: service = SnippetDslService(session=SimpleNamespace()) monkeypatch.setattr( @@ -127,12 +164,12 @@ def test_import_snippet_returns_failed_when_yaml_url_fetch_raises(monkeypatch: p def test_import_snippet_rejects_oversized_yaml_content(monkeypatch: pytest.MonkeyPatch) -> None: service = SnippetDslService(session=SimpleNamespace()) - monkeypatch.setattr("services.snippet_dsl_service.DSL_MAX_SIZE", 3) + monkeypatch.setattr("services.snippet_dsl_service.DSL_MAX_SIZE", 1) result = service.import_snippet( account=SimpleNamespace(current_tenant_id="tenant-1"), import_mode=ImportMode.YAML_CONTENT.value, - yaml_content="too large", + yaml_content="é", ) assert result.status == ImportStatus.FAILED From dbd3316615109157230691bb2c5638ce89c35967 Mon Sep 17 00:00:00 2001 From: Copilot <198982749+Copilot@users.noreply.github.com> Date: Tue, 7 Jul 2026 10:20:15 +0800 Subject: [PATCH 17/70] chore: set NEXT_PUBLIC_ENABLE_FEATURE_PREVIEW default to true (#38362) Co-authored-by: copilot-swe-agent[bot] <198982749+Copilot@users.noreply.github.com> --- docker/.env.example | 2 +- docker/envs/core-services/web.env.example | 2 +- web/.env.example | 2 +- web/docker/entrypoint.sh | 2 +- web/env.ts | 2 +- 5 files changed, 5 insertions(+), 5 deletions(-) diff --git a/docker/.env.example b/docker/.env.example index 746f40df56f..532c78e75ce 100644 --- a/docker/.env.example +++ b/docker/.env.example @@ -153,7 +153,7 @@ ENABLE_WEBSITE_WATERCRAWL=true NEXT_PUBLIC_ENABLE_SINGLE_DOLLAR_LATEX=false # Enable preview features still in development (currently the /create and # /refine slash commands in the "Go to Anything" command palette). -NEXT_PUBLIC_ENABLE_FEATURE_PREVIEW=false +NEXT_PUBLIC_ENABLE_FEATURE_PREVIEW=true ENABLE_AGENT_V2=false EXPERIMENTAL_ENABLE_VINEXT=false diff --git a/docker/envs/core-services/web.env.example b/docker/envs/core-services/web.env.example index bd788a1b16c..0b75ec7c5b8 100644 --- a/docker/envs/core-services/web.env.example +++ b/docker/envs/core-services/web.env.example @@ -24,7 +24,7 @@ ENABLE_WEBSITE_JINAREADER=true ENABLE_WEBSITE_FIRECRAWL=true ENABLE_WEBSITE_WATERCRAWL=true NEXT_PUBLIC_ENABLE_SINGLE_DOLLAR_LATEX=false -NEXT_PUBLIC_ENABLE_FEATURE_PREVIEW=false +NEXT_PUBLIC_ENABLE_FEATURE_PREVIEW=true ENABLE_AGENT_V2=false NEXT_PUBLIC_COOKIE_DOMAIN= NEXT_PUBLIC_BATCH_CONCURRENCY=5 diff --git a/web/.env.example b/web/.env.example index 7363ce628f3..ea80c352352 100644 --- a/web/.env.example +++ b/web/.env.example @@ -88,7 +88,7 @@ NEXT_PUBLIC_ENABLE_SINGLE_DOLLAR_LATEX=false # Enable preview features still in development (currently the /create and # /refine slash commands in the "Go to Anything" command palette) -NEXT_PUBLIC_ENABLE_FEATURE_PREVIEW=false +NEXT_PUBLIC_ENABLE_FEATURE_PREVIEW=true # Enable Agent v2 frontend entry points. NEXT_PUBLIC_ENABLE_AGENT_V2=false diff --git a/web/docker/entrypoint.sh b/web/docker/entrypoint.sh index 7fddc825610..2508af3bc61 100755 --- a/web/docker/entrypoint.sh +++ b/web/docker/entrypoint.sh @@ -57,7 +57,7 @@ export NEXT_PUBLIC_ENABLE_WEBSITE_JINAREADER=${ENABLE_WEBSITE_JINAREADER:-true} export NEXT_PUBLIC_ENABLE_WEBSITE_FIRECRAWL=${ENABLE_WEBSITE_FIRECRAWL:-true} export NEXT_PUBLIC_ENABLE_WEBSITE_WATERCRAWL=${ENABLE_WEBSITE_WATERCRAWL:-true} export NEXT_PUBLIC_ENABLE_SINGLE_DOLLAR_LATEX=${NEXT_PUBLIC_ENABLE_SINGLE_DOLLAR_LATEX:-false} -export NEXT_PUBLIC_ENABLE_FEATURE_PREVIEW=${NEXT_PUBLIC_ENABLE_FEATURE_PREVIEW:-false} +export NEXT_PUBLIC_ENABLE_FEATURE_PREVIEW=${NEXT_PUBLIC_ENABLE_FEATURE_PREVIEW:-true} export NEXT_PUBLIC_ENABLE_AGENT_V2=${NEXT_PUBLIC_ENABLE_AGENT_V2:-${ENABLE_AGENT_V2:-false}} export NEXT_PUBLIC_LOOP_NODE_MAX_COUNT=${LOOP_NODE_MAX_COUNT} export NEXT_PUBLIC_MAX_PARALLEL_LIMIT=${MAX_PARALLEL_LIMIT} diff --git a/web/env.ts b/web/env.ts index 15d93fa11ea..bcfd55e8a00 100644 --- a/web/env.ts +++ b/web/env.ts @@ -69,7 +69,7 @@ const clientSchema = { * Currently gates the `/create` and `/refine` slash commands in the * "Go to Anything" command palette (Cmd/Ctrl+K). */ - NEXT_PUBLIC_ENABLE_FEATURE_PREVIEW: coercedBoolean.default(false), + NEXT_PUBLIC_ENABLE_FEATURE_PREVIEW: coercedBoolean.default(true), /** * Cloud-only system-features defaults. From 5a342f9258148b1af4a63f0f0066351501f42974 Mon Sep 17 00:00:00 2001 From: Charles Yao Date: Tue, 7 Jul 2026 05:26:50 +0200 Subject: [PATCH 18/70] feat(mcp): support MCP protocol 2025-06-18 for workflow-as-MCP server (version negotiation + structured output) (#37892) Co-authored-by: Claude Sonnet 4.6 Co-authored-by: yunlu.wen --- api/controllers/mcp/mcp.py | 53 +++- api/core/mcp/server/streamable_http.py | 112 +++++++- api/core/mcp/types.py | 9 +- .../controllers/mcp/test_mcp.py | 222 ++++++++++++++- .../core/mcp/server/test_streamable_http.py | 258 +++++++++++++++++- api/tests/unit_tests/core/mcp/test_types.py | 6 +- 6 files changed, 628 insertions(+), 32 deletions(-) diff --git a/api/controllers/mcp/mcp.py b/api/controllers/mcp/mcp.py index 3830c9585a5..380ca470c47 100644 --- a/api/controllers/mcp/mcp.py +++ b/api/controllers/mcp/mcp.py @@ -1,6 +1,6 @@ from typing import Any, Union -from flask import Response +from flask import Response, request from flask_restx import Resource from pydantic import BaseModel, Field, ValidationError from sqlalchemy import select @@ -9,7 +9,7 @@ from sqlalchemy.orm import Session, sessionmaker from controllers.common.schema import register_schema_model from controllers.mcp import mcp_ns from core.mcp import types as mcp_types -from core.mcp.server.streamable_http import handle_mcp_request +from core.mcp.server.streamable_http import handle_mcp_request, negotiate_protocol_version from extensions.ext_database import db from graphon.variables.input_entities import VariableEntity, VariableEntityType from libs import helper @@ -68,6 +68,17 @@ class MCPAppApi(Resource): request_id: Union[int, str] | None = args.id mcp_request = self._parse_mcp_request(args.model_dump(exclude_none=True)) + # Resolve the negotiated protocol version from the MCP-Protocol-Version header. + is_initialize = isinstance(mcp_request.root, mcp_types.InitializeRequest) + header_value = request.headers.get("MCP-Protocol-Version") + protocol_version = negotiate_protocol_version(header_value, is_initialize) + if protocol_version is None: + # A notification never receives a response, even with an unsupported header. + if isinstance(mcp_request, mcp_types.ClientNotification): + protocol_version = mcp_types.DEFAULT_NEGOTIATED_VERSION + else: + return self._protocol_version_error_response(request_id, header_value) + with sessionmaker(db.engine, expire_on_commit=False).begin() as session: # Get MCP server and app mcp_server, app = self._get_mcp_server_and_app(server_code, session) @@ -77,7 +88,28 @@ class MCPAppApi(Resource): user_input_form = self._get_user_input_form(app) # Handle notification vs request differently - return self._process_mcp_message(mcp_request, request_id, app, mcp_server, user_input_form, session) + return self._process_mcp_message( + mcp_request, request_id, app, mcp_server, user_input_form, session, protocol_version + ) + + def _protocol_version_error_response( + self, request_id: Union[int, str] | None, header_value: str | None + ) -> Response: + """Return a JSON-RPC error for an unsupported MCP-Protocol-Version header. + + Per JSON-RPC 2.0, an error whose request id is unknown uses a null id, so we echo the + offending request's id directly (None -> null) instead of fabricating a placeholder. + """ + error_data = mcp_types.ErrorData( + code=mcp_types.INVALID_REQUEST, + message=f"Unsupported MCP-Protocol-Version: {header_value}", + ) + error_response = { + "jsonrpc": "2.0", + "id": request_id, + "error": error_data.model_dump(by_alias=True, mode="json", exclude_none=True), + } + return helper.compact_generate_response(error_response) def _get_mcp_server_and_app(self, server_code: str, session: Session) -> tuple[AppMCPServer, App]: """Get and validate MCP server and app in one query session""" @@ -104,12 +136,15 @@ class MCPAppApi(Resource): mcp_server: AppMCPServer, user_input_form: list[VariableEntity], session: Session, + protocol_version: str, ) -> Response: """Process MCP message (notification or request)""" if isinstance(mcp_request, mcp_types.ClientNotification): return self._handle_notification(mcp_request) else: - return self._handle_request(mcp_request, request_id, app, mcp_server, user_input_form, session) + return self._handle_request( + mcp_request, request_id, app, mcp_server, user_input_form, session, protocol_version + ) def _handle_notification(self, mcp_request: mcp_types.ClientNotification) -> Response: """Handle MCP notification""" @@ -127,12 +162,15 @@ class MCPAppApi(Resource): mcp_server: AppMCPServer, user_input_form: list[VariableEntity], session: Session, + protocol_version: str, ) -> Response: """Handle MCP request""" if request_id is None: raise MCPRequestError(mcp_types.INVALID_REQUEST, "Request ID is required") - result = self._handle_mcp_request(app, mcp_server, mcp_request, user_input_form, session, request_id) + result = self._handle_mcp_request( + app, mcp_server, mcp_request, user_input_form, session, request_id, protocol_version + ) if result is None: # This shouldn't happen for requests, but handle gracefully raise MCPRequestError(mcp_types.INTERNAL_ERROR, "No response generated for request") @@ -229,6 +267,7 @@ class MCPAppApi(Resource): user_input_form: list[VariableEntity], session: Session, request_id: Union[int, str], + protocol_version: str, ) -> mcp_types.JSONRPCResponse | mcp_types.JSONRPCError | None: """Handle MCP request and return response""" end_user = self._retrieve_end_user(mcp_server.tenant_id, mcp_server.id) @@ -238,4 +277,6 @@ class MCPAppApi(Resource): client_name = f"{client_info.name}@{client_info.version}" end_user = self._create_end_user(client_name, app.tenant_id, app.id, mcp_server.id, session) - return handle_mcp_request(session, app, mcp_request, user_input_form, mcp_server, end_user, request_id) + return handle_mcp_request( + session, app, mcp_request, user_input_form, mcp_server, end_user, request_id, protocol_version + ) diff --git a/api/core/mcp/server/streamable_http.py b/api/core/mcp/server/streamable_http.py index 3bb75e485a8..964f3211db0 100644 --- a/api/core/mcp/server/streamable_http.py +++ b/api/core/mcp/server/streamable_http.py @@ -15,6 +15,36 @@ from services.app_generate_service import AppGenerateService logger = logging.getLogger(__name__) +# Structured tool output (outputSchema + structuredContent) was introduced in MCP 2025-06-18. +STRUCTURED_OUTPUT_MIN_VERSION = "2025-06-18" + + +def _supports_structured_output(protocol_version: str) -> bool: + """Return True when the negotiated protocol version supports structured tool output. + + MCP protocol versions are YYYY-MM-DD strings, so lexical comparison equals chronological. + """ + return protocol_version >= STRUCTURED_OUTPUT_MIN_VERSION + + +def negotiate_protocol_version(header_value: str | None, is_initialize: bool) -> str | None: + """Resolve the negotiated protocol version for an incoming MCP request. + + The version is taken from the MCP-Protocol-Version header on post-initialize requests. + Returns the version to use for behavior gating, or None when the client sent an explicit + but unsupported header (the caller should reply with a JSON-RPC INVALID_REQUEST error). + Initialize requests negotiate via the request body, so they always receive + DEFAULT_NEGOTIATED_VERSION and their header is never validated or rejected. + """ + if is_initialize: + return mcp_types.DEFAULT_NEGOTIATED_VERSION + # Treat an absent or empty header as "not specified" -> default version. + if not header_value: + return mcp_types.DEFAULT_NEGOTIATED_VERSION + if header_value not in mcp_types.SERVER_SUPPORTED_PROTOCOL_VERSIONS: + return None + return header_value + class ToolParameterSchemaDict(TypedDict): type: str @@ -35,6 +65,7 @@ def handle_mcp_request( mcp_server: AppMCPServer, end_user: EndUser | None = None, request_id: int | str = 1, + protocol_version: str = mcp_types.DEFAULT_NEGOTIATED_VERSION, ) -> mcp_types.JSONRPCResponse | mcp_types.JSONRPCError: """ Handle MCP request and return JSON-RPC response @@ -77,15 +108,24 @@ def handle_mcp_request( # Dispatch request to appropriate handler based on instance type match request_root: case mcp_types.InitializeRequest(): - return create_success_response(handle_initialize(mcp_server.description)) + return create_success_response( + handle_initialize(mcp_server.description, request_root.params.protocolVersion) + ) case mcp_types.ListToolsRequest(): return create_success_response( handle_list_tools( - app.name, app.mode, user_input_form, mcp_server.description, mcp_server.parameters_dict + app.name, + app.mode, + user_input_form, + mcp_server.description, + mcp_server.parameters_dict, + protocol_version, ) ) case mcp_types.CallToolRequest(): - return create_success_response(handle_call_tool(session, app, request, user_input_form, end_user)) + return create_success_response( + handle_call_tool(session, app, request, user_input_form, end_user, protocol_version) + ) case mcp_types.PingRequest(): return create_success_response(handle_ping()) case _: @@ -104,14 +144,22 @@ def handle_ping() -> mcp_types.EmptyResult: return mcp_types.EmptyResult() -def handle_initialize(description: str) -> mcp_types.InitializeResult: - """Handle initialize request""" +def handle_initialize(description: str, requested_version: str | int) -> mcp_types.InitializeResult: + """Handle initialize request, negotiating the protocol version with the client. + + Echoes the client's requested version when the server supports it, otherwise returns the + server's latest supported version (per the MCP lifecycle spec). + """ + negotiated_version: str = mcp_types.SERVER_LATEST_PROTOCOL_VERSION + if isinstance(requested_version, str) and requested_version in mcp_types.SERVER_SUPPORTED_PROTOCOL_VERSIONS: + negotiated_version = requested_version + capabilities = mcp_types.ServerCapabilities( tools=mcp_types.ToolsCapability(listChanged=False), ) return mcp_types.InitializeResult( - protocolVersion=mcp_types.SERVER_LATEST_PROTOCOL_VERSION, + protocolVersion=negotiated_version, capabilities=capabilities, serverInfo=mcp_types.Implementation(name="Dify", version=dify_config.project.version), instructions=description, @@ -124,19 +172,23 @@ def handle_list_tools( user_input_form: list[VariableEntity], description: str, parameters_dict: dict[str, str], + protocol_version: str = mcp_types.DEFAULT_NEGOTIATED_VERSION, ) -> mcp_types.ListToolsResult: """Handle list tools request""" parameter_schema = build_parameter_schema(app_mode, user_input_form, parameters_dict) + supports_structured = _supports_structured_output(protocol_version) - return mcp_types.ListToolsResult( - tools=[ - mcp_types.Tool( - name=app_name, - description=description, - inputSchema=cast(dict[str, Any], parameter_schema), - ) - ], + # For 2025-06-18+ clients, expose an explicit display title and a permissive output + # schema. Both stay None (and are stripped by exclude_none serialization) for older + # clients, so their tool definition is unchanged. + tool = mcp_types.Tool( + name=app_name, + title=app_name if supports_structured else None, + description=description, + inputSchema=cast(dict[str, Any], parameter_schema), + outputSchema={"type": "object"} if supports_structured else None, ) + return mcp_types.ListToolsResult(tools=[tool]) def handle_call_tool( @@ -145,6 +197,7 @@ def handle_call_tool( request: mcp_types.ClientRequest, user_input_form: list[VariableEntity], end_user: EndUser | None, + protocol_version: str = mcp_types.DEFAULT_NEGOTIATED_VERSION, ) -> mcp_types.CallToolResult: """Handle call tool request""" request_obj = cast(mcp_types.CallToolRequest, request.root) @@ -163,7 +216,13 @@ def handle_call_tool( ) answer = extract_answer_from_response(app, response) - return mcp_types.CallToolResult(content=[mcp_types.TextContent(text=answer, type="text")]) + structured_content = None + if _supports_structured_output(protocol_version): + structured_content = extract_structured_output(app, response, answer) + return mcp_types.CallToolResult( + content=[mcp_types.TextContent(text=answer, type="text")], + structuredContent=structured_content, + ) def build_parameter_schema( @@ -204,6 +263,29 @@ def prepare_tool_arguments(app: App, arguments: dict[str, Any]) -> ToolArguments return {"query": query, "inputs": args_copy} +def extract_structured_output(app: App, response: Any, answer: str) -> dict[str, Any] | None: + """Build MCP structured tool output (2025-06-18) from the app response. + + WORKFLOW mode exposes the raw outputs mapping; chat/agent/completion modes expose the + answer string under an "answer" key. Returns None when no structured output is available. + """ + match app.mode: + case AppMode.WORKFLOW: + if isinstance(response, Mapping): + data = response.get("data") + if isinstance(data, Mapping): + outputs = data.get("outputs") + # All three guards use Mapping for consistency; coerce to a concrete dict + # because structuredContent must be a JSON object (dict[str, Any]). + if isinstance(outputs, Mapping): + return dict(outputs) + return None + case AppMode.ADVANCED_CHAT | AppMode.CHAT | AppMode.AGENT_CHAT | AppMode.COMPLETION: + return {"answer": answer} + case _: + return None + + def extract_answer_from_response(app: App, response: Any) -> str: """Extract answer from app generate response""" answer = "" diff --git a/api/core/mcp/types.py b/api/core/mcp/types.py index 9470d39f414..8d1e6587c5f 100644 --- a/api/core/mcp/types.py +++ b/api/core/mcp/types.py @@ -22,10 +22,13 @@ for reference. * Define additional model classes instead of using dictionaries. Do this even if they're not separate types in the schema. """ -# Client support both version, not support 2025-06-18 yet. +# Latest protocol version the Dify MCP client negotiates with upstream MCP servers. LATEST_PROTOCOL_VERSION = "2025-06-18" -# Server support 2024-11-05 to allow claude to use. -SERVER_LATEST_PROTOCOL_VERSION = "2024-11-05" +# Latest protocol version the Dify MCP server advertises to connecting clients. +SERVER_LATEST_PROTOCOL_VERSION = "2025-06-18" +# Protocol versions the Dify MCP server can negotiate down to (e.g. Claude on 2024-11-05). +SERVER_SUPPORTED_PROTOCOL_VERSIONS: frozenset[str] = frozenset({"2024-11-05", "2025-03-26", "2025-06-18"}) +# Version assumed when a client omits the MCP-Protocol-Version header on post-initialize requests. DEFAULT_NEGOTIATED_VERSION = "2025-03-26" ProgressToken = str | int Cursor = str diff --git a/api/tests/test_containers_integration_tests/controllers/mcp/test_mcp.py b/api/tests/test_containers_integration_tests/controllers/mcp/test_mcp.py index c281f071560..4c9d2e28116 100644 --- a/api/tests/test_containers_integration_tests/controllers/mcp/test_mcp.py +++ b/api/tests/test_containers_integration_tests/controllers/mcp/test_mcp.py @@ -8,7 +8,7 @@ from unittest.mock import MagicMock, patch from uuid import uuid4 import pytest -from flask import Response +from flask import Flask, Response from pydantic import ValidationError import controllers.mcp.mcp as module @@ -37,12 +37,15 @@ class DummyServer: self.app_id = app_id self.tenant_id = tenant_id self.id = server_id + self.description = "Test server" + self.parameters_dict = {} class DummyApp: def __init__(self, mode, workflow=None, app_model_config=None): self.id = _APP_ID self.tenant_id = _TENANT_ID + self.name = "test_app" self.mode = mode self.workflow = workflow self.app_model_config = app_model_config @@ -494,3 +497,220 @@ class TestMCPAppApi: with pytest.raises(module.MCPRequestError) as exc_info: post_fn("server-1") assert "Invalid user_input_form" in str(exc_info.value) + + +_UNSUPPORTED_VERSION = "1999-01-01" + + +def _initialize_payload(protocol_version: str = "2024-11-05") -> dict[str, object]: + return { + "jsonrpc": "2.0", + "method": "initialize", + "id": 1, + "params": { + "protocolVersion": protocol_version, + "capabilities": {}, + "clientInfo": {"name": "test-client", "version": "1.0"}, + }, + } + + +def _tools_list_payload(request_id: int | None = 1) -> dict[str, object]: + payload: dict[str, object] = {"jsonrpc": "2.0", "method": "tools/list", "params": {}} + if request_id is not None: + payload["id"] = request_id + return payload + + +def _tools_call_payload() -> dict[str, object]: + return { + "jsonrpc": "2.0", + "method": "tools/call", + "id": 1, + "params": {"name": "test_app", "arguments": {"query": "test question"}}, + } + + +class TestMCPProtocolVersionNegotiationApi: + """MCP protocol version negotiation exercised through the HTTP controller layer. + + Covers the MCP-Protocol-Version header contract (resolution, rejection, threading) + and the serialized JSON responses seen by modern (2025-06-18) vs legacy (2024-11-05) + clients, including the back-compat guarantee that legacy responses carry none of the + structured-output fields. + """ + + def _make_api(self) -> module.MCPAppApi: + server = DummyServer(status=module.AppMCPServerStatus.ACTIVE) + app = DummyApp(mode=module.AppMode.CHAT, app_model_config=DummyConfig()) + api = module.MCPAppApi() + api._get_mcp_server_and_app = MagicMock(return_value=(server, app)) + api._retrieve_end_user = MagicMock(return_value=MagicMock()) + return api + + def _post( + self, flask_app: Flask, api: module.MCPAppApi, payload: dict[str, object], headers: dict[str, str] | None = None + ) -> Response: + fake_payload(payload) + post_fn = unwrap(api.post) + with flask_app.test_request_context(headers=headers): + return post_fn("server-1") + + @pytest.mark.parametrize("version", sorted(module.mcp_types.SERVER_SUPPORTED_PROTOCOL_VERSIONS)) + def test_initialize_echoes_supported_body_version(self, flask_app_with_containers, version): + """Initialize echoes every supported client-requested version back unchanged.""" + api = self._make_api() + + response = self._post(flask_app_with_containers, api, _initialize_payload(version)) + + body = response.get_json() + assert body["result"]["protocolVersion"] == version + + def test_initialize_falls_back_for_unsupported_body_version(self, flask_app_with_containers): + """An unsupported requested version falls back to the server latest.""" + api = self._make_api() + + response = self._post(flask_app_with_containers, api, _initialize_payload(_UNSUPPORTED_VERSION)) + + body = response.get_json() + assert body["result"]["protocolVersion"] == module.mcp_types.SERVER_LATEST_PROTOCOL_VERSION + + def test_initialize_ignores_unsupported_header(self, flask_app_with_containers): + """Initialize negotiates via the request body, so its header is never rejected.""" + api = self._make_api() + + response = self._post( + flask_app_with_containers, + api, + _initialize_payload("2024-11-05"), + headers={"MCP-Protocol-Version": _UNSUPPORTED_VERSION}, + ) + + body = response.get_json() + assert "error" not in body + assert body["result"]["protocolVersion"] == "2024-11-05" + + @pytest.mark.parametrize("request_id", [5, None]) + def test_unsupported_header_returns_invalid_request_error(self, flask_app_with_containers, request_id): + """An unsupported header gets a JSON-RPC error echoing the request id (missing id -> null).""" + api = self._make_api() + + with patch.object(module, "handle_mcp_request", autospec=True) as mock_handle: + response = self._post( + flask_app_with_containers, + api, + _tools_list_payload(request_id=request_id), + headers={"MCP-Protocol-Version": _UNSUPPORTED_VERSION}, + ) + + body = response.get_json() + assert response.status_code == 200 + assert body["jsonrpc"] == "2.0" + assert body["id"] == request_id + assert body["error"]["code"] == module.mcp_types.INVALID_REQUEST + assert _UNSUPPORTED_VERSION in body["error"]["message"] + mock_handle.assert_not_called() + + def test_notification_with_unsupported_header_is_accepted(self, flask_app_with_containers): + """A notification is accepted (202, no body) even with an unsupported header.""" + api = self._make_api() + + response = self._post( + flask_app_with_containers, + api, + {"jsonrpc": "2.0", "method": "notifications/initialized", "params": {}}, + headers={"MCP-Protocol-Version": _UNSUPPORTED_VERSION}, + ) + + assert response.status_code == 202 + + @pytest.mark.parametrize("version", sorted(module.mcp_types.SERVER_SUPPORTED_PROTOCOL_VERSIONS)) + def test_supported_header_is_threaded_to_handler(self, flask_app_with_containers, version): + """Every supported header value is passed through to handle_mcp_request.""" + api = self._make_api() + + with patch.object(module, "handle_mcp_request", return_value=DummyResult(), autospec=True) as mock_handle: + self._post( + flask_app_with_containers, + api, + _tools_list_payload(), + headers={"MCP-Protocol-Version": version}, + ) + + assert mock_handle.call_args.args[-1] == version + + def test_absent_header_defaults_to_back_compat_version(self, flask_app_with_containers): + """An absent header resolves to the spec's default version (2025-03-26).""" + api = self._make_api() + + with patch.object(module, "handle_mcp_request", return_value=DummyResult(), autospec=True) as mock_handle: + self._post(flask_app_with_containers, api, _tools_list_payload()) + + assert mock_handle.call_args.args[-1] == module.mcp_types.DEFAULT_NEGOTIATED_VERSION + + def test_tools_list_json_advertises_structured_output_for_modern_client(self, flask_app_with_containers): + """A 2025-06-18 client sees outputSchema and title in the serialized tool JSON.""" + api = self._make_api() + + response = self._post( + flask_app_with_containers, + api, + _tools_list_payload(), + headers={"MCP-Protocol-Version": "2025-06-18"}, + ) + + tool = response.get_json()["result"]["tools"][0] + assert tool["outputSchema"] == {"type": "object"} + assert tool["title"] == "test_app" + + def test_tools_list_json_unchanged_for_legacy_client(self, flask_app_with_containers): + """A 2024-11-05 client sees exactly the pre-upgrade tool JSON keys.""" + api = self._make_api() + + response = self._post( + flask_app_with_containers, + api, + _tools_list_payload(), + headers={"MCP-Protocol-Version": "2024-11-05"}, + ) + + tool = response.get_json()["result"]["tools"][0] + assert set(tool) == {"name", "description", "inputSchema"} + + @patch("core.mcp.server.streamable_http.AppGenerateService") + def test_tools_call_json_includes_structured_content_for_modern_client( + self, mock_app_generate, flask_app_with_containers + ): + """A 2025-06-18 client receives structuredContent alongside the text content.""" + api = self._make_api() + mock_app_generate.generate.return_value = {"answer": "test answer"} + + response = self._post( + flask_app_with_containers, + api, + _tools_call_payload(), + headers={"MCP-Protocol-Version": "2025-06-18"}, + ) + + result = response.get_json()["result"] + assert result["structuredContent"] == {"answer": "test answer"} + assert result["content"][0]["text"] == "test answer" + + @patch("core.mcp.server.streamable_http.AppGenerateService") + def test_tools_call_json_omits_structured_content_for_legacy_client( + self, mock_app_generate, flask_app_with_containers + ): + """A 2024-11-05 client receives the pre-upgrade tools/call JSON without structuredContent.""" + api = self._make_api() + mock_app_generate.generate.return_value = {"answer": "test answer"} + + response = self._post( + flask_app_with_containers, + api, + _tools_call_payload(), + headers={"MCP-Protocol-Version": "2024-11-05"}, + ) + + result = response.get_json()["result"] + assert "structuredContent" not in result + assert result["content"][0]["text"] == "test answer" diff --git a/api/tests/unit_tests/core/mcp/server/test_streamable_http.py b/api/tests/unit_tests/core/mcp/server/test_streamable_http.py index 3f95736ef92..cc00e02252e 100644 --- a/api/tests/unit_tests/core/mcp/server/test_streamable_http.py +++ b/api/tests/unit_tests/core/mcp/server/test_streamable_http.py @@ -10,11 +10,13 @@ from core.mcp.server.streamable_http import ( build_parameter_schema, convert_input_form_to_parameters, extract_answer_from_response, + extract_structured_output, handle_call_tool, handle_initialize, handle_list_tools, handle_mcp_request, handle_ping, + negotiate_protocol_version, prepare_tool_arguments, process_mapping_response, ) @@ -64,6 +66,8 @@ class TestHandleMCPRequest: # Setup initialize request self.mock_request.root = Mock(spec=types.InitializeRequest) self.mock_request.root.id = 123 + self.mock_request.root.params = Mock() + self.mock_request.root.params.protocolVersion = "2025-06-18" request_type = Mock(return_value=types.InitializeRequest) with patch("core.mcp.server.streamable_http.type", request_type): @@ -91,6 +95,33 @@ class TestHandleMCPRequest: assert result.jsonrpc == "2.0" assert result.id == 123 + def test_handle_list_tools_request_threads_protocol_version(self): + """The negotiated version reaches handle_list_tools through the dispatcher.""" + self.mock_request.root = Mock(spec=types.ListToolsRequest) + self.mock_request.root.id = 123 + + result = handle_mcp_request( + Mock(), self.app, self.mock_request, self.user_input_form, self.mcp_server, self.end_user, 123, "2025-06-18" + ) + + assert isinstance(result, types.JSONRPCResponse) + tool = result.result["tools"][0] + assert tool["outputSchema"] == {"type": "object"} + assert tool["title"] == "test_app" + + def test_handle_list_tools_request_legacy_serialization_unchanged(self): + """A 2024-11-05 tools/list response serializes without any 2025-06-18 fields.""" + self.mock_request.root = Mock(spec=types.ListToolsRequest) + self.mock_request.root.id = 123 + + result = handle_mcp_request( + Mock(), self.app, self.mock_request, self.user_input_form, self.mcp_server, self.end_user, 123, "2024-11-05" + ) + + assert isinstance(result, types.JSONRPCResponse) + tool = result.result["tools"][0] + assert set(tool) == {"name", "description", "inputSchema"} + @patch("core.mcp.server.streamable_http.AppGenerateService") def test_handle_call_tool_request(self, mock_app_generate): """Test handling call tool request""" @@ -119,6 +150,43 @@ class TestHandleMCPRequest: # Verify AppGenerateService was called mock_app_generate.generate.assert_called_once() + @patch("core.mcp.server.streamable_http.AppGenerateService") + def test_handle_call_tool_request_threads_protocol_version(self, mock_app_generate): + """The negotiated version reaches handle_call_tool through the dispatcher.""" + mock_call_request = Mock(spec=types.CallToolRequest) + mock_call_request.params = Mock() + mock_call_request.params.arguments = {"query": "test question"} + mock_call_request.id = 123 + self.mock_request.root = mock_call_request + + mock_app_generate.generate.return_value = {"answer": "test answer"} + + result = handle_mcp_request( + Mock(), self.app, self.mock_request, self.user_input_form, self.mcp_server, self.end_user, 123, "2025-06-18" + ) + + assert isinstance(result, types.JSONRPCResponse) + assert result.result["structuredContent"] == {"answer": "test answer"} + + @patch("core.mcp.server.streamable_http.AppGenerateService") + def test_handle_call_tool_request_legacy_serialization_unchanged(self, mock_app_generate): + """A 2024-11-05 tools/call response serializes without structuredContent.""" + mock_call_request = Mock(spec=types.CallToolRequest) + mock_call_request.params = Mock() + mock_call_request.params.arguments = {"query": "test question"} + mock_call_request.id = 123 + self.mock_request.root = mock_call_request + + mock_app_generate.generate.return_value = {"answer": "test answer"} + + result = handle_mcp_request( + Mock(), self.app, self.mock_request, self.user_input_form, self.mcp_server, self.end_user, 123, "2024-11-05" + ) + + assert isinstance(result, types.JSONRPCResponse) + assert "structuredContent" not in result.result + assert result.result["content"][0]["text"] == "test answer" + def test_handle_unknown_request_type(self): """Test handling unknown request type""" @@ -183,18 +251,49 @@ class TestIndividualHandlers: result = handle_ping() assert isinstance(result, types.EmptyResult) - def test_handle_initialize(self): - """Test initialize handler""" - description = "Test server" - + def test_handle_initialize_echoes_supported_version(self): + """A supported requested version is echoed back unchanged.""" with patch("core.mcp.server.streamable_http.dify_config") as mock_config: mock_config.project.version = "1.0.0" - result = handle_initialize(description) + result = handle_initialize("Test server", "2024-11-05") assert isinstance(result, types.InitializeResult) - assert result.protocolVersion == types.SERVER_LATEST_PROTOCOL_VERSION + assert result.protocolVersion == "2024-11-05" assert result.instructions == "Test server" + def test_handle_initialize_echoes_intermediate_version(self): + """The intermediate supported version (2025-03-26) is echoed back.""" + with patch("core.mcp.server.streamable_http.dify_config") as mock_config: + mock_config.project.version = "1.0.0" + result = handle_initialize("Test server", "2025-03-26") + + assert result.protocolVersion == "2025-03-26" + + def test_handle_initialize_negotiates_latest_for_modern_client(self): + """A 2025-06-18 client gets 2025-06-18 back.""" + with patch("core.mcp.server.streamable_http.dify_config") as mock_config: + mock_config.project.version = "1.0.0" + result = handle_initialize("Test server", "2025-06-18") + + assert result.protocolVersion == "2025-06-18" + + def test_handle_initialize_falls_back_for_unknown_version(self): + """An unsupported requested version falls back to the server latest.""" + with patch("core.mcp.server.streamable_http.dify_config") as mock_config: + mock_config.project.version = "1.0.0" + result = handle_initialize("Test server", "1999-01-01") + + assert result.protocolVersion == types.SERVER_LATEST_PROTOCOL_VERSION + assert result.protocolVersion == "2025-06-18" + + def test_handle_initialize_non_string_version_falls_back(self): + """A malformed (non-string) requested version falls back to the server latest.""" + with patch("core.mcp.server.streamable_http.dify_config") as mock_config: + mock_config.project.version = "1.0.0" + result = handle_initialize("Test server", 20250618) + + assert result.protocolVersion == types.SERVER_LATEST_PROTOCOL_VERSION + def test_handle_list_tools(self): """Test list tools handler""" app_name = "test_app" @@ -210,6 +309,30 @@ class TestIndividualHandlers: assert result.tools[0].name == "test_app" assert result.tools[0].description == "Test server" + def test_handle_list_tools_adds_structured_output_for_modern_client(self): + """Tool advertises outputSchema and title when negotiated >= 2025-06-18.""" + result = handle_list_tools("test_app", AppMode.CHAT, [], "Test server", {}, "2025-06-18") + + tool = result.tools[0] + assert tool.outputSchema == {"type": "object"} + assert tool.title == "test_app" + + def test_handle_list_tools_omits_structured_output_for_legacy_client(self): + """Tool stays unchanged (no outputSchema/title) for 2024-11-05 clients.""" + result = handle_list_tools("test_app", AppMode.CHAT, [], "Test server", {}, "2024-11-05") + + tool = result.tools[0] + assert tool.outputSchema is None + assert tool.title is None + + def test_handle_list_tools_omits_structured_output_for_intermediate_client(self): + """The 2025-03-26 negotiated version is below the structured-output threshold.""" + result = handle_list_tools("test_app", AppMode.CHAT, [], "Test server", {}, "2025-03-26") + + tool = result.tools[0] + assert tool.outputSchema is None + assert tool.title is None + @patch("core.mcp.server.streamable_http.AppGenerateService") def test_handle_call_tool(self, mock_app_generate): """Test call tool handler""" @@ -239,6 +362,44 @@ class TestIndividualHandlers: assert hasattr(text_content, "text") assert text_content.text == "test answer" + @patch("core.mcp.server.streamable_http.AppGenerateService") + def test_handle_call_tool_structured_output_modern_client(self, mock_app_generate): + """structuredContent is attached alongside TextContent for >= 2025-06-18.""" + app = Mock(spec=App) + app.mode = AppMode.CHAT + + mock_request = Mock() + mock_call_request = Mock(spec=types.CallToolRequest) + mock_call_request.params = Mock() + mock_call_request.params.arguments = {"query": "test question"} + mock_request.root = mock_call_request + + mock_app_generate.generate.return_value = {"answer": "test answer"} + + result = handle_call_tool(Mock(), app, mock_request, [], Mock(spec=EndUser), "2025-06-18") + + assert result.structuredContent == {"answer": "test answer"} + assert result.content[0].text == "test answer" + + @patch("core.mcp.server.streamable_http.AppGenerateService") + def test_handle_call_tool_no_structured_output_legacy_client(self, mock_app_generate): + """structuredContent is omitted for 2024-11-05 clients.""" + app = Mock(spec=App) + app.mode = AppMode.CHAT + + mock_request = Mock() + mock_call_request = Mock(spec=types.CallToolRequest) + mock_call_request.params = Mock() + mock_call_request.params.arguments = {"query": "test question"} + mock_request.root = mock_call_request + + mock_app_generate.generate.return_value = {"answer": "test answer"} + + result = handle_call_tool(Mock(), app, mock_request, [], Mock(spec=EndUser), "2024-11-05") + + assert result.structuredContent is None + assert result.content[0].text == "test answer" + def test_handle_call_tool_no_end_user(self): """Test call tool handler without end user""" app = Mock(spec=App) @@ -375,6 +536,65 @@ class TestUtilityFunctions: assert result == "thinking...more thinking" + def test_extract_structured_output_workflow(self): + """Workflow mode exposes the raw outputs mapping as structured content.""" + app = Mock(spec=App) + app.mode = AppMode.WORKFLOW + + response = {"data": {"outputs": {"result": "test result"}}} + + assert extract_structured_output(app, response, "ignored") == {"result": "test result"} + + def test_extract_structured_output_chat(self): + """Chat mode wraps the answer string under an 'answer' key.""" + app = Mock(spec=App) + app.mode = AppMode.CHAT + + assert extract_structured_output(app, {"answer": "hi"}, "hi") == {"answer": "hi"} + + def test_extract_structured_output_workflow_missing_outputs(self): + """Missing or malformed outputs fall back to None.""" + app = Mock(spec=App) + app.mode = AppMode.WORKFLOW + + assert extract_structured_output(app, {"data": {}}, "ignored") is None + + def test_extract_structured_output_workflow_non_mapping_response(self): + """A non-mapping workflow response yields no structured output.""" + app = Mock(spec=App) + app.mode = AppMode.WORKFLOW + + assert extract_structured_output(app, None, "ignored") is None + + def test_extract_structured_output_workflow_non_mapping_data(self): + """A non-mapping 'data' entry yields no structured output.""" + app = Mock(spec=App) + app.mode = AppMode.WORKFLOW + + assert extract_structured_output(app, {"data": "not a mapping"}, "ignored") is None + + def test_extract_structured_output_workflow_non_mapping_outputs(self): + """A non-mapping 'outputs' entry yields no structured output.""" + app = Mock(spec=App) + app.mode = AppMode.WORKFLOW + + assert extract_structured_output(app, {"data": {"outputs": ["not", "a", "mapping"]}}, "ignored") is None + + @pytest.mark.parametrize("mode", [AppMode.ADVANCED_CHAT, AppMode.AGENT_CHAT, AppMode.COMPLETION]) + def test_extract_structured_output_other_answer_modes(self, mode): + """Every chat-style mode wraps the answer string under an 'answer' key.""" + app = Mock(spec=App) + app.mode = mode + + assert extract_structured_output(app, {"answer": "hi"}, "hi") == {"answer": "hi"} + + def test_extract_structured_output_unknown_mode(self): + """Modes outside the MCP surface produce no structured output.""" + app = Mock(spec=App) + app.mode = AppMode.CHANNEL + + assert extract_structured_output(app, {"answer": "hi"}, "hi") is None + def test_process_mapping_response_invalid_mode(self): """Test processing mapping response with invalid app mode""" app = Mock(spec=App) @@ -578,3 +798,29 @@ class TestUtilityFunctions: # Or validation should also raise SchemaError with pytest.raises(jsonschema.exceptions.SchemaError): jsonschema.validate(instance={"count": 1.23}, schema=bad_schema) + + +class TestNegotiateProtocolVersion: + """Test the MCP-Protocol-Version header resolver.""" + + def test_initialize_ignores_header(self): + """Initialize negotiates via the request body, so its header is ignored.""" + assert negotiate_protocol_version("anything", True) == types.DEFAULT_NEGOTIATED_VERSION + + def test_absent_header_defaults(self): + """An absent header defaults to 2025-03-26 per the spec back-compat rule.""" + assert negotiate_protocol_version(None, False) == types.DEFAULT_NEGOTIATED_VERSION + + def test_empty_header_treated_as_absent(self): + """An empty header value is treated as absent and defaults to 2025-03-26.""" + assert negotiate_protocol_version("", False) == types.DEFAULT_NEGOTIATED_VERSION + + def test_supported_header_passes_through(self): + """All supported header values are used as the negotiated version.""" + assert negotiate_protocol_version("2025-06-18", False) == "2025-06-18" + assert negotiate_protocol_version("2025-03-26", False) == "2025-03-26" + assert negotiate_protocol_version("2024-11-05", False) == "2024-11-05" + + def test_unsupported_header_returns_none(self): + """An explicit but unsupported header signals an error (None).""" + assert negotiate_protocol_version("1999-01-01", False) is None diff --git a/api/tests/unit_tests/core/mcp/test_types.py b/api/tests/unit_tests/core/mcp/test_types.py index d4fe353f0aa..ff16f1f9dc8 100644 --- a/api/tests/unit_tests/core/mcp/test_types.py +++ b/api/tests/unit_tests/core/mcp/test_types.py @@ -4,6 +4,7 @@ import pytest from pydantic import ValidationError from core.mcp.types import ( + DEFAULT_NEGOTIATED_VERSION, INTERNAL_ERROR, INVALID_PARAMS, INVALID_REQUEST, @@ -11,6 +12,7 @@ from core.mcp.types import ( METHOD_NOT_FOUND, PARSE_ERROR, SERVER_LATEST_PROTOCOL_VERSION, + SERVER_SUPPORTED_PROTOCOL_VERSIONS, Annotations, CallToolRequest, CallToolRequestParams, @@ -59,7 +61,9 @@ class TestConstants: def test_protocol_versions(self): """Test protocol version constants.""" assert LATEST_PROTOCOL_VERSION == "2025-06-18" - assert SERVER_LATEST_PROTOCOL_VERSION == "2024-11-05" + assert SERVER_LATEST_PROTOCOL_VERSION == "2025-06-18" + assert DEFAULT_NEGOTIATED_VERSION == "2025-03-26" + assert sorted(SERVER_SUPPORTED_PROTOCOL_VERSIONS) == ["2024-11-05", "2025-03-26", "2025-06-18"] def test_error_codes(self): """Test JSON-RPC error code constants.""" From 800c9f4fca005cd399383e87cc8b3865ee9a4102 Mon Sep 17 00:00:00 2001 From: Stephen Zhou Date: Tue, 7 Jul 2026 12:48:12 +0800 Subject: [PATCH 19/70] test(e2e): stabilize Agent v2 external runtime checks (#38493) Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- .../common/agent_app_parameters.py | 4 +- api/core/app/apps/agent_app/app_generator.py | 4 +- .../common/test_agent_app_parameters.py | 33 +++++++-- .../app/apps/agent_app/test_resolve_agent.py | 25 ++++++- .../agent-v2/access-point-web-app.steps.ts | 8 ++- .../agent-v2/build-draft.steps.ts | 5 +- e2e/features/support/hooks.ts | 70 +++++++++++++++---- 7 files changed, 122 insertions(+), 27 deletions(-) diff --git a/api/controllers/common/agent_app_parameters.py b/api/controllers/common/agent_app_parameters.py index 32e338957ea..8c2fbccd513 100644 --- a/api/controllers/common/agent_app_parameters.py +++ b/api/controllers/common/agent_app_parameters.py @@ -32,7 +32,9 @@ def get_published_agent_app_feature_dict_and_user_input_form( ) if agent is None: raise AgentAppGeneratorError("Agent App has no bound Agent") - if not agent.active_config_snapshot_id or not agent.active_config_is_published: + # active_config_is_published means the draft has no unpublished edits; the public app + # can still read parameters from the active snapshot while a newer draft is pending. + if not agent.active_config_snapshot_id: raise AgentAppNotPublishedError("Agent has not been published") snapshot = db.session.scalar( diff --git a/api/core/app/apps/agent_app/app_generator.py b/api/core/app/apps/agent_app/app_generator.py index a940ccf6ceb..9531f9092a4 100644 --- a/api/core/app/apps/agent_app/app_generator.py +++ b/api/core/app/apps/agent_app/app_generator.py @@ -611,7 +611,9 @@ class AgentAppGenerator(MessageBasedAppGenerator): "build_draft" if draft.draft_type == AgentConfigDraftType.DEBUG_BUILD else "draft" ) return agent, draft.id, config_version_kind, agent_soul - if not agent.active_config_snapshot_id or not agent.active_config_is_published: + # active_config_is_published tracks whether the editable draft matches the active snapshot. + # Public runtime must keep serving the active snapshot even when unpublished draft edits exist. + if not agent.active_config_snapshot_id: raise AgentAppNotPublishedError("Agent has not been published") _, snapshot, agent_soul = self._resolve_agent_by_id( tenant_id=app_model.tenant_id, diff --git a/api/tests/unit_tests/controllers/common/test_agent_app_parameters.py b/api/tests/unit_tests/controllers/common/test_agent_app_parameters.py index f963729c8a2..b4639d54d2f 100644 --- a/api/tests/unit_tests/controllers/common/test_agent_app_parameters.py +++ b/api/tests/unit_tests/controllers/common/test_agent_app_parameters.py @@ -85,15 +85,13 @@ def test_published_agent_app_parameters_requires_existing_active_agent(monkeypat @pytest.mark.parametrize( - ("active_config_snapshot_id", "active_config_is_published"), + "active_config_is_published", [ - (None, True), - ("snapshot-1", False), + True, + False, ], ) -def test_published_agent_app_parameters_requires_published_agent( - monkeypatch, active_config_snapshot_id, active_config_is_published -): +def test_published_agent_app_parameters_requires_published_agent(monkeypatch, active_config_is_published): app_model = SimpleNamespace( tenant_id="tenant-1", bound_agent_id="agent-1", @@ -101,7 +99,7 @@ def test_published_agent_app_parameters_requires_published_agent( ) agent = SimpleNamespace( id="agent-1", - active_config_snapshot_id=active_config_snapshot_id, + active_config_snapshot_id=None, active_config_is_published=active_config_is_published, ) monkeypatch.setattr(agent_app_parameters.db.session, "scalar", lambda _: agent) @@ -110,6 +108,27 @@ def test_published_agent_app_parameters_requires_published_agent( get_published_agent_app_feature_dict_and_user_input_form(app_model) +def test_published_agent_app_parameters_allows_unpublished_draft_with_active_snapshot(monkeypatch): + app_model = SimpleNamespace( + tenant_id="tenant-1", + bound_agent_id="agent-1", + app_model_config=None, + ) + agent = SimpleNamespace( + id="agent-1", + active_config_snapshot_id="snapshot-1", + active_config_is_published=False, + ) + snapshot = SimpleNamespace(config_snapshot_dict={}) + query_results = iter([agent, snapshot]) + monkeypatch.setattr(agent_app_parameters.db.session, "scalar", lambda _: next(query_results)) + + features_dict, user_input_form = get_published_agent_app_feature_dict_and_user_input_form(app_model) + + assert features_dict["file_upload"]["enabled"] is True + assert user_input_form == [] + + def test_published_agent_app_parameters_requires_published_snapshot(monkeypatch): app_model = SimpleNamespace( tenant_id="tenant-1", diff --git a/api/tests/unit_tests/core/app/apps/agent_app/test_resolve_agent.py b/api/tests/unit_tests/core/app/apps/agent_app/test_resolve_agent.py index 293c6676e3c..ec0ce7d23b1 100644 --- a/api/tests/unit_tests/core/app/apps/agent_app/test_resolve_agent.py +++ b/api/tests/unit_tests/core/app/apps/agent_app/test_resolve_agent.py @@ -97,12 +97,35 @@ class TestResolveAgent: assert config_version_kind == "snapshot" assert soul.model is not None - def test_unpublished_agent_raises_before_model_resolution(self, monkeypatch: pytest.MonkeyPatch): + def test_unpublished_draft_still_resolves_active_snapshot(self, monkeypatch: pytest.MonkeyPatch): bound_agent = SimpleNamespace( id="agent-1", active_config_snapshot_id="snap-1", active_config_is_published=False, ) + inner_agent = SimpleNamespace(id="agent-1") + snapshot = _snapshot() + _patch_session(monkeypatch, [bound_agent, inner_agent, snapshot]) + app_model = SimpleNamespace(id="app-1", tenant_id="t1") + + agent, config_id, config_version_kind, soul = AgentAppGenerator()._resolve_agent( + app_model, + invoke_from=InvokeFrom.WEB_APP, + draft_type=None, + user=SimpleNamespace(id="user-1"), + ) # type: ignore[arg-type] + + assert agent is bound_agent + assert config_id == snapshot.id + assert config_version_kind == "snapshot" + assert soul.prompt.system_prompt == "You are Iris." + + def test_agent_without_active_snapshot_raises_before_model_resolution(self, monkeypatch: pytest.MonkeyPatch): + bound_agent = SimpleNamespace( + id="agent-1", + active_config_snapshot_id=None, + active_config_is_published=False, + ) _patch_session(monkeypatch, [bound_agent]) app_model = SimpleNamespace(id="app-1", tenant_id="t1") diff --git a/e2e/features/step-definitions/agent-v2/access-point-web-app.steps.ts b/e2e/features/step-definitions/agent-v2/access-point-web-app.steps.ts index d9d85e83772..e5de7e1ac26 100644 --- a/e2e/features/step-definitions/agent-v2/access-point-web-app.steps.ts +++ b/e2e/features/step-definitions/agent-v2/access-point-web-app.steps.ts @@ -1,3 +1,4 @@ +import type { Page } from '@playwright/test' import type { DifyWorld } from '../../support/world' import { Given, Then, When } from '@cucumber/cucumber' import { expect } from '@playwright/test' @@ -12,6 +13,9 @@ import { const WEB_APP_RUNTIME_RESPONSE_STEP_TIMEOUT_MS = 180_000 +const getWebAppMessageInput = (webAppPage: Page) => + webAppPage.getByPlaceholder(/^Talk to /).last() + Then('I should see the Agent v2 Web app access URL', async function (this: DifyWorld) { const webAppCard = getWebAppCard(this) @@ -71,7 +75,7 @@ When('I send an E2E message in the Agent v2 Web app', async function (this: Dify if (!webAppPage) throw new Error('No Agent v2 Web app page was opened.') - const messageInput = webAppPage.getByRole('textbox').last() + const messageInput = getWebAppMessageInput(webAppPage) await expect(messageInput).toBeEditable({ timeout: 30_000 }) await messageInput.fill('Please reply with the test success marker.') await messageInput.press('Enter') @@ -84,7 +88,7 @@ Then('the Agent v2 Web app should open in a new tab', async function (this: Dify throw new Error('No Agent v2 Web app page was opened.') await expect(webAppPage).toHaveURL(webAppURL) - await expect(webAppPage.getByRole('textbox').last()).toBeEditable({ timeout: 30_000 }) + await expect(getWebAppMessageInput(webAppPage)).toBeEditable({ timeout: 30_000 }) await webAppPage.close() this.agentBuilder.accessPoint.webAppPage = undefined this.agentBuilder.accessPoint.webAppURL = undefined diff --git a/e2e/features/step-definitions/agent-v2/build-draft.steps.ts b/e2e/features/step-definitions/agent-v2/build-draft.steps.ts index 22e11f7bd63..87901434035 100644 --- a/e2e/features/step-definitions/agent-v2/build-draft.steps.ts +++ b/e2e/features/step-definitions/agent-v2/build-draft.steps.ts @@ -254,9 +254,8 @@ Then('I should see the Agent v2 Build mode confirmation state', async function ( const page = this.getPage() await expect(page.getByText('Build mode', { exact: true })).toBeVisible() - await expect( - page.getByText('You\'re in build mode. Shape this setup through the chat on the right, then Apply.'), - ).toBeVisible() + await expect(page.getByText('Configure can only be updated by the agent in this mode.')).toBeVisible() + await expect(page.getByText('Shape this setup through the chat on the right, then Apply.')).toBeVisible() }) Then('Agent v2 Build chat should be blocked until a model is configured', async function (this: DifyWorld) { diff --git a/e2e/features/support/hooks.ts b/e2e/features/support/hooks.ts index 8fd59934dae..658f3efc69b 100644 --- a/e2e/features/support/hooks.ts +++ b/e2e/features/support/hooks.ts @@ -1,4 +1,4 @@ -import type { Browser } from '@playwright/test' +import type { Browser, Page } from '@playwright/test' import type { Buffer } from 'node:buffer' import type { DifyWorld } from './world' import { mkdir, writeFile } from 'node:fs/promises' @@ -34,18 +34,50 @@ const sanitizeForPath = (value: string) => const writeArtifact = async ( scenarioName: string, + label: string, extension: 'html' | 'png', contents: Buffer | string, ) => { const artifactPath = path.join( artifactsDir, - `${Date.now()}-${sanitizeForPath(scenarioName || 'scenario')}.${extension}`, + `${Date.now()}-${sanitizeForPath(scenarioName || 'scenario')}-${sanitizeForPath(label)}.${extension}`, ) await writeFile(artifactPath, contents) return artifactPath } +const uniqueDiagnosticPages = (pages: { label: string, page: Page | undefined }[]) => { + const seen = new Set() + + return pages.filter(({ page }) => { + if (!page || page.isClosed() || seen.has(page)) + return false + + seen.add(page) + return true + }) as { label: string, page: Page }[] +} + +const captureDiagnosticPage = async ( + world: DifyWorld, + scenarioName: string, + label: string, + page: Page, +) => { + const screenshot = await page.screenshot({ + fullPage: true, + }) + const screenshotPath = await writeArtifact(scenarioName, label, 'png', screenshot) + world.attach(screenshot, 'image/png') + + const html = await page.content() + const htmlPath = await writeArtifact(scenarioName, label, 'html', html) + world.attach(html, 'text/html') + + return [screenshotPath, htmlPath] +} + const recordCleanup = async ( errors: string[], label: string, @@ -91,16 +123,26 @@ After(async function (this: DifyWorld, { pickle, result }) { const elapsedMs = this.scenarioStartedAt ? Date.now() - this.scenarioStartedAt : undefined const status = result?.status || Status.UNKNOWN - if (diagnosticArtifactStatuses.has(status) && this.page) { - const screenshot = await this.page.screenshot({ - fullPage: true, - }) - const screenshotPath = await writeArtifact(pickle.name, 'png', screenshot) - this.attach(screenshot, 'image/png') + if (diagnosticArtifactStatuses.has(status)) { + const artifactPaths: string[] = [] + const artifactErrors: string[] = [] + const diagnosticPages = uniqueDiagnosticPages([ + { label: 'main-page', page: this.page }, + { label: 'agent-v2-web-app', page: this.agentBuilder.accessPoint.webAppPage }, + { label: 'agent-v2-api-reference', page: this.agentBuilder.accessPoint.apiReferencePage }, + { label: 'agent-v2-workflow-reference', page: this.agentBuilder.accessPoint.workflowReferencePage }, + { label: 'agent-v2-concurrent-configure', page: this.agentBuilder.configure.concurrentPage }, + { label: 'agent-v2-workflow-console', page: this.agentBuilder.workflow.agentConsolePage }, + ]) - const html = await this.page.content() - const htmlPath = await writeArtifact(pickle.name, 'html', html) - this.attach(html, 'text/html') + for (const { label, page } of diagnosticPages) { + try { + artifactPaths.push(...await captureDiagnosticPage(this, pickle.name, label, page)) + } + catch (error) { + artifactErrors.push(`${label}: ${error instanceof Error ? error.message : String(error)}`) + } + } if (this.consoleErrors.length > 0) this.attach(`Console Errors:\n${this.consoleErrors.join('\n')}`, 'text/plain') @@ -108,7 +150,11 @@ After(async function (this: DifyWorld, { pickle, result }) { if (this.pageErrors.length > 0) this.attach(`Page Errors:\n${this.pageErrors.join('\n')}`, 'text/plain') - this.attach(`Artifacts:\n${[screenshotPath, htmlPath].join('\n')}`, 'text/plain') + if (artifactErrors.length > 0) + this.attach(`Artifact Errors:\n${artifactErrors.join('\n')}`, 'text/plain') + + if (artifactPaths.length > 0) + this.attach(`Artifacts:\n${artifactPaths.join('\n')}`, 'text/plain') } console.warn( From f3ba28463bebbc3762218990909ca7a6d667f9b2 Mon Sep 17 00:00:00 2001 From: Asuka Minato Date: Tue, 7 Jul 2026 14:04:44 +0900 Subject: [PATCH 20/70] chore: add sqlite3 to conftest (#38475) --- .../test_legacy_model_type_migration.py | 9 ------- api/tests/unit_tests/conftest.py | 27 +++++++++++++++++++ .../services/test_account_service.py | 13 --------- 3 files changed, 27 insertions(+), 22 deletions(-) diff --git a/api/tests/unit_tests/commands/test_legacy_model_type_migration.py b/api/tests/unit_tests/commands/test_legacy_model_type_migration.py index e14c8ed3243..9eb73b82516 100644 --- a/api/tests/unit_tests/commands/test_legacy_model_type_migration.py +++ b/api/tests/unit_tests/commands/test_legacy_model_type_migration.py @@ -32,15 +32,6 @@ from tests.helpers.legacy_model_type_migration import ( ) -@pytest.fixture -def sqlite_engine(tmp_path: Path) -> sa.Engine: - engine = sa.create_engine(f"sqlite:///{tmp_path / 'legacy_model_type_migration.sqlite'}") - try: - yield engine - finally: - engine.dispose() - - @pytest.fixture def dirty_fixture(sqlite_engine: sa.Engine): return seed_legacy_model_type_dirty_data(sqlite_engine) diff --git a/api/tests/unit_tests/conftest.py b/api/tests/unit_tests/conftest.py index 7174530e976..d95e5e8501d 100644 --- a/api/tests/unit_tests/conftest.py +++ b/api/tests/unit_tests/conftest.py @@ -1,9 +1,12 @@ import os +from collections.abc import Iterator from unittest.mock import MagicMock, patch import pytest from flask import Flask from sqlalchemy import create_engine +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session, sessionmaker # Getting the absolute path of the current file's directory ABS_PATH = os.path.dirname(os.path.abspath(__file__)) @@ -34,6 +37,7 @@ os.environ.setdefault("STORAGE_TYPE", "opendal") from core.db.session_factory import configure_session_factory, session_factory from extensions import ext_redis +from models.base import TypeBase def _patch_redis_clients_on_loaded_modules(): @@ -113,6 +117,29 @@ def _unit_test_engine(): engine.dispose() +@pytest.fixture +def sqlite_engine() -> Iterator[Engine]: + """Create an isolated in-memory SQLite engine for tests that need a disposable database.""" + + engine = create_engine("sqlite:///:memory:") + try: + yield engine + finally: + engine.dispose() + + +@pytest.fixture +def sqlite_session(request: pytest.FixtureRequest, sqlite_engine: Engine) -> Iterator[Session]: + """Yield a SQLite session after creating the model tables passed through ``request.param``.""" + + models: tuple[type[TypeBase], ...] = request.param + tables = [model.metadata.tables[model.__tablename__] for model in models] + TypeBase.metadata.create_all(sqlite_engine, tables=tables) + session_factory = sessionmaker(bind=sqlite_engine, expire_on_commit=False) + with session_factory() as session: + yield session + + @pytest.fixture(autouse=True) def _configure_session_factory(_unit_test_engine): try: diff --git a/api/tests/unit_tests/services/test_account_service.py b/api/tests/unit_tests/services/test_account_service.py index 498e295447e..233191ca0b7 100644 --- a/api/tests/unit_tests/services/test_account_service.py +++ b/api/tests/unit_tests/services/test_account_service.py @@ -1,16 +1,13 @@ import json -from collections.abc import Iterator from datetime import datetime, timedelta from types import SimpleNamespace from unittest.mock import MagicMock, patch import pytest -from sqlalchemy import create_engine from sqlalchemy.orm import Session from configs import dify_config from models.account import Account, AccountStatus, TenantAccountRole, TenantStatus -from models.base import TypeBase from services.account_service import AccountService, RegisterService, TenantService from services.errors.account import ( AccountAlreadyInTenantError, @@ -22,16 +19,6 @@ from services.errors.account import ( ) -@pytest.fixture -def sqlite_session(request: pytest.FixtureRequest) -> Iterator[Session]: - models: tuple[type[TypeBase], ...] = request.param - engine = create_engine("sqlite:///:memory:") - tables = [model.metadata.tables[model.__tablename__] for model in models] - Account.metadata.create_all(engine, tables=tables) - with Session(engine, expire_on_commit=False) as session: - yield session - - class TestAccountAssociatedDataFactory: """Factory class for creating test data and mock objects for account service tests.""" From dd0c4a229643204a2130b18edb861b304b506e7e Mon Sep 17 00:00:00 2001 From: Jashwanth Reddy Gummula Date: Tue, 7 Jul 2026 12:07:20 +0530 Subject: [PATCH 21/70] fix: resolve 36288 mypy errors (#37850) Co-authored-by: WH-2099 --- api/core/helper/code_executor/code_executor.py | 6 ++++-- .../code_executor/jinja2/jinja2_transformer.py | 2 +- .../code_executor/template_transformer.py | 8 ++++---- api/core/helper/creators.py | 2 +- api/core/helper/credential_utils.py | 2 +- api/core/helper/download.py | 5 ++++- api/core/helper/encrypter.py | 14 ++++++++------ api/core/helper/marketplace.py | 18 ++++++++++++++---- api/core/helper/model_provider_cache.py | 4 ++-- api/core/helper/module_import_helper.py | 11 ++++++----- api/core/helper/provider_cache.py | 16 ++++++++-------- api/core/helper/provider_encryption.py | 4 +++- api/core/helper/tool_parameter_cache.py | 4 ++-- api/core/helper/trace_id_helper.py | 2 +- api/core/plugin/entities/plugin_daemon.py | 2 +- .../src/dify_vdb_qdrant/qdrant_vector.py | 1 - .../tidb_on_qdrant_vector.py | 1 - .../code_executor/test_template_transformer.py | 3 +-- .../unit_tests/core/helper/test_creators.py | 12 ++++++++++++ .../unit_tests/core/helper/test_marketplace.py | 11 +++++++++++ 20 files changed, 84 insertions(+), 44 deletions(-) diff --git a/api/core/helper/code_executor/code_executor.py b/api/core/helper/code_executor/code_executor.py index 951e065b2cb..ab2a959a1fc 100644 --- a/api/core/helper/code_executor/code_executor.py +++ b/api/core/helper/code_executor/code_executor.py @@ -13,7 +13,7 @@ from core.helper.code_executor.jinja2.jinja2_transformer import Jinja2TemplateTr from core.helper.code_executor.python3.python3_transformer import Python3TemplateTransformer from core.helper.code_executor.template_transformer import TemplateTransformer from core.helper.http_client_pooling import get_pooled_http_client -from graphon.nodes.code.entities import CodeLanguage +from graphon.nodes.code.entities import CodeLanguage as CodeLanguage # noqa: PLC0414 logger = logging.getLogger(__name__) code_execution_endpoint_url = URL(str(dify_config.CODE_EXECUTION_ENDPOINT)) @@ -133,7 +133,9 @@ class CodeExecutor: return response_code.data.stdout or "" @classmethod - def execute_workflow_code_template(cls, language: CodeLanguage, code: str, inputs: Mapping[str, Any]): + def execute_workflow_code_template( + cls, language: CodeLanguage, code: str, inputs: Mapping[str, Any] + ) -> dict[str, Any]: """ Execute code :param language: code language diff --git a/api/core/helper/code_executor/jinja2/jinja2_transformer.py b/api/core/helper/code_executor/jinja2/jinja2_transformer.py index 9cf5089f7b5..d1c75c981b6 100644 --- a/api/core/helper/code_executor/jinja2/jinja2_transformer.py +++ b/api/core/helper/code_executor/jinja2/jinja2_transformer.py @@ -11,7 +11,7 @@ class Jinja2TemplateTransformer(TemplateTransformer): @classmethod @override - def transform_response(cls, response: str): + def transform_response(cls, response: str) -> dict[str, Any]: """ Transform response to dict :param response: response diff --git a/api/core/helper/code_executor/template_transformer.py b/api/core/helper/code_executor/template_transformer.py index 3a6c314159b..501f460ba95 100644 --- a/api/core/helper/code_executor/template_transformer.py +++ b/api/core/helper/code_executor/template_transformer.py @@ -36,14 +36,14 @@ class TemplateTransformer(ABC): return runner_script, preload_script @classmethod - def extract_result_str_from_response(cls, response: str): + def extract_result_str_from_response(cls, response: str) -> str: result = re.search(rf"{cls._result_tag}(.*){cls._result_tag}", response, re.DOTALL) if not result: raise ValueError(f"Failed to parse result: no result tag found in response. Response: {response[:200]}...") return result.group(1) @classmethod - def transform_response(cls, response: str) -> Mapping[str, Any]: + def transform_response(cls, response: str) -> dict[str, Any]: """ Transform response to dict :param response: response @@ -71,7 +71,7 @@ class TemplateTransformer(ABC): return result @classmethod - def _post_process_result(cls, result: dict[Any, Any]) -> dict[Any, Any]: + def _post_process_result(cls, result: dict[str, Any]) -> dict[str, Any]: """ Post-process the result to convert scientific notation strings back to numbers """ @@ -89,7 +89,7 @@ class TemplateTransformer(ABC): return [convert_scientific_notation(v) for v in value] return value - return convert_scientific_notation(result) + return {key: convert_scientific_notation(value) for key, value in result.items()} @classmethod @abstractmethod diff --git a/api/core/helper/creators.py b/api/core/helper/creators.py index b01e16f18a7..4ad61371512 100644 --- a/api/core/helper/creators.py +++ b/api/core/helper/creators.py @@ -24,7 +24,7 @@ def upload_dsl(dsl_file_bytes: bytes, filename: str = "template.yaml") -> str: response.raise_for_status() data = response.json() claim_code = data.get("data", {}).get("claim_code") - if not claim_code: + if not isinstance(claim_code, str) or not claim_code: raise ValueError("Creators Platform did not return a valid claim_code") return claim_code diff --git a/api/core/helper/credential_utils.py b/api/core/helper/credential_utils.py index e8f3ba0a547..a57474a8c12 100644 --- a/api/core/helper/credential_utils.py +++ b/api/core/helper/credential_utils.py @@ -45,7 +45,7 @@ def is_credential_exists(credential_id: str, credential_type: "PluginCredentialT def runtime_check_credential_policy_compliance( credential_id: str, provider: str, credential_type: "PluginCredentialType", check_existence: bool = True -): +) -> None: if dify_config.ENTERPRISE_DISABLE_RUNTIME_CREDENTIAL_CHECK: return check_credential_policy_compliance( diff --git a/api/core/helper/download.py b/api/core/helper/download.py index 364d45b1e9e..e99be256a01 100644 --- a/api/core/helper/download.py +++ b/api/core/helper/download.py @@ -1,4 +1,7 @@ -def download_with_size_limit(url, max_download_size: int, **kwargs): +from typing import Any + + +def download_with_size_limit(url: str, max_download_size: int, **kwargs: Any) -> bytes: from core.file import remote_fetcher response = remote_fetcher.make_request("GET", url, follow_redirects=True, **kwargs) diff --git a/api/core/helper/encrypter.py b/api/core/helper/encrypter.py index 20125ec6b30..f72f6c6be9f 100644 --- a/api/core/helper/encrypter.py +++ b/api/core/helper/encrypter.py @@ -1,5 +1,7 @@ import base64 +from Crypto.PublicKey import RSA + from libs import rsa @@ -11,13 +13,13 @@ def obfuscated_token(token: str) -> str: return token[:6] + "*" * 12 + token[-2:] -def full_mask_token(token_length=20): +def full_mask_token(token_length: int = 20) -> str: return "*" * token_length -def encrypt_token(tenant_id: str, token: str): - from extensions.ext_database import db +def encrypt_token(tenant_id: str, token: str) -> str: from models.account import Tenant + from models.engine import db if not (tenant := db.session.get(Tenant, tenant_id)): raise ValueError(f"Tenant with id {tenant_id} not found") @@ -30,15 +32,15 @@ def decrypt_token(tenant_id: str, token: str) -> str: return rsa.decrypt(base64.b64decode(token), tenant_id) -def batch_decrypt_token(tenant_id: str, tokens: list[str]): +def batch_decrypt_token(tenant_id: str, tokens: list[str]) -> list[str]: rsa_key, cipher_rsa = rsa.get_decrypt_decoding(tenant_id) return [rsa.decrypt_token_with_decoding(base64.b64decode(token), rsa_key, cipher_rsa) for token in tokens] -def get_decrypt_decoding(tenant_id: str): +def get_decrypt_decoding(tenant_id: str) -> tuple[RSA.RsaKey, object]: return rsa.get_decrypt_decoding(tenant_id) -def decrypt_token_with_decoding(token: str, rsa_key, cipher_rsa): +def decrypt_token_with_decoding(token: str, rsa_key: RSA.RsaKey, cipher_rsa: object) -> str: return rsa.decrypt_token_with_decoding(base64.b64decode(token), rsa_key, cipher_rsa) diff --git a/api/core/helper/marketplace.py b/api/core/helper/marketplace.py index 0b77891ce16..e6e1d565769 100644 --- a/api/core/helper/marketplace.py +++ b/api/core/helper/marketplace.py @@ -1,5 +1,6 @@ import logging from collections.abc import Sequence +from typing import Any from urllib.parse import urlencode import httpx @@ -21,7 +22,7 @@ def get_plugin_pkg_url(plugin_unique_identifier: str) -> str: return f"{marketplace_api_url / 'api/v1/plugins/download'}?{query}" -def download_plugin_pkg(plugin_unique_identifier: str): +def download_plugin_pkg(plugin_unique_identifier: str) -> bytes: return download_with_size_limit(get_plugin_pkg_url(plugin_unique_identifier), dify_config.PLUGIN_MAX_PACKAGE_SIZE) @@ -41,7 +42,7 @@ def batch_fetch_plugin_manifests(plugin_ids: list[str]) -> Sequence[MarketplaceP return [MarketplacePluginDeclaration.model_validate(plugin) for plugin in response.json()["data"]["plugins"]] -def batch_fetch_plugin_by_ids(plugin_ids: list[str]) -> list[dict]: +def batch_fetch_plugin_by_ids(plugin_ids: list[str]) -> list[dict[str, Any]]: if not plugin_ids: return [] @@ -55,10 +56,19 @@ def batch_fetch_plugin_by_ids(plugin_ids: list[str]) -> list[dict]: response.raise_for_status() data = response.json() - return data.get("data", {}).get("plugins", []) + plugins = data.get("data", {}).get("plugins", []) + if not isinstance(plugins, list): + raise ValueError("Marketplace did not return a valid plugins list") + + result: list[dict[str, Any]] = [] + for plugin in plugins: + if not isinstance(plugin, dict) or not all(isinstance(key, str) for key in plugin): + raise ValueError("Marketplace did not return a valid plugins list") + result.append(plugin) + return result -def record_install_plugin_event(plugin_unique_identifier: str): +def record_install_plugin_event(plugin_unique_identifier: str) -> None: url = str(marketplace_api_url / "api/v1/stats/plugins/install_count") response = httpx.post(url, json={"unique_identifier": plugin_unique_identifier}, timeout=MARKETPLACE_TIMEOUT) response.raise_for_status() diff --git a/api/core/helper/model_provider_cache.py b/api/core/helper/model_provider_cache.py index 10d79a82392..2b9e6613378 100644 --- a/api/core/helper/model_provider_cache.py +++ b/api/core/helper/model_provider_cache.py @@ -34,7 +34,7 @@ class ProviderCredentialsCache: else: return None - def set(self, credentials: dict[str, Any]): + def set(self, credentials: dict[str, Any]) -> None: """ Cache model provider credentials. @@ -43,7 +43,7 @@ class ProviderCredentialsCache: """ redis_client.setex(self.cache_key, 86400, json.dumps(credentials)) - def delete(self): + def delete(self) -> None: """ Delete cached model provider credentials. diff --git a/api/core/helper/module_import_helper.py b/api/core/helper/module_import_helper.py index 768210d899b..4d37abf4897 100644 --- a/api/core/helper/module_import_helper.py +++ b/api/core/helper/module_import_helper.py @@ -20,17 +20,18 @@ def import_module_from_source[T: (str, bytes)]( raise Exception(f"Failed to load module {module_name} from {py_file_path!r}") else: # Refer to: https://docs.python.org/3/library/importlib.html#importing-a-source-file-directly - # FIXME: mypy does not support the type of spec.loader - spec = importlib.util.spec_from_file_location(module_name, py_file_path) # type: ignore[assignment] - if not spec or not spec.loader: + new_spec = importlib.util.spec_from_file_location(module_name, py_file_path) + if not new_spec or not new_spec.loader: raise Exception(f"Failed to load module {module_name} from {py_file_path!r}") if use_lazy_loader: # Refer to: https://docs.python.org/3/library/importlib.html#implementing-lazy-imports - spec.loader = importlib.util.LazyLoader(spec.loader) + new_spec.loader = importlib.util.LazyLoader(new_spec.loader) + spec = new_spec module = importlib.util.module_from_spec(spec) if not existed_spec: sys.modules[module_name] = module - spec.loader.exec_module(module) + if spec.loader is not None: + spec.loader.exec_module(module) return module except Exception as e: logger.exception("Failed to load module %s from script file '%s'", module_name, repr(py_file_path)) diff --git a/api/core/helper/provider_cache.py b/api/core/helper/provider_cache.py index 6ad08dfe178..a3b61887892 100644 --- a/api/core/helper/provider_cache.py +++ b/api/core/helper/provider_cache.py @@ -9,11 +9,11 @@ from extensions.ext_redis import redis_client class ProviderCredentialsCache(ABC): """Base class for provider credentials cache""" - def __init__(self, **kwargs): + def __init__(self, **kwargs: Any) -> None: self.cache_key = self._generate_cache_key(**kwargs) @abstractmethod - def _generate_cache_key(self, **kwargs) -> str: + def _generate_cache_key(self, **kwargs: Any) -> str: """Generate cache key based on subclass implementation""" pass @@ -28,11 +28,11 @@ class ProviderCredentialsCache(ABC): return None return None - def set(self, config: dict[str, Any]): + def set(self, config: dict[str, Any]) -> None: """Cache provider credentials""" redis_client.setex(self.cache_key, 86400, json.dumps(config)) - def delete(self): + def delete(self) -> None: """Delete cached provider credentials""" redis_client.delete(self.cache_key) @@ -48,7 +48,7 @@ class SingletonProviderCredentialsCache(ProviderCredentialsCache): ) @override - def _generate_cache_key(self, **kwargs) -> str: + def _generate_cache_key(self, **kwargs: Any) -> str: tenant_id = kwargs["tenant_id"] provider_type = kwargs["provider_type"] identity_name = kwargs["provider_identity"] @@ -63,7 +63,7 @@ class ToolProviderCredentialsCache(ProviderCredentialsCache): super().__init__(tenant_id=tenant_id, provider=provider, credential_id=credential_id) @override - def _generate_cache_key(self, **kwargs) -> str: + def _generate_cache_key(self, **kwargs: Any) -> str: tenant_id = kwargs["tenant_id"] provider = kwargs["provider"] credential_id = kwargs["credential_id"] @@ -77,10 +77,10 @@ class NoOpProviderCredentialCache: """Get cached provider credentials""" return None - def set(self, config: dict[str, Any]): + def set(self, config: dict[str, Any]) -> None: """Cache provider credentials""" pass - def delete(self): + def delete(self) -> None: """Delete cached provider credentials""" pass diff --git a/api/core/helper/provider_encryption.py b/api/core/helper/provider_encryption.py index d392d22687b..de9d5e15be1 100644 --- a/api/core/helper/provider_encryption.py +++ b/api/core/helper/provider_encryption.py @@ -125,5 +125,7 @@ class ProviderConfigEncrypter: return data -def create_provider_encrypter(tenant_id: str, config: list[BasicProviderConfig], cache: ProviderConfigCache): +def create_provider_encrypter( + tenant_id: str, config: list[BasicProviderConfig], cache: ProviderConfigCache +) -> tuple[ProviderConfigEncrypter, ProviderConfigCache]: return ProviderConfigEncrypter(tenant_id=tenant_id, config=config, provider_config_cache=cache), cache diff --git a/api/core/helper/tool_parameter_cache.py b/api/core/helper/tool_parameter_cache.py index bf5bf9af03b..2650eb0c2c2 100644 --- a/api/core/helper/tool_parameter_cache.py +++ b/api/core/helper/tool_parameter_cache.py @@ -37,11 +37,11 @@ class ToolParameterCache: else: return None - def set(self, parameters: dict[str, Any]): + def set(self, parameters: dict[str, Any]) -> None: """Cache model provider credentials.""" redis_client.setex(self.cache_key, 86400, json.dumps(parameters)) - def delete(self): + def delete(self) -> None: """ Delete cached model provider credentials. diff --git a/api/core/helper/trace_id_helper.py b/api/core/helper/trace_id_helper.py index 8b022c1d065..e1ebd45e074 100644 --- a/api/core/helper/trace_id_helper.py +++ b/api/core/helper/trace_id_helper.py @@ -61,7 +61,7 @@ def get_external_trace_id(request: Any) -> str | None: return None -def extract_external_trace_id_from_args(args: Mapping[str, Any]): +def extract_external_trace_id_from_args(args: Mapping[str, Any]) -> dict[str, Any]: """ Extract 'external_trace_id' from args. diff --git a/api/core/plugin/entities/plugin_daemon.py b/api/core/plugin/entities/plugin_daemon.py index 507a6ea5cd3..4cf55ef8e3c 100644 --- a/api/core/plugin/entities/plugin_daemon.py +++ b/api/core/plugin/entities/plugin_daemon.py @@ -228,7 +228,7 @@ class CredentialType(enum.StrEnum): OAUTH2 = "oauth2" UNAUTHORIZED = "unauthorized" - def get_name(self): + def get_name(self) -> str: if self == CredentialType.API_KEY: return "API KEY" elif self == CredentialType.OAUTH2: diff --git a/api/providers/vdb/vdb-qdrant/src/dify_vdb_qdrant/qdrant_vector.py b/api/providers/vdb/vdb-qdrant/src/dify_vdb_qdrant/qdrant_vector.py index 2e6395bacc7..d52a3d15544 100644 --- a/api/providers/vdb/vdb-qdrant/src/dify_vdb_qdrant/qdrant_vector.py +++ b/api/providers/vdb/vdb-qdrant/src/dify_vdb_qdrant/qdrant_vector.py @@ -33,7 +33,6 @@ from models.dataset import Dataset, DatasetCollectionBinding if TYPE_CHECKING: from qdrant_client.conversions import common_types - from qdrant_client.http import models as rest type DictFilter = dict[str, str | int | bool | dict | list] type MetadataFilter = DictFilter | common_types.Filter diff --git a/api/providers/vdb/vdb-tidb-on-qdrant/src/dify_vdb_tidb_on_qdrant/tidb_on_qdrant_vector.py b/api/providers/vdb/vdb-tidb-on-qdrant/src/dify_vdb_tidb_on_qdrant/tidb_on_qdrant_vector.py index 9e6dc27203d..b352243b92a 100644 --- a/api/providers/vdb/vdb-tidb-on-qdrant/src/dify_vdb_tidb_on_qdrant/tidb_on_qdrant_vector.py +++ b/api/providers/vdb/vdb-tidb-on-qdrant/src/dify_vdb_tidb_on_qdrant/tidb_on_qdrant_vector.py @@ -41,7 +41,6 @@ from models.enums import TidbAuthBindingStatus if TYPE_CHECKING: from qdrant_client import grpc # noqa from qdrant_client.conversions import common_types - from qdrant_client.http import models as rest type DictFilter = dict[str, str | int | bool | dict | list] type MetadataFilter = DictFilter | common_types.Filter diff --git a/api/tests/unit_tests/core/helper/code_executor/test_template_transformer.py b/api/tests/unit_tests/core/helper/code_executor/test_template_transformer.py index 5b54b8e6474..d9d171efaea 100644 --- a/api/tests/unit_tests/core/helper/code_executor/test_template_transformer.py +++ b/api/tests/unit_tests/core/helper/code_executor/test_template_transformer.py @@ -1,6 +1,5 @@ import json from base64 import b64decode -from collections.abc import Mapping from typing import Any import pytest @@ -44,7 +43,7 @@ def test_serialize_inputs_encodes_payload() -> None: def test_transform_response_parses_json_result_and_converts_scientific_notation() -> None: response = '<>{"value": "1e+3", "nested": {"x": "2E-2"}, "arr": ["3e+1"]}<>' - result: Mapping[str, Any] = _DummyTransformer.transform_response(response) + result: dict[str, Any] = _DummyTransformer.transform_response(response) assert result == {"value": 1000.0, "nested": {"x": 0.02}, "arr": [30.0]} diff --git a/api/tests/unit_tests/core/helper/test_creators.py b/api/tests/unit_tests/core/helper/test_creators.py index 8750f6d9070..cac7ad35a71 100644 --- a/api/tests/unit_tests/core/helper/test_creators.py +++ b/api/tests/unit_tests/core/helper/test_creators.py @@ -46,6 +46,18 @@ class TestUploadDSL: with pytest.raises(ValueError, match="claim_code"): upload_dsl(b"app: demo") + @patch("core.helper.creators.httpx.post") + def test_raises_on_non_string_claim_code(self, mock_post): + mock_response = MagicMock(spec=httpx.Response) + mock_response.json.return_value = {"data": {"claim_code": 123}} + mock_response.raise_for_status = MagicMock() + mock_post.return_value = mock_response + + from core.helper.creators import upload_dsl + + with pytest.raises(ValueError, match="claim_code"): + upload_dsl(b"app: demo") + @patch("core.helper.creators.httpx.post") def test_raises_on_http_error(self, mock_post): mock_response = MagicMock(spec=httpx.Response) diff --git a/api/tests/unit_tests/core/helper/test_marketplace.py b/api/tests/unit_tests/core/helper/test_marketplace.py index eba1e4e5442..a14a9fe582d 100644 --- a/api/tests/unit_tests/core/helper/test_marketplace.py +++ b/api/tests/unit_tests/core/helper/test_marketplace.py @@ -1,6 +1,7 @@ from types import SimpleNamespace from unittest.mock import MagicMock +import pytest from pytest_mock import MockerFixture from core.helper.marketplace import ( @@ -53,6 +54,16 @@ def test_batch_fetch_plugin_by_ids_returns_plugins_from_response(mocker: MockerF response.raise_for_status.assert_called_once() +def test_batch_fetch_plugin_by_ids_rejects_invalid_plugins_response(mocker: MockerFixture) -> None: + response = MagicMock() + response.json.return_value = {"data": {"plugins": ["p1"]}} + response.raise_for_status.return_value = None + mocker.patch("core.helper.marketplace.httpx.post", return_value=response) + + with pytest.raises(ValueError, match="plugins list"): + batch_fetch_plugin_by_ids(["p1"]) + + def test_batch_fetch_plugin_manifests_returns_empty_for_empty_input(mocker: MockerFixture) -> None: post_mock = mocker.patch("core.helper.marketplace.httpx.post") From faaa4708a697e327ad0b5cc5f3828cb65c647585 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=9E=E6=B3=95=E6=93=8D=E4=BD=9C?= Date: Tue, 7 Jul 2026 15:39:24 +0800 Subject: [PATCH 22/70] fix: editor should not manage member (#38503) --- api/services/enterprise/rbac_service.py | 1 - 1 file changed, 1 deletion(-) diff --git a/api/services/enterprise/rbac_service.py b/api/services/enterprise/rbac_service.py index 450e4d0be33..b2e77156d3d 100644 --- a/api/services/enterprise/rbac_service.py +++ b/api/services/enterprise/rbac_service.py @@ -365,7 +365,6 @@ _LEGACY_WORKSPACE_ADMIN_KEYS: list[str] = [ ] _LEGACY_WORKSPACE_EDITOR_KEYS: list[str] = [ - "workspace.member.manage", "api_extension.manage", "plugin.install", "credential.use", From 6922c4548966c212563e09a485064feae58008b1 Mon Sep 17 00:00:00 2001 From: wangxiaolei Date: Tue, 7 Jul 2026 15:53:01 +0800 Subject: [PATCH 23/70] chore: update editor permission (#38505) From fb92e9a3470415bec35c50857870d38cf4a09a03 Mon Sep 17 00:00:00 2001 From: FFXN <31929997+FFXN@users.noreply.github.com> Date: Tue, 7 Jul 2026 15:53:14 +0800 Subject: [PATCH 24/70] chore: improve cherry pick missed message (#38496) --- .github/scripts/check-hotfix-cherry-picks.sh | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/scripts/check-hotfix-cherry-picks.sh b/.github/scripts/check-hotfix-cherry-picks.sh index 11dc024ccf8..cb97dc20699 100644 --- a/.github/scripts/check-hotfix-cherry-picks.sh +++ b/.github/scripts/check-hotfix-cherry-picks.sh @@ -45,7 +45,7 @@ while IFS= read -r commit_sha; do ) if [[ -z "$source_sha" ]]; then - error "Commit $commit_sha ($subject) is missing cherry-pick provenance. $REMEDIATION_HINT" + error "Commit $commit_sha ($subject) is missing cherry-pick provenance. $REMEDIATION_HINT If version differences prevent using git cherry-pick -x, manually add '(cherry picked from commit )' to the commit message." failed=1 continue fi From 31b17513c2ebc7ae54961fcf7bf5475e8499e1b6 Mon Sep 17 00:00:00 2001 From: yyh <92089059+lyzno1@users.noreply.github.com> Date: Tue, 7 Jul 2026 16:58:35 +0800 Subject: [PATCH 25/70] docs(component): document focus-visible guidance (#38509) --- .agents/skills/how-to-write-component/SKILL.md | 1 + 1 file changed, 1 insertion(+) diff --git a/.agents/skills/how-to-write-component/SKILL.md b/.agents/skills/how-to-write-component/SKILL.md index d38bf53529b..e66eec40c88 100644 --- a/.agents/skills/how-to-write-component/SKILL.md +++ b/.agents/skills/how-to-write-component/SKILL.md @@ -24,6 +24,7 @@ Use this as the component decision guide for Dify web. Existing code is referenc - Search before adding UI, hooks, helpers, query utilities, or styling patterns. Reuse existing base components, feature components, hooks, utilities, and design styles when they fit. - Follow Dify's CSS-first Tailwind v4 contract from `packages/dify-ui/README.md` and `packages/dify-ui/AGENTS.md`. Prefer design-system tokens, utilities, and radius mappings over generic Tailwind choices. +- Preserve visible keyboard focus states on the final focusable element. Prefer styled `@langgenius/dify-ui/*` controls when available, because components such as `Button` and form/control primitives carry the standard Dify UI `focus-visible` styling. Do not assume every Dify UI export provides visual focus styles: headless anatomy parts and direct Base UI re-exports such as dialog/popover/tooltip/drawer triggers usually only provide behavior and semantics. When using native `button` / `a`, custom trigger `render` props, clickable rows, icon buttons, menu-like items, or direct trigger parts, verify the rendered focusable element has a visible focus state. If it does not, add the standard Dify UI focus style: `outline-hidden focus-visible:ring-2 focus-visible:ring-state-accent-solid`. Do not hide outlines without an equivalent visible `focus-visible` indicator. Component-specific focus styles should follow an existing styled primitive pattern or a concrete design constraint, not a new ad hoc style. - Group feature code by workflow, route, or ownership area with route-aligned names: components, hooks, local types, query helpers, atoms, constants, tests, and small utilities should live near the code that changes with them. - For each feature module, keep a module-local `README.md` as a boundary note. Start with the module name, a brief one-sentence description, then split dependencies into `Internal Modules` and `External Modules` sections; keep both sections and write `None.` when one category is empty. `Internal Modules` lists modules inside the same overall feature using paths from that feature root, such as `shared/domain/runtime-status`; `External Modules` lists project modules outside the feature using paths from the web root without a `web/` prefix, such as `app/components/base/skeleton`. Omit npm packages, workspace package dependencies, and whitelisted plumbing modules. Do not copy caller-relative import paths into the README. - Module README whitelist: `@/service/client`, `@/next/*`. From 2c6ec1a761b92e07f995a6c46ea69e28c97d149f Mon Sep 17 00:00:00 2001 From: yyh <92089059+lyzno1@users.noreply.github.com> Date: Tue, 7 Jul 2026 17:47:30 +0800 Subject: [PATCH 26/70] refactor(web): move app context layout styles to shell (#38511) --- web/app/(commonLayout)/layout.tsx | 38 ++++++++++-------- web/app/account/(commonLayout)/layout.tsx | 40 ++++++++++--------- .../__tests__/maintenance-notice.spec.tsx | 17 ++++++++ .../components/header/maintenance-notice.tsx | 12 +++++- web/context/app-context-provider.tsx | 13 ++---- 5 files changed, 73 insertions(+), 47 deletions(-) diff --git a/web/app/(commonLayout)/layout.tsx b/web/app/(commonLayout)/layout.tsx index d4b9bc599e7..02302752df8 100644 --- a/web/app/(commonLayout)/layout.tsx +++ b/web/app/(commonLayout)/layout.tsx @@ -1,8 +1,9 @@ -import type { ReactNode } from 'react' +import * as React from 'react' import AmplitudeProvider from '@/app/components/base/amplitude' import { GoogleAnalyticsScripts } from '@/app/components/base/ga' import Zendesk from '@/app/components/base/zendesk' import { EducationVerifyActionRecorder } from '@/app/components/education-verify-action-recorder' +import MaintenanceNotice from '@/app/components/header/maintenance-notice' import MainNavLayout from '@/app/components/main-nav/layout' import { NextRouteStateBridge } from '@/app/components/next-route-state' import { OAuthRegistrationAnalytics } from '@/app/components/oauth-registration-analytics' @@ -17,32 +18,35 @@ export default async function Layout({ children, detailSidebar, }: { - children: ReactNode - detailSidebar: ReactNode + children: React.ReactNode + detailSidebar: React.ReactNode }) { return ( - <> + - - - - - - {children} - - - - - - +
    + + + + + + + {children} + + + + + + +
    - +
    ) } diff --git a/web/app/account/(commonLayout)/layout.tsx b/web/app/account/(commonLayout)/layout.tsx index a97588d2030..9420503c08b 100644 --- a/web/app/account/(commonLayout)/layout.tsx +++ b/web/app/account/(commonLayout)/layout.tsx @@ -1,10 +1,10 @@ -import type { ReactNode } from 'react' import * as React from 'react' import { CommonLayoutHydrationBoundary } from '@/app/(commonLayout)/hydration-boundary' import AmplitudeProvider from '@/app/components/base/amplitude' import { GoogleAnalyticsScripts } from '@/app/components/base/ga' import { EducationVerifyActionRecorder } from '@/app/components/education-verify-action-recorder' import HeaderWrapper from '@/app/components/header/header-wrapper' +import MaintenanceNotice from '@/app/components/header/maintenance-notice' import { OAuthRegistrationAnalytics } from '@/app/components/oauth-registration-analytics' import { AppContextProvider } from '@/context/app-context-provider' import { EventEmitterContextProvider } from '@/context/event-emitter-provider' @@ -12,30 +12,32 @@ import { ModalContextProvider } from '@/context/modal-context-provider' import { ProviderContextProvider } from '@/context/provider-context-provider' import Header from './header' -const Layout = async ({ children }: { children: ReactNode }) => { +export default async function Layout({ children }: { children: React.ReactNode }) { return ( - <> + - - - - - -
    - -
    - {children} -
    - - - - +
    + + + + + + +
    + +
    + {children} +
    + + + + +
    - + ) } -export default Layout diff --git a/web/app/components/header/__tests__/maintenance-notice.spec.tsx b/web/app/components/header/__tests__/maintenance-notice.spec.tsx index bab8c014fc8..203dd55fa91 100644 --- a/web/app/components/header/__tests__/maintenance-notice.spec.tsx +++ b/web/app/components/header/__tests__/maintenance-notice.spec.tsx @@ -4,10 +4,18 @@ import { useLanguage } from '@/app/components/header/account-setting/model-provi import { NOTICE_I18N } from '@/i18n-config/language' import MaintenanceNotice from '../maintenance-notice' +const mockEnv = vi.hoisted(() => ({ + NEXT_PUBLIC_MAINTENANCE_NOTICE: 'true', +})) + vi.mock('@/app/components/base/icons/src/vender/line/general', () => ({ X: (props: React.SVGProps) => , })) +vi.mock('@/env', () => ({ + env: mockEnv, +})) + vi.mock( '@/app/components/header/account-setting/model-provider-page/hooks', () => ({ @@ -44,6 +52,7 @@ describe('MaintenanceNotice', () => { beforeEach(() => { vi.clearAllMocks() localStorage.clear() + mockEnv.NEXT_PUBLIC_MAINTENANCE_NOTICE = 'true' vi.mocked(useLanguage).mockReturnValue('en_US') setNoticeHref('#') }) @@ -71,6 +80,14 @@ describe('MaintenanceNotice', () => { const { container } = render() expect(container.firstChild).toBeNull() }) + + it('should not render when the notice env flag is disabled', () => { + mockEnv.NEXT_PUBLIC_MAINTENANCE_NOTICE = '' + + const { container } = render() + + expect(container.firstChild).toBeNull() + }) }) describe('User Interactions', () => { diff --git a/web/app/components/header/maintenance-notice.tsx b/web/app/components/header/maintenance-notice.tsx index 918c3ce78f4..fcfbb39e100 100644 --- a/web/app/components/header/maintenance-notice.tsx +++ b/web/app/components/header/maintenance-notice.tsx @@ -1,11 +1,21 @@ +'use client' + import { useState } from 'react' import { useTranslation } from 'react-i18next' import { X } from '@/app/components/base/icons/src/vender/line/general' import { useLanguage } from '@/app/components/header/account-setting/model-provider-page/hooks' +import { env } from '@/env' import { NOTICE_I18N } from '@/i18n-config/language' import { useHideMaintenanceNotice } from './storage' -const MaintenanceNotice = () => { +function MaintenanceNotice() { + if (!env.NEXT_PUBLIC_MAINTENANCE_NOTICE) + return null + + return +} + +function MaintenanceNoticeContent() { const { t } = useTranslation() const locale = useLanguage() diff --git a/web/context/app-context-provider.tsx b/web/context/app-context-provider.tsx index 2091f0818c2..509d36e3e59 100644 --- a/web/context/app-context-provider.tsx +++ b/web/context/app-context-provider.tsx @@ -2,14 +2,13 @@ import type { GetAccountProfileResponse } from '@dify/contracts/api/console/account/types.gen' import type { PostWorkspacesCurrentResponse } from '@dify/contracts/api/console/workspaces/types.gen' -import type { FC, ReactNode } from 'react' +import type { ReactNode } from 'react' import type { ICurrentWorkspace, LangGeniusVersionResponse } from '@/models/common' import { useQuery, useQueryClient, useSuspenseQuery } from '@tanstack/react-query' import { useCallback, useEffect, useMemo } from 'react' import { setUserId, setUserProperties } from '@/app/components/base/amplitude' import { flushRegistrationSuccess } from '@/app/components/base/amplitude/registration-tracking' import { setZendeskConversationFields } from '@/app/components/base/zendesk/utils' -import MaintenanceNotice from '@/app/components/header/maintenance-notice' import { ZENDESK_FIELD_IDS } from '@/config' import { AppContext, @@ -18,7 +17,6 @@ import { userProfilePlaceholder, useSelector, } from '@/context/app-context' -import { env } from '@/env' import { userProfileQueryOptions } from '@/features/account-profile/client' import { systemFeaturesQueryOptions } from '@/features/system-features/client' import { useWorkspacePermissionKeys } from '@/service/access-control/use-permission-keys' @@ -66,7 +64,7 @@ const normalizeCurrentWorkspace = (workspace?: PostWorkspacesCurrentResponse): I } } -export const AppContextProvider: FC = ({ children }) => { +export function AppContextProvider({ children }: AppContextProviderProps) { const queryClient = useQueryClient() const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) const { data: userProfileResp } = useSuspenseQuery(userProfileQueryOptions()) @@ -192,12 +190,7 @@ export const AppContextProvider: FC = ({ children }) => workspacePermissionKeys: workspacePermissionKeysQuery.data?.workspace.permission_keys ?? emptyWorkspacePermissionKeys, }} > -
    - {env.NEXT_PUBLIC_MAINTENANCE_NOTICE && } -
    - {children} -
    -
    + {children} ) } From 3ddfba5ca5dc13176723b710f1b126b6b5b68a4d Mon Sep 17 00:00:00 2001 From: yyh <92089059+lyzno1@users.noreply.github.com> Date: Tue, 7 Jul 2026 21:02:30 +0800 Subject: [PATCH 27/70] fix(web): add backdrop blur to skip nav (#38517) --- web/app/components/main-nav/skip-nav.tsx | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/web/app/components/main-nav/skip-nav.tsx b/web/app/components/main-nav/skip-nav.tsx index 6bdb82c6563..ddf798e3178 100644 --- a/web/app/components/main-nav/skip-nav.tsx +++ b/web/app/components/main-nav/skip-nav.tsx @@ -26,7 +26,7 @@ export function SkipNav({ href={MAIN_CONTENT_HREF} onClick={handleClick} className={cn( - 'fixed top-2 left-2 z-60 inline-flex h-9 -translate-y-[calc(100%+0.75rem)] items-center justify-center rounded-lg border-[0.5px] border-components-button-secondary-border bg-components-button-secondary-bg px-3 system-sm-medium text-components-button-secondary-text outline-hidden transition-transform duration-150 focus-visible:translate-y-0 focus-visible:shadow-lg focus-visible:ring-2 focus-visible:shadow-shadow-shadow-5 focus-visible:ring-state-accent-solid motion-reduce:transition-none', + 'fixed top-2 left-2 z-60 inline-flex h-9 -translate-y-[calc(100%+0.75rem)] items-center justify-center rounded-lg border-[0.5px] border-components-button-secondary-border bg-components-button-secondary-bg px-3 system-sm-medium text-components-button-secondary-text outline-hidden backdrop-blur-[5px] transition-transform duration-150 focus-visible:translate-y-0 focus-visible:shadow-lg focus-visible:ring-2 focus-visible:shadow-shadow-shadow-5 focus-visible:ring-state-accent-solid motion-reduce:transition-none', className, )} {...props} From 6edce14e887fbbc94ae838002abea519a24acfbd Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=9E=E6=B3=95=E6=93=8D=E4=BD=9C?= Date: Tue, 7 Jul 2026 21:03:32 +0800 Subject: [PATCH 28/70] fix: can't debug model plugins (#38500) --- api/core/plugin/plugin_service.py | 126 +++++++- .../services/plugin/test_plugin_service.py | 150 +++++++++- .../__tests__/index.non-cloud.spec.tsx | 5 +- .../__tests__/index.spec.tsx | 275 +++++++++++++++++- .../model-provider-page/index.tsx | 47 ++- .../model-provider-page-body.tsx | 38 ++- .../__tests__/provider-card-actions.spec.tsx | 9 + .../provider-card-actions.tsx | 53 ++-- 8 files changed, 648 insertions(+), 55 deletions(-) diff --git a/api/core/plugin/plugin_service.py b/api/core/plugin/plugin_service.py index 6b306e2df86..89274b635ac 100644 --- a/api/core/plugin/plugin_service.py +++ b/api/core/plugin/plugin_service.py @@ -92,6 +92,7 @@ class PluginService: PLUGIN_MODEL_PROVIDERS_REDIS_KEY_PREFIX = "plugin_model_providers:tenant_id:" PLUGIN_MODEL_PROVIDERS_GENERATION_REDIS_KEY_PREFIX = "plugin_model_providers_generation:tenant_id:" PLUGIN_MODEL_PROVIDERS_LOCK_REDIS_KEY_PREFIX = "plugin_model_providers_refresh_lock:tenant_id:" + PLUGIN_MODEL_PROVIDERS_REMOTE_DEBUG_REDIS_KEY_PREFIX = "plugin_model_providers_remote_debug:tenant_id:" PLUGIN_MODEL_PROVIDERS_LOCK_TTL = 30 PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_TIMEOUT = 2.0 PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_INTERVAL = 0.05 @@ -117,6 +118,10 @@ class PluginService: def _get_plugin_model_providers_lock_key(cls, tenant_id: str, generation: int) -> str: return f"{cls.PLUGIN_MODEL_PROVIDERS_LOCK_REDIS_KEY_PREFIX}{tenant_id}:generation:{generation}" + @classmethod + def _get_plugin_model_providers_remote_debug_cache_key(cls, tenant_id: str) -> str: + return f"{cls.PLUGIN_MODEL_PROVIDERS_REMOTE_DEBUG_REDIS_KEY_PREFIX}{tenant_id}" + @staticmethod def _get_provider_short_name_alias(provider: PluginModelProviderEntity) -> str: """ @@ -259,6 +264,111 @@ class PluginService: except (RedisError, RuntimeError): logger.warning("Failed to cache plugin model providers for tenant %s.", tenant_id, exc_info=True) + @classmethod + def _get_remote_model_plugin_cache_marker(cls, plugins: Sequence[PluginEntity]) -> str | None: + remote_model_plugins = sorted( + f"{plugin.plugin_id}:{plugin.plugin_unique_identifier}" + for plugin in plugins + if plugin.source == PluginInstallationSource.Remote + ) + if not remote_model_plugins: + return None + + return "\n".join(remote_model_plugins) + + @classmethod + def _load_cached_remote_model_plugin_marker(cls, tenant_id: str) -> str | None: + cache_key = cls._get_plugin_model_providers_remote_debug_cache_key(tenant_id) + try: + cached_marker = redis_client.get(cache_key) + except (RedisError, RuntimeError): + logger.warning("Failed to read remote debug model plugin marker for tenant %s.", tenant_id, exc_info=True) + return None + + if cached_marker is None: + return None + if isinstance(cached_marker, bytes): + try: + return cached_marker.decode() + except UnicodeDecodeError: + logger.warning( + "Invalid remote debug model plugin marker for tenant %s; deleting cache marker.", + tenant_id, + exc_info=True, + ) + try: + redis_client.delete(cache_key) + except (RedisError, RuntimeError): + logger.warning( + "Failed to delete invalid remote debug model plugin marker for tenant %s.", + tenant_id, + exc_info=True, + ) + return None + if isinstance(cached_marker, str): + return cached_marker + + logger.warning("Invalid remote debug model plugin marker for tenant %s; deleting cache marker.", tenant_id) + try: + redis_client.delete(cache_key) + except (RedisError, RuntimeError): + logger.warning( + "Failed to delete invalid remote debug model plugin marker for tenant %s.", + tenant_id, + exc_info=True, + ) + return None + + @classmethod + def _store_cached_remote_model_plugin_marker(cls, tenant_id: str, marker: str | None) -> None: + cache_key = cls._get_plugin_model_providers_remote_debug_cache_key(tenant_id) + try: + if marker is None: + redis_client.delete(cache_key) + else: + redis_client.setex(cache_key, dify_config.PLUGIN_MODEL_PROVIDERS_CACHE_TTL, marker) + except (RedisError, RuntimeError): + logger.warning("Failed to cache remote debug model plugin marker for tenant %s.", tenant_id, exc_info=True) + + @classmethod + def _load_cached_plugin_model_provider_plugin_ids(cls, tenant_id: str) -> set[str] | None: + """Return plugin ids represented by the current provider cache, or None when no usable cache exists.""" + generation = cls._load_plugin_model_providers_generation(tenant_id) + cached_providers, _ = cls._load_cached_plugin_model_providers_for_generation(tenant_id, generation) + if cached_providers is None: + return None + + plugin_ids: set[str] = set() + for provider in cached_providers: + last_slash = provider.provider.rfind("/") + if last_slash > 0: + plugin_ids.add(provider.provider[:last_slash]) + + return plugin_ids + + @classmethod + def _should_invalidate_model_provider_cache_for_remote_model_plugins( + cls, + tenant_id: str, + plugins: Sequence[PluginEntity], + ) -> bool: + remote_model_plugin_marker = cls._get_remote_model_plugin_cache_marker(plugins) + cached_remote_model_plugin_marker = cls._load_cached_remote_model_plugin_marker(tenant_id) + if remote_model_plugin_marker is None: + return cached_remote_model_plugin_marker is not None + + if remote_model_plugin_marker != cached_remote_model_plugin_marker: + return True + + remote_model_plugin_ids = { + plugin.plugin_id for plugin in plugins if plugin.source == PluginInstallationSource.Remote + } + cached_plugin_ids = cls._load_cached_plugin_model_provider_plugin_ids(tenant_id) + if cached_plugin_ids is None: + return False + + return not remote_model_plugin_ids.issubset(cached_plugin_ids) + @classmethod @contextmanager def _plugin_model_providers_refresh_lock( @@ -571,7 +681,21 @@ class PluginService: This keeps pagination usable before category is persisted on installation rows. """ manager = PluginInstaller() - return manager.list_plugins_by_category(tenant_id, category, page, page_size) + plugins = manager.list_plugins_by_category(tenant_id, category, page, page_size) + if category == PluginCategory.Model: + should_invalidate_model_provider_cache = ( + PluginService._should_invalidate_model_provider_cache_for_remote_model_plugins( + tenant_id, + plugins.list, + ) + ) + if should_invalidate_model_provider_cache: + PluginService.invalidate_plugin_model_providers_cache(tenant_id) + + remote_model_plugin_marker = PluginService._get_remote_model_plugin_cache_marker(plugins.list) + PluginService._store_cached_remote_model_plugin_marker(tenant_id, remote_model_plugin_marker) + + return plugins @staticmethod def _normalize_endpoint_count(value: object) -> int: diff --git a/api/tests/unit_tests/services/plugin/test_plugin_service.py b/api/tests/unit_tests/services/plugin/test_plugin_service.py index a8922154a95..278898926b9 100644 --- a/api/tests/unit_tests/services/plugin/test_plugin_service.py +++ b/api/tests/unit_tests/services/plugin/test_plugin_service.py @@ -8,7 +8,7 @@ import zstandard from pydantic import TypeAdapter from redis import RedisError -from core.plugin.entities.plugin import PluginInstallationSource +from core.plugin.entities.plugin import PluginCategory, PluginInstallationSource from core.plugin.entities.plugin_daemon import PluginInstallTask, PluginInstallTaskStatus, PluginModelProviderEntity from graphon.model_runtime.entities.common_entities import I18nObject from graphon.model_runtime.entities.provider_entities import ConfigurateMethod, ProviderEntity @@ -71,6 +71,16 @@ def _build_install_task(*, task_id: str = "task-1", status: PluginInstallTaskSta ) +def _build_remote_model_plugin( + *, plugin_id: str = "langgenius/debug-model", plugin_unique_identifier: str = "langgenius/debug-model:1.0.0" +) -> SimpleNamespace: + return SimpleNamespace( + plugin_id=plugin_id, + plugin_unique_identifier=plugin_unique_identifier, + source=PluginInstallationSource.Remote, + ) + + def _provider_cache_key(tenant_id: str, generation: int | None = None) -> str: if generation is None: return f"plugin_model_providers:tenant_id:{tenant_id}" @@ -797,6 +807,144 @@ class TestPluginListEndpointCounts: class TestPluginModelProviderCacheInvalidation: + def test_get_debugging_key_does_not_invalidate_model_provider_cache(self) -> None: + """Reading a debug key does not mean a debug runtime has registered a model provider.""" + with ( + patch(f"{MODULE}.PluginDebuggingClient") as debugging_client_cls, + patch(f"{MODULE}.PluginService.invalidate_plugin_model_providers_cache") as invalidate_cache, + ): + debugging_client_cls.return_value.get_debugging_key.return_value = "debug-key" + + from core.plugin.plugin_service import PluginService + + result = PluginService.get_debugging_key("tenant-1") + + assert result == "debug-key" + debugging_client_cls.return_value.get_debugging_key.assert_called_once_with("tenant-1") + invalidate_cache.assert_not_called() + + def test_list_model_category_invalidates_when_remote_model_plugin_is_missing_from_provider_cache(self) -> None: + """Remote model plugins are daemon-registered, so category reads repair a stale provider cache.""" + remote_plugin = _build_remote_model_plugin() + remote_plugin_marker = "langgenius/debug-model:langgenius/debug-model:1.0.0" + plugins = SimpleNamespace(list=[remote_plugin], has_more=False) + + with ( + patch(f"{MODULE}.PluginInstaller") as installer_cls, + patch( + f"{MODULE}.PluginService._load_cached_remote_model_plugin_marker", + return_value=remote_plugin_marker, + ), + patch( + f"{MODULE}.PluginService._load_cached_plugin_model_provider_plugin_ids", + return_value={"langgenius/openai"}, + ), + patch(f"{MODULE}.PluginService.invalidate_plugin_model_providers_cache") as invalidate_cache, + patch(f"{MODULE}.PluginService._store_cached_remote_model_plugin_marker") as store_marker, + ): + installer_cls.return_value.list_plugins_by_category.return_value = plugins + + from core.plugin.plugin_service import PluginService + + result = PluginService.list_by_category("tenant-1", PluginCategory.Model, 1, 100) + + assert result is plugins + installer_cls.return_value.list_plugins_by_category.assert_called_once_with( + "tenant-1", PluginCategory.Model, 1, 100 + ) + invalidate_cache.assert_called_once_with("tenant-1") + store_marker.assert_called_once_with("tenant-1", remote_plugin_marker) + + def test_list_model_category_invalidates_when_remote_model_plugin_identity_changes(self) -> None: + """A debug model plugin can share plugin_id with an installed plugin, so identity changes bust cache too.""" + remote_plugin = _build_remote_model_plugin( + plugin_id="langgenius/openai", + plugin_unique_identifier="langgenius/openai:debug", + ) + remote_plugin_marker = "langgenius/openai:langgenius/openai:debug" + plugins = SimpleNamespace(list=[remote_plugin], has_more=False) + + with ( + patch(f"{MODULE}.PluginInstaller") as installer_cls, + patch( + f"{MODULE}.PluginService._load_cached_remote_model_plugin_marker", + return_value="langgenius/openai:langgenius/openai:1.0.0", + ), + patch( + f"{MODULE}.PluginService._load_cached_plugin_model_provider_plugin_ids", + return_value={"langgenius/openai"}, + ) as load_cached_provider_plugin_ids, + patch(f"{MODULE}.PluginService.invalidate_plugin_model_providers_cache") as invalidate_cache, + patch(f"{MODULE}.PluginService._store_cached_remote_model_plugin_marker") as store_marker, + ): + installer_cls.return_value.list_plugins_by_category.return_value = plugins + + from core.plugin.plugin_service import PluginService + + result = PluginService.list_by_category("tenant-1", PluginCategory.Model, 1, 100) + + assert result is plugins + invalidate_cache.assert_called_once_with("tenant-1") + load_cached_provider_plugin_ids.assert_not_called() + store_marker.assert_called_once_with("tenant-1", remote_plugin_marker) + + def test_list_model_category_keeps_provider_cache_when_remote_model_plugin_is_already_cached(self) -> None: + """A connected remote model plugin should not force provider cache churn once represented.""" + remote_plugin = _build_remote_model_plugin() + remote_plugin_marker = "langgenius/debug-model:langgenius/debug-model:1.0.0" + plugins = SimpleNamespace(list=[remote_plugin], has_more=False) + + with ( + patch(f"{MODULE}.PluginInstaller") as installer_cls, + patch( + f"{MODULE}.PluginService._load_cached_remote_model_plugin_marker", + return_value=remote_plugin_marker, + ), + patch( + f"{MODULE}.PluginService._load_cached_plugin_model_provider_plugin_ids", + return_value={"langgenius/debug-model"}, + ), + patch(f"{MODULE}.PluginService.invalidate_plugin_model_providers_cache") as invalidate_cache, + patch(f"{MODULE}.PluginService._store_cached_remote_model_plugin_marker") as store_marker, + ): + installer_cls.return_value.list_plugins_by_category.return_value = plugins + + from core.plugin.plugin_service import PluginService + + result = PluginService.list_by_category("tenant-1", PluginCategory.Model, 1, 100) + + assert result is plugins + invalidate_cache.assert_not_called() + store_marker.assert_called_once_with("tenant-1", remote_plugin_marker) + + def test_list_model_category_invalidates_when_remote_model_plugin_disconnects(self) -> None: + """The current model category result clears provider cache when the previous debug model disappears.""" + installed_plugin = SimpleNamespace( + plugin_id="langgenius/openai", + plugin_unique_identifier="langgenius/openai:1.0.0", + source=PluginInstallationSource.Marketplace, + ) + plugins = SimpleNamespace(list=[installed_plugin], has_more=True) + + with ( + patch(f"{MODULE}.PluginInstaller") as installer_cls, + patch( + f"{MODULE}.PluginService._load_cached_remote_model_plugin_marker", + return_value="langgenius/debug-model:langgenius/debug-model:1.0.0", + ), + patch(f"{MODULE}.PluginService.invalidate_plugin_model_providers_cache") as invalidate_cache, + patch(f"{MODULE}.PluginService._store_cached_remote_model_plugin_marker") as store_marker, + ): + installer_cls.return_value.list_plugins_by_category.return_value = plugins + + from core.plugin.plugin_service import PluginService + + result = PluginService.list_by_category("tenant-1", PluginCategory.Model, 1, 100) + + assert result is plugins + invalidate_cache.assert_called_once_with("tenant-1") + store_marker.assert_called_once_with("tenant-1", None) + def test_fetch_install_task_invalidates_model_provider_cache_when_finished(self) -> None: """Finished plugin install tasks invalidate tenant provider cache.""" task = _build_install_task(status=PluginInstallTaskStatus.Success) diff --git a/web/app/components/header/account-setting/model-provider-page/__tests__/index.non-cloud.spec.tsx b/web/app/components/header/account-setting/model-provider-page/__tests__/index.non-cloud.spec.tsx index f7bf0d1ce05..8896c4575c8 100644 --- a/web/app/components/header/account-setting/model-provider-page/__tests__/index.non-cloud.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/__tests__/index.non-cloud.spec.tsx @@ -41,6 +41,7 @@ vi.mock('@/context/provider-context', () => ({ vi.mock('../hooks', () => ({ useDefaultModel: () => ({ data: null, isLoading: false }), + useLanguage: () => 'en_US', })) vi.mock('../provider-added-card', () => ({ @@ -84,9 +85,11 @@ vi.mock('@/app/components/plugins/plugin-page/use-reference-setting', () => ({ })) vi.mock('@/service/use-plugins', () => ({ - useCheckInstalled: () => ({ + useInstalledPluginList: () => ({ data: { plugins: [] }, }), + useInvalidateInstalledPluginList: () => vi.fn(), + useInvalidateCheckInstalled: () => vi.fn(), usePluginAutoUpgradeSettings: () => ({ data: { category: 'model', diff --git a/web/app/components/header/account-setting/model-provider-page/__tests__/index.spec.tsx b/web/app/components/header/account-setting/model-provider-page/__tests__/index.spec.tsx index 426276cfc5f..a648f4288af 100644 --- a/web/app/components/header/account-setting/model-provider-page/__tests__/index.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/__tests__/index.spec.tsx @@ -1,7 +1,9 @@ import type { ReactNode } from 'react' +import type { PluginDeclaration, PluginDetail } from '@/app/components/plugins/types' import { act, fireEvent, screen } from '@testing-library/react' import { describe, expect, it, vi } from 'vitest' import { renderWithSystemFeatures } from '@/__tests__/utils/mock-system-features' +import { PluginCategoryEnum, PluginSource } from '@/app/components/plugins/types' import { CurrentSystemQuotaTypeEnum, CustomConfigurationStatusEnum, @@ -41,10 +43,18 @@ const { mockReferenceSetting, mockAutoUpgradeError } = vi.hoisted(() => ({ }, })) -const { mockProviderContextState } = vi.hoisted(() => ({ +const { mockProviderContextState, mockRefreshModelProviders } = vi.hoisted(() => ({ mockProviderContextState: { isLoadingModelProviders: false, }, + mockRefreshModelProviders: vi.fn(), +})) + +const { mockInstalledModelPlugins, mockUseInstalledPluginList } = vi.hoisted(() => ({ + mockInstalledModelPlugins: { + value: [] as PluginDetail[], + }, + mockUseInstalledPluginList: vi.fn(), })) const mockQuotaConfig = { @@ -79,6 +89,70 @@ const saveUpdateSettings = () => { fireEvent.click(screen.getByRole('button', { name: 'common.operation.save' })) } +const createPluginDeclaration = (overrides: Partial = {}): PluginDeclaration => ({ + plugin_unique_identifier: 'langgenius/debug-model:1.0.0', + version: '1.0.0', + author: 'langgenius', + icon: 'debug-model.png', + icon_dark: 'debug-model-dark.png', + name: 'debug-model', + category: PluginCategoryEnum.model, + label: { en_US: 'Debug Model' } as unknown as PluginDeclaration['label'], + description: { en_US: 'Debug model provider' } as unknown as PluginDeclaration['description'], + created_at: '2024-01-01', + resource: null, + plugins: null, + verified: false, + endpoint: null, + tool: undefined, + datasource: undefined, + model: {}, + tags: [], + agent_strategy: null, + meta: { + version: '1.0.0', + }, + trigger: {} as unknown as PluginDeclaration['trigger'], + ...overrides, +}) + +const createPluginDetail = (overrides: Partial = {}): PluginDetail => { + const { + declaration: overrideDeclaration, + plugin_id: overridePluginId, + ...restOverrides + } = overrides + const declaration = overrideDeclaration ?? createPluginDeclaration() + const pluginId = overridePluginId ?? 'langgenius/debug-model' + + return { + id: 'plugin-installation-id', + created_at: '2024-01-01', + updated_at: '2024-01-01', + name: declaration.name, + plugin_id: pluginId, + plugin_unique_identifier: declaration.plugin_unique_identifier, + declaration, + installation_id: 'plugin-installation-id', + tenant_id: 'tenant-id', + endpoints_setups: 0, + endpoints_active: 0, + version: '1.0.0', + latest_version: '1.0.0', + latest_unique_identifier: declaration.plugin_unique_identifier, + source: PluginSource.debugging, + meta: { + repo: '', + version: '1.0.0', + package: '', + }, + status: 'active', + deprecated_reason: '', + alternative_plugin_id: '', + ...restOverrides, + } +} + const mockProviders = [ { provider: 'openai', @@ -106,6 +180,7 @@ vi.mock('@/context/provider-context', () => ({ useProviderContext: () => ({ modelProviders: mockProviders, isLoadingModelProviders: mockProviderContextState.isLoadingModelProviders, + refreshModelProviders: mockRefreshModelProviders, }), })) @@ -119,6 +194,7 @@ const mockDefaultModels: Record = vi.mock('../hooks', () => ({ useDefaultModel: (type: string) => mockDefaultModels[type] ?? { data: null, isLoading: false }, + useLanguage: () => 'en_US', })) vi.mock('../install-from-marketplace', () => ({ @@ -126,7 +202,24 @@ vi.mock('../install-from-marketplace', () => ({ })) vi.mock('../provider-added-card', () => ({ - default: ({ provider }: { provider: { provider: string } }) =>
    {provider.provider}
    , + default: ({ + notConfigured, + provider, + pluginDetail, + }: { + notConfigured?: boolean + provider: { provider: string } + pluginDetail?: { plugin_id: string, source?: string } + }) => ( +
    + {provider.provider} +
    + ), })) vi.mock('../provider-added-card/quota-panel', () => ({ @@ -160,12 +253,12 @@ vi.mock('@/app/components/plugins/plugin-page/use-reference-setting', () => ({ })) vi.mock('@/service/use-plugins', () => ({ - useInstalledPluginList: () => ({ - data: { plugins: [] }, - }), - useCheckInstalled: () => ({ - data: { plugins: [] }, - }), + useInstalledPluginList: (...args: unknown[]) => { + mockUseInstalledPluginList(...args) + return { + data: { plugins: mockInstalledModelPlugins.value }, + } + }, usePluginAutoUpgradeSettings: () => ({ data: mockReferenceSetting.auto_upgrade ? { @@ -280,6 +373,9 @@ describe('ModelProviderPage', () => { beforeEach(() => { vi.useFakeTimers() vi.clearAllMocks() + mockUseInstalledPluginList.mockClear() + mockRefreshModelProviders.mockClear() + mockInstalledModelPlugins.value = [] mockProviderContextState.isLoadingModelProviders = false mockAutoUpgradeError.value = undefined mockReferenceSetting.auto_upgrade = { @@ -418,6 +514,107 @@ describe('ModelProviderPage', () => { expect(screen.getByText('anthropic')).toBeInTheDocument() }) + it('should use the model plugin installation list to attach plugin detail to provider cards', () => { + mockProviders.splice(0, mockProviders.length, { + provider: 'langgenius/openai/openai', + label: { en_US: 'OpenAI' }, + custom_configuration: { status: CustomConfigurationStatusEnum.active }, + system_configuration: { + enabled: false, + current_quota_type: CurrentSystemQuotaTypeEnum.free, + quota_configurations: [mockQuotaConfig], + }, + }) + mockInstalledModelPlugins.value = [ + createPluginDetail({ + plugin_id: 'langgenius/openai', + declaration: createPluginDeclaration({ + plugin_unique_identifier: 'langgenius/openai:1.0.0', + name: 'openai', + label: { en_US: 'OpenAI Plugin' } as unknown as PluginDeclaration['label'], + }), + }), + ] + + renderModelProviderPage() + + expect(mockUseInstalledPluginList).toHaveBeenCalledWith(false, 100, { category: PluginCategoryEnum.model }) + expect(screen.getByTestId('provider-card')).toHaveAttribute('data-plugin-id', 'langgenius/openai') + expect(screen.queryByText('OpenAI Plugin')).not.toBeInTheDocument() + }) + + it('should not render installed model plugins that are not registered as model providers', () => { + mockInstalledModelPlugins.value = [ + createPluginDetail({ + plugin_id: 'langgenius/debug-model', + declaration: createPluginDeclaration({ + label: { en_US: 'Debug Model' } as unknown as PluginDeclaration['label'], + description: { en_US: 'Debug model provider' } as unknown as PluginDeclaration['description'], + }), + }), + ] + + renderModelProviderPage() + + expect(screen.queryByText('Debug Model')).not.toBeInTheDocument() + expect(screen.queryByText('langgenius/debug-model')).not.toBeInTheDocument() + expect(screen.queryByRole('button', { name: 'plugin actions langgenius/debug-model' })).not.toBeInTheDocument() + }) + + it('should refresh model providers once when a debugging model plugin is missing from providers', () => { + mockInstalledModelPlugins.value = [ + createPluginDetail({ + plugin_id: 'langgenius/debug-model', + declaration: createPluginDeclaration({ + label: { en_US: 'Debug Model' } as unknown as PluginDeclaration['label'], + }), + }), + ] + + renderModelProviderPage() + + expect(mockRefreshModelProviders).toHaveBeenCalledTimes(1) + }) + + it('should prefer debugging plugin detail when an installed model plugin shares the same plugin id', () => { + mockProviders.splice(0, mockProviders.length, { + provider: 'langgenius/openai/openai', + label: { en_US: 'OpenAI' }, + custom_configuration: { status: CustomConfigurationStatusEnum.active }, + system_configuration: { + enabled: false, + current_quota_type: CurrentSystemQuotaTypeEnum.free, + quota_configurations: [mockQuotaConfig], + }, + }) + mockInstalledModelPlugins.value = [ + createPluginDetail({ + plugin_id: 'langgenius/openai', + declaration: createPluginDeclaration({ + plugin_unique_identifier: 'langgenius/openai:debug', + name: 'openai', + label: { en_US: 'OpenAI Debug Plugin' } as unknown as PluginDeclaration['label'], + }), + source: PluginSource.debugging, + }), + createPluginDetail({ + plugin_id: 'langgenius/openai', + declaration: createPluginDeclaration({ + plugin_unique_identifier: 'langgenius/openai:1.0.0', + name: 'openai', + label: { en_US: 'OpenAI Installed Plugin' } as unknown as PluginDeclaration['label'], + }), + source: PluginSource.marketplace, + }), + ] + + renderModelProviderPage() + + expect(screen.getByTestId('provider-card')).toHaveAttribute('data-plugin-id', 'langgenius/openai') + expect(screen.getByTestId('provider-card')).toHaveAttribute('data-plugin-source', PluginSource.debugging) + expect(mockRefreshModelProviders).toHaveBeenCalledTimes(1) + }) + it('should show provider placeholders while model providers are loading', () => { mockProviderContextState.isLoadingModelProviders = true @@ -572,4 +769,66 @@ describe('ModelProviderPage', () => { ]) expect(screen.queryByText('common.modelProvider.toBeConfigured')).not.toBeInTheDocument() }) + + it('should prioritize debugging model plugins within their provider section', () => { + mockProviders.splice(0, mockProviders.length, { + provider: 'langgenius/openai/openai', + label: { en_US: 'OpenAI Fixed' }, + custom_configuration: { status: CustomConfigurationStatusEnum.active }, + system_configuration: { + enabled: false, + current_quota_type: CurrentSystemQuotaTypeEnum.free, + quota_configurations: [mockQuotaConfig], + }, + }, { + provider: 'zeta-provider', + label: { en_US: 'Zeta Provider' }, + custom_configuration: { status: CustomConfigurationStatusEnum.active }, + system_configuration: { + enabled: false, + current_quota_type: CurrentSystemQuotaTypeEnum.free, + quota_configurations: [mockQuotaConfig], + }, + }, { + provider: 'langgenius/normal-model/normal-model', + label: { en_US: 'Normal Model' }, + custom_configuration: { status: CustomConfigurationStatusEnum.noConfigure }, + system_configuration: { + enabled: false, + current_quota_type: CurrentSystemQuotaTypeEnum.free, + quota_configurations: [mockQuotaConfig], + }, + }, { + provider: 'langgenius/debug-model/debug-model', + label: { en_US: 'Debug Model' }, + custom_configuration: { status: CustomConfigurationStatusEnum.noConfigure }, + system_configuration: { + enabled: false, + current_quota_type: CurrentSystemQuotaTypeEnum.free, + quota_configurations: [mockQuotaConfig], + }, + }) + mockInstalledModelPlugins.value = [ + createPluginDetail({ + plugin_id: 'langgenius/debug-model', + declaration: createPluginDeclaration({ + plugin_unique_identifier: 'langgenius/debug-model:1.0.0', + name: 'debug-model', + label: { en_US: 'Debug Model' } as unknown as PluginDeclaration['label'], + }), + }), + ] + + renderModelProviderPage() + + const renderedProviders = screen.getAllByTestId('provider-card').map(item => item.textContent) + expect(renderedProviders).toEqual([ + 'langgenius/openai/openai', + 'zeta-provider', + 'langgenius/debug-model/debug-model', + 'langgenius/normal-model/normal-model', + ]) + expect(screen.getAllByTestId('provider-card')[2]).toHaveAttribute('data-not-configured', 'true') + expect(screen.getByText('common.modelProvider.toBeConfigured')).toBeInTheDocument() + }) }) diff --git a/web/app/components/header/account-setting/model-provider-page/index.tsx b/web/app/components/header/account-setting/model-provider-page/index.tsx index 5d8ddfd8f13..9bc43082349 100644 --- a/web/app/components/header/account-setting/model-provider-page/index.tsx +++ b/web/app/components/header/account-setting/model-provider-page/index.tsx @@ -6,15 +6,15 @@ import type { PluginDetail } from '@/app/components/plugins/types' import { useSuspenseQuery } from '@tanstack/react-query' import { useDebounce } from 'ahooks' import { noop } from 'es-toolkit/function' -import { useMemo } from 'react' +import { useEffect, useMemo, useRef } from 'react' import { useTranslation } from 'react-i18next' import { SearchInput } from '@/app/components/base/search-input' import { usePluginsWithLatestVersion } from '@/app/components/plugins/hooks' import { usePluginSettingsAccess } from '@/app/components/plugins/plugin-page/use-reference-setting' -import { PluginCategoryEnum } from '@/app/components/plugins/types' +import { PluginCategoryEnum, PluginSource } from '@/app/components/plugins/types' import { useProviderContext } from '@/context/provider-context' import { systemFeaturesQueryOptions } from '@/features/system-features/client' -import { useCheckInstalled } from '@/service/use-plugins' +import { useInstalledPluginList } from '@/service/use-plugins' import UpdateSettingDialog from '../update-setting-dialog' import { CustomConfigurationStatusEnum, @@ -25,7 +25,6 @@ import { } from './hooks' import ModelProviderPageBody from './model-provider-page-body' import SystemModelSelector from './system-model-selector' -import { providerToPluginId } from './utils' type SystemModelConfigStatus = 'no-provider' | 'none-configured' | 'partially-configured' | 'fully-configured' @@ -58,23 +57,43 @@ const ModelProviderPage = ({ const { data: rerankDefaultModel, isLoading: isRerankDefaultModelLoading } = useDefaultModel(ModelTypeEnum.rerank) const { data: speech2textDefaultModel, isLoading: isSpeech2textDefaultModelLoading } = useDefaultModel(ModelTypeEnum.speech2text) const { data: ttsDefaultModel, isLoading: isTTSDefaultModelLoading } = useDefaultModel(ModelTypeEnum.tts) - const { modelProviders: providers, isLoadingModelProviders } = useProviderContext() + const { modelProviders: providers, isLoadingModelProviders, refreshModelProviders } = useProviderContext() const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) - const allPluginIds = useMemo(() => { - return [...new Set(providers.map(p => providerToPluginId(p.provider)).filter(Boolean))] - }, [providers]) - const { data: installedPlugins } = useCheckInstalled({ - pluginIds: allPluginIds, - enabled: allPluginIds.length > 0, + const { data: installedModelPlugins } = useInstalledPluginList(false, 100, { + category: PluginCategoryEnum.model, }) - const enrichedPlugins = usePluginsWithLatestVersion(installedPlugins?.plugins) + const enrichedPlugins = usePluginsWithLatestVersion(installedModelPlugins?.plugins) const pluginDetailMap = useMemo(() => { const map = new Map() - for (const plugin of enrichedPlugins) - map.set(plugin.plugin_id, plugin) + for (const plugin of enrichedPlugins) { + const existingPlugin = map.get(plugin.plugin_id) + if (!existingPlugin || plugin.source === PluginSource.debugging) + map.set(plugin.plugin_id, plugin) + } return map }, [enrichedPlugins]) + const debuggingModelPluginKey = useMemo(() => { + const debuggingModelPluginIds = enrichedPlugins + .filter(plugin => plugin.source === PluginSource.debugging) + .map(plugin => `${plugin.plugin_id}:${plugin.plugin_unique_identifier}`) + .sort() + + return debuggingModelPluginIds.join(',') + }, [enrichedPlugins]) + const refreshedDebuggingModelPluginKeyRef = useRef('') + useEffect(() => { + if (!debuggingModelPluginKey) { + refreshedDebuggingModelPluginKeyRef.current = '' + return + } + + if (refreshedDebuggingModelPluginKeyRef.current === debuggingModelPluginKey) + return + + refreshedDebuggingModelPluginKeyRef.current = debuggingModelPluginKey + refreshModelProviders?.() + }, [debuggingModelPluginKey, refreshModelProviders]) const enableMarketplace = systemFeatures.enable_marketplace const isDefaultModelLoading = isTextGenerationDefaultModelLoading || isEmbeddingsDefaultModelLoading diff --git a/web/app/components/header/account-setting/model-provider-page/model-provider-page-body.tsx b/web/app/components/header/account-setting/model-provider-page/model-provider-page-body.tsx index 625731d29a3..e614b2c917c 100644 --- a/web/app/components/header/account-setting/model-provider-page/model-provider-page-body.tsx +++ b/web/app/components/header/account-setting/model-provider-page/model-provider-page-body.tsx @@ -3,6 +3,7 @@ import type { ModelProvider } from './declarations' import type { PluginDetail } from '@/app/components/plugins/types' import { Trans, useTranslation } from 'react-i18next' import { SkeletonContainer, SkeletonRectangle, SkeletonRow } from '@/app/components/base/skeleton' +import { PluginSource } from '@/app/components/plugins/types' import { IS_CLOUD_EDITION } from '@/config' import InstallFromMarketplace from './install-from-marketplace' import ProviderAddedCard from './provider-added-card' @@ -99,21 +100,40 @@ type ProviderCardListProps = { notConfigured?: boolean } +function isDebuggingProvider(provider: ModelProvider, pluginDetailMap: Map) { + return pluginDetailMap.get(providerToPluginId(provider.provider))?.source === PluginSource.debugging +} + function ProviderCardList({ providers, pluginDetailMap, notConfigured, }: ProviderCardListProps) { + const sortedProviders = [...providers] + .sort((a, b) => { + const aIsDebuggingPlugin = isDebuggingProvider(a, pluginDetailMap) + const bIsDebuggingPlugin = isDebuggingProvider(b, pluginDetailMap) + + if (aIsDebuggingPlugin === bIsDebuggingPlugin) + return 0 + + return aIsDebuggingPlugin ? -1 : 1 + }) + return (
    - {providers.map(provider => ( - - ))} + {sortedProviders.map((provider) => { + const pluginDetail = pluginDetailMap.get(providerToPluginId(provider.provider)) + + return ( + + ) + })}
    ) } @@ -157,8 +177,8 @@ const ModelProviderPageBody: FC = ({
    {t('modelProvider.toBeConfigured', { ns: 'common' })}
    diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/provider-card-actions.spec.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/provider-card-actions.spec.tsx index b62b56a8a6f..935292a599a 100644 --- a/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/provider-card-actions.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/provider-card-actions.spec.tsx @@ -158,6 +158,15 @@ describe('ProviderCardActions', () => { expect(mockHandleUpdate).toHaveBeenCalledWith(true) }) + it('should show a compact debug badge after the version for debugging plugins', () => { + render() + + const version = screen.getByText('1.0.0') + const debugBadge = screen.getByText('appDebug.operation.debugConfig') + + expect(version.compareDocumentPosition(debugBadge) & Node.DOCUMENT_POSITION_FOLLOWING).toBeTruthy() + }) + it('should trigger the latest marketplace update when clicking the update button', () => { render() diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/provider-card-actions.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/provider-card-actions.tsx index b2511b18670..4538cd9d16e 100644 --- a/web/app/components/header/account-setting/model-provider-page/provider-added-card/provider-card-actions.tsx +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/provider-card-actions.tsx @@ -30,6 +30,7 @@ const ProviderCardActions: FC = ({ detail, onUpdate }) => { const { source, version, latest_version, latest_unique_identifier, meta } = detail const author = detail.declaration?.author ?? '' const name = detail.declaration?.name ?? detail.name + const isDebuggingPlugin = source === PluginSource.debugging const { modalStates, @@ -80,31 +81,41 @@ const ProviderCardActions: FC = ({ detail, onUpdate }) => { return ( <> {!!version && ( - + + {version} + {canUpdatePlugin && isFromMarketplace && } + + )} + hasRedCornerMark={hasNewVersion} + /> + )} + /> + {isDebuggingPlugin && ( - {version} - {canUpdatePlugin && isFromMarketplace && } - - )} - hasRedCornerMark={hasNewVersion} + text={t('operation.debugConfig', { ns: 'appDebug' })} /> )} - /> + )} {canUpdatePlugin && (hasNewVersion || isFromGitHub) && ( From 56f3d0a11ef1127f04cb03848b2a4731be9c1115 Mon Sep 17 00:00:00 2001 From: Stephen Zhou Date: Tue, 7 Jul 2026 22:08:38 +0800 Subject: [PATCH 29/70] refactor(web): clarify app context bootstrap graph (#38516) Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- eslint-suppressions.json | 21 - .../__tests__/app-context-provider.spec.tsx | 446 +++++++++++++++--- web/context/app-context-effects.ts | 82 ++++ web/context/app-context-normalizers.ts | 78 +++ web/context/app-context-provider.tsx | 206 ++------ web/context/app-context-state.ts | 116 +++++ .../__tests__/use-permission-keys.spec.tsx | 33 +- .../access-control/use-permission-keys.ts | 14 +- web/service/lang-genius-version.ts | 13 + web/service/use-common.ts | 14 +- 10 files changed, 721 insertions(+), 302 deletions(-) create mode 100644 web/context/app-context-effects.ts create mode 100644 web/context/app-context-normalizers.ts create mode 100644 web/context/app-context-state.ts create mode 100644 web/service/lang-genius-version.ts diff --git a/eslint-suppressions.json b/eslint-suppressions.json index e4b92e2db43..176341002bf 100644 --- a/eslint-suppressions.json +++ b/eslint-suppressions.json @@ -7031,11 +7031,6 @@ "count": 1 } }, - "web/service/access-control/__tests__/use-permission-keys.spec.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "web/service/access-control/__tests__/use-workspace-access-rules.spec.tsx": { "no-restricted-imports": { "count": 1 @@ -7056,11 +7051,6 @@ "count": 1 } }, - "web/service/access-control/use-permission-keys.ts": { - "no-restricted-imports": { - "count": 1 - } - }, "web/service/access-control/use-workspace-access-rules.ts": { "no-restricted-imports": { "count": 1 @@ -7228,17 +7218,6 @@ "count": 1 } }, - "web/service/use-common.ts": { - "no-restricted-imports": { - "count": 1 - }, - "ts/no-empty-object-type": { - "count": 1 - }, - "ts/no-explicit-any": { - "count": 1 - } - }, "web/service/use-datasource.ts": { "no-restricted-imports": { "count": 1 diff --git a/web/context/__tests__/app-context-provider.spec.tsx b/web/context/__tests__/app-context-provider.spec.tsx index 0e55cd8459e..7b574823025 100644 --- a/web/context/__tests__/app-context-provider.spec.tsx +++ b/web/context/__tests__/app-context-provider.spec.tsx @@ -1,8 +1,18 @@ -import { render, screen } from '@testing-library/react' -import { useAppContext, useSelector } from '../app-context' +import type { ReactNode } from 'react' +import { QueryClient, QueryClientProvider } from '@tanstack/react-query' +import { fireEvent, render, screen, waitFor } from '@testing-library/react' +import { Provider as JotaiProvider } from 'jotai' +import { queryClientAtom } from 'jotai-tanstack-query' +import { useHydrateAtoms } from 'jotai/react/utils' +import { Suspense } from 'react' +import { setUserId, setUserProperties } from '@/app/components/base/amplitude' +import { flushRegistrationSuccess } from '@/app/components/base/amplitude/registration-tracking' +import { setZendeskConversationFields } from '@/app/components/base/zendesk/utils' +import { ZENDESK_FIELD_IDS } from '@/config' +import { initialWorkspaceInfo, useAppContext, useSelector } from '../app-context' import { AppContextProvider } from '../app-context-provider' -const mockInvalidateQueries = vi.hoisted(() => vi.fn()) +const mockGetRequest = vi.hoisted(() => vi.fn()) const mockPermissionKeysState = vi.hoisted(() => ({ isPending: false, permissionKeys: ['app.create_and_management'], @@ -19,55 +29,84 @@ const mockCurrentWorkspaceResponse = vi.hoisted(() => ({ next_credit_reset_date: 1706745600, custom_config: {}, })) - -vi.mock('@tanstack/react-query', () => ({ - useQueryClient: () => ({ - invalidateQueries: mockInvalidateQueries, - }), - useSuspenseQuery: (options: { queryKey?: readonly unknown[] }) => { - if (options.queryKey?.[0] === 'system-features') { - return { - data: { - branding: { - enabled: false, - }, - }, - } +const mockCurrentWorkspaceQueryState = vi.hoisted(() => ({ + data: mockCurrentWorkspaceResponse as typeof mockCurrentWorkspaceResponse | undefined, + isPending: false, +})) +const mockUserProfileResponseState = vi.hoisted(() => ({ + data: { + profile: { + id: 'user-1', + name: 'User', + email: 'user@example.com', + avatar: '', + avatar_url: '', + is_password_set: true, + }, + meta: { + currentVersion: '1.0.0', + currentEnv: 'cloud', + }, + } as { + profile?: { + id: string + name: string + email: string + avatar: string + avatar_url: string + is_password_set: boolean } - - return { - data: { - profile: { - id: 'user-1', - name: 'User', - email: 'user@example.com', - avatar: '', - avatar_url: '', - is_password_set: true, - }, - meta: { - currentVersion: '1.0.0', - currentEnv: 'cloud', - }, - }, + meta: { + currentVersion: string | null + currentEnv: string | null } }, - useQuery: (options: { select?: (workspace: typeof mockCurrentWorkspaceResponse) => unknown }) => ({ - data: options.select ? options.select(mockCurrentWorkspaceResponse) : mockCurrentWorkspaceResponse, - isFetching: false, - isPending: false, - }), })) +const mockSystemFeaturesState = vi.hoisted(() => ({ + data: { + branding: { + enabled: false, + }, + }, +})) +const mockLangGeniusVersionState = vi.hoisted(() => ({ + data: { + version: '1.0.1', + release_date: '', + release_notes: '', + can_auto_update: false, + } as { + version: string + release_date: string + release_notes: string + can_auto_update: boolean + } | undefined, +})) + +vi.mock('@/config', async (importOriginal) => { + const actual = await importOriginal() + return { + ...actual, + ZENDESK_FIELD_IDS: { + ENVIRONMENT: 'environment-field', + VERSION: 'version-field', + EMAIL: 'email-field', + WORKSPACE_ID: 'workspace-id-field', + }, + } +}) vi.mock('@/features/system-features/client', () => ({ systemFeaturesQueryOptions: () => ({ queryKey: ['system-features'], + queryFn: async () => mockSystemFeaturesState.data, }), })) vi.mock('@/features/account-profile/client', () => ({ userProfileQueryOptions: () => ({ queryKey: ['user-profile'], + queryFn: async () => mockUserProfileResponseState.data, }), })) @@ -77,8 +116,16 @@ vi.mock('@/service/client', () => ({ current: { post: { key: () => ['current-workspace'], - queryOptions: (options: Record) => ({ + queryOptions: (options: { + select?: (workspace?: typeof mockCurrentWorkspaceResponse) => unknown + }) => ({ queryKey: ['current-workspace'], + queryFn: async () => { + if (mockCurrentWorkspaceQueryState.isPending) + return new Promise(() => {}) + + return mockCurrentWorkspaceQueryState.data + }, ...options, }), }, @@ -87,34 +134,9 @@ vi.mock('@/service/client', () => ({ }, })) -vi.mock('@/service/access-control/use-permission-keys', () => ({ - useWorkspacePermissionKeys: () => ({ - data: { - workspace: { - permission_keys: mockPermissionKeysState.permissionKeys, - }, - app: { - default_permission_keys: [], - overrides: [], - }, - dataset: { - default_permission_keys: [], - overrides: [], - }, - }, - isPending: mockPermissionKeysState.isPending, - }), -})) - -vi.mock('@/service/use-common', () => ({ - useLangGeniusVersion: () => ({ - data: { - version: '1.0.1', - release_date: '', - release_notes: '', - can_auto_update: false, - }, - }), +vi.mock('@/service/base', () => ({ + get: mockGetRequest, + post: vi.fn(), })) vi.mock('@/app/components/base/amplitude', () => ({ @@ -145,35 +167,303 @@ function AppContextProbe() { {selectedWorkspacePermissionKeys.join(',')} - loading: + permission loading: {String(context.isLoadingWorkspacePermissionKeys)} + + workspace loading: + {String(context.isLoadingCurrentWorkspace)} + + + workspace validating: + {String(context.isValidatingCurrentWorkspace)} + + + user: + {context.userProfile.email} + + + workspace: + {context.currentWorkspace.name} + role: {context.currentWorkspace.role} + + manager: + {String(context.isCurrentWorkspaceManager)} + + + owner: + {String(context.isCurrentWorkspaceOwner)} + + + editor: + {String(context.isCurrentWorkspaceEditor)} + + + dataset operator: + {String(context.isCurrentWorkspaceDatasetOperator)} + + + version: + {context.langGeniusVersionInfo.current_version} + / + {context.langGeniusVersionInfo.latest_version} + / + {context.langGeniusVersionInfo.current_env} + + + ) } +function TestQueryClientHydrator({ + children, + queryClient, +}: { + children: ReactNode + queryClient: QueryClient +}) { + useHydrateAtoms(new Map([[queryClientAtom, queryClient]])) + + return children +} + +function createTestQueryClient() { + return new QueryClient({ + defaultOptions: { + queries: { + retry: false, + staleTime: 0, + }, + }, + }) +} + +function renderProvider() { + const queryClient = createTestQueryClient() + const view = render( + + + + loading}> + + + + + + + , + ) + + return { + ...view, + queryClient, + } +} + describe('AppContextProvider', () => { beforeEach(() => { vi.clearAllMocks() mockPermissionKeysState.isPending = false mockPermissionKeysState.permissionKeys = ['app.create_and_management'] + mockCurrentWorkspaceQueryState.data = mockCurrentWorkspaceResponse + mockCurrentWorkspaceQueryState.isPending = false + mockUserProfileResponseState.data = { + profile: { + id: 'user-1', + name: 'User', + email: 'user@example.com', + avatar: '', + avatar_url: '', + is_password_set: true, + }, + meta: { + currentVersion: '1.0.0', + currentEnv: 'cloud', + }, + } + mockSystemFeaturesState.data = { + branding: { + enabled: false, + }, + } + mockLangGeniusVersionState.data = { + version: '1.0.1', + release_date: '', + release_notes: '', + can_auto_update: false, + } + mockGetRequest.mockImplementation((url: string) => { + if (url === '/workspaces/current/rbac/my-permissions') { + if (mockPermissionKeysState.isPending) + return new Promise(() => {}) + + return Promise.resolve({ + workspace: { + permission_keys: mockPermissionKeysState.permissionKeys, + }, + app: { + default_permission_keys: [], + overrides: [], + }, + dataset: { + default_permission_keys: [], + overrides: [], + }, + }) + } + + if (url === '/version') + return Promise.resolve(mockLangGeniusVersionState.data) + + return Promise.reject(new Error(`Unexpected GET ${url}`)) + }) }) - describe('Workspace Permission Keys', () => { - it('should provide current workspace permission keys from my-permissions', () => { - render( - - - , - ) + describe('Context compatibility values', () => { + it('should provide profile, workspace, permissions, loading state, and version metadata', async () => { + renderProvider() - expect(screen.getByText('keys:app.create_and_management')).toBeInTheDocument() - expect(screen.getByText('loading:false')).toBeInTheDocument() - expect(screen.getByText('role:editor')).toBeInTheDocument() + expect(await screen.findByText('user:user@example.com')).toBeInTheDocument() + expect(await screen.findByText('workspace:Workspace')).toBeInTheDocument() + expect(await screen.findByText('keys:app.create_and_management')).toBeInTheDocument() + expect(screen.getByText('permission loading:false')).toBeInTheDocument() + expect(screen.getByText('workspace loading:false')).toBeInTheDocument() + expect(screen.getByText('workspace validating:false')).toBeInTheDocument() + expect(await screen.findByText('version:1.0.0/1.0.1/cloud')).toBeInTheDocument() + }) + + it('should fall back to placeholder values when profile, workspace, permission, or version data is missing', async () => { + mockUserProfileResponseState.data = { + meta: { + currentVersion: null, + currentEnv: null, + }, + } + mockCurrentWorkspaceQueryState.data = undefined + mockPermissionKeysState.permissionKeys = [] + mockLangGeniusVersionState.data = undefined + + renderProvider() + + expect(await screen.findByText('user:')).toBeInTheDocument() + expect(screen.getByText(`workspace:${initialWorkspaceInfo.name}`)).toBeInTheDocument() + expect(screen.getByText(`role:${initialWorkspaceInfo.role}`)).toBeInTheDocument() + expect(screen.getByText('keys:')).toBeInTheDocument() + expect(screen.getByText('version://')).toBeInTheDocument() + }) + + it('should normalize invalid workspace roles to the initial workspace role', async () => { + mockCurrentWorkspaceQueryState.data = { + ...mockCurrentWorkspaceResponse, + role: 'unsupported-role', + } + + renderProvider() + + expect(await screen.findByText(`role:${initialWorkspaceInfo.role}`)).toBeInTheDocument() + }) + + it('should derive role flags from the current workspace role', async () => { + mockCurrentWorkspaceQueryState.data = { + ...mockCurrentWorkspaceResponse, + role: 'owner', + } + + renderProvider() + + expect(await screen.findByText('manager:true')).toBeInTheDocument() + expect(screen.getByText('owner:true')).toBeInTheDocument() + expect(screen.getByText('editor:true')).toBeInTheDocument() + expect(screen.getByText('dataset operator:false')).toBeInTheDocument() + }) + + it('should expose query loading and validating state', async () => { + mockPermissionKeysState.isPending = true + mockCurrentWorkspaceQueryState.isPending = true + + renderProvider() + + expect(await screen.findByText('workspace loading:true')).toBeInTheDocument() + expect(screen.getByText('workspace validating:true')).toBeInTheDocument() + expect(screen.getByText('permission loading:true')).toBeInTheDocument() + }) + }) + + describe('Refresh actions', () => { + it('should invalidate the source queries when refresh actions are called', async () => { + const { queryClient } = renderProvider() + const invalidateQueriesSpy = vi.spyOn(queryClient, 'invalidateQueries') + + fireEvent.click(await screen.findByRole('button', { name: /refresh user/i })) + fireEvent.click(screen.getByRole('button', { name: /refresh workspace/i })) + + expect(invalidateQueriesSpy).toHaveBeenCalledWith({ queryKey: ['user-profile'] }) + expect(invalidateQueriesSpy).toHaveBeenCalledWith({ queryKey: ['current-workspace'] }) + }) + }) + + describe('External side effects', () => { + it('should sync Zendesk fields and Amplitude identity when bootstrap data is available', async () => { + renderProvider() + + await waitFor(() => { + expect(setZendeskConversationFields).toHaveBeenCalledWith([{ + id: ZENDESK_FIELD_IDS.ENVIRONMENT, + value: 'cloud', + }]) + }) + expect(setZendeskConversationFields).toHaveBeenCalledWith([{ + id: ZENDESK_FIELD_IDS.VERSION, + value: '1.0.1', + }]) + expect(setZendeskConversationFields).toHaveBeenCalledWith([{ + id: ZENDESK_FIELD_IDS.EMAIL, + value: 'user@example.com', + }]) + await waitFor(() => { + expect(setZendeskConversationFields).toHaveBeenCalledWith([{ + id: ZENDESK_FIELD_IDS.WORKSPACE_ID, + value: 'workspace-1', + }]) + }) + await waitFor(() => { + expect(setUserId).toHaveBeenCalledWith('user@example.com') + expect(setUserProperties).toHaveBeenCalledWith(expect.objectContaining({ + email: 'user@example.com', + workspace_id: 'workspace-1', + workspace_role: 'editor', + })) + expect(flushRegistrationSuccess).toHaveBeenCalled() + }) + }) + + it('should not sync Amplitude identity when user id is missing', async () => { + mockUserProfileResponseState.data = { + profile: { + id: '', + name: '', + email: '', + avatar: '', + avatar_url: '', + is_password_set: false, + }, + meta: { + currentVersion: '1.0.0', + currentEnv: 'cloud', + }, + } + + renderProvider() + + await screen.findByText('user:') + expect(setUserId).not.toHaveBeenCalled() + expect(setUserProperties).not.toHaveBeenCalled() + expect(flushRegistrationSuccess).not.toHaveBeenCalled() }) }) }) diff --git a/web/context/app-context-effects.ts b/web/context/app-context-effects.ts new file mode 100644 index 00000000000..875cc2ec638 --- /dev/null +++ b/web/context/app-context-effects.ts @@ -0,0 +1,82 @@ +'use client' + +import { useAtomValue } from 'jotai' +import { useEffect } from 'react' +import { setUserId, setUserProperties } from '@/app/components/base/amplitude' +import { flushRegistrationSuccess } from '@/app/components/base/amplitude/registration-tracking' +import { setZendeskConversationFields } from '@/app/components/base/zendesk/utils' +import { ZENDESK_FIELD_IDS } from '@/config' +import { + currentWorkspaceAtom, + langGeniusVersionInfoAtom, + userProfileAtom, +} from './app-context-state' + +export function useSyncZendeskFields() { + const userProfile = useAtomValue(userProfileAtom) + const currentWorkspace = useAtomValue(currentWorkspaceAtom) + const langGeniusVersionInfo = useAtomValue(langGeniusVersionInfoAtom) + + useEffect(() => { + if (ZENDESK_FIELD_IDS.ENVIRONMENT && langGeniusVersionInfo?.current_env) { + setZendeskConversationFields([{ + id: ZENDESK_FIELD_IDS.ENVIRONMENT, + value: langGeniusVersionInfo.current_env.toLowerCase(), + }]) + } + }, [langGeniusVersionInfo?.current_env]) + + useEffect(() => { + if (ZENDESK_FIELD_IDS.VERSION && langGeniusVersionInfo?.version) { + setZendeskConversationFields([{ + id: ZENDESK_FIELD_IDS.VERSION, + value: langGeniusVersionInfo.version, + }]) + } + }, [langGeniusVersionInfo?.version]) + + useEffect(() => { + if (ZENDESK_FIELD_IDS.EMAIL && userProfile?.email) { + setZendeskConversationFields([{ + id: ZENDESK_FIELD_IDS.EMAIL, + value: userProfile.email, + }]) + } + }, [userProfile?.email]) + + useEffect(() => { + if (ZENDESK_FIELD_IDS.WORKSPACE_ID && currentWorkspace?.id) { + setZendeskConversationFields([{ + id: ZENDESK_FIELD_IDS.WORKSPACE_ID, + value: currentWorkspace.id, + }]) + } + }, [currentWorkspace?.id]) +} + +export function useSyncAmplitudeIdentity() { + const userProfile = useAtomValue(userProfileAtom) + const currentWorkspace = useAtomValue(currentWorkspaceAtom) + + useEffect(() => { + if (userProfile?.id) { + setUserId(userProfile.email) + const properties: Record = { + email: userProfile.email, + name: userProfile.name, + has_password: userProfile.is_password_set, + } + + if (currentWorkspace?.id) { + properties.workspace_id = currentWorkspace.id + properties.workspace_name = currentWorkspace.name + properties.workspace_plan = currentWorkspace.plan + properties.workspace_status = currentWorkspace.status + properties.workspace_role = currentWorkspace.role + } + + setUserProperties(properties) + flushRegistrationSuccess() + } + }, [userProfile, currentWorkspace]) +} diff --git a/web/context/app-context-normalizers.ts b/web/context/app-context-normalizers.ts new file mode 100644 index 00000000000..244d88fa682 --- /dev/null +++ b/web/context/app-context-normalizers.ts @@ -0,0 +1,78 @@ +import type { PostWorkspacesCurrentResponse } from '@dify/contracts/api/console/workspaces/types.gen' +import type { ICurrentWorkspace, LangGeniusVersionResponse } from '@/models/common' +import { initialLangGeniusVersionInfo, initialWorkspaceInfo } from './app-context' + +const workspaceRoles = new Set(['owner', 'admin', 'editor', 'dataset_operator', 'normal']) + +export const emptyWorkspacePermissionKeys: string[] = [] + +export type WorkspaceRoleFlags = { + isCurrentWorkspaceManager: boolean + isCurrentWorkspaceOwner: boolean + isCurrentWorkspaceEditor: boolean + isCurrentWorkspaceDatasetOperator: boolean +} + +export type ProfileMeta = { + currentVersion: string | null + currentEnv: string | null +} + +function resolveWorkspaceRole(role: PostWorkspacesCurrentResponse['role']): ICurrentWorkspace['role'] { + if (role && workspaceRoles.has(role as ICurrentWorkspace['role'])) + return role as ICurrentWorkspace['role'] + + return initialWorkspaceInfo.role +} + +export function normalizeCurrentWorkspace(workspace?: PostWorkspacesCurrentResponse): ICurrentWorkspace { + if (!workspace) + return initialWorkspaceInfo + + return { + id: workspace.id, + name: workspace.name ?? initialWorkspaceInfo.name, + plan: workspace.plan ?? initialWorkspaceInfo.plan, + status: workspace.status ?? initialWorkspaceInfo.status, + created_at: workspace.created_at ?? initialWorkspaceInfo.created_at, + role: resolveWorkspaceRole(workspace.role), + providers: initialWorkspaceInfo.providers, + trial_credits: workspace.trial_credits ?? initialWorkspaceInfo.trial_credits, + trial_credits_used: workspace.trial_credits_used ?? initialWorkspaceInfo.trial_credits_used, + next_credit_reset_date: workspace.next_credit_reset_date ?? initialWorkspaceInfo.next_credit_reset_date, + trial_end_reason: workspace.trial_end_reason ?? undefined, + custom_config: workspace.custom_config + ? { + remove_webapp_brand: workspace.custom_config.remove_webapp_brand ?? undefined, + replace_webapp_logo: workspace.custom_config.replace_webapp_logo ?? undefined, + } + : undefined, + } +} + +export function getWorkspaceRoleFlags(currentWorkspace: ICurrentWorkspace): WorkspaceRoleFlags { + return { + isCurrentWorkspaceManager: ['owner', 'admin'].includes(currentWorkspace.role), + isCurrentWorkspaceOwner: currentWorkspace.role === 'owner', + isCurrentWorkspaceEditor: ['owner', 'admin', 'editor'].includes(currentWorkspace.role), + isCurrentWorkspaceDatasetOperator: currentWorkspace.role === 'dataset_operator', + } +} + +export function getLangGeniusVersionInfo({ + meta, + versionData, +}: { + meta: ProfileMeta + versionData?: Omit +}): LangGeniusVersionResponse { + if (!meta.currentVersion || !versionData) + return initialLangGeniusVersionInfo + + return { + ...versionData, + current_version: meta.currentVersion, + latest_version: versionData.version, + current_env: meta.currentEnv || '', + } +} diff --git a/web/context/app-context-provider.tsx b/web/context/app-context-provider.tsx index 509d36e3e59..64dda2cc753 100644 --- a/web/context/app-context-provider.tsx +++ b/web/context/app-context-provider.tsx @@ -1,193 +1,65 @@ 'use client' -import type { GetAccountProfileResponse } from '@dify/contracts/api/console/account/types.gen' -import type { PostWorkspacesCurrentResponse } from '@dify/contracts/api/console/workspaces/types.gen' import type { ReactNode } from 'react' -import type { ICurrentWorkspace, LangGeniusVersionResponse } from '@/models/common' -import { useQuery, useQueryClient, useSuspenseQuery } from '@tanstack/react-query' -import { useCallback, useEffect, useMemo } from 'react' -import { setUserId, setUserProperties } from '@/app/components/base/amplitude' -import { flushRegistrationSuccess } from '@/app/components/base/amplitude/registration-tracking' -import { setZendeskConversationFields } from '@/app/components/base/zendesk/utils' -import { ZENDESK_FIELD_IDS } from '@/config' +import { useAtomValue, useSetAtom } from 'jotai' import { AppContext, - initialLangGeniusVersionInfo, - initialWorkspaceInfo, - userProfilePlaceholder, useSelector, } from '@/context/app-context' -import { userProfileQueryOptions } from '@/features/account-profile/client' -import { systemFeaturesQueryOptions } from '@/features/system-features/client' -import { useWorkspacePermissionKeys } from '@/service/access-control/use-permission-keys' -import { consoleQuery } from '@/service/client' import { - useLangGeniusVersion, -} from '@/service/use-common' + currentWorkspaceAtom, + currentWorkspaceLoadingAtom, + currentWorkspaceValidatingAtom, + langGeniusVersionInfoAtom, + refreshCurrentWorkspaceAtom, + refreshUserProfileAtom, + userProfileAtom, + workspacePermissionKeysAtom, + workspacePermissionKeysLoadingAtom, + workspaceRoleFlagsAtom, +} from '@/context/app-context-state' +import { + useSyncAmplitudeIdentity, + useSyncZendeskFields, +} from './app-context-effects' type AppContextProviderProps = { children: ReactNode } -const workspaceRoles = new Set(['owner', 'admin', 'editor', 'dataset_operator', 'normal']) -const emptyWorkspacePermissionKeys: string[] = [] - -const resolveWorkspaceRole = (role: PostWorkspacesCurrentResponse['role']): ICurrentWorkspace['role'] => { - if (role && workspaceRoles.has(role as ICurrentWorkspace['role'])) - return role as ICurrentWorkspace['role'] - - return initialWorkspaceInfo.role -} - -const normalizeCurrentWorkspace = (workspace?: PostWorkspacesCurrentResponse): ICurrentWorkspace => { - if (!workspace) - return initialWorkspaceInfo - - return { - id: workspace.id, - name: workspace.name ?? initialWorkspaceInfo.name, - plan: workspace.plan ?? initialWorkspaceInfo.plan, - status: workspace.status ?? initialWorkspaceInfo.status, - created_at: workspace.created_at ?? initialWorkspaceInfo.created_at, - role: resolveWorkspaceRole(workspace.role), - providers: initialWorkspaceInfo.providers, - trial_credits: workspace.trial_credits ?? initialWorkspaceInfo.trial_credits, - trial_credits_used: workspace.trial_credits_used ?? initialWorkspaceInfo.trial_credits_used, - next_credit_reset_date: workspace.next_credit_reset_date ?? initialWorkspaceInfo.next_credit_reset_date, - trial_end_reason: workspace.trial_end_reason ?? undefined, - custom_config: workspace.custom_config - ? { - remove_webapp_brand: workspace.custom_config.remove_webapp_brand ?? undefined, - replace_webapp_logo: workspace.custom_config.replace_webapp_logo ?? undefined, - } - : undefined, - } -} - export function AppContextProvider({ children }: AppContextProviderProps) { - const queryClient = useQueryClient() - const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) - const { data: userProfileResp } = useSuspenseQuery(userProfileQueryOptions()) - const currentWorkspaceQuery = useQuery(consoleQuery.workspaces.current.post.queryOptions({ - select: normalizeCurrentWorkspace, - })) - const workspacePermissionKeysQuery = useWorkspacePermissionKeys() - const langGeniusVersionQuery = useLangGeniusVersion( - userProfileResp?.meta.currentVersion, - !systemFeatures.branding.enabled, - ) + const userProfile = useAtomValue(userProfileAtom) + const currentWorkspace = useAtomValue(currentWorkspaceAtom) + const roleFlags = useAtomValue(workspaceRoleFlagsAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) + const langGeniusVersionInfo = useAtomValue(langGeniusVersionInfoAtom) + const isLoadingCurrentWorkspace = useAtomValue(currentWorkspaceLoadingAtom) + const isValidatingCurrentWorkspace = useAtomValue(currentWorkspaceValidatingAtom) + const isLoadingWorkspacePermissionKeys = useAtomValue(workspacePermissionKeysLoadingAtom) - const userProfile = useMemo(() => userProfileResp?.profile || userProfilePlaceholder, [userProfileResp?.profile]) - const currentWorkspace = currentWorkspaceQuery.data ?? initialWorkspaceInfo - const langGeniusVersionInfo = useMemo(() => { - if (!userProfileResp?.meta?.currentVersion || !langGeniusVersionQuery.data) - return initialLangGeniusVersionInfo + const refreshUserProfile = useSetAtom(refreshUserProfileAtom) + const refreshCurrentWorkspace = useSetAtom(refreshCurrentWorkspaceAtom) - const current_version = userProfileResp.meta.currentVersion - const current_env = userProfileResp.meta.currentEnv || '' - const versionData = langGeniusVersionQuery.data - return { - ...versionData, - current_version, - latest_version: versionData.version, - current_env, - } - }, [langGeniusVersionQuery.data, userProfileResp?.meta]) - - const isCurrentWorkspaceManager = useMemo(() => ['owner', 'admin'].includes(currentWorkspace.role), [currentWorkspace.role]) - const isCurrentWorkspaceOwner = useMemo(() => currentWorkspace.role === 'owner', [currentWorkspace.role]) - const isCurrentWorkspaceEditor = useMemo(() => ['owner', 'admin', 'editor'].includes(currentWorkspace.role), [currentWorkspace.role]) - const isCurrentWorkspaceDatasetOperator = useMemo(() => currentWorkspace.role === 'dataset_operator', [currentWorkspace.role]) - - const mutateUserProfile = useCallback(() => { - queryClient.invalidateQueries({ queryKey: userProfileQueryOptions().queryKey }) - }, [queryClient]) - - const mutateCurrentWorkspace = useCallback(() => { - queryClient.invalidateQueries({ queryKey: consoleQuery.workspaces.current.post.key() }) - }, [queryClient]) - - // #region Zendesk conversation fields - useEffect(() => { - if (ZENDESK_FIELD_IDS.ENVIRONMENT && langGeniusVersionInfo?.current_env) { - setZendeskConversationFields([{ - id: ZENDESK_FIELD_IDS.ENVIRONMENT, - value: langGeniusVersionInfo.current_env.toLowerCase(), - }]) - } - }, [langGeniusVersionInfo?.current_env]) - - useEffect(() => { - if (ZENDESK_FIELD_IDS.VERSION && langGeniusVersionInfo?.version) { - setZendeskConversationFields([{ - id: ZENDESK_FIELD_IDS.VERSION, - value: langGeniusVersionInfo.version, - }]) - } - }, [langGeniusVersionInfo?.version]) - - useEffect(() => { - if (ZENDESK_FIELD_IDS.EMAIL && userProfile?.email) { - setZendeskConversationFields([{ - id: ZENDESK_FIELD_IDS.EMAIL, - value: userProfile.email, - }]) - } - }, [userProfile?.email]) - - useEffect(() => { - if (ZENDESK_FIELD_IDS.WORKSPACE_ID && currentWorkspace?.id) { - setZendeskConversationFields([{ - id: ZENDESK_FIELD_IDS.WORKSPACE_ID, - value: currentWorkspace.id, - }]) - } - }, [currentWorkspace?.id]) - // #endregion Zendesk conversation fields - - useEffect(() => { - // Report user and workspace info to Amplitude when loaded - if (userProfile?.id) { - setUserId(userProfile.email) - const properties: Record = { - email: userProfile.email, - name: userProfile.name, - has_password: userProfile.is_password_set, - } - - if (currentWorkspace?.id) { - properties.workspace_id = currentWorkspace.id - properties.workspace_name = currentWorkspace.name - properties.workspace_plan = currentWorkspace.plan - properties.workspace_status = currentWorkspace.status - properties.workspace_role = currentWorkspace.role - } - - setUserProperties(properties) - - // The user ID is now attached, so replay any registration success event captured - // at signup time. This makes it land on the identified Amplitude profile instead - // of an anonymous one (no-op when nothing was deferred). - flushRegistrationSuccess() - } - }, [userProfile, currentWorkspace]) + useSyncZendeskFields() + useSyncAmplitudeIdentity() return ( { + refreshUserProfile() + }, langGeniusVersionInfo, useSelector, currentWorkspace, - isCurrentWorkspaceManager, - isCurrentWorkspaceOwner, - isCurrentWorkspaceEditor, - isCurrentWorkspaceDatasetOperator, - mutateCurrentWorkspace, - isLoadingCurrentWorkspace: currentWorkspaceQuery.isPending, - isLoadingWorkspacePermissionKeys: workspacePermissionKeysQuery.isPending, - isValidatingCurrentWorkspace: currentWorkspaceQuery.isFetching, - workspacePermissionKeys: workspacePermissionKeysQuery.data?.workspace.permission_keys ?? emptyWorkspacePermissionKeys, + ...roleFlags, + mutateCurrentWorkspace: () => { + refreshCurrentWorkspace() + }, + isLoadingCurrentWorkspace, + isLoadingWorkspacePermissionKeys, + isValidatingCurrentWorkspace, + workspacePermissionKeys, }} > {children} diff --git a/web/context/app-context-state.ts b/web/context/app-context-state.ts new file mode 100644 index 00000000000..780138f4d0f --- /dev/null +++ b/web/context/app-context-state.ts @@ -0,0 +1,116 @@ +'use client' + +import type { GetAccountProfileResponse } from '@dify/contracts/api/console/account/types.gen' +import type { GetSystemFeaturesResponse } from '@dify/contracts/api/console/system-features/types.gen' +import type { DefinedQueryObserverResult } from '@tanstack/react-query' +import type { UserProfileWithMeta } from '@/features/account-profile/client' +import { atom } from 'jotai' +import { atomWithQuery, atomWithSuspenseQuery, queryClientAtom } from 'jotai-tanstack-query' +import { userProfileQueryOptions } from '@/features/account-profile/client' +import { systemFeaturesQueryOptions } from '@/features/system-features/client' +import { workspacePermissionKeysQueryOptions } from '@/service/access-control/use-permission-keys' +import { consoleQuery } from '@/service/client' +import { langGeniusVersionQueryOptions } from '@/service/lang-genius-version' +import { + initialLangGeniusVersionInfo, + initialWorkspaceInfo, + userProfilePlaceholder, +} from './app-context' +import { + emptyWorkspacePermissionKeys, + getLangGeniusVersionInfo, + getWorkspaceRoleFlags, + normalizeCurrentWorkspace, +} from './app-context-normalizers' + +type SuspenseQueryResult = Omit, 'isPlaceholderData'> + +const accountProfileQueryAtom = atomWithSuspenseQuery(() => userProfileQueryOptions()) + +const systemFeaturesQueryAtom = atomWithSuspenseQuery(() => systemFeaturesQueryOptions()) + +export const userProfileAtom = atom((get): GetAccountProfileResponse => { + const accountProfileQuery = get(accountProfileQueryAtom) as SuspenseQueryResult + + return accountProfileQuery.data?.profile || userProfilePlaceholder +}) + +const profileMetaAtom = atom((get) => { + const accountProfileQuery = get(accountProfileQueryAtom) as SuspenseQueryResult + + return accountProfileQuery.data?.meta ?? { + currentVersion: null, + currentEnv: null, + } +}) + +const currentWorkspaceQueryAtom = atomWithQuery(() => { + return consoleQuery.workspaces.current.post.queryOptions({ + select: normalizeCurrentWorkspace, + }) +}) + +const normalizedCurrentWorkspaceAtom = atom((get) => { + return get(currentWorkspaceQueryAtom).data ?? initialWorkspaceInfo +}) + +export const currentWorkspaceAtom = atom((get) => { + return get(normalizedCurrentWorkspaceAtom) +}) + +export const workspaceRoleFlagsAtom = atom((get) => { + return getWorkspaceRoleFlags(get(currentWorkspaceAtom)) +}) + +const workspacePermissionKeysQueryAtom = atomWithQuery((get) => { + const workspaceId = get(currentWorkspaceAtom).id + + return workspacePermissionKeysQueryOptions(workspaceId) +}) + +export const workspacePermissionKeysAtom = atom((get) => { + return get(workspacePermissionKeysQueryAtom).data?.workspace.permission_keys ?? emptyWorkspacePermissionKeys +}) + +export const workspacePermissionKeysLoadingAtom = atom((get) => { + return get(workspacePermissionKeysQueryAtom).isPending +}) + +export const currentWorkspaceLoadingAtom = atom((get) => { + return get(currentWorkspaceQueryAtom).isPending +}) + +export const currentWorkspaceValidatingAtom = atom((get) => { + return get(currentWorkspaceQueryAtom).isFetching +}) + +const versionQueryAtom = atomWithQuery((get) => { + const meta = get(profileMetaAtom) + const systemFeaturesQuery = get(systemFeaturesQueryAtom) as SuspenseQueryResult + const enabled = Boolean(meta.currentVersion && !systemFeaturesQuery.data?.branding.enabled) + + return langGeniusVersionQueryOptions(meta.currentVersion, enabled) +}) + +export const langGeniusVersionInfoAtom = atom((get) => { + const meta = get(profileMetaAtom) + const versionData = get(versionQueryAtom).data + + if (!versionData) + return initialLangGeniusVersionInfo + + return getLangGeniusVersionInfo({ + meta, + versionData, + }) +}) + +export const refreshUserProfileAtom = atom(null, (get) => { + const queryClient = get(queryClientAtom) + queryClient.invalidateQueries({ queryKey: userProfileQueryOptions().queryKey }) +}) + +export const refreshCurrentWorkspaceAtom = atom(null, (get) => { + const queryClient = get(queryClientAtom) + queryClient.invalidateQueries({ queryKey: consoleQuery.workspaces.current.post.key() }) +}) diff --git a/web/service/access-control/__tests__/use-permission-keys.spec.tsx b/web/service/access-control/__tests__/use-permission-keys.spec.tsx index 1849a48a5ed..428dbfe3ad9 100644 --- a/web/service/access-control/__tests__/use-permission-keys.spec.tsx +++ b/web/service/access-control/__tests__/use-permission-keys.spec.tsx @@ -1,26 +1,13 @@ -import type { ReactNode } from 'react' -import { QueryClient, QueryClientProvider } from '@tanstack/react-query' -import { renderHook, waitFor } from '@testing-library/react' +import { QueryClient } from '@tanstack/react-query' +// eslint-disable-next-line no-restricted-imports import { get } from '@/service/base' -import { useWorkspacePermissionKeys } from '../use-permission-keys' +import { workspacePermissionKeysQueryOptions } from '../use-permission-keys' vi.mock('@/service/base', () => ({ get: vi.fn(), })) -const createWrapper = () => { - const queryClient = new QueryClient({ - defaultOptions: { - queries: { retry: false }, - }, - }) - - return ({ children }: { children: ReactNode }) => ( - {children} - ) -} - -describe('useWorkspacePermissionKeys', () => { +describe('workspacePermissionKeysQueryOptions', () => { beforeEach(() => { vi.clearAllMocks() vi.mocked(get).mockResolvedValue({ @@ -33,11 +20,15 @@ describe('useWorkspacePermissionKeys', () => { // Current-user permissions come from the my-permissions RBAC endpoint. describe('Queries', () => { it('should fetch workspace permission keys', async () => { - renderHook(() => useWorkspacePermissionKeys(), { wrapper: createWrapper() }) - - await waitFor(() => { - expect(get).toHaveBeenCalledWith('/workspaces/current/rbac/my-permissions') + const queryClient = new QueryClient({ + defaultOptions: { + queries: { retry: false }, + }, }) + + await queryClient.fetchQuery(workspacePermissionKeysQueryOptions()) + + expect(get).toHaveBeenCalledWith('/workspaces/current/rbac/my-permissions') }) }) }) diff --git a/web/service/access-control/use-permission-keys.ts b/web/service/access-control/use-permission-keys.ts index af19cf3e840..9953b904ce7 100644 --- a/web/service/access-control/use-permission-keys.ts +++ b/web/service/access-control/use-permission-keys.ts @@ -1,12 +1,18 @@ import type { PermissionKeysResponse } from '@/models/access-control' -import { useQuery } from '@tanstack/react-query' +import { queryOptions } from '@tanstack/react-query' +// eslint-disable-next-line no-restricted-imports import { get } from '../base' const NAME_SPACE = 'workspace-permission-keys' -export const useWorkspacePermissionKeys = () => { - return useQuery({ - queryKey: [NAME_SPACE], +const workspacePermissionKeysQueryKey = (workspaceId?: string) => { + return workspaceId ? [NAME_SPACE, workspaceId] as const : [NAME_SPACE] as const +} + +export const workspacePermissionKeysQueryOptions = (workspaceId?: string) => { + return queryOptions({ + queryKey: workspacePermissionKeysQueryKey(workspaceId), queryFn: () => get('/workspaces/current/rbac/my-permissions'), + enabled: workspaceId === undefined || Boolean(workspaceId), }) } diff --git a/web/service/lang-genius-version.ts b/web/service/lang-genius-version.ts new file mode 100644 index 00000000000..6799fcd1ad4 --- /dev/null +++ b/web/service/lang-genius-version.ts @@ -0,0 +1,13 @@ +import type { LangGeniusVersionResponse } from '@/models/common' +import { queryOptions } from '@tanstack/react-query' +// eslint-disable-next-line no-restricted-imports +import { get } from './base' +import { commonQueryKeys } from './use-common' + +export const langGeniusVersionQueryOptions = (currentVersion?: string | null, enabled?: boolean) => { + return queryOptions({ + queryKey: commonQueryKeys.langGeniusVersion(currentVersion || undefined), + queryFn: () => get('/version', { params: { current_version: currentVersion } }), + enabled: !!currentVersion && (enabled ?? true), + }) +} diff --git a/web/service/use-common.ts b/web/service/use-common.ts index 36644a3ef37..f02d1f403dd 100644 --- a/web/service/use-common.ts +++ b/web/service/use-common.ts @@ -10,13 +10,13 @@ import type { CodeBasedExtension, CommonResponse, FileUploadConfigResponse, - LangGeniusVersionResponse, Member, StructuredOutputRulesRequestBody, StructuredOutputRulesResponse, } from '@/models/common' import type { RETRIEVE_METHOD } from '@/types/app' import { useMutation, useQuery, useQueryClient } from '@tanstack/react-query' +// eslint-disable-next-line no-restricted-imports import { get, post } from './base' const NAME_SPACE = 'common' @@ -54,14 +54,6 @@ export const useFileUploadConfig = () => { }) } -export const useLangGeniusVersion = (currentVersion?: string | null, enabled?: boolean) => { - return useQuery({ - queryKey: commonQueryKeys.langGeniusVersion(currentVersion || undefined), - queryFn: () => get('/version', { params: { current_version: currentVersion } }), - enabled: !!currentVersion && (enabled ?? true), - }) -} - export const useGenerateStructuredOutputRules = () => { return useMutation({ mutationKey: [NAME_SPACE, 'generate-structured-output-rules'], @@ -95,7 +87,7 @@ export const useMailValidity = () => { }) } -export type MailRegisterResponse = { result: string, data: {} } +export type MailRegisterResponse = { result: string, data: Record } export const useMailRegister = () => { return useMutation({ @@ -149,7 +141,7 @@ export const useFilePreview = (fileID: string) => { export type SchemaTypeDefinition = { name: string schema: { - properties: Record + properties: Record } } From 64aa1426815005c5fcdc3180d7b86439b450979a Mon Sep 17 00:00:00 2001 From: chariri Date: Tue, 7 Jul 2026 23:30:28 +0900 Subject: [PATCH 30/70] chore(api): cache the setup status to cut down DB access (#36966) Co-authored-by: Asuka Minato Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> Co-authored-by: Byron.wang --- api/controllers/console/setup.py | 3 +- api/controllers/console/wraps.py | 62 ++++++++++++++++++- .../console/test_fastopenapi_setup.py | 2 + .../controllers/console/test_wraps.py | 40 ++++++++++++ 4 files changed, 103 insertions(+), 4 deletions(-) diff --git a/api/controllers/console/setup.py b/api/controllers/console/setup.py index 3b5c1bbe18f..2b99693a9ca 100644 --- a/api/controllers/console/setup.py +++ b/api/controllers/console/setup.py @@ -13,7 +13,7 @@ from services.account_service import RegisterService, TenantService from .error import AlreadySetupError, NotInitValidateError from .init_validate import get_init_validate_status -from .wraps import only_edition_self_hosted +from .wraps import mark_setup_completed, only_edition_self_hosted class SetupRequestPayload(BaseModel): @@ -96,6 +96,7 @@ def setup_system(payload: SetupRequestPayload) -> SetupResponse: language=payload.language, session=db.session, ) + mark_setup_completed() return SetupResponse(result="success") diff --git a/api/controllers/console/wraps.py b/api/controllers/console/wraps.py index 017793ffe0b..37d7239170c 100644 --- a/api/controllers/console/wraps.py +++ b/api/controllers/console/wraps.py @@ -4,7 +4,7 @@ import os import time from collections.abc import Callable from functools import wraps -from typing import Any, Concatenate, overload +from typing import Any, Concatenate, Protocol, cast, overload from flask import abort, request from pydantic import BaseModel, ValidationError @@ -46,6 +46,60 @@ ERROR_MSG_INVALID_ENCRYPTED_DATA = "Invalid encrypted data" ERROR_MSG_INVALID_ENCRYPTED_CODE = "Invalid encrypted code" +class OnceTrueCallable[**P](Protocol): + def __call__(self, *args: P.args, **kwargs: P.kwargs) -> bool: ... + + def mark_success(self) -> None: ... + + def reset_success(self) -> None: ... + + +def once_true[**P](func: Callable[P, bool]) -> OnceTrueCallable[P]: + """Wrap a predicate so only a strict True result is memoized.""" + has_success = False + + def mark_success() -> None: + nonlocal has_success + + has_success = True + + def reset_success() -> None: + nonlocal has_success + + has_success = False + + @wraps(func) + def wrapper(*args: P.args, **kwargs: P.kwargs) -> bool: + nonlocal has_success + + if has_success: + return True + + result = func(*args, **kwargs) + if result is True: + has_success = True + + return result + + wrapper.mark_success = mark_success # type: ignore[attr-defined] + wrapper.reset_success = reset_success # type: ignore[attr-defined] + return cast(OnceTrueCallable[P], wrapper) + + +def mark_setup_completed() -> None: + """Remember in this process that one-time self-hosted setup has completed.""" + _is_setup_completed.mark_success() + + +@once_true +def _is_setup_completed() -> bool: + """Check whether setup exists, caching only successful observations. + + Use `once_true` instead of `@cache` because a pre-setup False result must not be memoized. + """ + return db.session.scalar(select(DifySetup).limit(1)) is not None + + @overload def account_initialization_required[T, **P, R]( view: Callable[Concatenate[T, P], R], @@ -246,7 +300,9 @@ def setup_required[T, **P, R]( @overload -def setup_required[**P, R](view: Callable[P, R]) -> Callable[P, R]: ... +def setup_required[**P, R](view: Callable[P, R]) -> Callable[P, R]: + """Require self-hosted bootstrap setup before serving protected routes.""" + ... def setup_required[R](view: Callable[..., R]) -> Callable[..., R]: @@ -255,7 +311,7 @@ def setup_required[R](view: Callable[..., R]) -> Callable[..., R]: # The overloads keep Resource methods method-aware for pyrefly while # preserving support for plain functions used in tests and utilities. # check setup - if dify_config.EDITION == "SELF_HOSTED" and not db.session.scalar(select(DifySetup).limit(1)): + if dify_config.EDITION == "SELF_HOSTED" and not _is_setup_completed(): if os.environ.get("INIT_PASSWORD"): raise NotInitValidateError() raise NotSetupError() diff --git a/api/tests/unit_tests/controllers/console/test_fastopenapi_setup.py b/api/tests/unit_tests/controllers/console/test_fastopenapi_setup.py index 385539b6f30..2b385304d32 100644 --- a/api/tests/unit_tests/controllers/console/test_fastopenapi_setup.py +++ b/api/tests/unit_tests/controllers/console/test_fastopenapi_setup.py @@ -48,9 +48,11 @@ def test_console_setup_fastopenapi_post_success(app: Flask): patch("controllers.console.setup.TenantService.get_tenant_count", return_value=0), patch("controllers.console.setup.get_init_validate_status", return_value=True), patch("controllers.console.setup.RegisterService.setup"), + patch("controllers.console.setup.mark_setup_completed") as mark_setup_completed, ): client = app.test_client() response = client.post("/console/api/setup", json=payload) assert response.status_code == 201 assert response.get_json() == {"result": "success"} + mark_setup_completed.assert_called_once_with() diff --git a/api/tests/unit_tests/controllers/console/test_wraps.py b/api/tests/unit_tests/controllers/console/test_wraps.py index 618a1f52180..172125f4635 100644 --- a/api/tests/unit_tests/controllers/console/test_wraps.py +++ b/api/tests/unit_tests/controllers/console/test_wraps.py @@ -13,6 +13,7 @@ from controllers.console.workspace.error import AccountNotInitializedError from controllers.console.wraps import ( RBACPermission, RBACResourceScope, + _is_setup_completed, account_initialization_required, cloud_edition_billing_enabled, cloud_edition_billing_rate_limit_check, @@ -35,6 +36,12 @@ from models.account import AccountStatus, TenantAccountRole from services.feature_service import LicenseStatus +@pytest.fixture(autouse=True) +def reset_setup_required_cache(): + """Keep setup_required's process cache isolated across unit tests.""" + _is_setup_completed.reset_success() + + class MockUser(UserMixin): """Simple User class for testing.""" @@ -735,6 +742,39 @@ class TestSystemSetup: # Assert assert result == "admin_success" + @patch("controllers.console.wraps.db") + def test_should_cache_completed_setup(self, mock_db): + """Test that completed setup skips repeated DB reads in this process""" + mock_db.session.scalar.return_value = MagicMock() + + @setup_required + def admin_view(): + return "admin_success" + + with patch("controllers.console.wraps.dify_config.EDITION", "SELF_HOSTED"): + assert admin_view() == "admin_success" + assert admin_view() == "admin_success" + + assert mock_db.session.scalar.call_count == 1 + + @patch("controllers.console.wraps.db") + @patch("controllers.console.wraps.os.environ.get") + def test_should_not_cache_missing_setup(self, mock_environ_get, mock_db): + """Test that first-time bootstrap completion can be observed later in the same process""" + mock_db.session.scalar.side_effect = [None, MagicMock()] + mock_environ_get.return_value = None + + @setup_required + def admin_view(): + return "admin_success" + + with patch("controllers.console.wraps.dify_config.EDITION", "SELF_HOSTED"): + with pytest.raises(NotSetupError): + admin_view() + assert admin_view() == "admin_success" + + assert mock_db.session.scalar.call_count == 2 + @patch("controllers.console.wraps.db") @patch("controllers.console.wraps.os.environ.get") def test_should_raise_not_init_validate_error_with_init_password(self, mock_environ_get, mock_db: MagicMock): From 09c5c5e5ed3f5764f7293586b08aaf3ccbef6b46 Mon Sep 17 00:00:00 2001 From: "ojasarora.eth" Date: Tue, 7 Jul 2026 15:36:17 +0100 Subject: [PATCH 31/70] refactor(test): replace SimpleNamespace with typed mocks in schedule service tests (#38393) --- .../services/test_schedule_service.py | 23 +++++++++++++++---- 1 file changed, 19 insertions(+), 4 deletions(-) diff --git a/api/tests/unit_tests/services/test_schedule_service.py b/api/tests/unit_tests/services/test_schedule_service.py index 0f8f7ffab58..d5006bd5f26 100644 --- a/api/tests/unit_tests/services/test_schedule_service.py +++ b/api/tests/unit_tests/services/test_schedule_service.py @@ -1,7 +1,7 @@ +import json import unittest from datetime import UTC, datetime -from types import SimpleNamespace -from typing import Any, cast +from typing import Any from unittest.mock import MagicMock, Mock import pytest @@ -11,7 +11,7 @@ from core.trigger.constants import TRIGGER_SCHEDULE_NODE_TYPE from core.workflow.nodes.trigger_schedule.entities import VisualConfig from core.workflow.nodes.trigger_schedule.exc import ScheduleConfigError from libs.schedule_utils import calculate_next_run_at, convert_12h_to_24h -from models.workflow import Workflow +from models.workflow import Workflow, WorkflowType from services.trigger.schedule_service import ScheduleService @@ -503,7 +503,22 @@ def session_mock() -> MagicMock: def _workflow(**kwargs: Any) -> Workflow: - return cast(Workflow, SimpleNamespace(**kwargs)) + graph_dict = kwargs.pop("graph_dict", {}) + workflow = Workflow.new( + tenant_id="tenant-1", + app_id="app-1", + type=WorkflowType.WORKFLOW, + version="draft", + graph=json.dumps(graph_dict), + features="{}", + created_by="account-1", + environment_variables=[], + conversation_variables=[], + rag_pipeline_variables=[], + ) + for key, value in kwargs.items(): + setattr(workflow, key, value) + return workflow def test_to_schedule_config_should_build_from_cron_mode() -> None: From 5308b95aff8703524ebc23f9397863c527a2bc39 Mon Sep 17 00:00:00 2001 From: Stephen Zhou Date: Tue, 7 Jul 2026 23:10:58 +0800 Subject: [PATCH 32/70] refactor(web): reduce query atom subscriptions (#38521) --- .../skills/how-to-write-component/SKILL.md | 1 + .../deployments/create-guide/state/queries.ts | 41 ++++++++++++------- .../deployments/create-guide/state/source.ts | 21 +++++++--- .../create-guide/state/submission.ts | 4 +- .../deployments/create-guide/state/target.ts | 18 ++++---- .../ui/__tests__/source-step.spec.tsx | 22 ++++++++-- .../create-guide/ui/source-step.tsx | 34 +++++++++++---- .../create-guide/ui/target-step.tsx | 36 ++++++++++------ .../deployments/create-release/state/index.ts | 13 ++++++ .../ui/__tests__/source-app-picker.spec.tsx | 23 ++++++++++- .../create-release/ui/source-app-picker.tsx | 41 ++++++++++++------- .../deployments/deploy-drawer/state/index.ts | 28 ++++++++----- .../deployments/deploy-drawer/ui/form.tsx | 11 ++--- .../channels/__tests__/section.spec.tsx | 22 ++++++---- .../detail/access/channels/section.tsx | 18 ++++---- .../__tests__/permissions.spec.tsx | 18 ++++---- .../detail/access/permissions/section.tsx | 14 ++++--- .../deployments/detail/access/state.ts | 5 +++ .../api-token-management/section.tsx | 22 ++++++---- .../developer-api-header-switch.tsx | 16 +++++--- .../deployments/detail/api-tokens/state.ts | 5 +++ .../deployments/detail/deployment-sidebar.tsx | 16 +++++--- .../header-actions/new-deployment-button.tsx | 8 ++-- .../deployments/detail/instances/index.tsx | 8 ++-- .../deployments/detail/overview/index.tsx | 15 ++++--- .../deployments/detail/overview/state.ts | 5 +++ .../__tests__/deploy-release-menu.spec.tsx | 24 ++--------- .../release-actions/deploy-release-menu.tsx | 24 ++++++----- .../detail/releases/release-actions/state.ts | 20 ++++++++- .../release-history/release-history-table.tsx | 11 ++--- .../deployments/detail/releases/state.ts | 5 +++ web/features/deployments/detail/state.ts | 11 ++++- web/features/deployments/list/state/index.ts | 28 +++++++++---- web/features/deployments/list/ui/shell.tsx | 25 +++++++++-- 34 files changed, 420 insertions(+), 193 deletions(-) diff --git a/.agents/skills/how-to-write-component/SKILL.md b/.agents/skills/how-to-write-component/SKILL.md index e66eec40c88..9d013876dd9 100644 --- a/.agents/skills/how-to-write-component/SKILL.md +++ b/.agents/skills/how-to-write-component/SKILL.md @@ -54,6 +54,7 @@ Use this as the component decision guide for Dify web. Existing code is referenc - Treat `useParams`, route args, and `nuqs` query state as framework-owned state. When atom logic needs those values, hydrate primitive atoms at the route or surface boundary, such as with `useHydrateAtoms(..., { dangerouslyForceHydrate: true })`; keep URL updates in the route/query-state APIs instead of write atoms. - Within a route-owned feature, choose one source for route identity. If route params are bridged into feature atoms, use that bridge consistently for route-derived queries and actions instead of also threading the same route id through page, tab, and section props. - For async work tied to atom state, use `atomWithQuery` or `atomWithMutation`; write atoms should update only the inputs that drive those atoms. This applies to pure frontend async work as well as network requests, so do not hand-roll loading/error/in-flight state with `useState` or `useRef` for atom-orchestrated async behavior. For component-owned remote work, use `useQuery` or `useMutation` directly. +- `jotai-tanstack-query` query atoms do not support TanStack Query tracked properties. A component that reads `useAtomValue(queryAtom)` subscribes to the whole query result, even if it only accesses `data`, `isLoading`, or `isError`. Export field-specific derived atoms and have components read the exact fields they render; use `selectAtom(queryAtom, result => result.field)` for query-result fields so unchanged selections do not notify subscribers. Keep direct `useAtomValue(queryAtom)` only when the component or hook genuinely needs the full observer result. - Row-local async state belongs to the row owner unless it participates in a shared Jotai workflow or needs atom-scoped reset semantics. - Leave query and mutation atoms unscoped so they keep shared QueryClient cache and invalidation behavior. Scope resettable primitives and explicit hydration tuples; scope a derived atom only when every dependency should be private to that surface. - For scoped primitives that are always hydrated by `ScopeProvider`, prefer `atomWithLazy(() => { throw new Error(...) })` when consumers should see a non-null type. diff --git a/web/features/deployments/create-guide/state/queries.ts b/web/features/deployments/create-guide/state/queries.ts index 22541b25b74..b9c099b9efc 100644 --- a/web/features/deployments/create-guide/state/queries.ts +++ b/web/features/deployments/create-guide/state/queries.ts @@ -4,6 +4,7 @@ import type { Getter } from 'jotai/vanilla' import { keepPreviousData, skipToken } from '@tanstack/react-query' import { atom } from 'jotai' import { atomWithInfiniteQuery, atomWithQuery } from 'jotai-tanstack-query' +import { selectAtom } from 'jotai/utils' import { encodeDslContent } from '@/features/deployments/shared/domain/dsl' import { consoleQuery } from '@/service/client' import { effectiveMethodAtom, instanceNameAtom, submissionUnsupportedDslNodesAtom } from './primitives' @@ -53,6 +54,11 @@ export const deployableEnvironmentsQueryAtom = atomWithQuery((get) => { }) }) +export const deployableEnvironmentsDataAtom = selectAtom(deployableEnvironmentsQueryAtom, query => query.data) +export const deployableEnvironmentsIsErrorAtom = selectAtom(deployableEnvironmentsQueryAtom, query => query.isError) +export const deployableEnvironmentsIsLoadingAtom = selectAtom(deployableEnvironmentsQueryAtom, query => query.isLoading) +export const deployableEnvironmentsIsFetchingAtom = selectAtom(deployableEnvironmentsQueryAtom, query => query.isFetching) + const precheckReleaseQueryAtom = atomWithQuery((get) => { const method = get(effectiveMethodAtom) const effectiveSelectedApp = get(effectiveSelectedAppAtom) @@ -88,17 +94,22 @@ const precheckReleaseQueryAtom = atomWithQuery((get) => { return precheckReleaseQueryOptions }) +const precheckReleaseDataAtom = selectAtom(precheckReleaseQueryAtom, query => query.data) +const precheckReleaseIsSuccessAtom = selectAtom(precheckReleaseQueryAtom, query => query.isSuccess) +const precheckReleaseIsLoadingAtom = selectAtom(precheckReleaseQueryAtom, query => query.isLoading) +const precheckReleaseIsFetchingAtom = selectAtom(precheckReleaseQueryAtom, query => query.isFetching) + function precheckReleaseReady(get: Getter) { - const precheckReleaseQuery = get(precheckReleaseQueryAtom) + const precheckRelease = get(precheckReleaseDataAtom) return sourceReady(get) - && precheckReleaseQuery.isSuccess - && Boolean(precheckReleaseQuery.data?.canCreate) - && (precheckReleaseQuery.data?.unsupportedNodes.length ?? 0) === 0 + && get(precheckReleaseIsSuccessAtom) + && Boolean(precheckRelease?.canCreate) + && (precheckRelease?.unsupportedNodes.length ?? 0) === 0 && get(submissionUnsupportedDslNodesAtom).length === 0 } -export const deploymentOptionsQueryAtom = atomWithQuery((get) => { +const deploymentOptionsQueryAtom = atomWithQuery((get) => { const method = get(effectiveMethodAtom) const effectiveSelectedApp = get(effectiveSelectedAppAtom) const dslContent = get(dslContentAtom) @@ -133,6 +144,12 @@ export const deploymentOptionsQueryAtom = atomWithQuery((get) => { return deploymentOptionsQueryOptions }) +export const deploymentOptionsDataAtom = selectAtom(deploymentOptionsQueryAtom, query => query.data) +export const deploymentOptionsIsErrorAtom = selectAtom(deploymentOptionsQueryAtom, query => query.isError) +export const deploymentOptionsIsLoadingAtom = selectAtom(deploymentOptionsQueryAtom, query => query.isLoading) +export const deploymentOptionsIsFetchingAtom = selectAtom(deploymentOptionsQueryAtom, query => query.isFetching) +const deploymentOptionsIsSuccessAtom = selectAtom(deploymentOptionsQueryAtom, query => query.isSuccess) + export const unsupportedDslNodesAtom = atom((get) => { const submissionUnsupportedDslNodes = get(submissionUnsupportedDslNodesAtom) if (submissionUnsupportedDslNodes.length > 0) @@ -141,7 +158,7 @@ export const unsupportedDslNodesAtom = atom((get) => { if (!sourceReady(get)) return [] - return get(precheckReleaseQueryAtom).data?.unsupportedNodes ?? [] + return get(precheckReleaseDataAtom)?.unsupportedNodes ?? [] }) const precheckReleaseReadyAtom = atom((get) => { @@ -149,21 +166,17 @@ const precheckReleaseReadyAtom = atom((get) => { }) export const deploymentOptionsReadyAtom = atom((get) => { - const deploymentOptionsQuery = get(deploymentOptionsQueryAtom) - return sourceReady(get) && get(precheckReleaseReadyAtom) - && deploymentOptionsQuery.isSuccess + && get(deploymentOptionsIsSuccessAtom) }) export const deploymentOptionsContentCheckedAtom = atom((get) => { - const deploymentOptionsQuery = get(deploymentOptionsQueryAtom) - const precheckReleaseQuery = get(precheckReleaseQueryAtom) - const isLoadingOptions = deploymentOptionsQuery.isLoading || (deploymentOptionsQuery.isFetching && !deploymentOptionsQuery.data) - const isCheckingReleaseContent = precheckReleaseQuery.isLoading || (precheckReleaseQuery.isFetching && !precheckReleaseQuery.data) + const isLoadingOptions = get(deploymentOptionsIsLoadingAtom) || (get(deploymentOptionsIsFetchingAtom) && !get(deploymentOptionsDataAtom)) + const isCheckingReleaseContent = get(precheckReleaseIsLoadingAtom) || (get(precheckReleaseIsFetchingAtom) && !get(precheckReleaseDataAtom)) if (!sourceReady(get) || isCheckingReleaseContent || isLoadingOptions) return false - return get(precheckReleaseReadyAtom) && deploymentOptionsQuery.isSuccess + return get(precheckReleaseReadyAtom) && get(deploymentOptionsIsSuccessAtom) }) diff --git a/web/features/deployments/create-guide/state/source.ts b/web/features/deployments/create-guide/state/source.ts index 39bb834b8f1..274fcd3a401 100644 --- a/web/features/deployments/create-guide/state/source.ts +++ b/web/features/deployments/create-guide/state/source.ts @@ -5,6 +5,7 @@ import type { WorkflowSourceApp } from './types' import { keepPreviousData, queryOptions } from '@tanstack/react-query' import { atom } from 'jotai' import { atomWithInfiniteQuery, atomWithQuery } from 'jotai-tanstack-query' +import { selectAtom } from 'jotai/utils' import { dslAppName, isWorkflowDsl } from '@/features/deployments/shared/domain/dsl' import { consoleQuery } from '@/service/client' import { normalizeAppPagination } from '@/service/use-apps' @@ -92,18 +93,28 @@ export const sourceAppsQueryAtom = atomWithInfiniteQuery((get) => { }) }) +const sourceAppsDataAtom = selectAtom(sourceAppsQueryAtom, query => query.data) +export const sourceAppsErrorAtom = selectAtom(sourceAppsQueryAtom, query => query.error) +export const sourceAppsFetchNextPageAtom = selectAtom(sourceAppsQueryAtom, query => query.fetchNextPage) +export const sourceAppsHasNextPageAtom = selectAtom(sourceAppsQueryAtom, query => query.hasNextPage) +export const sourceAppsIsFetchingAtom = selectAtom(sourceAppsQueryAtom, query => query.isFetching) +export const sourceAppsIsFetchingNextPageAtom = selectAtom(sourceAppsQueryAtom, query => query.isFetchingNextPage) +export const sourceAppsIsLoadingAtom = selectAtom(sourceAppsQueryAtom, query => query.isLoading) +export const sourceAppsIsPlaceholderDataAtom = selectAtom(sourceAppsQueryAtom, query => query.isPlaceholderData) + +export const sourceAppsAtom = atom((get) => { + return (get(sourceAppsDataAtom)?.pages.flatMap(page => page.data) ?? []) as WorkflowSourceApp[] +}) + export const effectiveSelectedAppAtom = atom((get) => { const selectedApp = get(selectedAppAtom) if (selectedApp) return selectedApp - const sourceAppsQuery = get(sourceAppsQueryAtom) - if (sourceAppsQuery.isPlaceholderData) + if (get(sourceAppsIsPlaceholderDataAtom)) return undefined - const sourceApps = (sourceAppsQuery.data?.pages.flatMap(page => page.data) ?? []) as WorkflowSourceApp[] - - return sourceApps[0] + return get(sourceAppsAtom)[0] }) export function sourceReady(get: Getter) { diff --git a/web/features/deployments/create-guide/state/submission.ts b/web/features/deployments/create-guide/state/submission.ts index 675598e6307..9f13688c0b4 100644 --- a/web/features/deployments/create-guide/state/submission.ts +++ b/web/features/deployments/create-guide/state/submission.ts @@ -22,7 +22,7 @@ import { selectedEnvironmentIdAtom, submissionUnsupportedDslNodesAtom, } from './primitives' -import { deployableEnvironmentsQueryAtom, deploymentOptionsQueryAtom } from './queries' +import { deployableEnvironmentsQueryAtom, deploymentOptionsDataAtom } from './queries' import { submittedReleaseReadyAtom } from './release' import { dslContentAtom, effectiveSelectedAppAtom } from './source' import { @@ -75,7 +75,7 @@ export const createDeploymentGuideSubmissionAtom = atom(null, async (get, set, { const effectiveSelectedApp = get(effectiveSelectedAppAtom) const deployableEnvironmentsQuery = get(deployableEnvironmentsQueryAtom) - const deploymentOptions = get(deploymentOptionsQueryAtom).data?.options + const deploymentOptions = get(deploymentOptionsDataAtom)?.options const envVarSlots = get(deploymentTargetEnvVarSlotsAtom) const envVarValues = get(envVarValuesAtom) const bindingSlots = get(deploymentTargetBindingSlotsAtom) diff --git a/web/features/deployments/create-guide/state/target.ts b/web/features/deployments/create-guide/state/target.ts index 36b3fee69d9..57f75389f2d 100644 --- a/web/features/deployments/create-guide/state/target.ts +++ b/web/features/deployments/create-guide/state/target.ts @@ -11,23 +11,21 @@ import { import { dslEnvVarSlots } from '@/features/deployments/shared/domain/dsl' import { environmentMatchesIdentifier } from './environment' import { effectiveMethodAtom, envVarValuesAtom, manualBindingSelectionsAtom, selectedEnvironmentIdAtom } from './primitives' -import { deployableEnvironmentsQueryAtom, deploymentOptionsQueryAtom, deploymentOptionsReadyAtom } from './queries' +import { deployableEnvironmentsDataAtom, deploymentOptionsDataAtom, deploymentOptionsReadyAtom } from './queries' import { submittedReleaseReadyAtom } from './release' import { dslContentAtom, sourceReady } from './source' import { envVarSelectionReady } from './utils' export const deployableEnvironmentsAtom = atom((get) => { - const deployableEnvironmentsQuery = get(deployableEnvironmentsQueryAtom) + const deployableEnvironments = get(deployableEnvironmentsDataAtom) return sourceReady(get) - ? deployableEnvironmentsQuery.data?.environments ?? [] + ? deployableEnvironments?.environments ?? [] : [] }) const deployableEnvironmentsReadyAtom = atom((get) => { - const deployableEnvironmentsQuery = get(deployableEnvironmentsQueryAtom) - - return sourceReady(get) && deployableEnvironmentsQuery.isSuccess + return sourceReady(get) && Boolean(get(deployableEnvironmentsDataAtom)) }) export const effectiveSelectedEnvironmentIdAtom = atom((get) => { @@ -35,10 +33,10 @@ export const effectiveSelectedEnvironmentIdAtom = atom((get) => { }) export const deploymentTargetBindingSlotsAtom = atom((get) => { - const deploymentOptionsQuery = get(deploymentOptionsQueryAtom) + const deploymentOptions = get(deploymentOptionsDataAtom) return sourceReady(get) - ? deploymentOptionsQuery.data?.options?.credentialSlots?.filter(slot => runtimeCredentialSlotKey(slot)) ?? [] + ? deploymentOptions?.options?.credentialSlots?.filter(slot => runtimeCredentialSlotKey(slot)) ?? [] : [] }) @@ -59,8 +57,8 @@ export const requiredBindingsReadyAtom = atom((get) => { export const deploymentTargetEnvVarSlotsAtom = atom((get) => { const method = get(effectiveMethodAtom) - const deploymentOptionsQuery = get(deploymentOptionsQueryAtom) - const slots = sourceReady(get) ? deploymentOptionsQuery.data?.options?.envVarSlots : undefined + const deploymentOptions = get(deploymentOptionsDataAtom) + const slots = sourceReady(get) ? deploymentOptions?.options?.envVarSlots : undefined const dslContent = get(dslContentAtom) // Deployment options own the canonical slot list; DSL metadata only enriches import-DSL defaults. diff --git a/web/features/deployments/create-guide/ui/__tests__/source-step.spec.tsx b/web/features/deployments/create-guide/ui/__tests__/source-step.spec.tsx index 46fea37411f..cb620975f8d 100644 --- a/web/features/deployments/create-guide/ui/__tests__/source-step.spec.tsx +++ b/web/features/deployments/create-guide/ui/__tests__/source-step.spec.tsx @@ -5,12 +5,13 @@ import { SourceStepContent } from '../source-step' const mocks = vi.hoisted(() => { const sourceAppsQuery = { data: { pages: [{ data: [] }] }, + error: null, + fetchNextPage: vi.fn(), hasNextPage: false, isFetching: false, isFetchingNextPage: false, isLoading: false, isPlaceholderData: false, - fetchNextPage: vi.fn(), } return { @@ -46,6 +47,14 @@ vi.mock('@/features/deployments/create-guide/state/source', async () => { dslUnsupportedModeAtom: atom(false), effectiveSelectedAppAtom: atom(undefined), isReadingDslAtom: atom(false), + sourceAppsAtom: atom(() => mocks.sourceAppsQuery.data.pages.flatMap(page => page.data)), + sourceAppsErrorAtom: atom(() => mocks.sourceAppsQuery.error), + sourceAppsFetchNextPageAtom: atom(() => mocks.sourceAppsQuery.fetchNextPage), + sourceAppsHasNextPageAtom: atom(() => mocks.sourceAppsQuery.hasNextPage), + sourceAppsIsFetchingAtom: atom(() => mocks.sourceAppsQuery.isFetching), + sourceAppsIsFetchingNextPageAtom: atom(() => mocks.sourceAppsQuery.isFetchingNextPage), + sourceAppsIsLoadingAtom: atom(() => mocks.sourceAppsQuery.isLoading), + sourceAppsIsPlaceholderDataAtom: atom(() => mocks.sourceAppsQuery.isPlaceholderData), sourceAppsQueryAtom: atom(mocks.sourceAppsQuery), } }) @@ -80,12 +89,13 @@ describe('SourceStepContent', () => { vi.clearAllMocks() Object.assign(mocks.sourceAppsQuery, { data: { pages: [{ data: [] }] }, + error: null, + fetchNextPage: vi.fn(), hasNextPage: false, isFetching: false, isFetchingNextPage: false, isLoading: false, isPlaceholderData: false, - fetchNextPage: vi.fn(), }) }) @@ -114,7 +124,13 @@ describe('SourceStepContent', () => { render() expect(mocks.useInfiniteScroll).toHaveBeenCalledWith( - mocks.sourceAppsQuery, + expect.objectContaining({ + fetchNextPage: expect.any(Function), + hasNextPage: expect.any(Boolean), + isFetching: false, + isFetchingNextPage: false, + isLoading: false, + }), expect.objectContaining({ rootMargin: '0px 0px 160px 0px', threshold: 0.1, diff --git a/web/features/deployments/create-guide/ui/source-step.tsx b/web/features/deployments/create-guide/ui/source-step.tsx index 56debc4530a..50c3b938c16 100644 --- a/web/features/deployments/create-guide/ui/source-step.tsx +++ b/web/features/deployments/create-guide/ui/source-step.tsx @@ -21,7 +21,14 @@ import { dslUnsupportedModeAtom, effectiveSelectedAppAtom, isReadingDslAtom, - sourceAppsQueryAtom, + sourceAppsAtom, + sourceAppsErrorAtom, + sourceAppsFetchNextPageAtom, + sourceAppsHasNextPageAtom, + sourceAppsIsFetchingAtom, + sourceAppsIsFetchingNextPageAtom, + sourceAppsIsLoadingAtom, + sourceAppsIsPlaceholderDataAtom, } from '@/features/deployments/create-guide/state/source' import { continueFromSourceAtom, @@ -189,10 +196,23 @@ function SourceAppList() { const { t } = useTranslation('deployments') const selectSourceApp = useSetAtom(selectSourceAppAtom) const effectiveSelectedApp = useAtomValue(effectiveSelectedAppAtom) - const sourceAppsQuery = useAtomValue(sourceAppsQueryAtom) - const sourceApps = (sourceAppsQuery.data?.pages.flatMap(page => page.data) ?? []) as WorkflowSourceApp[] - const sourceAppsLoading = sourceAppsQuery.isLoading || sourceAppsQuery.isPlaceholderData || (sourceAppsQuery.isFetching && sourceApps.length === 0) - const { rootRef, sentinelRef } = useInfiniteScroll(sourceAppsQuery, { + const sourceApps = useAtomValue(sourceAppsAtom) + const sourceAppsError = useAtomValue(sourceAppsErrorAtom) + const sourceAppsFetchNextPage = useAtomValue(sourceAppsFetchNextPageAtom) + const sourceAppsHasNextPage = useAtomValue(sourceAppsHasNextPageAtom) + const sourceAppsIsFetching = useAtomValue(sourceAppsIsFetchingAtom) + const sourceAppsIsFetchingNextPage = useAtomValue(sourceAppsIsFetchingNextPageAtom) + const sourceAppsIsLoading = useAtomValue(sourceAppsIsLoadingAtom) + const sourceAppsIsPlaceholderData = useAtomValue(sourceAppsIsPlaceholderDataAtom) + const sourceAppsLoading = sourceAppsIsLoading || sourceAppsIsPlaceholderData || (sourceAppsIsFetching && sourceApps.length === 0) + const { rootRef, sentinelRef } = useInfiniteScroll({ + error: sourceAppsError, + fetchNextPage: sourceAppsFetchNextPage, + hasNextPage: sourceAppsHasNextPage, + isFetching: sourceAppsIsFetching, + isFetchingNextPage: sourceAppsIsFetchingNextPage, + isLoading: sourceAppsIsLoading, + }, { enabled: !sourceAppsLoading, rootMargin: '0px 0px 160px 0px', threshold: 0.1, @@ -218,12 +238,12 @@ function SourceAppList() { onSelect={() => selectSourceApp(app)} /> ))} - {sourceAppsQuery.isFetchingNextPage && ( + {sourceAppsIsFetchingNextPage && (
    {t('createModal.loadingApps')}
    )} - {sourceAppsQuery.hasNextPage && diff --git a/web/features/deployments/create-guide/ui/target-step.tsx b/web/features/deployments/create-guide/ui/target-step.tsx index e915aa146b9..f543a5ee969 100644 --- a/web/features/deployments/create-guide/ui/target-step.tsx +++ b/web/features/deployments/create-guide/ui/target-step.tsx @@ -16,8 +16,13 @@ import { stepAtom, } from '@/features/deployments/create-guide/state/primitives' import { - deployableEnvironmentsQueryAtom, - deploymentOptionsQueryAtom, + deployableEnvironmentsIsErrorAtom, + deployableEnvironmentsIsFetchingAtom, + deployableEnvironmentsIsLoadingAtom, + deploymentOptionsDataAtom, + deploymentOptionsIsErrorAtom, + deploymentOptionsIsFetchingAtom, + deploymentOptionsIsLoadingAtom, unsupportedDslNodesAtom, } from '@/features/deployments/create-guide/state/queries' import { @@ -72,11 +77,12 @@ export function TargetStepContent() { function TargetEnvironmentSection() { const { t } = useTranslation('deployments') - const environmentsQuery = useAtomValue(deployableEnvironmentsQueryAtom) + const environmentsIsError = useAtomValue(deployableEnvironmentsIsErrorAtom) + const environmentsIsFetching = useAtomValue(deployableEnvironmentsIsFetchingAtom) + const environmentsIsLoading = useAtomValue(deployableEnvironmentsIsLoadingAtom) const environments = useAtomValue(deployableEnvironmentsAtom) const effectiveSelectedEnvironmentId = useAtomValue(effectiveSelectedEnvironmentIdAtom) - const isEnvironmentError = environmentsQuery.isError - const isEnvironmentLoading = environmentsQuery.isLoading || (environmentsQuery.isFetching && !environmentsQuery.data) + const isEnvironmentLoading = environmentsIsLoading || (environmentsIsFetching && environments.length === 0) const selectEnvironment = useSetAtom(selectedEnvironmentIdAtom) const hasEnvironmentOptions = environments.length > 0 @@ -102,7 +108,7 @@ function TargetEnvironmentSection() { ? : (
    - {isEnvironmentError + {environmentsIsError ? t('createGuide.target.loadEnvironmentsFailed') : t('createGuide.target.noEnvironmentOptions')}
    @@ -144,11 +150,14 @@ function EnvironmentOptionRow({ environment }: { function TargetBindingSection() { const { t } = useTranslation('deployments') - const deploymentOptionsQuery = useAtomValue(deploymentOptionsQueryAtom) + const deploymentOptions = useAtomValue(deploymentOptionsDataAtom) + const deploymentOptionsIsError = useAtomValue(deploymentOptionsIsErrorAtom) + const deploymentOptionsIsFetching = useAtomValue(deploymentOptionsIsFetchingAtom) + const deploymentOptionsIsLoading = useAtomValue(deploymentOptionsIsLoadingAtom) const bindingSlots = useAtomValue(deploymentTargetBindingSlotsAtom) const bindingSelections = useAtomValue(deploymentTargetBindingSelectionsAtom) - const isBindingError = deploymentOptionsQuery.isError - const isBindingLoading = deploymentOptionsQuery.isLoading || (deploymentOptionsQuery.isFetching && !deploymentOptionsQuery.data) + const isBindingError = deploymentOptionsIsError + const isBindingLoading = deploymentOptionsIsLoading || (deploymentOptionsIsFetching && !deploymentOptions) const selectBinding = useSetAtom(selectBindingAtom) const unsupportedDslNodes = useAtomValue(unsupportedDslNodesAtom) const shouldRender = !(isBindingError && unsupportedDslNodes.length > 0) @@ -196,10 +205,13 @@ function TargetEnvVarSection() { const { t } = useTranslation('deployments') const setEnvVar = useSetAtom(setEnvVarAtom) const envVarValues = useAtomValue(envVarValuesAtom) - const deploymentOptionsQuery = useAtomValue(deploymentOptionsQueryAtom) + const deploymentOptions = useAtomValue(deploymentOptionsDataAtom) + const deploymentOptionsIsError = useAtomValue(deploymentOptionsIsErrorAtom) + const deploymentOptionsIsFetching = useAtomValue(deploymentOptionsIsFetchingAtom) + const deploymentOptionsIsLoading = useAtomValue(deploymentOptionsIsLoadingAtom) const envVarSlots = useAtomValue(deploymentTargetEnvVarSlotsAtom) - const isBindingError = deploymentOptionsQuery.isError - const isBindingLoading = deploymentOptionsQuery.isLoading || (deploymentOptionsQuery.isFetching && !deploymentOptionsQuery.data) + const isBindingError = deploymentOptionsIsError + const isBindingLoading = deploymentOptionsIsLoading || (deploymentOptionsIsFetching && !deploymentOptions) if (isBindingLoading || isBindingError) return null diff --git a/web/features/deployments/create-release/state/index.ts b/web/features/deployments/create-release/state/index.ts index b47ea3243be..67a670b351a 100644 --- a/web/features/deployments/create-release/state/index.ts +++ b/web/features/deployments/create-release/state/index.ts @@ -20,6 +20,7 @@ import { atomWithQuery, queryClientAtom, } from 'jotai-tanstack-query' +import { selectAtom } from 'jotai/utils' import * as z from 'zod' import { consoleQuery } from '@/service/client' import { normalizeAppPagination } from '@/service/use-apps' @@ -270,6 +271,18 @@ export const createReleaseSourceAppsQueryAtom = atomWithInfiniteQuery((get) => { }) }) +const createReleaseSourceAppsDataAtom = selectAtom(createReleaseSourceAppsQueryAtom, query => query.data) +export const createReleaseSourceAppsErrorAtom = selectAtom(createReleaseSourceAppsQueryAtom, query => query.error) +export const createReleaseSourceAppsFetchNextPageAtom = selectAtom(createReleaseSourceAppsQueryAtom, query => query.fetchNextPage) +export const createReleaseSourceAppsHasNextPageAtom = selectAtom(createReleaseSourceAppsQueryAtom, query => query.hasNextPage) +export const createReleaseSourceAppsIsFetchingAtom = selectAtom(createReleaseSourceAppsQueryAtom, query => query.isFetching) +export const createReleaseSourceAppsIsFetchingNextPageAtom = selectAtom(createReleaseSourceAppsQueryAtom, query => query.isFetchingNextPage) +export const createReleaseSourceAppsIsLoadingAtom = selectAtom(createReleaseSourceAppsQueryAtom, query => query.isLoading) + +export const createReleaseSourceAppsAtom = atom((get) => { + return get(createReleaseSourceAppsDataAtom)?.pages.flatMap(page => page.data) ?? [] +}) + export const createReleaseDslContentAtom = atom((get) => { return get(createReleaseDslFileContentQueryAtom).data ?? '' }) diff --git a/web/features/deployments/create-release/ui/__tests__/source-app-picker.spec.tsx b/web/features/deployments/create-release/ui/__tests__/source-app-picker.spec.tsx index a33f2664042..d7bbc829df1 100644 --- a/web/features/deployments/create-release/ui/__tests__/source-app-picker.spec.tsx +++ b/web/features/deployments/create-release/ui/__tests__/source-app-picker.spec.tsx @@ -35,6 +35,13 @@ vi.mock('@/features/deployments/create-release/state', async () => { const { atom } = await import('jotai') return { + createReleaseSourceAppsAtom: atom(() => mocks.sourceAppsQuery.data.pages.flatMap(page => page.data)), + createReleaseSourceAppsErrorAtom: atom(() => mocks.sourceAppsQuery.error), + createReleaseSourceAppsFetchNextPageAtom: atom(() => mocks.sourceAppsQuery.fetchNextPage), + createReleaseSourceAppsHasNextPageAtom: atom(() => mocks.sourceAppsQuery.hasNextPage), + createReleaseSourceAppsIsFetchingAtom: atom(() => mocks.sourceAppsQuery.isFetching), + createReleaseSourceAppsIsFetchingNextPageAtom: atom(() => mocks.sourceAppsQuery.isFetchingNextPage), + createReleaseSourceAppsIsLoadingAtom: atom(() => mocks.sourceAppsQuery.isLoading), createReleaseSourceAppSearchTextAtom: atom(''), createReleaseSourceAppsQueryAtom: atom(mocks.sourceAppsQuery), } @@ -98,7 +105,13 @@ describe('SourceAppPicker', () => { renderSourceAppPicker(false) expect(mocks.useInfiniteScroll).toHaveBeenCalledWith( - mocks.sourceAppsQuery, + expect.objectContaining({ + fetchNextPage: expect.any(Function), + hasNextPage: true, + isFetching: false, + isFetchingNextPage: false, + isLoading: false, + }), expect.objectContaining({ enabled: false, rootMargin: '0px 0px 160px 0px', @@ -110,7 +123,13 @@ describe('SourceAppPicker', () => { await waitFor(() => { expect(mocks.useInfiniteScroll).toHaveBeenLastCalledWith( - mocks.sourceAppsQuery, + expect.objectContaining({ + fetchNextPage: expect.any(Function), + hasNextPage: true, + isFetching: false, + isFetchingNextPage: false, + isLoading: false, + }), expect.objectContaining({ enabled: true, rootMargin: '0px 0px 160px 0px', diff --git a/web/features/deployments/create-release/ui/source-app-picker.tsx b/web/features/deployments/create-release/ui/source-app-picker.tsx index 414d12a60f6..e495a8b3ebf 100644 --- a/web/features/deployments/create-release/ui/source-app-picker.tsx +++ b/web/features/deployments/create-release/ui/source-app-picker.tsx @@ -21,8 +21,14 @@ import { SkeletonRectangle, SkeletonRow } from '@/app/components/base/skeleton' import { useInfiniteScroll } from '@/features/deployments/shared/hooks/use-infinite-scroll' import { TitleTooltip } from '../../shared/components/title-tooltip' import { + createReleaseSourceAppsAtom, createReleaseSourceAppSearchTextAtom, - createReleaseSourceAppsQueryAtom, + createReleaseSourceAppsErrorAtom, + createReleaseSourceAppsFetchNextPageAtom, + createReleaseSourceAppsHasNextPageAtom, + createReleaseSourceAppsIsFetchingAtom, + createReleaseSourceAppsIsFetchingNextPageAtom, + createReleaseSourceAppsIsLoadingAtom, } from '../state' const SOURCE_APP_PICKER_SKELETON_KEYS = ['first-source-app', 'second-source-app', 'third-source-app'] @@ -134,21 +140,26 @@ export function SourceAppPicker({ value, onChange, disabled = false }: { const [isShow, setIsShow] = useState(false) const searchText = useAtomValue(createReleaseSourceAppSearchTextAtom) const setSearchText = useSetAtom(createReleaseSourceAppSearchTextAtom) - const sourceAppsQuery = useAtomValue(createReleaseSourceAppsQueryAtom) - const { - data, - isLoading, - isFetchingNextPage, - hasNextPage, - } = sourceAppsQuery - const { rootRef, sentinelRef } = useInfiniteScroll(sourceAppsQuery, { + const apps = useAtomValue(createReleaseSourceAppsAtom) + const sourceAppsError = useAtomValue(createReleaseSourceAppsErrorAtom) + const sourceAppsFetchNextPage = useAtomValue(createReleaseSourceAppsFetchNextPageAtom) + const sourceAppsHasNextPage = useAtomValue(createReleaseSourceAppsHasNextPageAtom) + const sourceAppsIsFetching = useAtomValue(createReleaseSourceAppsIsFetchingAtom) + const sourceAppsIsFetchingNextPage = useAtomValue(createReleaseSourceAppsIsFetchingNextPageAtom) + const sourceAppsIsLoading = useAtomValue(createReleaseSourceAppsIsLoadingAtom) + const { rootRef, sentinelRef } = useInfiniteScroll({ + error: sourceAppsError, + fetchNextPage: sourceAppsFetchNextPage, + hasNextPage: sourceAppsHasNextPage, + isFetching: sourceAppsIsFetching, + isFetchingNextPage: sourceAppsIsFetchingNextPage, + isLoading: sourceAppsIsLoading, + }, { enabled: isShow && !disabled, rootMargin: '0px 0px 160px 0px', threshold: 0.1, }) - const apps = data?.pages.flatMap(page => page.data) ?? [] - return ( items={apps} @@ -208,23 +219,23 @@ export function SourceAppPicker({ value, onChange, disabled = false }: {
    - {(isLoading || isFetchingNextPage) && apps.length === 0 && } + {(sourceAppsIsLoading || sourceAppsIsFetchingNextPage) && apps.length === 0 && } {(app: App) => ( )} - {!(isLoading || isFetchingNextPage) && ( + {!(sourceAppsIsLoading || sourceAppsIsFetchingNextPage) && ( {t('createModal.appSearchEmpty')} )} - {isFetchingNextPage && apps.length > 0 && ( + {sourceAppsIsFetchingNextPage && apps.length > 0 && (
    {t('createModal.loadingApps')}
    )} - {hasNextPage && diff --git a/web/features/deployments/deploy-drawer/state/index.ts b/web/features/deployments/deploy-drawer/state/index.ts index f0d37466567..9e5bbb50263 100644 --- a/web/features/deployments/deploy-drawer/state/index.ts +++ b/web/features/deployments/deploy-drawer/state/index.ts @@ -18,6 +18,7 @@ import { toast } from '@langgenius/dify-ui/toast' import { skipToken } from '@tanstack/react-query' import { atom } from 'jotai' import { atomWithMutation, atomWithQuery } from 'jotai-tanstack-query' +import { selectAtom } from 'jotai/utils' import { consoleQuery } from '@/service/client' import { envVarBindingSlotFromContract } from '../../shared/components/env-var-bindings-utils' import { @@ -81,6 +82,10 @@ export const releaseDeploymentViewQueryAtom = atomWithQuery((get) => { }) }) +export const releaseDeploymentViewAtom = selectAtom(releaseDeploymentViewQueryAtom, query => query.data) +export const releaseDeploymentViewIsLoadingAtom = selectAtom(releaseDeploymentViewQueryAtom, query => query.isLoading) +export const releaseDeploymentViewIsErrorAtom = selectAtom(releaseDeploymentViewQueryAtom, query => query.isError) + const selectedEnvIdAtom = atom(undefined) const selectedReleaseIdAtom = atom(undefined) const manualBindingsAtom = atom({}) @@ -223,44 +228,47 @@ const releaseDeploymentOptionsQueryAtom = atomWithQuery((get) => { }) }) -export const deployBindingSlotsAtom = atom((get) => { - const deploymentOptionsQuery = get(releaseDeploymentOptionsQueryAtom) +const releaseDeploymentOptionsAtom = selectAtom(releaseDeploymentOptionsQueryAtom, query => query.data) +const releaseDeploymentOptionsIsLoadingAtom = selectAtom(releaseDeploymentOptionsQueryAtom, query => query.isLoading) +const releaseDeploymentOptionsIsFetchingAtom = selectAtom(releaseDeploymentOptionsQueryAtom, query => query.isFetching) +const releaseDeploymentOptionsIsErrorAtom = selectAtom(releaseDeploymentOptionsQueryAtom, query => query.isError) - return deploymentOptionsQuery.data?.options.credentialSlots.filter(slot => runtimeCredentialSlotKey(slot)) ?? [] +export const deployBindingSlotsAtom = atom((get) => { + const deploymentOptions = get(releaseDeploymentOptionsAtom) + + return deploymentOptions?.options.credentialSlots.filter(slot => runtimeCredentialSlotKey(slot)) ?? [] }) export const deployEnvVarSlotsAtom = atom((get): EnvVarBindingSlot[] => { - const deploymentOptionsQuery = get(releaseDeploymentOptionsQueryAtom) + const deploymentOptions = get(releaseDeploymentOptionsAtom) - return deploymentOptionsQuery.data?.options.envVarSlots.flatMap((slot): EnvVarBindingSlot[] => { + return deploymentOptions?.options.envVarSlots.flatMap((slot): EnvVarBindingSlot[] => { const bindingSlot = envVarBindingSlotFromContract(slot) return bindingSlot ? [bindingSlot] : [] }) ?? [] }) export const deployIsBindingOptionsLoadingAtom = atom((get) => { - const deploymentOptionsQuery = get(releaseDeploymentOptionsQueryAtom) const releaseId = get(deployTargetReleaseIdAtom) return Boolean( releaseId && get(deployHasSelectedEnvironmentAtom) - && (deploymentOptionsQuery.isLoading || deploymentOptionsQuery.isFetching), + && (get(releaseDeploymentOptionsIsLoadingAtom) || get(releaseDeploymentOptionsIsFetchingAtom)), ) }) export const deployHasBindingOptionsErrorAtom = atom((get) => { - return get(releaseDeploymentOptionsQueryAtom).isError + return get(releaseDeploymentOptionsIsErrorAtom) }) const deployIsBindingOptionsReadyAtom = atom((get) => { - const deploymentOptionsQuery = get(releaseDeploymentOptionsQueryAtom) const releaseId = get(deployTargetReleaseIdAtom) return Boolean( releaseId && get(deployHasSelectedEnvironmentAtom) - && deploymentOptionsQuery.data + && get(releaseDeploymentOptionsAtom) && !get(deployIsBindingOptionsLoadingAtom) && !get(deployHasBindingOptionsErrorAtom), ) diff --git a/web/features/deployments/deploy-drawer/ui/form.tsx b/web/features/deployments/deploy-drawer/ui/form.tsx index 6940c122380..8d11d9584b9 100644 --- a/web/features/deployments/deploy-drawer/ui/form.tsx +++ b/web/features/deployments/deploy-drawer/ui/form.tsx @@ -11,7 +11,7 @@ import { ScopeProvider } from 'jotai-scope' import { useTranslation } from 'react-i18next' import { EnvVarBindingsPanel } from '../../shared/components/env-var-bindings' import { isAvailableDeploymentTarget } from '../../shared/domain/runtime-status' -import { canAttemptDeployAtom, canSubmitDeployAtom, closeDeployDrawerAtom, deployBindingSlotsAtom, deployEnvVarSlotsAtom, deployEnvVarValuesAtom, deployFormAppInstanceIdAtom, deployHasBindingOptionsErrorAtom, deployHasSelectedEnvironmentAtom, deployIsBindingOptionsLoadingAtom, deployReadyFormConfigAtom, deployReadyFormLocalAtoms, deployReleaseSubmissionAtom, deploySelectedBindingsAtom, deployShowValidationErrorsAtom, deployTargetReleaseIdAtom, isDeployReleaseSubmittingAtom, releaseDeploymentViewQueryAtom, selectDeployBindingAtom, setDeployEnvVarAtom, showDeployValidationErrorsAtom } from '../state' +import { canAttemptDeployAtom, canSubmitDeployAtom, closeDeployDrawerAtom, deployBindingSlotsAtom, deployEnvVarSlotsAtom, deployEnvVarValuesAtom, deployFormAppInstanceIdAtom, deployHasBindingOptionsErrorAtom, deployHasSelectedEnvironmentAtom, deployIsBindingOptionsLoadingAtom, deployReadyFormConfigAtom, deployReadyFormLocalAtoms, deployReleaseSubmissionAtom, deploySelectedBindingsAtom, deployShowValidationErrorsAtom, deployTargetReleaseIdAtom, isDeployReleaseSubmittingAtom, releaseDeploymentViewAtom, releaseDeploymentViewIsErrorAtom, releaseDeploymentViewIsLoadingAtom, selectDeployBindingAtom, setDeployEnvVarAtom, showDeployValidationErrorsAtom } from '../state' import { currentReleaseIdForEnvironment, selectableDeployReleases, @@ -203,13 +203,15 @@ function DeployFormContent({ presetReleaseId, }: DeployFormProps) { const { t } = useTranslation('deployments') - const releaseDeploymentViewQuery = useAtomValue(releaseDeploymentViewQueryAtom) + const deploymentView = useAtomValue(releaseDeploymentViewAtom) + const isLoading = useAtomValue(releaseDeploymentViewIsLoadingAtom) + const isError = useAtomValue(releaseDeploymentViewIsErrorAtom) - if (releaseDeploymentViewQuery.isLoading) { + if (isLoading) { return } - if (releaseDeploymentViewQuery.isError) { + if (isError) { return (
    {t('common.loadFailed')} @@ -217,7 +219,6 @@ function DeployFormContent({ ) } - const deploymentView = releaseDeploymentViewQuery.data if (!deploymentView) { return (
    diff --git a/web/features/deployments/detail/access/channels/__tests__/section.spec.tsx b/web/features/deployments/detail/access/channels/__tests__/section.spec.tsx index ab08477cc94..d5c7c45eb17 100644 --- a/web/features/deployments/detail/access/channels/__tests__/section.spec.tsx +++ b/web/features/deployments/detail/access/channels/__tests__/section.spec.tsx @@ -2,7 +2,11 @@ import type { AccessChannels, AccessEndpoint } from '@dify/contracts/enterprise/ import { render, screen } from '@testing-library/react' import { beforeEach, describe, expect, it, vi } from 'vitest' import { deploymentRouteAppInstanceIdAtom } from '../../../../route-state' -import { accessSettingsQueryAtom } from '../../state' +import { + accessSettingsAtom, + accessSettingsIsErrorAtom, + accessSettingsIsLoadingAtom, +} from '../../state' import { AccessChannelsSection } from '../section' const mockToggleAccessChannel = vi.hoisted(() => vi.fn()) @@ -68,17 +72,17 @@ describe('AccessChannelsSection', () => { mockUseAtomValue.mockImplementation((atom) => { if (atom === deploymentRouteAppInstanceIdAtom) return 'app-instance-1' - if (atom === accessSettingsQueryAtom) { + if (atom === accessSettingsAtom) { return { - data: { - accessChannels: createAccessChannels(), - webAppEndpoints: [createEndpoint('https://app.example.com/webapp')], - cliEndpoint: createEndpoint('https://cli.example.com/entry'), - }, - isLoading: false, - isError: false, + accessChannels: createAccessChannels(), + webAppEndpoints: [createEndpoint('https://app.example.com/webapp')], + cliEndpoint: createEndpoint('https://cli.example.com/entry'), } } + if (atom === accessSettingsIsLoadingAtom) + return false + if (atom === accessSettingsIsErrorAtom) + return false return undefined }) }) diff --git a/web/features/deployments/detail/access/channels/section.tsx b/web/features/deployments/detail/access/channels/section.tsx index affd031b790..e772a1ca8e0 100644 --- a/web/features/deployments/detail/access/channels/section.tsx +++ b/web/features/deployments/detail/access/channels/section.tsx @@ -12,7 +12,11 @@ import { deploymentRouteAppInstanceIdAtom } from '../../../route-state' import { DeploymentEmptyState, DeploymentNoticeState, DeploymentStateMessage } from '../../../shared/components/empty-state' import { CopyPill, EndpointRow } from '../../../shared/components/endpoint' import { Section } from '../../../shared/components/section' -import { accessSettingsQueryAtom } from '../state' +import { + accessSettingsAtom, + accessSettingsIsErrorAtom, + accessSettingsIsLoadingAtom, +} from '../state' import { getUrlOrigin } from './url' const ACCESS_CHANNEL_SKELETON_SECTIONS = [ @@ -115,12 +119,12 @@ function ChannelRow({ info, children }: { export function AccessChannelsSection() { const { t } = useTranslation('deployments') const appInstanceId = useAtomValue(deploymentRouteAppInstanceIdAtom) - const accessSettingsQuery = useAtomValue(accessSettingsQueryAtom) - const accessChannels = accessSettingsQuery.data?.accessChannels - const webAppEndpoints: AccessEndpoint[] | undefined = accessSettingsQuery.data?.webAppEndpoints - const cliEndpoint: AccessEndpoint | undefined = accessSettingsQuery.data?.cliEndpoint - const isLoading = accessSettingsQuery.isLoading - const isError = accessSettingsQuery.isError + const accessSettings = useAtomValue(accessSettingsAtom) + const isLoading = useAtomValue(accessSettingsIsLoadingAtom) + const isError = useAtomValue(accessSettingsIsErrorAtom) + const accessChannels = accessSettings?.accessChannels + const webAppEndpoints: AccessEndpoint[] | undefined = accessSettings?.webAppEndpoints + const cliEndpoint: AccessEndpoint | undefined = accessSettings?.cliEndpoint const runEnabled = accessChannels?.webAppEnabled ?? false const webappRows = webAppEndpoints?.flatMap((endpoint) => { const endpointUrl = endpoint.endpointUrl diff --git a/web/features/deployments/detail/access/permissions/__tests__/permissions.spec.tsx b/web/features/deployments/detail/access/permissions/__tests__/permissions.spec.tsx index ecdbc625769..577be5ce086 100644 --- a/web/features/deployments/detail/access/permissions/__tests__/permissions.spec.tsx +++ b/web/features/deployments/detail/access/permissions/__tests__/permissions.spec.tsx @@ -5,7 +5,11 @@ import { fireEvent, render, screen } from '@testing-library/react' import { createStore, Provider as JotaiProvider } from 'jotai' import { describe, expect, it, vi } from 'vitest' import { deploymentRouteAppInstanceIdAtom } from '../../../../route-state' -import { accessSettingsQueryAtom } from '../../state' +import { + accessSettingsAtom, + accessSettingsIsErrorAtom, + accessSettingsIsLoadingAtom, +} from '../../state' import { EnvironmentPermissionRow } from '../environment-permission-row' import { AccessPermissionsSection } from '../section' @@ -236,15 +240,15 @@ describe('AccessPermissionsSection', () => { mockUseAtomValue.mockImplementation((atom) => { if (atom === deploymentRouteAppInstanceIdAtom) return 'app-instance-1' - if (atom === accessSettingsQueryAtom) { + if (atom === accessSettingsAtom) { return { - data: { - environmentPolicies: [createEnvironmentAccessPolicy()], - }, - isLoading: false, - isError: false, + environmentPolicies: [createEnvironmentAccessPolicy()], } } + if (atom === accessSettingsIsLoadingAtom) + return false + if (atom === accessSettingsIsErrorAtom) + return false return undefined }) }) diff --git a/web/features/deployments/detail/access/permissions/section.tsx b/web/features/deployments/detail/access/permissions/section.tsx index c1153c36bc6..8ff5e5d2d99 100644 --- a/web/features/deployments/detail/access/permissions/section.tsx +++ b/web/features/deployments/detail/access/permissions/section.tsx @@ -7,7 +7,11 @@ import { SkeletonRectangle } from '@/app/components/base/skeleton' import { deploymentRouteAppInstanceIdAtom } from '../../../route-state' import { DeploymentEmptyState, DeploymentStateMessage } from '../../../shared/components/empty-state' import { Section } from '../../../shared/components/section' -import { accessSettingsQueryAtom } from '../state' +import { + accessSettingsAtom, + accessSettingsIsErrorAtom, + accessSettingsIsLoadingAtom, +} from '../state' import { EnvironmentPermissionRow } from './environment-permission-row' const ACCESS_PERMISSIONS_SKELETON_KEYS = ['production', 'staging', 'development'] @@ -28,10 +32,10 @@ function AccessPermissionsSkeleton() { export function AccessPermissionsSection() { const { t } = useTranslation('deployments') const appInstanceId = useAtomValue(deploymentRouteAppInstanceIdAtom) - const accessSettingsQuery = useAtomValue(accessSettingsQueryAtom) - const environmentPolicies: EnvironmentAccessPolicy[] | undefined = accessSettingsQuery.data?.environmentPolicies - const isLoading = accessSettingsQuery.isLoading - const isError = accessSettingsQuery.isError + const accessSettings = useAtomValue(accessSettingsAtom) + const isLoading = useAtomValue(accessSettingsIsLoadingAtom) + const isError = useAtomValue(accessSettingsIsErrorAtom) + const environmentPolicies: EnvironmentAccessPolicy[] | undefined = accessSettings?.environmentPolicies const policyRows = environmentPolicies ?? [] return ( diff --git a/web/features/deployments/detail/access/state.ts b/web/features/deployments/detail/access/state.ts index 620678f9496..fea3e11ccf7 100644 --- a/web/features/deployments/detail/access/state.ts +++ b/web/features/deployments/detail/access/state.ts @@ -2,6 +2,7 @@ import { skipToken } from '@tanstack/react-query' import { atomWithQuery } from 'jotai-tanstack-query' +import { selectAtom } from 'jotai/utils' import { consoleQuery } from '@/service/client' import { deploymentRouteAppInstanceIdAtom } from '../../route-state' @@ -17,3 +18,7 @@ export const accessSettingsQueryAtom = atomWithQuery((get) => { enabled: Boolean(appInstanceId), }) }) + +export const accessSettingsAtom = selectAtom(accessSettingsQueryAtom, query => query.data) +export const accessSettingsIsLoadingAtom = selectAtom(accessSettingsQueryAtom, query => query.isLoading) +export const accessSettingsIsErrorAtom = selectAtom(accessSettingsQueryAtom, query => query.isError) diff --git a/web/features/deployments/detail/api-tokens/api-token-management/section.tsx b/web/features/deployments/detail/api-tokens/api-token-management/section.tsx index 34eb480a9de..5938d2822e9 100644 --- a/web/features/deployments/detail/api-tokens/api-token-management/section.tsx +++ b/web/features/deployments/detail/api-tokens/api-token-management/section.tsx @@ -16,7 +16,11 @@ import { ApiKeyGenerateMenu } from '../api-keys/api-key-generate-menu' import { ApiKeyList } from '../api-keys/api-key-list' import { CreatedApiTokenDialog } from '../api-keys/created-token-dialog' import { DeveloperApiDocsDrawer } from '../docs/docs-drawer' -import { developerApiSettingsQueryAtom } from '../state' +import { + developerApiSettingsAtom, + developerApiSettingsIsErrorAtom, + developerApiSettingsIsLoadingAtom, +} from '../state' import { DeveloperApiSkeleton } from './skeleton' type CreatedApiToken = { @@ -86,21 +90,23 @@ export function DeveloperApiSection() { const { t } = useTranslation('deployments') const appInstanceId = useAtomValue(deploymentRouteAppInstanceIdAtom) const [createdApiToken, setCreatedApiToken] = useState() - const developerApiSettingsQuery = useAtomValue(developerApiSettingsQueryAtom) - const accessChannels = developerApiSettingsQuery.data?.accessChannels + const developerApiSettings = useAtomValue(developerApiSettingsAtom) + const isLoading = useAtomValue(developerApiSettingsIsLoadingAtom) + const isError = useAtomValue(developerApiSettingsIsErrorAtom) + const accessChannels = developerApiSettings?.accessChannels const apiEnabled = accessChannels?.developerApiEnabled ?? false - const apiUrl = developerApiSettingsQuery.data?.developerApiUrl.apiUrl - const apiKeys: ApiKey[] = developerApiSettingsQuery.data?.apiKeys ?? [] - const environments = developerApiSettingsQuery.data?.environments ?? [] + const apiUrl = developerApiSettings?.developerApiUrl.apiUrl + const apiKeys: ApiKey[] = developerApiSettings?.apiKeys ?? [] + const environments = developerApiSettings?.environments ?? [] const visibleCreatedApiToken = createdApiToken && createdApiToken.appInstanceId === appInstanceId ? createdApiToken.token : undefined const hasSelectableEnvironment = environments.some(environment => Boolean(environment.id)) - if (developerApiSettingsQuery.isLoading) + if (isLoading) return - if (developerApiSettingsQuery.isError || !appInstanceId) + if (isError || !appInstanceId) return {t('common.loadFailed')} if (!apiEnabled) { diff --git a/web/features/deployments/detail/api-tokens/developer-api-header-switch.tsx b/web/features/deployments/detail/api-tokens/developer-api-header-switch.tsx index 9415754ab5e..348d2d34d79 100644 --- a/web/features/deployments/detail/api-tokens/developer-api-header-switch.tsx +++ b/web/features/deployments/detail/api-tokens/developer-api-header-switch.tsx @@ -7,7 +7,11 @@ import { useAtomValue } from 'jotai' import { useTranslation } from 'react-i18next' import { consoleQuery } from '@/service/client' import { deploymentRouteAppInstanceIdAtom } from '../../route-state' -import { developerApiSettingsQueryAtom } from './state' +import { + developerApiSettingsAtom, + developerApiSettingsIsErrorAtom, + developerApiSettingsIsLoadingAtom, +} from './state' function DeveloperApiSwitch({ checked, accessChannels, disabled }: { checked: boolean @@ -43,11 +47,13 @@ function DeveloperApiSwitch({ checked, accessChannels, disabled }: { export function DeveloperApiHeaderSwitch() { const { t } = useTranslation('deployments') - const developerApiSettingsQuery = useAtomValue(developerApiSettingsQueryAtom) - const accessChannels = developerApiSettingsQuery.data?.accessChannels + const developerApiSettings = useAtomValue(developerApiSettingsAtom) + const isLoading = useAtomValue(developerApiSettingsIsLoadingAtom) + const isError = useAtomValue(developerApiSettingsIsErrorAtom) + const accessChannels = developerApiSettings?.accessChannels const apiEnabled = accessChannels?.developerApiEnabled ?? false - if (developerApiSettingsQuery.isLoading) + if (isLoading) return return ( @@ -58,7 +64,7 @@ export function DeveloperApiHeaderSwitch() {
    ) diff --git a/web/features/deployments/detail/api-tokens/state.ts b/web/features/deployments/detail/api-tokens/state.ts index 9eece7967e2..7fde83b0d13 100644 --- a/web/features/deployments/detail/api-tokens/state.ts +++ b/web/features/deployments/detail/api-tokens/state.ts @@ -2,6 +2,7 @@ import { skipToken } from '@tanstack/react-query' import { atomWithQuery } from 'jotai-tanstack-query' +import { selectAtom } from 'jotai/utils' import { consoleQuery } from '@/service/client' import { deploymentRouteAppInstanceIdAtom } from '../../route-state' @@ -17,3 +18,7 @@ export const developerApiSettingsQueryAtom = atomWithQuery((get) => { enabled: Boolean(appInstanceId), }) }) + +export const developerApiSettingsAtom = selectAtom(developerApiSettingsQueryAtom, query => query.data) +export const developerApiSettingsIsLoadingAtom = selectAtom(developerApiSettingsQueryAtom, query => query.isLoading) +export const developerApiSettingsIsErrorAtom = selectAtom(developerApiSettingsQueryAtom, query => query.isError) diff --git a/web/features/deployments/detail/deployment-sidebar.tsx b/web/features/deployments/detail/deployment-sidebar.tsx index 95f50c768d6..74a98d606b2 100644 --- a/web/features/deployments/detail/deployment-sidebar.tsx +++ b/web/features/deployments/detail/deployment-sidebar.tsx @@ -20,7 +20,11 @@ import { usePathname } from '@/next/navigation' import { DeploymentActionsMenu } from '../deployment-actions' import { deploymentRouteAppInstanceIdAtom } from '../route-state' import { TitleTooltip } from '../shared/components/title-tooltip' -import { deploymentDetailAppInstanceQueryAtom } from './state' +import { + deploymentDetailAppInstanceAtom, + deploymentDetailAppInstanceIsErrorAtom, + deploymentDetailAppInstanceIsLoadingAtom, +} from './state' type TabDef = { key: InstanceDetailTabKey @@ -95,10 +99,12 @@ function DeploymentDetailInstanceInfo({ appInstanceId, expand }: { expand: boolean }) { const { t } = useTranslation('deployments') - const overviewQuery = useAtomValue(deploymentDetailAppInstanceQueryAtom) - const app = overviewQuery.data?.appInstance - const isLoading = !app && overviewQuery.isLoading - const isUnavailable = !app || overviewQuery.isError + const overview = useAtomValue(deploymentDetailAppInstanceAtom) + const isOverviewLoading = useAtomValue(deploymentDetailAppInstanceIsLoadingAtom) + const isOverviewError = useAtomValue(deploymentDetailAppInstanceIsErrorAtom) + const app = overview?.appInstance + const isLoading = !app && isOverviewLoading + const isUnavailable = !app || isOverviewError const instanceName = app ? app.displayName : appInstanceId return ( diff --git a/web/features/deployments/detail/instances/header-actions/new-deployment-button.tsx b/web/features/deployments/detail/instances/header-actions/new-deployment-button.tsx index 30531fde96b..d35d5afffc7 100644 --- a/web/features/deployments/detail/instances/header-actions/new-deployment-button.tsx +++ b/web/features/deployments/detail/instances/header-actions/new-deployment-button.tsx @@ -6,7 +6,8 @@ import { useTranslation } from 'react-i18next' import { openDeployDrawerAtom } from '../../../deploy-drawer/state' import { deploymentRouteAppInstanceIdAtom } from '../../../route-state' import { - deploymentEnvironmentDeploymentsQueryAtom, + deploymentEnvironmentDeploymentsIsErrorAtom, + deploymentEnvironmentDeploymentsIsLoadingAtom, deploymentRuntimeInstanceRowsAtom, } from '../../state' @@ -34,10 +35,11 @@ export function NewDeploymentButton() { } export function NewDeploymentHeaderAction() { - const environmentDeploymentsQuery = useAtomValue(deploymentEnvironmentDeploymentsQueryAtom) + const isLoading = useAtomValue(deploymentEnvironmentDeploymentsIsLoadingAtom) + const hasError = useAtomValue(deploymentEnvironmentDeploymentsIsErrorAtom) const rows = useAtomValue(deploymentRuntimeInstanceRowsAtom) - if (environmentDeploymentsQuery.isLoading || environmentDeploymentsQuery.isError || rows.length === 0) + if (isLoading || hasError || rows.length === 0) return null return diff --git a/web/features/deployments/detail/instances/index.tsx b/web/features/deployments/detail/instances/index.tsx index e33ce938d43..8df9101f916 100644 --- a/web/features/deployments/detail/instances/index.tsx +++ b/web/features/deployments/detail/instances/index.tsx @@ -14,7 +14,8 @@ import { } from '../../shared/components/detail-table' import { DeploymentEmptyState, DeploymentStateMessage } from '../../shared/components/empty-state' import { - deploymentEnvironmentDeploymentsQueryAtom, + deploymentEnvironmentDeploymentsIsErrorAtom, + deploymentEnvironmentDeploymentsIsLoadingAtom, deploymentRuntimeInstanceRowsAtom, } from '../state' import { DeploymentEnvironmentList } from './environment-list/deployment-environment-list' @@ -89,10 +90,9 @@ function DeploymentEnvironmentListSkeleton() { export function DeploymentInstances() { const { t } = useTranslation('deployments') - const environmentDeploymentsQuery = useAtomValue(deploymentEnvironmentDeploymentsQueryAtom) + const isLoading = useAtomValue(deploymentEnvironmentDeploymentsIsLoadingAtom) + const hasError = useAtomValue(deploymentEnvironmentDeploymentsIsErrorAtom) const rows = useAtomValue(deploymentRuntimeInstanceRowsAtom) - const isLoading = environmentDeploymentsQuery.isLoading - const hasError = environmentDeploymentsQuery.isError return (
    diff --git a/web/features/deployments/detail/overview/index.tsx b/web/features/deployments/detail/overview/index.tsx index da486fe10ca..fed571c6543 100644 --- a/web/features/deployments/detail/overview/index.tsx +++ b/web/features/deployments/detail/overview/index.tsx @@ -9,7 +9,11 @@ import { hasRuntimeInstanceDeployment } from '../../shared/domain/runtime-status import { AccessStatusSection, AccessStatusSectionSkeleton, ApiTokenSummarySection, ApiTokenSummarySectionSkeleton } from './access-summary/access-status-section' import { EnvironmentStrip, EnvironmentStripSkeleton } from './environment-status/environment-strip' import { ReleaseHero, ReleaseHeroSkeleton } from './release-summary/release-hero' -import { deploymentOverviewQueryAtom } from './state' +import { + deploymentOverviewAtom, + deploymentOverviewIsErrorAtom, + deploymentOverviewIsLoadingAtom, +} from './state' function OverviewLayout({ children }: { children: React.ReactNode }) { return ( @@ -63,13 +67,14 @@ function OverviewLoadingSkeleton() { export function DeploymentOverview() { const { t } = useTranslation('deployments') - const overviewQuery = useAtomValue(deploymentOverviewQueryAtom) - const overview = overviewQuery.data + const overview = useAtomValue(deploymentOverviewAtom) + const isLoading = useAtomValue(deploymentOverviewIsLoadingAtom) + const isError = useAtomValue(deploymentOverviewIsErrorAtom) - if (overviewQuery.isLoading) + if (isLoading) return - if (overviewQuery.isError) { + if (isError) { return ( {t('common.loadFailed')} diff --git a/web/features/deployments/detail/overview/state.ts b/web/features/deployments/detail/overview/state.ts index c9b1c9b7bc9..ee7cb68c3ce 100644 --- a/web/features/deployments/detail/overview/state.ts +++ b/web/features/deployments/detail/overview/state.ts @@ -2,6 +2,7 @@ import { skipToken } from '@tanstack/react-query' import { atomWithQuery } from 'jotai-tanstack-query' +import { selectAtom } from 'jotai/utils' import { consoleQuery } from '@/service/client' import { deploymentRouteAppInstanceIdAtom } from '../../route-state' @@ -17,3 +18,7 @@ export const deploymentOverviewQueryAtom = atomWithQuery((get) => { enabled: Boolean(appInstanceId), }) }) + +export const deploymentOverviewAtom = selectAtom(deploymentOverviewQueryAtom, query => query.data) +export const deploymentOverviewIsLoadingAtom = selectAtom(deploymentOverviewQueryAtom, query => query.isLoading) +export const deploymentOverviewIsErrorAtom = selectAtom(deploymentOverviewQueryAtom, query => query.isError) diff --git a/web/features/deployments/detail/releases/release-actions/__tests__/deploy-release-menu.spec.tsx b/web/features/deployments/detail/releases/release-actions/__tests__/deploy-release-menu.spec.tsx index aa22da2ed76..a0f0ff76305 100644 --- a/web/features/deployments/detail/releases/release-actions/__tests__/deploy-release-menu.spec.tsx +++ b/web/features/deployments/detail/releases/release-actions/__tests__/deploy-release-menu.spec.tsx @@ -41,8 +41,10 @@ vi.mock('../state', async (importOriginal) => { return { ...actual, - deployReleaseMenuEnvironmentDeploymentsQueryAtom: atom(environmentDeploymentsErrorResult()), - deployReleaseMenuAppInstanceQueryAtom: atom(appInstanceResult()), + deployReleaseMenuEnvironmentDeploymentsAtom: atom(undefined), + deployReleaseMenuEnvironmentDeploymentsIsErrorAtom: atom(true), + deployReleaseMenuEnvironmentDeploymentsIsLoadingAtom: atom(false), + deployReleaseMenuAppInstanceNameAtom: atom('Deployment 1'), } }) @@ -82,24 +84,6 @@ function createRelease(): Release { } } -function environmentDeploymentsErrorResult() { - return { - isError: true, - isLoading: false, - data: undefined, - } -} - -function appInstanceResult() { - return { - data: { - appInstance: { - displayName: 'Deployment 1', - }, - }, - } -} - describe('DeployReleaseMenu', () => { beforeEach(() => { vi.clearAllMocks() diff --git a/web/features/deployments/detail/releases/release-actions/deploy-release-menu.tsx b/web/features/deployments/detail/releases/release-actions/deploy-release-menu.tsx index ce97ff261b8..a0a125ab2b4 100644 --- a/web/features/deployments/detail/releases/release-actions/deploy-release-menu.tsx +++ b/web/features/deployments/detail/releases/release-actions/deploy-release-menu.tsx @@ -28,8 +28,10 @@ import { EditReleaseDialog } from './edit-release-dialog' import { exportReleaseDsl } from './release-dsl-export' import { deleteReleaseDialogOpenAtom, - deployReleaseMenuAppInstanceQueryAtom, - deployReleaseMenuEnvironmentDeploymentsQueryAtom, + deployReleaseMenuAppInstanceNameAtom, + deployReleaseMenuEnvironmentDeploymentsAtom, + deployReleaseMenuEnvironmentDeploymentsIsErrorAtom, + deployReleaseMenuEnvironmentDeploymentsIsLoadingAtom, deployReleaseMenuOpenAtom, openDeleteReleaseDialogAtom, openEditReleaseDialogAtom, @@ -54,19 +56,21 @@ function DeployReleaseMenuContent({ onDeleted }: { const setDeleteReleaseDialogOpen = useSetAtom(deleteReleaseDialogOpenAtom) const openEditReleaseDialog = useSetAtom(openEditReleaseDialogAtom) const openDeleteReleaseDialog = useSetAtom(openDeleteReleaseDialogAtom) - const environmentDeploymentsQuery = useAtomValue(deployReleaseMenuEnvironmentDeploymentsQueryAtom) - const appInstanceQuery = useAtomValue(deployReleaseMenuAppInstanceQueryAtom) + const environmentDeployments = useAtomValue(deployReleaseMenuEnvironmentDeploymentsAtom) + const environmentDeploymentsIsLoading = useAtomValue(deployReleaseMenuEnvironmentDeploymentsIsLoadingAtom) + const environmentDeploymentsIsError = useAtomValue(deployReleaseMenuEnvironmentDeploymentsIsErrorAtom) + const appInstanceName = useAtomValue(deployReleaseMenuAppInstanceNameAtom) const deleteRelease = useMutation(consoleQuery.enterprise.releaseService.deleteRelease.mutationOptions()) const exportReleaseDslMutation = useMutation(mutationOptions({ mutationKey: ['deployments', 'release-dsl-export'], mutationFn: (input: ExportReleaseDslInput) => exportReleaseDsl(input), })) - const environments = (environmentDeploymentsQuery.data?.environmentDeployments ?? []) + const deploymentEnvironmentRows = environmentDeployments?.environmentDeployments ?? [] + const environments = deploymentEnvironmentRows .map(row => row.environment) - const deploymentRows = environmentDeploymentsQuery.data?.environmentDeployments.filter(row => !isUndeployedDeploymentRow(row)) ?? [] + const deploymentRows = deploymentEnvironmentRows.filter(row => !isUndeployedDeploymentRow(row)) const targetRelease = releaseRows.find(release => release.id === releaseId) - const appInstanceName = appInstanceQuery.data?.appInstance.displayName if (!targetRelease) return null @@ -74,8 +78,8 @@ function DeployReleaseMenuContent({ onDeleted }: { const release = targetRelease const targetReleaseName = release.displayName const deleteUsageCount = releaseUsageCount(releaseId, deploymentRows) - const isCheckingDeleteUsage = open && environmentDeploymentsQuery.isLoading - const hasDeleteUsageCheckFailed = open && environmentDeploymentsQuery.isError + const isCheckingDeleteUsage = open && environmentDeploymentsIsLoading + const hasDeleteUsageCheckFailed = open && environmentDeploymentsIsError const isReleaseInUse = deleteUsageCount > 0 const isDeletingRelease = deleteRelease.isPending const isExportingDsl = exportReleaseDslMutation.isPending @@ -130,7 +134,7 @@ function DeployReleaseMenuContent({ onDeleted }: { const groupedRows = buildDeployMenuSections({ environments, - environmentDeployments: environmentDeploymentsQuery.data?.environmentDeployments ?? [], + environmentDeployments: deploymentEnvironmentRows, releaseRows, releaseId, targetRelease: release, diff --git a/web/features/deployments/detail/releases/release-actions/state.ts b/web/features/deployments/detail/releases/release-actions/state.ts index de9630800f9..a4e164ccedf 100644 --- a/web/features/deployments/detail/releases/release-actions/state.ts +++ b/web/features/deployments/detail/releases/release-actions/state.ts @@ -4,7 +4,7 @@ import type { Release } from '@dify/contracts/enterprise/types.gen' import { skipToken } from '@tanstack/react-query' import { atom } from 'jotai' import { atomWithQuery } from 'jotai-tanstack-query' -import { atomWithLazy } from 'jotai/utils' +import { atomWithLazy, selectAtom } from 'jotai/utils' import { consoleQuery } from '@/service/client' import { deploymentRouteAppInstanceIdAtom } from '../../../route-state' @@ -35,6 +35,19 @@ export const deployReleaseMenuEnvironmentDeploymentsQueryAtom = atomWithQuery((g }) }) +export const deployReleaseMenuEnvironmentDeploymentsAtom = selectAtom( + deployReleaseMenuEnvironmentDeploymentsQueryAtom, + query => query.data, +) +export const deployReleaseMenuEnvironmentDeploymentsIsLoadingAtom = selectAtom( + deployReleaseMenuEnvironmentDeploymentsQueryAtom, + query => query.isLoading, +) +export const deployReleaseMenuEnvironmentDeploymentsIsErrorAtom = selectAtom( + deployReleaseMenuEnvironmentDeploymentsQueryAtom, + query => query.isError, +) + export const deployReleaseMenuAppInstanceQueryAtom = atomWithQuery((get) => { const appInstanceId = get(deploymentRouteAppInstanceIdAtom) const menuOpen = get(deployReleaseMenuOpenAtom) @@ -49,6 +62,11 @@ export const deployReleaseMenuAppInstanceQueryAtom = atomWithQuery((get) => { }) }) +export const deployReleaseMenuAppInstanceNameAtom = selectAtom( + deployReleaseMenuAppInstanceQueryAtom, + query => query.data?.appInstance.displayName, +) + export const openEditReleaseDialogAtom = atom(null, (_get, set) => { set(deployReleaseMenuOpenAtom, false) set(deleteReleaseDialogOpenAtom, false) diff --git a/web/features/deployments/detail/releases/release-history/release-history-table.tsx b/web/features/deployments/detail/releases/release-history/release-history-table.tsx index 07debb96df1..c154c2bf1d2 100644 --- a/web/features/deployments/detail/releases/release-history/release-history-table.tsx +++ b/web/features/deployments/detail/releases/release-history/release-history-table.tsx @@ -7,8 +7,10 @@ import { DeploymentEmptyState, DeploymentStateMessage } from '../../../shared/co import { adjustReleaseHistoryPageAfterDeleteAtom, RELEASE_HISTORY_PAGE_SIZE, + releaseHistoryAtom, releaseHistoryCurrentPageAtom, - releaseHistoryQueryAtom, + releaseHistoryIsErrorAtom, + releaseHistoryIsLoadingAtom, setReleaseHistoryCurrentPageAtom, } from '../state' import { ReleaseHistoryRows } from './release-history-rows' @@ -20,9 +22,9 @@ export function ReleaseHistoryTable() { const currentPage = useAtomValue(releaseHistoryCurrentPageAtom) const setCurrentPage = useSetAtom(setReleaseHistoryCurrentPageAtom) const adjustPageAfterDelete = useSetAtom(adjustReleaseHistoryPageAfterDeleteAtom) - const releaseHistoryQuery = useAtomValue(releaseHistoryQueryAtom) - const isLoading = releaseHistoryQuery.isLoading - const hasError = releaseHistoryQuery.isError + const releaseHistory = useAtomValue(releaseHistoryAtom) + const isLoading = useAtomValue(releaseHistoryIsLoadingAtom) + const hasError = useAtomValue(releaseHistoryIsErrorAtom) if (isLoading) return @@ -35,7 +37,6 @@ export function ReleaseHistoryTable() { ) } - const releaseHistory = releaseHistoryQuery.data if (!releaseHistory) { return ( diff --git a/web/features/deployments/detail/releases/state.ts b/web/features/deployments/detail/releases/state.ts index add79f6e2a0..43500b45ea5 100644 --- a/web/features/deployments/detail/releases/state.ts +++ b/web/features/deployments/detail/releases/state.ts @@ -3,6 +3,7 @@ import { keepPreviousData, skipToken } from '@tanstack/react-query' import { atom } from 'jotai' import { atomWithQuery } from 'jotai-tanstack-query' +import { selectAtom } from 'jotai/utils' import { consoleQuery } from '@/service/client' import { deploymentRouteAppInstanceIdAtom } from '../../route-state' @@ -29,6 +30,10 @@ export const releaseHistoryQueryAtom = atomWithQuery((get) => { }) }) +export const releaseHistoryAtom = selectAtom(releaseHistoryQueryAtom, query => query.data) +export const releaseHistoryIsLoadingAtom = selectAtom(releaseHistoryQueryAtom, query => query.isLoading) +export const releaseHistoryIsErrorAtom = selectAtom(releaseHistoryQueryAtom, query => query.isError) + export const setReleaseHistoryCurrentPageAtom = atom(null, (_get, set, page: number) => { set(releaseHistoryCurrentPageAtom, Math.max(page, 0)) }) diff --git a/web/features/deployments/detail/state.ts b/web/features/deployments/detail/state.ts index 6356964f797..cc521f15804 100644 --- a/web/features/deployments/detail/state.ts +++ b/web/features/deployments/detail/state.ts @@ -3,6 +3,7 @@ import { skipToken } from '@tanstack/react-query' import { atom } from 'jotai' import { atomWithQuery } from 'jotai-tanstack-query' +import { selectAtom } from 'jotai/utils' import { nextPathnameAtom } from '@/app/components/next-route-state/atoms' import { consoleQuery } from '@/service/client' import { deploymentRouteAppInstanceIdAtom } from '../route-state' @@ -35,6 +36,10 @@ export const deploymentDetailAppInstanceQueryAtom = atomWithQuery((get) => { }) }) +export const deploymentDetailAppInstanceAtom = selectAtom(deploymentDetailAppInstanceQueryAtom, query => query.data) +export const deploymentDetailAppInstanceIsLoadingAtom = selectAtom(deploymentDetailAppInstanceQueryAtom, query => query.isLoading) +export const deploymentDetailAppInstanceIsErrorAtom = selectAtom(deploymentDetailAppInstanceQueryAtom, query => query.isError) + export const deploymentEnvironmentDeploymentsQueryAtom = atomWithQuery((get) => { const appInstanceId = get(deploymentRouteAppInstanceIdAtom) @@ -49,6 +54,10 @@ export const deploymentEnvironmentDeploymentsQueryAtom = atomWithQuery((get) => }) }) +export const deploymentEnvironmentDeploymentsAtom = selectAtom(deploymentEnvironmentDeploymentsQueryAtom, query => query.data) +export const deploymentEnvironmentDeploymentsIsLoadingAtom = selectAtom(deploymentEnvironmentDeploymentsQueryAtom, query => query.isLoading) +export const deploymentEnvironmentDeploymentsIsErrorAtom = selectAtom(deploymentEnvironmentDeploymentsQueryAtom, query => query.isError) + export const deploymentRuntimeInstanceRowsAtom = atom((get) => { - return get(deploymentEnvironmentDeploymentsQueryAtom).data?.environmentDeployments.filter(hasRuntimeInstanceDeployment) ?? [] + return get(deploymentEnvironmentDeploymentsAtom)?.environmentDeployments.filter(hasRuntimeInstanceDeployment) ?? [] }) diff --git a/web/features/deployments/list/state/index.ts b/web/features/deployments/list/state/index.ts index f33a2525640..db52a0f329b 100644 --- a/web/features/deployments/list/state/index.ts +++ b/web/features/deployments/list/state/index.ts @@ -4,7 +4,7 @@ import type { ReactNode } from 'react' import { keepPreviousData } from '@tanstack/react-query' import { atom } from 'jotai' import { atomWithInfiniteQuery, atomWithQuery } from 'jotai-tanstack-query' -import { useHydrateAtoms } from 'jotai/utils' +import { selectAtom, useHydrateAtoms } from 'jotai/utils' import { parseAsString, useQueryState } from 'nuqs' import { consoleQuery } from '@/service/client' import { deploymentStatusPollingInterval } from '../../shared/domain/runtime-status' @@ -52,8 +52,10 @@ const deploymentsListEnvironmentsQueryAtom = atomWithQuery(() => { }) }) +const deploymentsListEnvironmentsDataAtom = selectAtom(deploymentsListEnvironmentsQueryAtom, query => query.data) + export const deploymentsListEnvironmentFilterOptionsAtom = atom((get): DeploymentsListEnvironmentFilterOption[] => { - const environments = get(deploymentsListEnvironmentsQueryAtom).data?.environments ?? [] + const environments = get(deploymentsListEnvironmentsDataAtom)?.environments ?? [] return [ { @@ -83,7 +85,7 @@ export const deploymentsListSelectedEnvironmentFilterOptionAtom = atom((get): De : allOption) }) -export const deploymentsListQueryAtom = atomWithInfiniteQuery((get) => { +const deploymentsListQueryAtom = atomWithInfiniteQuery((get) => { const queryKeywords = get(deploymentsListKeywordsAtom).trim() const queryEnvironmentId = get(deploymentsListEnvironmentIdAtom) ?? undefined @@ -114,26 +116,34 @@ export const deploymentsListQueryAtom = atomWithInfiniteQuery((get) => { }) }) +const deploymentsListDataAtom = selectAtom(deploymentsListQueryAtom, query => query.data) +export const deploymentsListErrorAtom = selectAtom(deploymentsListQueryAtom, query => query.error) +export const deploymentsListFetchNextPageAtom = selectAtom(deploymentsListQueryAtom, query => query.fetchNextPage) +export const deploymentsListHasNextPageAtom = selectAtom(deploymentsListQueryAtom, query => query.hasNextPage) +export const deploymentsListIsFetchingAtom = selectAtom(deploymentsListQueryAtom, query => query.isFetching) +export const deploymentsListIsFetchingNextPageAtom = selectAtom(deploymentsListQueryAtom, query => query.isFetchingNextPage) +export const deploymentsListIsLoadingAtom = selectAtom(deploymentsListQueryAtom, query => query.isLoading) +const deploymentsListIsErrorAtom = selectAtom(deploymentsListQueryAtom, query => query.isError) + export const deploymentsListRowsAtom = atom((get) => { - return get(deploymentsListQueryAtom).data?.pages.flatMap(page => page.appInstanceSummaries) ?? [] + return get(deploymentsListDataAtom)?.pages.flatMap(page => page.appInstanceSummaries) ?? [] }) export const deploymentsListShowSkeletonAtom = atom((get) => { - const deploymentsListQuery = get(deploymentsListQueryAtom) - const pages = deploymentsListQuery.data?.pages ?? [] + const pages = get(deploymentsListDataAtom)?.pages ?? [] - return deploymentsListQuery.isLoading || (deploymentsListQuery.isFetching && pages.length === 0) + return get(deploymentsListIsLoadingAtom) || (get(deploymentsListIsFetchingAtom) && pages.length === 0) }) export const deploymentsListShowEmptyStateAtom = atom((get) => { return !get(deploymentsListShowSkeletonAtom) - && !get(deploymentsListQueryAtom).isError + && !get(deploymentsListIsErrorAtom) && get(deploymentsListRowsAtom).length === 0 }) export const deploymentsListShowErrorStateAtom = atom((get) => { return !get(deploymentsListShowSkeletonAtom) - && get(deploymentsListQueryAtom).isError + && get(deploymentsListIsErrorAtom) }) export const deploymentsListHasFilterAtom = atom((get) => { diff --git a/web/features/deployments/list/ui/shell.tsx b/web/features/deployments/list/ui/shell.tsx index e94c870eaa5..97313e60039 100644 --- a/web/features/deployments/list/ui/shell.tsx +++ b/web/features/deployments/list/ui/shell.tsx @@ -12,8 +12,13 @@ import { SkeletonRectangle } from '@/app/components/base/skeleton' import { DeploymentEmptyState, DeploymentStateMessage } from '../../shared/components/empty-state' import { useInfiniteScroll } from '../../shared/hooks/use-infinite-scroll' import { + deploymentsListErrorAtom, + deploymentsListFetchNextPageAtom, deploymentsListHasFilterAtom, - deploymentsListQueryAtom, + deploymentsListHasNextPageAtom, + deploymentsListIsFetchingAtom, + deploymentsListIsFetchingNextPageAtom, + deploymentsListIsLoadingAtom, deploymentsListRowsAtom, deploymentsListShowEmptyStateAtom, deploymentsListShowErrorStateAtom, @@ -157,13 +162,25 @@ function DeploymentsListControls() { export function DeploymentsListShell() { const { t } = useTranslation('deployments') - const deploymentsListQuery = useAtomValue(deploymentsListQueryAtom) + const deploymentsListError = useAtomValue(deploymentsListErrorAtom) + const deploymentsListFetchNextPage = useAtomValue(deploymentsListFetchNextPageAtom) + const deploymentsListHasNextPage = useAtomValue(deploymentsListHasNextPageAtom) + const deploymentsListIsFetching = useAtomValue(deploymentsListIsFetchingAtom) + const deploymentsListIsFetchingNextPage = useAtomValue(deploymentsListIsFetchingNextPageAtom) + const deploymentsListIsLoading = useAtomValue(deploymentsListIsLoadingAtom) const appInstanceSummaries = useAtomValue(deploymentsListRowsAtom) const showSkeleton = useAtomValue(deploymentsListShowSkeletonAtom) const showErrorState = useAtomValue(deploymentsListShowErrorStateAtom) const showEmptyState = useAtomValue(deploymentsListShowEmptyStateAtom) - const { rootRef, sentinelRef } = useInfiniteScroll(deploymentsListQuery) + const { rootRef, sentinelRef } = useInfiniteScroll({ + error: deploymentsListError, + fetchNextPage: deploymentsListFetchNextPage, + hasNextPage: deploymentsListHasNextPage, + isFetching: deploymentsListIsFetching, + isFetchingNextPage: deploymentsListIsFetchingNextPage, + isLoading: deploymentsListIsLoading, + }) return (
    @@ -185,7 +202,7 @@ export function DeploymentsListShell() { summary={summary} /> ))} - {deploymentsListQuery.isFetchingNextPage && } + {deploymentsListIsFetchingNextPage && }
    From abd720146d09e71bf8f153b4fddbf1c78d1af038 Mon Sep 17 00:00:00 2001 From: Nian <11332799+Lillian68@users.noreply.github.com> Date: Tue, 7 Jul 2026 23:31:04 +0800 Subject: [PATCH 33/70] test(services): cover DSL import and plugin migration regressions (#36072) Co-authored-by: WH-2099 --- api/services/plugin/plugin_migration.py | 7 ++-- .../services/plugin/test_plugin_migration.py | 32 ++++++++++++++++- .../test_rag_pipeline_dsl_service.py | 13 +++++++ .../services/test_app_dsl_service.py | 35 +++++++++---------- 4 files changed, 64 insertions(+), 23 deletions(-) diff --git a/api/services/plugin/plugin_migration.py b/api/services/plugin/plugin_migration.py index d6f154df812..5c2ddb77e0f 100644 --- a/api/services/plugin/plugin_migration.py +++ b/api/services/plugin/plugin_migration.py @@ -416,10 +416,9 @@ class PluginMigration: data = _tenant_plugin_adapter.validate_json(line) tenant_id = data["tenant_id"] plugin_ids = data["plugins"] - plugin_not_exist: list[str] = [] - for plugin_id in plugin_ids: - if plugin_id not in package_identifier_by_plugin_id: - plugin_not_exist.append(plugin_id) + plugin_not_exist = [ + plugin_id for plugin_id in plugin_ids if plugin_id not in package_identifier_by_plugin_id + ] if plugin_not_exist: not_installed.append( diff --git a/api/tests/unit_tests/services/plugin/test_plugin_migration.py b/api/tests/unit_tests/services/plugin/test_plugin_migration.py index fa10d775aa6..d94ab540dfd 100644 --- a/api/tests/unit_tests/services/plugin/test_plugin_migration.py +++ b/api/tests/unit_tests/services/plugin/test_plugin_migration.py @@ -148,5 +148,35 @@ class TestHandlePluginInstanceInstall: PluginMigration.install_plugins(str(extracted_plugins), str(output_file), workers=1) assert json.loads(output_file.read_text())["not_installed"] == [ - {"tenant_id": "tenant1", "plugin_not_exist": ["langgenius/missing"]} + { + "tenant_id": "tenant1", + "plugin_not_exist": ["langgenius/missing"], + } ] + mock_installer.install_from_identifiers.assert_called_once() + + def test_install_plugins_skips_unresolved_plugins(self, tmp_path) -> None: + extracted_plugins = tmp_path / "plugins.jsonl" + output_file = tmp_path / "output.json" + extracted_plugins.write_text('{"tenant_id":"tenant1","plugins":["langgenius/missing"]}\n') + + with ( + patch( + f"{MIGRATION_MODULE}.PluginMigration.extract_unique_plugins", + return_value={ + "plugins": {}, + "plugin_not_exist": ["langgenius/missing"], + }, + ), + patch(f"{MIGRATION_MODULE}.PluginMigration.handle_plugin_instance_install", return_value={}), + patch(f"{MIGRATION_MODULE}.PluginInstaller") as mock_installer_cls, + ): + mock_installer = MagicMock() + mock_installer.list_plugins.return_value = [] + mock_installer_cls.return_value = mock_installer + + PluginMigration.install_plugins(str(extracted_plugins), str(output_file), workers=1) + + output = json.loads(output_file.read_text()) + assert output["not_installed"] == [{"tenant_id": "tenant1", "plugin_not_exist": ["langgenius/missing"]}] + mock_installer.install_from_identifiers.assert_not_called() diff --git a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_dsl_service.py b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_dsl_service.py index 93884d07c5c..5cdb2afd093 100644 --- a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_dsl_service.py +++ b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_dsl_service.py @@ -633,6 +633,19 @@ def test_import_rag_pipeline_yaml_content_requires_content() -> None: assert "yaml_content is required" in result.error +def test_import_rag_pipeline_rejects_oversized_yaml_content_before_parsing( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr("services.rag_pipeline.rag_pipeline_dsl_service.DSL_MAX_SIZE", 3) + service = RagPipelineDslService(session=Mock()) + account = Mock(current_tenant_id="t1") + + result = service.import_rag_pipeline(account=account, import_mode="yaml-content", yaml_content="你你") + + assert result.status == ImportStatus.FAILED + assert result.error == "File size exceeds the limit of 10MB" + + def test_import_rag_pipeline_yaml_content_requires_mapping() -> None: service = RagPipelineDslService(session=Mock()) account = Mock(current_tenant_id="t1") diff --git a/api/tests/unit_tests/services/test_app_dsl_service.py b/api/tests/unit_tests/services/test_app_dsl_service.py index 64236ea5a90..ae621bdcf52 100644 --- a/api/tests/unit_tests/services/test_app_dsl_service.py +++ b/api/tests/unit_tests/services/test_app_dsl_service.py @@ -1,50 +1,49 @@ -from types import SimpleNamespace from unittest.mock import Mock -from services.app_dsl_service import AppDslService, ImportStatus +import pytest + +from services.app_dsl_service import AppDslService +from services.entities.dsl_entities import ImportStatus -def test_import_app_rejects_oversized_yaml_content_by_bytes(monkeypatch) -> None: - monkeypatch.setattr("services.app_dsl_service.DSL_MAX_SIZE", 1) - service = AppDslService(session=SimpleNamespace()) +def test_import_app_rejects_oversized_yaml_content_before_parsing(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr("services.app_dsl_service.DSL_MAX_SIZE", 3) + service = AppDslService(session=Mock()) + account = Mock(current_tenant_id="tenant-1") - result = service.import_app( - account=SimpleNamespace(current_tenant_id="tenant-1"), - import_mode="yaml-content", - yaml_content="é", - ) + result = service.import_app(account=account, import_mode="yaml-content", yaml_content="你你") assert result.status == ImportStatus.FAILED - assert "10MB" in result.error + assert result.error == "File size exceeds the limit of 10MB" -def test_import_app_rejects_oversized_yaml_url_bytes_before_decode(monkeypatch) -> None: +def test_import_app_rejects_oversized_yaml_url_bytes_before_decode(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr("services.app_dsl_service.DSL_MAX_SIZE", 1) response = Mock() response.raise_for_status.return_value = None response.content = b"\xff\xff" monkeypatch.setattr("services.app_dsl_service.remote_fetcher.make_request", Mock(return_value=response)) - service = AppDslService(session=SimpleNamespace()) + service = AppDslService(session=Mock()) result = service.import_app( - account=SimpleNamespace(current_tenant_id="tenant-1"), + account=Mock(current_tenant_id="tenant-1"), import_mode="yaml-url", yaml_url="https://example.com/app.yaml", ) assert result.status == ImportStatus.FAILED - assert "10MB" in result.error + assert result.error == "File size exceeds the limit of 10MB" -def test_import_app_returns_decode_error_for_invalid_yaml_url_bytes(monkeypatch) -> None: +def test_import_app_returns_decode_error_for_invalid_yaml_url_bytes(monkeypatch: pytest.MonkeyPatch) -> None: response = Mock() response.raise_for_status.return_value = None response.content = b"\xff" monkeypatch.setattr("services.app_dsl_service.remote_fetcher.make_request", Mock(return_value=response)) - service = AppDslService(session=SimpleNamespace()) + service = AppDslService(session=Mock()) result = service.import_app( - account=SimpleNamespace(current_tenant_id="tenant-1"), + account=Mock(current_tenant_id="tenant-1"), import_mode="yaml-url", yaml_url="https://example.com/app.yaml", ) From 915655683c596e7cbb0e200d78a29310468333b7 Mon Sep 17 00:00:00 2001 From: L1nSn0w Date: Wed, 8 Jul 2026 10:22:59 +0800 Subject: [PATCH 34/70] refactor(openapi): resource-oriented paths for /openapi/v1 + difyctl version gate (#38367) Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- api/configs/packaging/pyproject.py | 16 ++ api/controllers/openapi/__init__.py | 2 + api/controllers/openapi/_errors.py | 1 + api/controllers/openapi/_models.py | 8 +- api/controllers/openapi/_version_gate.py | 69 +++++ api/controllers/openapi/app_dsl.py | 8 +- api/controllers/openapi/app_run.py | 6 +- api/controllers/openapi/apps.py | 2 +- .../openapi/apps_permitted_external.py | 2 +- api/controllers/openapi/files.py | 4 +- api/controllers/openapi/human_input_form.py | 9 +- api/controllers/openapi/workspaces.py | 18 +- api/openapi/markdown/openapi-openapi.md | 118 ++++---- api/pyproject.toml | 5 + .../controllers/openapi/test_app_run.py | 24 +- .../controllers/openapi/test_apps.py | 16 +- .../controllers/openapi/test_app_dsl.py | 8 +- .../controllers/openapi/test_app_run.py | 2 +- .../controllers/openapi/test_apps.py | 6 +- .../controllers/openapi/test_files.py | 2 +- .../controllers/openapi/test_workspaces.py | 4 +- .../openapi/test_app_run_streaming.py | 8 +- .../openapi/test_human_input_form.py | 32 +-- .../controllers/openapi/test_version_gate.py | 100 +++++++ .../openapi/test_workspaces_members.py | 42 ++- cli/package.json | 6 +- cli/scripts/release-naming.test.ts | 36 +-- cli/scripts/release-r2-edge.test.ts | 2 +- cli/src/api/app-dsl.test.ts | 6 +- cli/src/api/app-dsl.ts | 4 +- cli/src/api/app-run.ts | 4 +- cli/src/api/apps.test.ts | 6 +- cli/src/api/apps.ts | 2 +- cli/src/api/file-upload.test.ts | 4 +- cli/src/api/file-upload.ts | 2 +- cli/src/api/members.test.ts | 10 +- cli/src/api/members.ts | 2 +- cli/src/api/permitted-external-apps.test.ts | 6 +- cli/src/api/permitted-external-apps.ts | 2 +- cli/src/api/workspaces.ts | 2 +- cli/src/cache/compat-store.test.ts | 57 ++++ cli/src/cache/compat-store.ts | 71 +++++ cli/src/commands/_shared/authed-command.ts | 5 + cli/src/commands/auth/login/index.ts | 4 + cli/src/commands/auth/login/login.ts | 7 + cli/src/commands/resume/app/run.test.ts | 2 +- cli/src/commands/resume/app/run.ts | 2 +- cli/src/commands/run/app/run.ts | 2 +- cli/src/commands/use/workspace/use.ts | 2 +- cli/src/commands/version/version.test.ts | 6 +- cli/src/http/client.test.ts | 4 +- cli/src/http/error-mapper.test.ts | 27 +- cli/src/http/error-mapper.ts | 11 + cli/src/store/manager.ts | 1 + cli/src/version/compat.test.ts | 32 ++- cli/src/version/compat.ts | 40 +-- cli/src/version/enforce.test.ts | 86 ++++++ cli/src/version/enforce.ts | 69 +++++ cli/src/version/nudge.ts | 4 +- cli/src/version/probe.test.ts | 4 +- cli/src/version/render.test.ts | 4 +- cli/src/version/render.ts | 6 +- .../suites/discovery/get-app-single.e2e.ts | 2 +- cli/test/fixtures/dify-mock/server.test.ts | 24 +- cli/test/fixtures/dify-mock/server.ts | 25 +- .../generated/api/openapi/orpc.gen.ts | 267 ++++++++---------- .../generated/api/openapi/types.gen.ts | 197 +++++++------ .../generated/api/openapi/zod.gen.ts | 87 +++--- packages/contracts/openapi-ts.api.config.ts | 13 +- 69 files changed, 1097 insertions(+), 570 deletions(-) create mode 100644 api/controllers/openapi/_version_gate.py create mode 100644 api/tests/unit_tests/controllers/openapi/test_version_gate.py create mode 100644 cli/src/cache/compat-store.test.ts create mode 100644 cli/src/cache/compat-store.ts create mode 100644 cli/src/version/enforce.test.ts create mode 100644 cli/src/version/enforce.ts diff --git a/api/configs/packaging/pyproject.py b/api/configs/packaging/pyproject.py index 90b1ecba065..c21c02082c5 100644 --- a/api/configs/packaging/pyproject.py +++ b/api/configs/packaging/pyproject.py @@ -6,6 +6,17 @@ class PyProjectConfig(BaseModel): version: str = Field(description="Dify version", default="") +class DifyToolConfig(BaseModel): + min_difyctl_version: str = Field( + description="Oldest difyctl version served on /openapi/v1", + default="0.0.0", + ) + + +class ToolConfig(BaseModel): + dify: DifyToolConfig = Field(default=DifyToolConfig()) + + class PyProjectTomlConfig(BaseSettings): """ configs in api/pyproject.toml @@ -15,3 +26,8 @@ class PyProjectTomlConfig(BaseSettings): description="configs in the project section of pyproject.toml", default=PyProjectConfig(), ) + + tool: ToolConfig = Field( + description="configs in the [tool.*] section of pyproject.toml", + default=ToolConfig(), + ) diff --git a/api/controllers/openapi/__init__.py b/api/controllers/openapi/__init__.py index 81c65ca03be..0260422ec1c 100644 --- a/api/controllers/openapi/__init__.py +++ b/api/controllers/openapi/__init__.py @@ -2,11 +2,13 @@ from flask import Blueprint from flask_restx import Namespace from controllers.openapi._errors import ErrorBody, OpenApiErrorCode, OpenApiErrorFormatter +from controllers.openapi._version_gate import attach_version_gate from libs.device_flow_security import attach_anti_framing from libs.external_api import ExternalApi bp = Blueprint("openapi", __name__, url_prefix="/openapi/v1") attach_anti_framing(bp) +attach_version_gate(bp) api = ExternalApi( bp, diff --git a/api/controllers/openapi/_errors.py b/api/controllers/openapi/_errors.py index 31d577665a3..92884dfcd50 100644 --- a/api/controllers/openapi/_errors.py +++ b/api/controllers/openapi/_errors.py @@ -45,6 +45,7 @@ class OpenApiErrorCode(StrEnum): TOO_MANY_REQUESTS = "too_many_requests" INTERNAL_ERROR = "internal_server_error" BAD_GATEWAY = "bad_gateway" + UPGRADE_REQUIRED = "upgrade_required" UNKNOWN = "unknown" # domain codes (must match the error_code attribute of the exception # classes raised on the openapi surface) diff --git a/api/controllers/openapi/_models.py b/api/controllers/openapi/_models.py index 6e8a9c9d439..5337612e7b6 100644 --- a/api/controllers/openapi/_models.py +++ b/api/controllers/openapi/_models.py @@ -279,7 +279,7 @@ def _csv_string_query_schema(schema: dict[str, Any]) -> None: class AppDescribeQuery(BaseModel): - """`?fields=` allow-list for GET /apps//describe. + """`?fields=` allow-list for GET /apps/. Empty / omitted → all blocks. Unknown member → ValidationError → 422. """ @@ -441,7 +441,7 @@ class MemberActionResponse(BaseModel): class TaskStopResponse(BaseModel): - """200 body for POST /apps//tasks//stop. The handler always returns + """200 body for POST /apps//tasks/:stop. The handler always returns {"result": "success"}, so `result` is required (no default) — the generated contract types it as a required `'success'` rather than an optional field.""" @@ -473,7 +473,7 @@ class AppDslImportPayload(BaseModel): class AppDslExportQuery(BaseModel): - """Query parameters for GET /apps//export.""" + """Query parameters for GET /apps//dsl.""" include_secret: bool = Field(False, description="Include encrypted secret values in the exported DSL") workflow_id: UUIDStr | None = Field( @@ -488,7 +488,7 @@ class AppDslExportResponse(BaseModel): class FormSubmitResponse(BaseModel): - """Empty 200 body for POST /apps//form/human_input/. `extra='forbid'` + """Empty 200 body for POST /apps//human-input-forms/:submit. `extra='forbid'` pins `additionalProperties: false` so the generated contract is an exact `{}` rather than an under-annotated open object.""" diff --git a/api/controllers/openapi/_version_gate.py b/api/controllers/openapi/_version_gate.py new file mode 100644 index 00000000000..785617fd950 --- /dev/null +++ b/api/controllers/openapi/_version_gate.py @@ -0,0 +1,69 @@ +"""Version gate: reject outdated difyctl clients on /openapi/v1 with HTTP 426. + +difyctl and the ``/openapi/v1`` surface ship in lockstep. A breaking path change +(resource-oriented paths) means an outdated difyctl would call removed paths and +get a bare 404; this gate returns ``426 Upgrade Required`` with an upgrade hint +instead. +""" + +from __future__ import annotations + +import re +from typing import Final + +from flask import Blueprint, Response, request +from packaging.version import InvalidVersion, Version + +from configs import dify_config +from controllers.openapi._errors import ErrorBody, OpenApiErrorCode + +_UPGRADE_HINT: Final = "Upgrade difyctl: https://docs.dify.ai/en/cli/install" + +# difyctl sends `User-Agent: difyctl/ (; ; )`. +_DIFYCTL_UA_RE = re.compile(r"^difyctl/(\d+\.\d+\.\d+(?:-[\w.]+)?)") + +_PREFIX: Final = "/openapi/v1/" + +# Paths a too-old client must still reach to discover that it is outdated. +_ALLOWLIST: Final = frozenset({"/openapi/v1/_version", "/openapi/v1/_health"}) + + +def _upgrade_required_response(client_version: str, min_version: str) -> Response: + body = ErrorBody( + code=OpenApiErrorCode.UPGRADE_REQUIRED, + message=f"difyctl {client_version} is no longer supported; upgrade to >= {min_version}.", + status=426, + hint=_UPGRADE_HINT, + ) + return Response(body.model_dump_json(exclude_none=True), status=426, mimetype="application/json") + + +def attach_version_gate(bp: Blueprint) -> None: + """Reject difyctl clients older than ``[tool.dify] min_difyctl_version`` with 426. + + Registered app-wide (``before_app_request``) rather than blueprint-scoped so it + also fires for requests to *removed* paths — those no longer match an openapi + route and would 404 before a blueprint-scoped ``before_request`` ever runs. The + prefix guard scopes it back to ``/openapi/v1``. Fails open for non-difyctl or + unparseable User-Agents (only a confidently-too-old difyctl is blocked). + """ + + @bp.before_app_request + def _enforce_min_client_version() -> Response | None: # pyright: ignore[reportUnusedFunction] + if not request.path.startswith(_PREFIX): + return None + if request.path in _ALLOWLIST: + return None + match = _DIFYCTL_UA_RE.match(request.headers.get("User-Agent", "")) + if match is None: + return None + try: + client_version = Version(match.group(1)) + except InvalidVersion: + return None + # Compare the numeric core (major.minor.patch) only — a pre-release build + # like 0.2.0-rc.1 must not sort below the 0.2.0 floor. + min_version = dify_config.tool.dify.min_difyctl_version + if client_version.release[:3] < Version(min_version).release[:3]: + return _upgrade_required_response(match.group(1), min_version) + return None diff --git a/api/controllers/openapi/app_dsl.py b/api/controllers/openapi/app_dsl.py index 9b1abd24bac..cea7127bd07 100644 --- a/api/controllers/openapi/app_dsl.py +++ b/api/controllers/openapi/app_dsl.py @@ -30,7 +30,7 @@ class AppDslImportApi(Resource): a new app. Returns 202 when the DSL version requires an explicit confirmation step - (major version mismatch). Callers must then POST to the confirm endpoint. + (major version mismatch). Callers must then POST to the imports :confirm method. Returns 400 when the import failed due to invalid DSL or a business error. """ @@ -79,7 +79,7 @@ class AppDslImportApi(Resource): return result, 200 -@openapi_ns.route("/workspaces//apps/imports//confirm") +@openapi_ns.route("/workspaces//apps/imports/:confirm") class AppDslImportConfirmApi(Resource): """Confirm a pending DSL import identified by ``import_id``. @@ -119,7 +119,7 @@ class AppDslImportConfirmApi(Resource): return result, 200 -@openapi_ns.route("/apps//export") +@openapi_ns.route("/apps//dsl") class AppDslExportApi(Resource): """Export an app's current draft configuration as a DSL YAML string. @@ -153,7 +153,7 @@ class AppDslExportApi(Resource): return AppDslExportResponse(data=data), 200 -@openapi_ns.route("/apps//check-dependencies") +@openapi_ns.route("/apps//dependencies:check") class AppDslCheckDependenciesApi(Resource): """Check for leaked plugin dependencies after a DSL import. diff --git a/api/controllers/openapi/app_run.py b/api/controllers/openapi/app_run.py index 6074c7c0e02..772513ad417 100644 --- a/api/controllers/openapi/app_run.py +++ b/api/controllers/openapi/app_run.py @@ -1,4 +1,4 @@ -"""POST /openapi/v1/apps//run — mode-agnostic runner.""" +"""POST /openapi/v1/apps/:run — mode-agnostic runner.""" from __future__ import annotations @@ -138,7 +138,7 @@ _DISPATCH: dict[AppMode, Callable[[App, Any, AppRunRequest, Session], Any]] = { } -@openapi_ns.route("/apps//run") +@openapi_ns.route("/apps/:run") class AppRunApi(Resource): @auth_router.guard( scope=Scope.APPS_RUN, @@ -174,7 +174,7 @@ class AppRunApi(Resource): return helper.compact_generate_response(stream_obj) -@openapi_ns.route("/apps//tasks//stop") +@openapi_ns.route("/apps//tasks/:stop") class AppRunTaskStopApi(Resource): @auth_router.guard( scope=Scope.APPS_RUN, diff --git a/api/controllers/openapi/apps.py b/api/controllers/openapi/apps.py index 8d5c9670e77..d4cb175ba5e 100644 --- a/api/controllers/openapi/apps.py +++ b/api/controllers/openapi/apps.py @@ -129,7 +129,7 @@ def build_app_describe_response(app: App, fields: set[str] | None) -> AppDescrib return AppDescribeResponse(info=info, parameters=parameters, input_schema=input_schema) -@openapi_ns.route("/apps//describe") +@openapi_ns.route("/apps/") class AppDescribeApi(AppReadResource): @auth_router.guard( scope=Scope.APPS_READ, diff --git a/api/controllers/openapi/apps_permitted_external.py b/api/controllers/openapi/apps_permitted_external.py index 718d3dbd169..5c6fdce5141 100644 --- a/api/controllers/openapi/apps_permitted_external.py +++ b/api/controllers/openapi/apps_permitted_external.py @@ -87,7 +87,7 @@ class PermittedExternalAppsListApi(Resource): return env -@openapi_ns.route("/permitted-external-apps//describe") +@openapi_ns.route("/permitted-external-apps/") class PermittedExternalAppDescribeApi(Resource): @auth_router.guard( scope=Scope.APPS_READ_PERMITTED_EXTERNAL, diff --git a/api/controllers/openapi/files.py b/api/controllers/openapi/files.py index 7326a4a922e..3b3f68fa36b 100644 --- a/api/controllers/openapi/files.py +++ b/api/controllers/openapi/files.py @@ -1,4 +1,4 @@ -"""POST /openapi/v1/apps//files/upload — upload a file for use in app inputs.""" +"""POST /openapi/v1/apps//files — upload a file for use in app inputs.""" from __future__ import annotations @@ -26,7 +26,7 @@ from libs.oauth_bearer import Scope from services.file_service import FileService -@openapi_ns.route("/apps//files/upload") +@openapi_ns.route("/apps//files") class AppFileUploadApi(Resource): @openapi_ns.doc("upload_file_for_app_input") @openapi_ns.doc(description="Upload a file to use as an input variable when running the app") diff --git a/api/controllers/openapi/human_input_form.py b/api/controllers/openapi/human_input_form.py index 998dd669836..593887ffcfc 100644 --- a/api/controllers/openapi/human_input_form.py +++ b/api/controllers/openapi/human_input_form.py @@ -1,8 +1,8 @@ """ OpenAPI bearer-authed human input form endpoints. -GET /apps//form/human_input/ — fetch paused form definition -POST /apps//form/human_input/ — submit form response +GET /apps//human-input-forms/ — fetch paused form definition +POST /apps//human-input-forms/:submit — submit form response """ from __future__ import annotations @@ -60,7 +60,7 @@ def _ensure_form_is_allowed_for_openapi(form) -> None: raise RecipientSurfaceMismatch() -@openapi_ns.route("/apps//form/human_input/") +@openapi_ns.route("/apps//human-input-forms/") class OpenApiWorkflowHumanInputFormApi(Resource): @openapi_ns.response(200, "Form definition", openapi_ns.models[HumanInputFormDefinitionResponse.__name__]) @auth_router.guard( @@ -79,6 +79,9 @@ class OpenApiWorkflowHumanInputFormApi(Resource): service.ensure_form_active(form) return _jsonify_form_definition(form) + +@openapi_ns.route("/apps//human-input-forms/:submit") +class OpenApiWorkflowHumanInputFormSubmitApi(Resource): @auth_router.guard( scope=Scope.APPS_RUN, rbac=RBACRequirement(resource_type=RBACResourceScope.APP, scene=RBACPermission.APP_TEST_AND_RUN), diff --git a/api/controllers/openapi/workspaces.py b/api/controllers/openapi/workspaces.py index 49f8fb9656f..c45c02e54d3 100644 --- a/api/controllers/openapi/workspaces.py +++ b/api/controllers/openapi/workspaces.py @@ -113,7 +113,7 @@ class WorkspaceByIdApi(Resource): return _workspace_detail(tenant, membership) -@openapi_ns.route("/workspaces//switch") +@openapi_ns.route("/workspaces/:switch") class WorkspaceSwitchApi(Resource): """Server-side switch — equivalent to the console's POST /workspaces/switch. @@ -212,11 +212,12 @@ class WorkspaceMembersApi(Resource): @openapi_ns.route("/workspaces//members/") class WorkspaceMemberApi(Resource): - """Remove a member. + """Remove a member (DELETE) or change a member's role (PATCH). Self-removal and owner-removal are explicitly rejected by the service layer (CannotOperateSelfError, NoPermissionError) — both surface as - 400 per the spec, with the service's message preserved. + 400 per the spec, with the service's message preserved. Owner can never be + assigned via PATCH (closed enum); admin cannot demote the standing owner. """ @auth_router.guard_workspace( @@ -243,15 +244,6 @@ class WorkspaceMemberApi(Resource): return MemberActionResponse() - -@openapi_ns.route("/workspaces//members//role") -class WorkspaceMemberRoleApi(Resource): - """Change a member's role. - - Owner cannot be assigned here (closed enum). Admin cannot demote the - standing owner (service NoPermissionError → 400, per spec). - """ - @auth_router.guard_workspace( scope=Scope.WORKSPACE_WRITE, allowed_token_types=frozenset({TokenType.OAUTH_ACCOUNT}), @@ -259,7 +251,7 @@ class WorkspaceMemberRoleApi(Resource): ) @returns(200, MemberActionResponse, description="Role updated") @accepts(body=MemberRoleUpdatePayload) - def put(self, workspace_id: str, member_id: str, *, auth_data: AuthData, body: MemberRoleUpdatePayload): + def patch(self, workspace_id: str, member_id: str, *, auth_data: AuthData, body: MemberRoleUpdatePayload): operator = _load_account(auth_data.account_id) tenant = _load_tenant(workspace_id) member = AccountService.get_account_by_id(db.session, member_id) diff --git a/api/openapi/markdown/openapi-openapi.md b/api/openapi/markdown/openapi-openapi.md index a649c2fa74f..c5ec6c3a7db 100644 --- a/api/openapi/markdown/openapi-openapi.md +++ b/api/openapi/markdown/openapi-openapi.md @@ -93,21 +93,7 @@ User-scoped operations | 422 | Validation error | **application/json**: [ErrorBody](#errorbody)
    | | default | Error | **application/json**: [ErrorBody](#errorbody)
    | -### [GET] /apps/{app_id}/check-dependencies -#### Parameters - -| Name | Located in | Description | Required | Schema | -| ---- | ---------- | ----------- | -------- | ------ | -| app_id | path | | Yes | string | - -#### Responses - -| Code | Description | Schema | -| ---- | ----------- | ------ | -| 200 | Dependencies checked | **application/json**: [CheckDependenciesResult](#checkdependenciesresult)
    | -| default | Error | **application/json**: [ErrorBody](#errorbody)
    | - -### [GET] /apps/{app_id}/describe +### [GET] /apps/{app_id} #### Parameters | Name | Located in | Description | Required | Schema | @@ -123,7 +109,21 @@ User-scoped operations | 422 | Validation error | **application/json**: [ErrorBody](#errorbody)
    | | default | Error | **application/json**: [ErrorBody](#errorbody)
    | -### [GET] /apps/{app_id}/export +### [GET] /apps/{app_id}/dependencies:check +#### Parameters + +| Name | Located in | Description | Required | Schema | +| ---- | ---------- | ----------- | -------- | ------ | +| app_id | path | | Yes | string | + +#### Responses + +| Code | Description | Schema | +| ---- | ----------- | ------ | +| 200 | Dependencies checked | **application/json**: [CheckDependenciesResult](#checkdependenciesresult)
    | +| default | Error | **application/json**: [ErrorBody](#errorbody)
    | + +### [GET] /apps/{app_id}/dsl #### Parameters | Name | Located in | Description | Required | Schema | @@ -140,7 +140,7 @@ User-scoped operations | 422 | Validation error | **application/json**: [ErrorBody](#errorbody)
    | | default | Error | **application/json**: [ErrorBody](#errorbody)
    | -### [POST] /apps/{app_id}/files/upload +### [POST] /apps/{app_id}/files Upload a file to use as an input variable when running the app #### Parameters @@ -160,7 +160,7 @@ Upload a file to use as an input variable when running the app | 415 | Unsupported file type or blocked extension | | | default | Error | **application/json**: [ErrorBody](#errorbody)
    | -### [GET] /apps/{app_id}/form/human_input/{form_token} +### [GET] /apps/{app_id}/human-input-forms/{form_token} #### Parameters | Name | Located in | Description | Required | Schema | @@ -174,7 +174,7 @@ Upload a file to use as an input variable when running the app | ---- | ----------- | ------ | | 200 | Form definition | **application/json**: [HumanInputFormDefinitionResponse](#humaninputformdefinitionresponse)
    | -### [POST] /apps/{app_id}/form/human_input/{form_token} +### [POST] /apps/{app_id}/human-input-forms/{form_token}:submit #### Parameters | Name | Located in | Description | Required | Schema | @@ -196,7 +196,38 @@ Upload a file to use as an input variable when running the app | 422 | Validation error | **application/json**: [ErrorBody](#errorbody)
    | | default | Error | **application/json**: [ErrorBody](#errorbody)
    | -### [POST] /apps/{app_id}/run +### [GET] /apps/{app_id}/tasks/{task_id}/events +#### Parameters + +| Name | Located in | Description | Required | Schema | +| ---- | ---------- | ----------- | -------- | ------ | +| continue_on_pause | query | Whether to keep the event stream open on pause | No | boolean | +| include_state_snapshot | query | Whether to include workflow state snapshots | No | boolean | +| app_id | path | | Yes | string | +| task_id | path | | Yes | string | + +#### Responses + +| Code | Description | Schema | +| ---- | ----------- | ------ | +| 200 | SSE event stream | **application/json**: [EventStreamResponse](#eventstreamresponse)
    | + +### [POST] /apps/{app_id}/tasks/{task_id}:stop +#### Parameters + +| Name | Located in | Description | Required | Schema | +| ---- | ---------- | ----------- | -------- | ------ | +| app_id | path | | Yes | string | +| task_id | path | | Yes | string | + +#### Responses + +| Code | Description | Schema | +| ---- | ----------- | ------ | +| 200 | Task stopped | **application/json**: [TaskStopResponse](#taskstopresponse)
    | +| default | Error | **application/json**: [ErrorBody](#errorbody)
    | + +### [POST] /apps/{app_id}:run #### Parameters | Name | Located in | Description | Required | Schema | @@ -216,37 +247,6 @@ Upload a file to use as an input variable when running the app | 200 | Run result (SSE stream) | **application/json**: [EventStreamResponse](#eventstreamresponse)
    | | 422 | Validation error | **application/json**: [ErrorBody](#errorbody)
    | -### [GET] /apps/{app_id}/tasks/{task_id}/events -#### Parameters - -| Name | Located in | Description | Required | Schema | -| ---- | ---------- | ----------- | -------- | ------ | -| continue_on_pause | query | Whether to keep the event stream open on pause | No | boolean | -| include_state_snapshot | query | Whether to include workflow state snapshots | No | boolean | -| app_id | path | | Yes | string | -| task_id | path | | Yes | string | - -#### Responses - -| Code | Description | Schema | -| ---- | ----------- | ------ | -| 200 | SSE event stream | **application/json**: [EventStreamResponse](#eventstreamresponse)
    | - -### [POST] /apps/{app_id}/tasks/{task_id}/stop -#### Parameters - -| Name | Located in | Description | Required | Schema | -| ---- | ---------- | ----------- | -------- | ------ | -| app_id | path | | Yes | string | -| task_id | path | | Yes | string | - -#### Responses - -| Code | Description | Schema | -| ---- | ----------- | ------ | -| 200 | Task stopped | **application/json**: [TaskStopResponse](#taskstopresponse)
    | -| default | Error | **application/json**: [ErrorBody](#errorbody)
    | - ### [POST] /oauth/device/approve #### Request Body @@ -330,7 +330,7 @@ Upload a file to use as an input variable when running the app | 422 | Validation error | **application/json**: [ErrorBody](#errorbody)
    | | default | Error | **application/json**: [ErrorBody](#errorbody)
    | -### [GET] /permitted-external-apps/{app_id}/describe +### [GET] /permitted-external-apps/{app_id} #### Parameters | Name | Located in | Description | Required | Schema | @@ -391,7 +391,7 @@ Upload a file to use as an input variable when running the app | 422 | Validation error | **application/json**: [ErrorBody](#errorbody)
    | | default | Error | **application/json**: [ErrorBody](#errorbody)
    | -### [POST] /workspaces/{workspace_id}/apps/imports/{import_id}/confirm +### [POST] /workspaces/{workspace_id}/apps/imports/{import_id}:confirm #### Parameters | Name | Located in | Description | Required | Schema | @@ -460,7 +460,7 @@ Upload a file to use as an input variable when running the app | 200 | Member removed | **application/json**: [MemberActionResponse](#memberactionresponse)
    | | default | Error | **application/json**: [ErrorBody](#errorbody)
    | -### [PUT] /workspaces/{workspace_id}/members/{member_id}/role +### [PATCH] /workspaces/{workspace_id}/members/{member_id} #### Parameters | Name | Located in | Description | Required | Schema | @@ -482,7 +482,7 @@ Upload a file to use as an input variable when running the app | 422 | Validation error | **application/json**: [ErrorBody](#errorbody)
    | | default | Error | **application/json**: [ErrorBody](#errorbody)
    | -### [POST] /workspaces/{workspace_id}/switch +### [POST] /workspaces/{workspace_id}:switch #### Parameters | Name | Located in | Description | Required | Schema | @@ -532,7 +532,7 @@ Upload a file to use as an input variable when running the app #### AppDescribeQuery -`?fields=` allow-list for GET /apps//describe. +`?fields=` allow-list for GET /apps/. Empty / omitted → all blocks. Unknown member → ValidationError → 422. @@ -550,7 +550,7 @@ Empty / omitted → all blocks. Unknown member → ValidationError → 422. #### AppDslExportQuery -Query parameters for GET /apps//export. +Query parameters for GET /apps//dsl. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | @@ -762,7 +762,7 @@ future server adds a code. Formatter tests pin emitted values to the enum. #### FormSubmitResponse -Empty 200 body for POST /apps//form/human_input/. `extra='forbid'` +Empty 200 body for POST /apps//human-input-forms/:submit. `extra='forbid'` pins `additionalProperties: false` so the generated contract is an exact `{}` rather than an under-annotated open object. @@ -1022,7 +1022,7 @@ generated CLI whitelist all derive from it. #### TaskStopResponse -200 body for POST /apps//tasks//stop. The handler always returns +200 body for POST /apps//tasks/:stop. The handler always returns {"result": "success"}, so `result` is required (no default) — the generated contract types it as a required `'success'` rather than an optional field. diff --git a/api/pyproject.toml b/api/pyproject.toml index d3fcb59d694..519e400c15f 100644 --- a/api/pyproject.toml +++ b/api/pyproject.toml @@ -52,6 +52,11 @@ dependencies = [ # Before adding new dependency, consider place it in # alphabet order (a-z) and suitable group. +[tool.dify] +# Oldest difyctl served on /openapi/v1. Bump in lockstep with breaking /openapi/v1 +# changes (paired with difyctl's own version + its compat.minDify). +min_difyctl_version = "0.2.0" + [tool.setuptools] packages = [] diff --git a/api/tests/integration_tests/controllers/openapi/test_app_run.py b/api/tests/integration_tests/controllers/openapi/test_app_run.py index b4f383a7cea..fbfedd0f269 100644 --- a/api/tests/integration_tests/controllers/openapi/test_app_run.py +++ b/api/tests/integration_tests/controllers/openapi/test_app_run.py @@ -1,4 +1,4 @@ -"""Integration tests for POST /openapi/v1/apps//run.""" +"""Integration tests for POST /openapi/v1/apps/:run.""" from __future__ import annotations @@ -36,7 +36,7 @@ def test_run_chat_dispatches_to_chat_handler( monkeypatch.setattr("controllers.openapi.app_run.AppGenerateService.generate", staticmethod(_fake_generate)) client = flask_app.test_client() res = client.post( - f"/openapi/v1/apps/{app_in_workspace.id}/run", + f"/openapi/v1/apps/{app_in_workspace.id}:run", json={"inputs": {}, "query": "hi", "response_mode": "blocking", "user": "spoof@x.com"}, headers={"Authorization": f"Bearer {account_token}"}, ) @@ -85,7 +85,7 @@ def test_run_chat_without_query_returns_422( ): client = flask_app.test_client() res = client.post( - f"/openapi/v1/apps/{app_in_workspace.id}/run", + f"/openapi/v1/apps/{app_in_workspace.id}:run", json={"inputs": {}, "response_mode": "blocking"}, headers={"Authorization": f"Bearer {account_token}"}, ) @@ -116,7 +116,7 @@ def test_run_completion_dispatches_to_completion_handler( monkeypatch.setattr("controllers.openapi.app_run.AppGenerateService.generate", staticmethod(_fake_generate)) client = flask_app.test_client() res = client.post( - f"/openapi/v1/apps/{app.id}/run", + f"/openapi/v1/apps/{app.id}:run", json={"inputs": {}, "response_mode": "blocking"}, headers={"Authorization": f"Bearer {account_token}"}, ) @@ -131,7 +131,7 @@ def test_run_workflow_with_query_returns_422( app = app_with_mode("workflow") client = flask_app.test_client() res = client.post( - f"/openapi/v1/apps/{app.id}/run", + f"/openapi/v1/apps/{app.id}:run", json={"inputs": {}, "query": "hi", "response_mode": "blocking"}, headers={"Authorization": f"Bearer {account_token}"}, ) @@ -154,7 +154,7 @@ def test_run_workflow_no_query_dispatches_to_workflow_handler( monkeypatch.setattr("controllers.openapi.app_run.AppGenerateService.generate", staticmethod(_fake_generate)) client = flask_app.test_client() res = client.post( - f"/openapi/v1/apps/{app.id}/run", + f"/openapi/v1/apps/{app.id}:run", json={"inputs": {}, "response_mode": "blocking"}, headers={"Authorization": f"Bearer {account_token}"}, ) @@ -170,7 +170,7 @@ def test_run_unsupported_mode_returns_422( app = app_with_mode("channel") client = flask_app.test_client() res = client.post( - f"/openapi/v1/apps/{app.id}/run", + f"/openapi/v1/apps/{app.id}:run", json={"inputs": {}, "response_mode": "blocking"}, headers={"Authorization": f"Bearer {account_token}"}, ) @@ -181,7 +181,7 @@ def test_run_unsupported_mode_returns_422( def test_run_without_bearer_returns_401(flask_app: Flask, app_in_workspace): client = flask_app.test_client() res = client.post( - f"/openapi/v1/apps/{app_in_workspace.id}/run", + f"/openapi/v1/apps/{app_in_workspace.id}:run", json={"inputs": {}, "query": "hi"}, ) assert res.status_code == 401 @@ -205,7 +205,7 @@ def test_run_with_insufficient_scope_returns_403( client = flask_app.test_client() res = client.post( - f"/openapi/v1/apps/{app_in_workspace.id}/run", + f"/openapi/v1/apps/{app_in_workspace.id}:run", json={"inputs": {}, "query": "hi"}, headers={"Authorization": f"Bearer {account_token}"}, ) @@ -215,7 +215,7 @@ def test_run_with_insufficient_scope_returns_403( def test_run_with_unknown_app_returns_404(flask_app: Flask, account_token): client = flask_app.test_client() res = client.post( - f"/openapi/v1/apps/{uuid.uuid4()}/run", + f"/openapi/v1/apps/{uuid.uuid4()}:run", json={"inputs": {}, "query": "hi"}, headers={"Authorization": f"Bearer {account_token}"}, ) @@ -235,7 +235,7 @@ def test_run_streaming_returns_event_stream( client = flask_app.test_client() res = client.post( - f"/openapi/v1/apps/{app_in_workspace.id}/run", + f"/openapi/v1/apps/{app_in_workspace.id}:run", json={"inputs": {}, "query": "hi", "response_mode": "streaming"}, headers={"Authorization": f"Bearer {account_token}"}, ) @@ -247,7 +247,7 @@ def test_run_streaming_returns_event_stream( def test_run_without_inputs_returns_422(flask_app: Flask, account_token, app_in_workspace): client = flask_app.test_client() res = client.post( - f"/openapi/v1/apps/{app_in_workspace.id}/run", + f"/openapi/v1/apps/{app_in_workspace.id}:run", json={"query": "hi"}, headers={"Authorization": f"Bearer {account_token}"}, ) diff --git a/api/tests/integration_tests/controllers/openapi/test_apps.py b/api/tests/integration_tests/controllers/openapi/test_apps.py index 20ac46fbbde..ab992d1c5eb 100644 --- a/api/tests/integration_tests/controllers/openapi/test_apps.py +++ b/api/tests/integration_tests/controllers/openapi/test_apps.py @@ -37,7 +37,7 @@ def test_apps_describe_returns_merged_shape( account_token: str, ): res = test_client.get( - f"/openapi/v1/apps/{app_in_workspace.id}/describe", + f"/openapi/v1/apps/{app_in_workspace.id}", headers={"Authorization": f"Bearer {account_token}"}, ) assert res.status_code == 200 @@ -53,7 +53,7 @@ def test_apps_describe_full_includes_input_schema( account_token: str, ): res = test_client.get( - f"/openapi/v1/apps/{app_in_workspace.id}/describe", + f"/openapi/v1/apps/{app_in_workspace.id}", headers={"Authorization": f"Bearer {account_token}"}, ) assert res.status_code == 200 @@ -70,7 +70,7 @@ def test_apps_describe_fields_info_only( account_token: str, ): res = test_client.get( - f"/openapi/v1/apps/{app_in_workspace.id}/describe?fields=info", + f"/openapi/v1/apps/{app_in_workspace.id}?fields=info", headers={"Authorization": f"Bearer {account_token}"}, ) assert res.status_code == 200 @@ -86,7 +86,7 @@ def test_apps_describe_fields_parameters_only( account_token: str, ): res = test_client.get( - f"/openapi/v1/apps/{app_in_workspace.id}/describe?fields=parameters", + f"/openapi/v1/apps/{app_in_workspace.id}?fields=parameters", headers={"Authorization": f"Bearer {account_token}"}, ) assert res.status_code == 200 @@ -102,7 +102,7 @@ def test_apps_describe_fields_input_schema_only( account_token: str, ): res = test_client.get( - f"/openapi/v1/apps/{app_in_workspace.id}/describe?fields=input_schema", + f"/openapi/v1/apps/{app_in_workspace.id}?fields=input_schema", headers={"Authorization": f"Bearer {account_token}"}, ) assert res.status_code == 200 @@ -118,7 +118,7 @@ def test_apps_describe_fields_combined( account_token: str, ): res = test_client.get( - f"/openapi/v1/apps/{app_in_workspace.id}/describe?fields=info,input_schema", + f"/openapi/v1/apps/{app_in_workspace.id}?fields=info,input_schema", headers={"Authorization": f"Bearer {account_token}"}, ) assert res.status_code == 200 @@ -134,7 +134,7 @@ def test_apps_describe_fields_unknown_returns_422( account_token: str, ): res = test_client.get( - f"/openapi/v1/apps/{app_in_workspace.id}/describe?fields=garbage", + f"/openapi/v1/apps/{app_in_workspace.id}?fields=garbage", headers={"Authorization": f"Bearer {account_token}"}, ) assert res.status_code == 422 @@ -146,7 +146,7 @@ def test_apps_describe_fields_extra_param_returns_422( account_token: str, ): res = test_client.get( - f"/openapi/v1/apps/{app_in_workspace.id}/describe?fields=info&page=1", + f"/openapi/v1/apps/{app_in_workspace.id}?fields=info&page=1", headers={"Authorization": f"Bearer {account_token}"}, ) assert res.status_code == 422 diff --git a/api/tests/test_containers_integration_tests/controllers/openapi/test_app_dsl.py b/api/tests/test_containers_integration_tests/controllers/openapi/test_app_dsl.py index 93e8927cfef..2b9feeede14 100644 --- a/api/tests/test_containers_integration_tests/controllers/openapi/test_app_dsl.py +++ b/api/tests/test_containers_integration_tests/controllers/openapi/test_app_dsl.py @@ -167,7 +167,7 @@ class TestDslImportConfirm: api = AppDslImportConfirmApi() with app.test_request_context( - f"/openapi/v1/workspaces/{tenant.id}/apps/imports/{import_id}/confirm", method="POST" + f"/openapi/v1/workspaces/{tenant.id}/apps/imports/{import_id}:confirm", method="POST" ): result, code = unwrap(api.post)( api, workspace_id=tenant.id, import_id=import_id, auth_data=auth_for(account) @@ -198,7 +198,7 @@ class TestDslExport: db_session_with_containers.commit() api = AppDslExportApi() - with app.test_request_context(f"/openapi/v1/apps/{app_model.id}/export"): + with app.test_request_context(f"/openapi/v1/apps/{app_model.id}/dsl"): response, code = unwrap(api.get)( api, app_id=app_model.id, auth_data=auth_for(account, app_model=app_model), query=AppDslExportQuery() ) @@ -216,7 +216,7 @@ class TestDslExport: app_model, account = _app_and_account(db_session_with_containers, mode="workflow") api = AppDslExportApi() - with app.test_request_context(f"/openapi/v1/apps/{app_model.id}/export"): + with app.test_request_context(f"/openapi/v1/apps/{app_model.id}/dsl"): result, code = unwrap(api.get)( api, app_id=app_model.id, auth_data=auth_for(account, app_model=app_model), query=AppDslExportQuery() ) @@ -232,7 +232,7 @@ class TestDslCheckDependencies: app_model, account = _app_and_account(db_session_with_containers, mode="chat") api = AppDslCheckDependenciesApi() - with app.test_request_context(f"/openapi/v1/apps/{app_model.id}/check-dependencies"): + with app.test_request_context(f"/openapi/v1/apps/{app_model.id}/dependencies:check"): result, code = unwrap(api.get)(api, app_id=app_model.id, auth_data=auth_for(account, app_model=app_model)) assert code == 200 diff --git a/api/tests/test_containers_integration_tests/controllers/openapi/test_app_run.py b/api/tests/test_containers_integration_tests/controllers/openapi/test_app_run.py index c6fde623677..8e9278ad244 100644 --- a/api/tests/test_containers_integration_tests/controllers/openapi/test_app_run.py +++ b/api/tests/test_containers_integration_tests/controllers/openapi/test_app_run.py @@ -38,7 +38,7 @@ class TestAppRunTaskStop: task_id = str(uuid4()) api = AppRunTaskStopApi() - with app.test_request_context(f"/openapi/v1/apps/{app_model.id}/tasks/{task_id}/stop", method="POST"): + with app.test_request_context(f"/openapi/v1/apps/{app_model.id}/tasks/{task_id}:stop", method="POST"): result = unwrap(api.post)( api, app_id=app_model.id, diff --git a/api/tests/test_containers_integration_tests/controllers/openapi/test_apps.py b/api/tests/test_containers_integration_tests/controllers/openapi/test_apps.py index 22f812e125b..24580ae0a0e 100644 --- a/api/tests/test_containers_integration_tests/controllers/openapi/test_apps.py +++ b/api/tests/test_containers_integration_tests/controllers/openapi/test_apps.py @@ -120,7 +120,7 @@ class TestAppDescribe: app_model = _create_app(db_session_with_containers, account, name="Describe Me", enable_api=True) api = AppDescribeApi() - with app.test_request_context(f"/openapi/v1/apps/{app_model.id}/describe?fields=info"): + with app.test_request_context(f"/openapi/v1/apps/{app_model.id}?fields=info"): result = unwrap(api.get)( api, app_id=app_model.id, auth_data=auth_for(account), query=AppDescribeQuery(fields="info") ) @@ -138,7 +138,7 @@ class TestAppDescribe: missing_id = str(uuid4()) api = AppDescribeApi() - with app.test_request_context(f"/openapi/v1/apps/{missing_id}/describe"): + with app.test_request_context(f"/openapi/v1/apps/{missing_id}"): with pytest.raises(NotFound): unwrap(api.get)(api, app_id=missing_id, auth_data=auth_for(account), query=AppDescribeQuery()) @@ -151,6 +151,6 @@ class TestAppDescribe: hidden = _create_app(db_session_with_containers, account, name="Hidden", enable_api=False) api = AppDescribeApi() - with app.test_request_context(f"/openapi/v1/apps/{hidden.id}/describe"): + with app.test_request_context(f"/openapi/v1/apps/{hidden.id}"): with pytest.raises(NotFound): unwrap(api.get)(api, app_id=hidden.id, auth_data=auth_for(account), query=AppDescribeQuery()) diff --git a/api/tests/test_containers_integration_tests/controllers/openapi/test_files.py b/api/tests/test_containers_integration_tests/controllers/openapi/test_files.py index b90d5ab907c..86cf70613c9 100644 --- a/api/tests/test_containers_integration_tests/controllers/openapi/test_files.py +++ b/api/tests/test_containers_integration_tests/controllers/openapi/test_files.py @@ -43,7 +43,7 @@ class TestAppFileUpload: api = AppFileUploadApi() data = {"file": (BytesIO(content), "note.txt", "text/plain")} with app.test_request_context( - f"/openapi/v1/apps/{app_model.id}/files/upload", + f"/openapi/v1/apps/{app_model.id}/files", method="POST", data=data, content_type="multipart/form-data", diff --git a/api/tests/test_containers_integration_tests/controllers/openapi/test_workspaces.py b/api/tests/test_containers_integration_tests/controllers/openapi/test_workspaces.py index 18075704325..5e794ae1982 100644 --- a/api/tests/test_containers_integration_tests/controllers/openapi/test_workspaces.py +++ b/api/tests/test_containers_integration_tests/controllers/openapi/test_workspaces.py @@ -95,7 +95,7 @@ class TestWorkspaceSwitch: ) api = WorkspaceSwitchApi() - with app.test_request_context(f"/openapi/v1/workspaces/{target.id}/switch", method="POST"): + with app.test_request_context(f"/openapi/v1/workspaces/{target.id}:switch", method="POST"): detail = unwrap(api.post)(api, workspace_id=target.id, auth_data=auth_for(account)) # Response reflects the post-switch state. @@ -118,6 +118,6 @@ class TestWorkspaceSwitch: assert outsider_ws is not None api = WorkspaceSwitchApi() - with app.test_request_context(f"/openapi/v1/workspaces/{outsider_ws.id}/switch", method="POST"): + with app.test_request_context(f"/openapi/v1/workspaces/{outsider_ws.id}:switch", method="POST"): with pytest.raises(NotFound): unwrap(api.post)(api, workspace_id=outsider_ws.id, auth_data=auth_for(account)) diff --git a/api/tests/unit_tests/controllers/openapi/test_app_run_streaming.py b/api/tests/unit_tests/controllers/openapi/test_app_run_streaming.py index b82ab254d45..7672e6414fd 100644 --- a/api/tests/unit_tests/controllers/openapi/test_app_run_streaming.py +++ b/api/tests/unit_tests/controllers/openapi/test_app_run_streaming.py @@ -72,7 +72,7 @@ def test_run_chat_always_calls_generate_with_streaming_true( "AppGenerateService", GenerateService, ) - with app.test_request_context(f"/openapi/v1/apps/{_TEST_APP_ID}/run", method="POST"): + with app.test_request_context(f"/openapi/v1/apps/{_TEST_APP_ID}:run", method="POST"): _run_chat( _make_app(), _make_account(), @@ -84,9 +84,9 @@ def test_run_chat_always_calls_generate_with_streaming_true( def test_stop_task_endpoint_registered(openapi_app): - """POST /openapi/v1/apps//tasks//stop must be registered.""" + """POST /openapi/v1/apps//tasks/:stop must be registered.""" rules = {r.rule for r in openapi_app.url_map.iter_rules()} - assert "/openapi/v1/apps//tasks//stop" in rules + assert "/openapi/v1/apps//tasks/:stop" in rules def test_stop_task_calls_queue_manager_and_graph_engine(app: Flask, bypass_pipeline, monkeypatch: pytest.MonkeyPatch): @@ -117,7 +117,7 @@ def test_stop_task_calls_queue_manager_and_graph_engine(app: Flask, bypass_pipel ) api = AppRunTaskStopApi() - with app.test_request_context("/openapi/v1/apps/app-1/tasks/task-1/stop", method="POST"): + with app.test_request_context("/openapi/v1/apps/app-1/tasks/task-1:stop", method="POST"): result = api.post.__wrapped__( api, app_id="app-1", diff --git a/api/tests/unit_tests/controllers/openapi/test_human_input_form.py b/api/tests/unit_tests/controllers/openapi/test_human_input_form.py index 5659cd6eeff..c4d3d21d84a 100644 --- a/api/tests/unit_tests/controllers/openapi/test_human_input_form.py +++ b/api/tests/unit_tests/controllers/openapi/test_human_input_form.py @@ -62,7 +62,7 @@ class TestOpenApiHumanInputFormGet: app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1") caller = SimpleNamespace(id="acct-1") - with app.test_request_context("/openapi/v1/apps/app-1/form/human_input/tok-1"): + with app.test_request_context("/openapi/v1/apps/app-1/human-input-forms/tok-1"): resp = api.get.__wrapped__( api, app_id="app-1", @@ -89,7 +89,7 @@ class TestOpenApiHumanInputFormGet: app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1") caller = SimpleNamespace(id="acct-1") - with app.test_request_context("/openapi/v1/apps/app-1/form/human_input/bad"): + with app.test_request_context("/openapi/v1/apps/app-1/human-input-forms/bad"): with pytest.raises(HumanInputFormNotFound): api.get.__wrapped__( api, @@ -117,7 +117,7 @@ class TestOpenApiHumanInputFormGet: app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1") caller = SimpleNamespace(id="acct-1") - with app.test_request_context("/openapi/v1/apps/app-1/form/human_input/tok-1"): + with app.test_request_context("/openapi/v1/apps/app-1/human-input-forms/tok-1"): with pytest.raises(HumanInputFormNotFound): api.get.__wrapped__( api, @@ -145,7 +145,7 @@ class TestOpenApiHumanInputFormGet: app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1") caller = SimpleNamespace(id="acct-1") - with app.test_request_context("/openapi/v1/apps/app-1/form/human_input/tok-1"): + with app.test_request_context("/openapi/v1/apps/app-1/human-input-forms/tok-1"): with pytest.raises(RecipientSurfaceMismatch): api.get.__wrapped__( api, @@ -165,7 +165,7 @@ class TestOpenApiHumanInputFormPost: ) def test_post_account_caller_uses_user_id(self, app: Flask, bypass_pipeline, monkeypatch: pytest.MonkeyPatch): - from controllers.openapi.human_input_form import OpenApiWorkflowHumanInputFormApi + from controllers.openapi.human_input_form import OpenApiWorkflowHumanInputFormSubmitApi form = self._make_form() service_mock = Mock() @@ -175,12 +175,12 @@ class TestOpenApiHumanInputFormPost: monkeypatch.setattr(module, "HumanInputService", lambda _engine: service_mock) monkeypatch.setattr(module, "db", SimpleNamespace(engine=object())) - api = OpenApiWorkflowHumanInputFormApi() + api = OpenApiWorkflowHumanInputFormSubmitApi() app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1") caller = SimpleNamespace(id="acct-42") with app.test_request_context( - "/openapi/v1/apps/app-1/form/human_input/tok-1", + "/openapi/v1/apps/app-1/human-input-forms/tok-1:submit", method="POST", json={"action": "approve", "inputs": {"field1": "val"}}, ): @@ -202,7 +202,7 @@ class TestOpenApiHumanInputFormPost: assert result == ({}, 200) def test_post_end_user_caller_uses_end_user_id(self, app: Flask, bypass_pipeline, monkeypatch: pytest.MonkeyPatch): - from controllers.openapi.human_input_form import OpenApiWorkflowHumanInputFormApi + from controllers.openapi.human_input_form import OpenApiWorkflowHumanInputFormSubmitApi form = self._make_form() service_mock = Mock() @@ -212,12 +212,12 @@ class TestOpenApiHumanInputFormPost: monkeypatch.setattr(module, "HumanInputService", lambda _engine: service_mock) monkeypatch.setattr(module, "db", SimpleNamespace(engine=object())) - api = OpenApiWorkflowHumanInputFormApi() + api = OpenApiWorkflowHumanInputFormSubmitApi() app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1") caller = SimpleNamespace(id="eu-7") with app.test_request_context( - "/openapi/v1/apps/app-1/form/human_input/tok-1", + "/openapi/v1/apps/app-1/human-input-forms/tok-1:submit", method="POST", json={"action": "approve", "inputs": {}}, ): @@ -241,7 +241,7 @@ class TestOpenApiHumanInputFormPost: def test_post_standalone_web_app_recipient_submits( self, app: Flask, bypass_pipeline, monkeypatch: pytest.MonkeyPatch ): - from controllers.openapi.human_input_form import OpenApiWorkflowHumanInputFormApi + from controllers.openapi.human_input_form import OpenApiWorkflowHumanInputFormSubmitApi form = self._make_form(recipient_type=RecipientType.STANDALONE_WEB_APP) service_mock = Mock() @@ -251,12 +251,12 @@ class TestOpenApiHumanInputFormPost: monkeypatch.setattr(module, "HumanInputService", lambda _engine: service_mock) monkeypatch.setattr(module, "db", SimpleNamespace(engine=object())) - api = OpenApiWorkflowHumanInputFormApi() + api = OpenApiWorkflowHumanInputFormSubmitApi() app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1") caller = SimpleNamespace(id="anyone") with app.test_request_context( - "/openapi/v1/apps/app-1/form/human_input/tok-1", + "/openapi/v1/apps/app-1/human-input-forms/tok-1:submit", method="POST", json={"action": "approve", "inputs": {}}, ): @@ -272,14 +272,14 @@ class TestOpenApiHumanInputFormPost: def test_post_rejects_invalid_body_with_422(self, app: Flask, bypass_pipeline): """Malformed body → 422 via @accepts (was an unmapped pydantic error → 500).""" - from controllers.openapi.human_input_form import OpenApiWorkflowHumanInputFormApi + from controllers.openapi.human_input_form import OpenApiWorkflowHumanInputFormSubmitApi - api = OpenApiWorkflowHumanInputFormApi() + api = OpenApiWorkflowHumanInputFormSubmitApi() app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1") caller = SimpleNamespace(id="acct-42") with app.test_request_context( - "/openapi/v1/apps/app-1/form/human_input/tok-1", + "/openapi/v1/apps/app-1/human-input-forms/tok-1:submit", method="POST", json={"inputs": {"field1": "val"}}, # missing required "action" ): diff --git a/api/tests/unit_tests/controllers/openapi/test_version_gate.py b/api/tests/unit_tests/controllers/openapi/test_version_gate.py new file mode 100644 index 00000000000..e2b16259952 --- /dev/null +++ b/api/tests/unit_tests/controllers/openapi/test_version_gate.py @@ -0,0 +1,100 @@ +"""Tests for the difyctl version gate on /openapi/v1 (HTTP 426 Upgrade Required). + +The gate is an app-level ``before_app_request`` hook: it must fire before routing, +so requests to *removed* paths (which no longer match a route) become 426 rather +than a bare 404. It reads the difyctl version from the User-Agent and fails open +for anything it can't confidently identify as an outdated difyctl. +""" + +from __future__ import annotations + +import uuid + +import pytest +from flask import Flask + +# Floor is [tool.dify] min_difyctl_version = "0.2.0". Comparison is on the numeric +# core (major.minor.patch), so 0.2.0-alpha passes (core 0.2.0 == floor) while +# 0.1.0 (core 0.1.0 < 0.2.0) is blocked. +OLD_UA = "difyctl/0.1.0 (darwin; arm64; stable)" +CURRENT_UA = "difyctl/0.2.0-alpha (darwin; arm64; stable)" + + +@pytest.fixture +def client(openapi_app: Flask): + return openapi_app.test_client() + + +def _gated_path() -> str: + """An existing, auth-guarded route on the surface (GET /apps/).""" + return f"/openapi/v1/apps/{uuid.uuid4()}" + + +class TestVersionGate: + def test_old_client_gets_426_with_upgrade_body(self, client): + res = client.get(_gated_path(), headers={"User-Agent": OLD_UA}) + + assert res.status_code == 426 + body = res.get_json() + assert body["code"] == "upgrade_required" + assert body["status"] == 426 + assert "0.1.0" in body["message"] + assert "0.2.0" in body["message"] + assert "docs.dify.ai" in body["hint"] + + def test_removed_old_path_gets_426_not_404(self, client): + # /apps//run was renamed to /apps/:run — the old path matches no + # route. The app-level gate must still turn it into 426, not a bare 404. + res = client.post( + f"/openapi/v1/apps/{uuid.uuid4()}/run", + headers={"User-Agent": OLD_UA}, + json={"inputs": {}}, + ) + + assert res.status_code == 426 + assert res.get_json()["code"] == "upgrade_required" + + def test_current_client_passes_gate(self, client): + # Gate passes → normal dispatch (auth rejects, never the gate's 426). + # 0.2.0-alpha == floor on the numeric core, so it passes despite the suffix. + res = client.get(_gated_path(), headers={"User-Agent": CURRENT_UA}) + + assert res.status_code != 426 + + def test_prerelease_at_floor_passes(self, client): + # Numeric-core comparison: a pre-release of the floor version (0.2.0-rc.1, + # core 0.2.0) passes, even though 0.2.0-rc.1 < 0.2.0 under naive ordering. + res = client.get(_gated_path(), headers={"User-Agent": "difyctl/0.2.0-rc.1 (darwin; arm64; rc)"}) + + assert res.status_code != 426 + + def test_prerelease_below_floor_gets_426(self, client): + # 0.1.9-rc.1 has core 0.1.9 < 0.2.0 floor → still blocked. + res = client.get(_gated_path(), headers={"User-Agent": "difyctl/0.1.9-rc.1 (darwin; arm64; rc)"}) + + assert res.status_code == 426 + + def test_non_difyctl_ua_passes(self, client): + res = client.get(_gated_path(), headers={"User-Agent": "curl/8.4.0"}) + + assert res.status_code != 426 + + def test_missing_ua_passes(self, client): + res = client.get(_gated_path()) + + assert res.status_code != 426 + + def test_unparseable_version_passes(self, client): + res = client.get(_gated_path(), headers={"User-Agent": "difyctl/notaversion (x; y; z)"}) + + assert res.status_code != 426 + + def test_version_probe_allowlisted(self, client): + res = client.get("/openapi/v1/_version", headers={"User-Agent": OLD_UA}) + + assert res.status_code == 200 + + def test_health_allowlisted(self, client): + res = client.get("/openapi/v1/_health", headers={"User-Agent": OLD_UA}) + + assert res.status_code == 200 diff --git a/api/tests/unit_tests/controllers/openapi/test_workspaces_members.py b/api/tests/unit_tests/controllers/openapi/test_workspaces_members.py index cf9fa671987..b78473fadda 100644 --- a/api/tests/unit_tests/controllers/openapi/test_workspaces_members.py +++ b/api/tests/unit_tests/controllers/openapi/test_workspaces_members.py @@ -1,7 +1,7 @@ """Member endpoints under /openapi/v1/workspaces//... Coverage: -- Route registration (5 endpoints across 4 URL patterns) +- Route registration (5 endpoints across 3 URL patterns) - Body validation lands at 400 (per spec — not Pydantic's default 422) - Domain exception → HTTP code mapping is preserved with the service's original message (so CLI users see what the console user sees) @@ -37,7 +37,6 @@ from controllers.openapi._models import MemberInvitePayload, MemberRoleUpdatePay from controllers.openapi.auth.data import AuthData from controllers.openapi.workspaces import ( WorkspaceMemberApi, - WorkspaceMemberRoleApi, WorkspaceMembersApi, WorkspaceSwitchApi, ) @@ -175,7 +174,7 @@ def _account_service(**overrides) -> SimpleNamespace: def test_switch_route_registered(openapi_app: Flask): - rule = _rule(openapi_app, "/openapi/v1/workspaces//switch") + rule = _rule(openapi_app, "/openapi/v1/workspaces/:switch") assert openapi_app.view_functions[rule.endpoint].view_class is WorkspaceSwitchApi assert "POST" in rule.methods @@ -191,12 +190,7 @@ def test_member_by_id_route_registered(openapi_app: Flask): rule = _rule(openapi_app, "/openapi/v1/workspaces//members/") assert openapi_app.view_functions[rule.endpoint].view_class is WorkspaceMemberApi assert "DELETE" in rule.methods - - -def test_member_role_route_registered(openapi_app: Flask): - rule = _rule(openapi_app, "/openapi/v1/workspaces//members//role") - assert openapi_app.view_functions[rule.endpoint].view_class is WorkspaceMemberRoleApi - assert "PUT" in rule.methods + assert "PATCH" in rule.methods # --------------------------------------------------------------------------- @@ -250,17 +244,17 @@ def test_update_role_rejects_invalid_body_with_422(app: Flask, bypass_pipeline): """Invalid role-update body surfaces as 422 through @accepts (was 400).""" ws_id, member_id = str(uuid.uuid4()), str(uuid.uuid4()) acct_id = uuid.uuid4() - api = WorkspaceMemberRoleApi() + api = WorkspaceMemberApi() with app.test_request_context( - f"/openapi/v1/workspaces/{ws_id}/members/{member_id}/role", - method="PUT", + f"/openapi/v1/workspaces/{ws_id}/members/{member_id}", + method="PATCH", data=json.dumps({"role": "owner"}), # closed enum rejects owner content_type="application/json", ): _seed(_auth_ctx(account_id=acct_id)) with pytest.raises(UnprocessableEntity): - api.put.__wrapped__(api, workspace_id=ws_id, member_id=member_id, auth_data=_auth_data(acct_id)) + api.patch.__wrapped__(api, workspace_id=ws_id, member_id=member_id, auth_data=_auth_data(acct_id)) # --------------------------------------------------------------------------- @@ -291,7 +285,7 @@ def test_switch_returns_workspace_detail_with_current_true( ) monkeypatch.setattr(sys.modules["controllers.openapi.workspaces"], "db", mock_db) - with app.test_request_context(f"/openapi/v1/workspaces/{ws_id}/switch", method="POST"): + with app.test_request_context(f"/openapi/v1/workspaces/{ws_id}:switch", method="POST"): _seed(_auth_ctx(account_id=acct_id)) body, status = api.post.__wrapped__(api, workspace_id=ws_id, auth_data=_auth_data(acct_id)) @@ -320,7 +314,7 @@ def test_switch_404s_when_service_raises_account_not_link_tenant( ) monkeypatch.setattr(sys.modules["controllers.openapi.workspaces"], "db", mock_db) - with app.test_request_context(f"/openapi/v1/workspaces/{ws_id}/switch", method="POST"): + with app.test_request_context(f"/openapi/v1/workspaces/{ws_id}:switch", method="POST"): _seed(_auth_ctx(account_id=acct_id)) with pytest.raises(NotFound): api.post.__wrapped__(api, workspace_id=ws_id, auth_data=_auth_data(acct_id)) @@ -767,7 +761,7 @@ def test_delete_member_404_when_member_missing(app: Flask, bypass_pipeline, monk def test_update_role_happy_path(app: Flask, bypass_pipeline, monkeypatch: pytest.MonkeyPatch): ws_id, member_id = str(uuid.uuid4()), str(uuid.uuid4()) acct_id = uuid.uuid4() - api = WorkspaceMemberRoleApi() + api = WorkspaceMemberApi() mock_db = MagicMock() mock_db.session.get.side_effect = [ @@ -785,13 +779,15 @@ def test_update_role_happy_path(app: Flask, bypass_pipeline, monkeypatch: pytest monkeypatch.setattr(sys.modules["controllers.openapi.workspaces"], "db", mock_db) with app.test_request_context( - f"/openapi/v1/workspaces/{ws_id}/members/{member_id}/role", - method="PUT", + f"/openapi/v1/workspaces/{ws_id}/members/{member_id}", + method="PATCH", data=json.dumps({"role": "admin"}), content_type="application/json", ): _seed(_auth_ctx(account_id=acct_id)) - body, status = api.put.__wrapped__(api, workspace_id=ws_id, member_id=member_id, auth_data=_auth_data(acct_id)) + body, status = api.patch.__wrapped__( + api, workspace_id=ws_id, member_id=member_id, auth_data=_auth_data(acct_id) + ) assert status == 200 assert body == {"result": "success"} @@ -811,7 +807,7 @@ def test_update_role_happy_path(app: Flask, bypass_pipeline, monkeypatch: pytest def test_update_role_exception_mapping(app: Flask, bypass_pipeline, monkeypatch, exc, expected): ws_id, member_id = str(uuid.uuid4()), str(uuid.uuid4()) acct_id = uuid.uuid4() - api = WorkspaceMemberRoleApi() + api = WorkspaceMemberApi() mock_db = MagicMock() mock_db.session.get.side_effect = [ @@ -828,14 +824,14 @@ def test_update_role_exception_mapping(app: Flask, bypass_pipeline, monkeypatch, monkeypatch.setattr(sys.modules["controllers.openapi.workspaces"], "db", mock_db) with app.test_request_context( - f"/openapi/v1/workspaces/{ws_id}/members/{member_id}/role", - method="PUT", + f"/openapi/v1/workspaces/{ws_id}/members/{member_id}", + method="PATCH", data=json.dumps({"role": "admin"}), content_type="application/json", ): _seed(_auth_ctx(account_id=acct_id)) with pytest.raises(expected): - api.put.__wrapped__( + api.patch.__wrapped__( api, workspace_id=ws_id, member_id=member_id, diff --git a/cli/package.json b/cli/package.json index 5121daf3624..0e87f58ef2f 100644 --- a/cli/package.json +++ b/cli/package.json @@ -1,13 +1,13 @@ { "name": "@langgenius/difyctl", "type": "module", - "version": "0.1.0-alpha", + "version": "0.2.0-alpha", "description": "Dify command-line interface", "difyctl": { "channel": "alpha", "compat": { - "minDify": "1.15.0", - "maxDify": "1.15.0" + "minDify": "1.16.0", + "maxDify": "1.16.0" }, "release": { "tagPrefix": "difyctl-v", diff --git a/cli/scripts/release-naming.test.ts b/cli/scripts/release-naming.test.ts index 559d15ad759..7641b531282 100644 --- a/cli/scripts/release-naming.test.ts +++ b/cli/scripts/release-naming.test.ts @@ -15,41 +15,41 @@ function run(args: string[]): { code: number, stdout: string, stderr: string } { } } -describe('release-naming compat-check (compat 1.15.0..1.15.0)', () => { +describe('release-naming compat-check (compat 1.16.0..1.16.0)', () => { it('accepts a version inside the window', () => { - expect(run(['compat-check', '1.15.0']).code).toBe(0) + expect(run(['compat-check', '1.16.0']).code).toBe(0) }) it('accepts the inclusive lower bound', () => { - expect(run(['compat-check', '1.15.0']).code).toBe(0) + expect(run(['compat-check', '1.16.0']).code).toBe(0) }) it('accepts the inclusive upper bound', () => { - expect(run(['compat-check', '1.15.0']).code).toBe(0) + expect(run(['compat-check', '1.16.0']).code).toBe(0) }) it('accepts a v-prefixed tag', () => { - expect(run(['compat-check', 'v1.15.0']).code).toBe(0) + expect(run(['compat-check', 'v1.16.0']).code).toBe(0) }) it('rejects a version below the lower bound', () => { - expect(run(['compat-check', '1.14.9']).code).not.toBe(0) + expect(run(['compat-check', '1.15.9']).code).not.toBe(0) }) it('rejects a version above the upper bound', () => { - expect(run(['compat-check', '1.15.1']).code).not.toBe(0) + expect(run(['compat-check', '1.16.1']).code).not.toBe(0) }) - it('treats a prerelease of the bound as below it (1.15.0-rc1 < 1.15.0)', () => { - expect(run(['compat-check', '1.15.0-rc1']).code).not.toBe(0) + it('treats a prerelease of the bound as below it (1.16.0-rc1 < 1.16.0)', () => { + expect(run(['compat-check', '1.16.0-rc1']).code).not.toBe(0) }) - it('ignores build metadata on the bound (1.15.0+build == 1.15.0)', () => { - expect(run(['compat-check', '1.15.0+build123']).code).toBe(0) + it('ignores build metadata on the bound (1.16.0+build == 1.16.0)', () => { + expect(run(['compat-check', '1.16.0+build123']).code).toBe(0) }) - it('ignores build metadata when out of range (1.15.1+build still rejected)', () => { - expect(run(['compat-check', '1.15.1+build123']).code).not.toBe(0) + it('ignores build metadata when out of range (1.16.1+build still rejected)', () => { + expect(run(['compat-check', '1.16.1+build123']).code).not.toBe(0) }) it('requires a version argument', () => { @@ -60,7 +60,7 @@ describe('release-naming compat-check (compat 1.15.0..1.15.0)', () => { describe('release-naming github-env', () => { it('emits difyctlTag = tagPrefix + version', () => { const { stdout } = run(['github-env']) - expect(stdout).toMatch(/^difyctlTag=difyctl-v0\.1\.0-alpha$/m) + expect(stdout).toMatch(/^difyctlTag=difyctl-v0\.2\.0-alpha$/m) }) it('still emits the existing trace fields', () => { @@ -75,14 +75,14 @@ describe('release-naming edge channel', () => { expect(run(['channels']).stdout).toMatch(/^edge$/m) }) - it('edge-version derives -edge. stripping the alpha prerelease', () => { - // package.json version is 0.1.0-alpha -> core 0.1.0 - expect(run(['edge-version', '2fd7b82']).stdout.trim()).toBe('0.1.0-edge.2fd7b82') + it('edge-version derives -edge. from the package version', () => { + // package.json version is 0.2.0-alpha -> core 0.2.0 + expect(run(['edge-version', '2fd7b82']).stdout.trim()).toBe('0.2.0-edge.2fd7b82') }) it('edge-version accepts a 40-char sha', () => { const sha = '2fd7b829e1f0aaaabbbbccccddddeeeeffff0000' - expect(run(['edge-version', sha]).stdout.trim()).toBe(`0.1.0-edge.${sha}`) + expect(run(['edge-version', sha]).stdout.trim()).toBe(`0.2.0-edge.${sha}`) }) it('edge-version rejects a non-hex sha', () => { diff --git a/cli/scripts/release-r2-edge.test.ts b/cli/scripts/release-r2-edge.test.ts index 12e3ee1822c..9db0c0ad2cb 100644 --- a/cli/scripts/release-r2-edge.test.ts +++ b/cli/scripts/release-r2-edge.test.ts @@ -81,7 +81,7 @@ describe('release-r2-edge manifest', () => { it('carries the compat window from package.json', () => { const { json } = buildManifest() - expect(json.compat).toEqual({ minDify: '1.15.0', maxDify: '1.15.0' }) + expect(json.compat).toEqual({ minDify: '1.16.0', maxDify: '1.16.0' }) }) it('lists all 5 targets with asset name + sha256 from the checksums file', () => { diff --git a/cli/src/api/app-dsl.test.ts b/cli/src/api/app-dsl.test.ts index 66d4507317e..1afff1fcbce 100644 --- a/cli/src/api/app-dsl.test.ts +++ b/cli/src/api/app-dsl.test.ts @@ -26,7 +26,7 @@ describe('AppDslClient.exportDsl', () => { const yaml = await makeClient(stub.url).exportDsl('app-1') expect(stub.captured.method).toBe('GET') - expect(stub.captured.url?.split('?')[0]).toBe('/openapi/v1/apps/app-1/export') + expect(stub.captured.url?.split('?')[0]).toBe('/openapi/v1/apps/app-1/dsl') expect(yaml).toBe(DSL_YAML) }) @@ -90,7 +90,7 @@ describe('AppDslClient.confirmImport', () => { const result = await makeClient(stub.url).confirmImport('ws-1', 'imp-1') expect(stub.captured.method).toBe('POST') - expect(stub.captured.url).toBe('/openapi/v1/workspaces/ws-1/apps/imports/imp-1/confirm') + expect(stub.captured.url).toBe('/openapi/v1/workspaces/ws-1/apps/imports/imp-1:confirm') expect(result.status).toBe('completed') }) }) @@ -107,7 +107,7 @@ describe('AppDslClient.checkDependencies', () => { const result = await makeClient(stub.url).checkDependencies('app-1') - expect(stub.captured.url?.split('?')[0]).toBe('/openapi/v1/apps/app-1/check-dependencies') + expect(stub.captured.url?.split('?')[0]).toBe('/openapi/v1/apps/app-1/dependencies:check') expect(result.leaked_dependencies).toEqual([]) }) }) diff --git a/cli/src/api/app-dsl.ts b/cli/src/api/app-dsl.ts index acb10a95351..19c26f5d7fb 100644 --- a/cli/src/api/app-dsl.ts +++ b/cli/src/api/app-dsl.ts @@ -33,7 +33,7 @@ export class AppDslClient { } async exportDsl(appId: string, query?: ExportQuery): Promise { - const resp = await this.orpc.apps.byAppId.export.get({ + const resp = await this.orpc.apps.byAppId.dsl.get({ params: { app_id: appId }, query: query !== undefined ? { @@ -52,7 +52,7 @@ export class AppDslClient { } async checkDependencies(appId: string): Promise { - return this.orpc.apps.byAppId.checkDependencies.get({ + return this.orpc.apps.byAppId.dependencies.check.get({ params: { app_id: appId }, }) } diff --git a/cli/src/api/app-run.ts b/cli/src/api/app-run.ts index cbd36de2049..a75b3f0c63a 100644 --- a/cli/src/api/app-run.ts +++ b/cli/src/api/app-run.ts @@ -54,7 +54,7 @@ export class AppRunClient { body: Record, opts: StreamOptions = {}, ): Promise> { - const res = await this.http.stream(`apps/${encodeURIComponent(appId)}/run`, { + const res = await this.http.stream(`apps/${encodeURIComponent(appId)}:run`, { method: 'POST', json: body, headers: { Accept: 'text/event-stream' }, @@ -79,7 +79,7 @@ export class AppRunClient { action: string, inputs: Record, ): Promise { - await this.orpc.apps.byAppId.form.humanInput.byFormToken.post({ + await this.orpc.apps.byAppId.humanInputForms.byFormToken.submit.post({ params: { app_id: appId, form_token: formToken }, body: { action, inputs }, }) diff --git a/cli/src/api/apps.test.ts b/cli/src/api/apps.test.ts index 861f60feb26..ba9126b7b98 100644 --- a/cli/src/api/apps.test.ts +++ b/cli/src/api/apps.test.ts @@ -82,12 +82,12 @@ describe('AppsClient.describe', () => { await stub?.stop() }) - it('hits /apps//describe, omits workspace_id and fields when not given', async () => { + it('hits /apps/, omits workspace_id and fields when not given', async () => { stub = await startStubServer(cap => jsonResponder(200, DESCRIBE_BODY, cap)) const res = await makeClient(stub.url).describe('app-1') - expect(stub.captured.url?.split('?')[0]).toBe('/openapi/v1/apps/app-1/describe') + expect(stub.captured.url?.split('?')[0]).toBe('/openapi/v1/apps/app-1') const q = queryOf(stub.captured.url) expect(q.has('workspace_id')).toBe(false) expect(q.has('fields')).toBe(false) @@ -107,6 +107,6 @@ describe('AppsClient.describe', () => { await makeClient(stub.url).describe('app/with space') - expect(stub.captured.url?.split('?')[0]).toBe('/openapi/v1/apps/app%2Fwith%20space/describe') + expect(stub.captured.url?.split('?')[0]).toBe('/openapi/v1/apps/app%2Fwith%20space') }) }) diff --git a/cli/src/api/apps.ts b/cli/src/api/apps.ts index 1189fdeaa06..fa29d66fd14 100644 --- a/cli/src/api/apps.ts +++ b/cli/src/api/apps.ts @@ -37,7 +37,7 @@ export class AppsClient implements AppReader { } async describe(appId: string, fields?: readonly string[]): Promise { - return this.orpc.apps.byAppId.describe.get({ + return this.orpc.apps.byAppId.get({ params: { app_id: appId }, query: { fields: fields !== undefined && fields.length > 0 ? fields.join(',') : undefined, diff --git a/cli/src/api/file-upload.test.ts b/cli/src/api/file-upload.test.ts index 018389916b0..602cf351e60 100644 --- a/cli/src/api/file-upload.test.ts +++ b/cli/src/api/file-upload.test.ts @@ -41,7 +41,7 @@ describe('FileUploadClient.upload', () => { const result = await makeClient(stub.url).upload('app-1', filePath) expect(stub.captured.method).toBe('POST') - expect(stub.captured.url).toBe('/openapi/v1/apps/app-1/files/upload') + expect(stub.captured.url).toBe('/openapi/v1/apps/app-1/files') // The client must let fetch own the multipart Content-Type + boundary; it // must NOT coerce this to application/json the way a json body would. const contentType = stub.captured.headers?.['content-type'] ?? '' @@ -61,7 +61,7 @@ describe('FileUploadClient.upload', () => { await makeClient(stub.url).upload('app/with space', filePath) - expect(stub.captured.url).toBe('/openapi/v1/apps/app%2Fwith%20space/files/upload') + expect(stub.captured.url).toBe('/openapi/v1/apps/app%2Fwith%20space/files') }) it('propagates a server 413 as a classified BaseError', async () => { diff --git a/cli/src/api/file-upload.ts b/cli/src/api/file-upload.ts index 7a032737a59..011f898c74e 100644 --- a/cli/src/api/file-upload.ts +++ b/cli/src/api/file-upload.ts @@ -65,7 +65,7 @@ export class FileUploadClient { form.append('file', blob, filename) return this.http.post( - `apps/${encodeURIComponent(appId)}/files/upload`, + `apps/${encodeURIComponent(appId)}/files`, { body: form, timeoutMs: 60_000 }, ) } diff --git a/cli/src/api/members.test.ts b/cli/src/api/members.test.ts index b4e01b76b24..a8fdc633f77 100644 --- a/cli/src/api/members.test.ts +++ b/cli/src/api/members.test.ts @@ -154,13 +154,13 @@ describe('MembersClient.updateRole', () => { await stub?.stop() }) - it('PUTs role payload to /role subresource', async () => { + it('PATCHes role payload to the member resource', async () => { stub = await startStubServer(cap => jsonResponder(200, { result: 'success' }, cap)) const result = await makeClient(stub.url).updateRole('ws-1', 'm-1', { role: 'admin' }) - expect(stub.captured.method).toBe('PUT') - expect(stub.captured.url).toBe('/openapi/v1/workspaces/ws-1/members/m-1/role') + expect(stub.captured.method).toBe('PATCH') + expect(stub.captured.url).toBe('/openapi/v1/workspaces/ws-1/members/m-1') expect(JSON.parse(stub.captured.body ?? '{}')).toEqual({ role: 'admin' }) expect(result.result).toBe('success') }) @@ -181,7 +181,7 @@ describe('WorkspacesClient.switch (integration with stub)', () => { await stub?.stop() }) - it('POSTs /workspaces//switch and returns workspace detail', async () => { + it('POSTs /workspaces/:switch and returns workspace detail', async () => { stub = await startStubServer(cap => jsonResponder( 200, @@ -200,7 +200,7 @@ describe('WorkspacesClient.switch (integration with stub)', () => { const result = await client.switch('ws-1') expect(stub.captured.method).toBe('POST') - expect(stub.captured.url).toBe('/openapi/v1/workspaces/ws-1/switch') + expect(stub.captured.url).toBe('/openapi/v1/workspaces/ws-1:switch') expect(result.current).toBe(true) }) diff --git a/cli/src/api/members.ts b/cli/src/api/members.ts index 7b1f80c08ff..8a9bc13081f 100644 --- a/cli/src/api/members.ts +++ b/cli/src/api/members.ts @@ -47,7 +47,7 @@ export class MembersClient { memberId: string, payload: MemberRoleUpdatePayload, ): Promise { - return this.orpc.workspaces.byWorkspaceId.members.byMemberId.role.put({ + return this.orpc.workspaces.byWorkspaceId.members.byMemberId.patch({ params: { workspace_id: workspaceId, member_id: memberId }, body: payload, }) diff --git a/cli/src/api/permitted-external-apps.test.ts b/cli/src/api/permitted-external-apps.test.ts index f6fa38cb3eb..58f47b4566f 100644 --- a/cli/src/api/permitted-external-apps.test.ts +++ b/cli/src/api/permitted-external-apps.test.ts @@ -12,15 +12,15 @@ describe('PermittedExternalAppsClient', () => { it('list calls permittedExternalApps.get with paging/filter query', async () => { const c = new PermittedExternalAppsClient(fakeHttp()) const get = vi.fn().mockResolvedValue({ page: 1, limit: 20, total: 0, has_more: false, data: [] }) - ;(c as unknown as WithOrpc).orpc = { permittedExternalApps: { get, byAppId: { describe: { get: vi.fn() } } } } + ;(c as unknown as WithOrpc).orpc = { permittedExternalApps: { get, byAppId: { get: vi.fn() } } } await c.list({ workspaceId: '', page: 2, limit: 5, mode: undefined, name: 'a' }) expect(get).toHaveBeenCalledWith({ query: { page: 2, limit: 5, mode: undefined, name: 'a' } }) }) - it('describe calls permittedExternalApps.byAppId.describe.get with app_id + fields', async () => { + it('describe calls permittedExternalApps.byAppId.get with app_id + fields', async () => { const c = new PermittedExternalAppsClient(fakeHttp()) const dget = vi.fn().mockResolvedValue({ info: null, parameters: null, input_schema: null }) - ;(c as unknown as WithOrpc).orpc = { permittedExternalApps: { get: vi.fn(), byAppId: { describe: { get: dget } } } } + ;(c as unknown as WithOrpc).orpc = { permittedExternalApps: { get: vi.fn(), byAppId: { get: dget } } } await c.describe('app-1', ['info']) expect(dget).toHaveBeenCalledWith({ params: { app_id: 'app-1' }, query: { fields: 'info' } }) }) diff --git a/cli/src/api/permitted-external-apps.ts b/cli/src/api/permitted-external-apps.ts index 497c398d0ba..c0164d3c536 100644 --- a/cli/src/api/permitted-external-apps.ts +++ b/cli/src/api/permitted-external-apps.ts @@ -26,7 +26,7 @@ export class PermittedExternalAppsClient implements AppReader { } async describe(appId: string, fields?: readonly string[]): Promise { - return this.orpc.permittedExternalApps.byAppId.describe.get({ + return this.orpc.permittedExternalApps.byAppId.get({ params: { app_id: appId }, query: { fields: fields !== undefined && fields.length > 0 ? fields.join(',') : undefined }, }) diff --git a/cli/src/api/workspaces.ts b/cli/src/api/workspaces.ts index 3ef587b574d..08495a01fce 100644 --- a/cli/src/api/workspaces.ts +++ b/cli/src/api/workspaces.ts @@ -19,7 +19,7 @@ export class WorkspacesClient { /** * Server-side workspace switch via OpenAPI POST - * `/workspaces/{id}/switch` — the bearer-authed equivalent of the + * `/workspaces/{id}:switch` — the bearer-authed equivalent of the * console's POST `/workspaces/switch`. The server updates the caller's * `current` tenant_account_join row. Callers MUST refresh their local * `hosts.yml` only after this resolves — never fall back to a local diff --git a/cli/src/cache/compat-store.test.ts b/cli/src/cache/compat-store.test.ts new file mode 100644 index 00000000000..25ad11db06f --- /dev/null +++ b/cli/src/cache/compat-store.test.ts @@ -0,0 +1,57 @@ +import { mkdtemp, rm } from 'node:fs/promises' +import { tmpdir } from 'node:os' +import { join } from 'node:path' +import { afterEach, beforeEach, describe, expect, it } from 'vitest' +import { ENV_CACHE_DIR } from '@/store/dir' +import { CACHE_COMPAT, getCache } from '@/store/manager' +import { loadCompatStore } from './compat-store' + +const HOST = 'https://cloud.dify.ai' +const NOW = new Date('2026-05-20T12:00:00.000Z') + +describe('compat-store', () => { + let dir: string + let prev: string | undefined + + beforeEach(async () => { + dir = await mkdtemp(join(tmpdir(), 'difyctl-compat-')) + prev = process.env[ENV_CACHE_DIR] + process.env[ENV_CACHE_DIR] = dir + }) + afterEach(async () => { + if (prev === undefined) + delete process.env[ENV_CACHE_DIR] + else + process.env[ENV_CACHE_DIR] = prev + await rm(dir, { recursive: true, force: true }) + }) + + const store = (now: Date = NOW) => loadCompatStore({ store: getCache(CACHE_COMPAT), now: () => now }) + + it('is not fresh before anything is marked', async () => { + expect((await store()).isFreshCompatible(HOST)).toBe(false) + }) + + it('is fresh right after markCompatible, and persists across loads', async () => { + await (await store()).markCompatible(HOST) + expect((await store()).isFreshCompatible(HOST)).toBe(true) + }) + + it('stays fresh within the 1h TTL', async () => { + const past = new Date(NOW.getTime() - 30 * 60 * 1000) + await (await store(past)).markCompatible(HOST) + expect((await store(NOW)).isFreshCompatible(HOST)).toBe(true) + }) + + it('expires after the 1h TTL', async () => { + const past = new Date(NOW.getTime() - 61 * 60 * 1000) + await (await store(past)).markCompatible(HOST) + expect((await store(NOW)).isFreshCompatible(HOST)).toBe(false) + }) + + it('tracks hosts independently', async () => { + const s = await store() + await s.markCompatible(HOST) + expect(s.isFreshCompatible('https://other.dify.ai')).toBe(false) + }) +}) diff --git a/cli/src/cache/compat-store.ts b/cli/src/cache/compat-store.ts new file mode 100644 index 00000000000..2df6dcf3377 --- /dev/null +++ b/cli/src/cache/compat-store.ts @@ -0,0 +1,71 @@ +import type { Store } from '@/store/store' +import { CACHE_COMPAT, getCache } from '@/store/manager' + +// How long a host stays "known compatible" before difyctl re-probes /_version. +export const COMPAT_TTL_MS = 60 * 60 * 1000 + +// Only *positive* (compatible) verdicts are cached — never "too old". A host that +// was too old is re-probed every time, so a just-upgraded server clears a previous +// block immediately instead of staying locked out for the whole TTL. +const COMPATIBLE_KEY = { key: 'compatible', default: {} as Record } as const + +export type CompatStore = { + readonly isFreshCompatible: (host: string, now?: Date) => boolean + readonly markCompatible: (host: string, now?: Date) => Promise +} + +export type CompatStoreOptions = { + readonly store?: Store + readonly now?: () => Date + readonly ttlMs?: number +} + +export async function loadCompatStore(opts: CompatStoreOptions = {}): Promise { + const store = opts.store ?? getCache(CACHE_COMPAT) + const ttlMs = opts.ttlMs ?? COMPAT_TTL_MS + const clock = opts.now ?? (() => new Date()) + const memory = await readCompatible(store) + + return { + isFreshCompatible: (host, now) => { + const last = memory.get(host) + if (last === undefined) + return false + const elapsed = Math.max(0, (now ?? clock()).getTime() - last) + return elapsed < ttlMs + }, + markCompatible: async (host, now) => { + const stamp = (now ?? clock()).getTime() + memory.set(host, stamp) + // Re-read disk inside the write cycle so concurrent processes touching + // different hosts don't clobber each other's stamps. + const onDisk = await readCompatible(store) + onDisk.set(host, stamp) + await writeCompatible(store, onDisk) + }, + } +} + +async function readCompatible(store: Store): Promise> { + const out = new Map() + let raw: Record + try { + raw = await store.get(COMPATIBLE_KEY) + } + catch { + return out + } + for (const [host, iso] of Object.entries(raw)) { + const t = Date.parse(iso) + if (!Number.isNaN(t)) + out.set(host, t) + } + return out +} + +async function writeCompatible(store: Store, state: Map): Promise { + const compatible: Record = {} + for (const [host, t] of state) + compatible[host] = new Date(t).toISOString() + await store.set(COMPATIBLE_KEY, compatible) +} diff --git a/cli/src/commands/_shared/authed-command.ts b/cli/src/commands/_shared/authed-command.ts index 8ef0b381e9e..1b0d6803180 100644 --- a/cli/src/commands/_shared/authed-command.ts +++ b/cli/src/commands/_shared/authed-command.ts @@ -14,6 +14,7 @@ import { createHttpClient } from '@/http/client' import { getTokenStore } from '@/store/manager' import { realStreams } from '@/sys/io/streams' import { hostWithScheme, openAPIBase } from '@/util/host' +import { enforceDifyVersion } from '@/version/enforce' import { versionInfo } from '@/version/info' import { maybeNudgeCompat } from '@/version/nudge' import { resolveRetryAttempts } from './global-flags.js' @@ -55,6 +56,10 @@ export async function buildAuthedContext( const cache = opts.withCache === true ? await loadAppInfoCache() : undefined + // Hard gate: refuse a server too old for this difyctl (throws → exit 6). + // Cached per host (1h) so most commands don't re-probe. Then the soft nudge + // handles the "server too new" direction. + await enforceDifyVersion(host) await runCompatNudge({ host, io }) return { reg, active, store, http, host, io, cache } diff --git a/cli/src/commands/auth/login/index.ts b/cli/src/commands/auth/login/index.ts index d9b6dd27fc3..214f38c3697 100644 --- a/cli/src/commands/auth/login/index.ts +++ b/cli/src/commands/auth/login/index.ts @@ -2,6 +2,7 @@ import type { CommandEffect } from '@/framework/command' import { DifyCommand } from '@/commands/_shared/dify-command' import { Flags } from '@/framework/flags' import { realStreams } from '@/sys/io/streams' +import { enforceDifyVersion } from '@/version/enforce' import { agentGuide } from './guide' import { runLogin } from './login' @@ -38,6 +39,9 @@ export default class Login extends DifyCommand { host: flags.host, noBrowser: flags['no-browser'], insecure: flags.insecure, + verifyServer: async (host) => { + await enforceDifyVersion(host, { forceFresh: true }) + }, }) } diff --git a/cli/src/commands/auth/login/login.ts b/cli/src/commands/auth/login/login.ts index 2c1ba5b95a9..f0eb14dc8e1 100644 --- a/cli/src/commands/auth/login/login.ts +++ b/cli/src/commands/auth/login/login.ts @@ -31,6 +31,10 @@ export type LoginOptions = { readonly browserEnv?: BrowserEnv readonly browserOpener?: BrowserOpener readonly clock?: Clock + // Version guard for the freshly-authenticated host; wired to enforceDifyVersion + // at the command boundary. Runs before the session is persisted so we never + // save credentials for a server too old for this difyctl. Defaults to a no-op. + readonly verifyServer?: (host: string) => Promise } export async function runLogin(opts: LoginOptions): Promise { @@ -70,6 +74,9 @@ export async function runLogin(opts: LoginOptions): Promise { spinner.stop() } + // Refuse to persist a session to a server too old for this difyctl. + await (opts.verifyServer ?? (async () => {}))(host) + const storeBundle = opts.store ?? await detectTokenStore() const display = bareHost(host) const email = accountEmail(success) diff --git a/cli/src/commands/resume/app/run.test.ts b/cli/src/commands/resume/app/run.test.ts index a72b0a93b25..8e1882906ae 100644 --- a/cli/src/commands/resume/app/run.test.ts +++ b/cli/src/commands/resume/app/run.test.ts @@ -41,7 +41,7 @@ describe('resumeApp pre-flight subject strategy', () => { const http = { baseURL: 'http://localhost', request: vi.fn().mockImplementation((opts: { path: string }) => { - if (typeof opts.path === 'string' && opts.path.includes('form/human_input')) { + if (typeof opts.path === 'string' && opts.path.includes('human-input-forms')) { return Promise.resolve(FORM_RESP) } // reconnect stream — return an async iterable that ends immediately diff --git a/cli/src/commands/resume/app/run.ts b/cli/src/commands/resume/app/run.ts index 1dc2855a118..0610f68d5e6 100644 --- a/cli/src/commands/resume/app/run.ts +++ b/cli/src/commands/resume/app/run.ts @@ -50,7 +50,7 @@ export async function resumeApp(opts: ResumeAppOptions, deps: ResumeAppDeps): Pr let action = opts.action if (action === undefined) { const formResp = await deps.http.get<{ user_actions: { id: string }[] }>( - `apps/${encodeURIComponent(opts.appId)}/form/human_input/${encodeURIComponent(opts.formToken)}`, + `apps/${encodeURIComponent(opts.appId)}/human-input-forms/${encodeURIComponent(opts.formToken)}`, ) if (formResp.user_actions.length === 1) { action = formResp.user_actions[0]?.id ?? '' diff --git a/cli/src/commands/run/app/run.ts b/cli/src/commands/run/app/run.ts index ab468678a72..c0598b0eab0 100644 --- a/cli/src/commands/run/app/run.ts +++ b/cli/src/commands/run/app/run.ts @@ -65,7 +65,7 @@ async function executeRun( const m = await meta.get(opts.appId, [FieldInfo]) const mode = m.info?.mode ?? '' if (mode === '') - throw new Error(`app ${opts.appId}: mode missing from /describe`) + throw new Error(`app ${opts.appId}: mode missing from app metadata`) if (mode === RUN_MODES.Workflow && opts.message !== undefined && opts.message !== '') { throw new BaseError({ diff --git a/cli/src/commands/use/workspace/use.ts b/cli/src/commands/use/workspace/use.ts index 3070a76aa3e..94f335c940e 100644 --- a/cli/src/commands/use/workspace/use.ts +++ b/cli/src/commands/use/workspace/use.ts @@ -27,7 +27,7 @@ export type UseWorkspaceDeps = { * workspace list and let the caller pick one interactively (TTY only). * * The server-side switch is the source of truth: if POST - * `/workspaces//switch` fails we abort before touching `hosts.yml`, so + * `/workspaces/:switch` fails we abort before touching `hosts.yml`, so * local state never diverges from the server. */ export async function runUseWorkspace( diff --git a/cli/src/commands/version/version.test.ts b/cli/src/commands/version/version.test.ts index 33b7c17b769..be1e0f48917 100644 --- a/cli/src/commands/version/version.test.ts +++ b/cli/src/commands/version/version.test.ts @@ -109,7 +109,7 @@ describe('Version command', () => { } it('--check-compat exits with COMPAT_FAIL_EXIT_CODE when compat is unsupported', async () => { - vi.spyOn(probe, 'runVersionProbe').mockResolvedValue(fakeReport({ status: 'unsupported' })) + vi.spyOn(probe, 'runVersionProbe').mockResolvedValue(fakeReport({ status: 'too_new' })) const exitSpy = stubProcessExit() const stderrSpy = vi.spyOn(process.stderr, 'write').mockImplementation(() => true) @@ -119,7 +119,7 @@ describe('Version command', () => { }) it('--check-compat -o json emits the JSON envelope on stdout before exiting', async () => { - vi.spyOn(probe, 'runVersionProbe').mockResolvedValue(fakeReport({ status: 'unsupported' })) + vi.spyOn(probe, 'runVersionProbe').mockResolvedValue(fakeReport({ status: 'too_new' })) const exitSpy = stubProcessExit() const stdoutSpy = vi.spyOn(process.stdout, 'write').mockImplementation(() => true) vi.spyOn(process.stderr, 'write').mockImplementation(() => true) @@ -131,7 +131,7 @@ describe('Version command', () => { expect(stdoutSpy).toHaveBeenCalled() const written = stdoutSpy.mock.calls.map(c => String(c[0])).join('') const parsed = JSON.parse(written) as { compat: { status: string } } - expect(parsed.compat.status).toBe('unsupported') + expect(parsed.compat.status).toBe('too_new') expect(exitSpy).toHaveBeenCalledWith(COMPAT_FAIL_EXIT_CODE) }) diff --git a/cli/src/http/client.test.ts b/cli/src/http/client.test.ts index ae4448843a4..fa397852c98 100644 --- a/cli/src/http/client.test.ts +++ b/cli/src/http/client.test.ts @@ -184,7 +184,7 @@ describe('http client', () => { const client = createHttpClient({ baseURL: base(mock.url), bearer: 'dfoa_test' }) let caught: unknown try { - await client.get('apps/nope/describe') + await client.get('apps/nope') } catch (err) { caught = err } expect(isHttpClientError(caught)).toBe(true) @@ -545,7 +545,7 @@ describe('empty / No-Content bodies', () => { }) try { const client = createHttpClient({ baseURL: stub.url, bearer: 'dfoa_test' }) - await expect(client.post('apps/app-1/tasks/t-1/stop', { json: {} })).resolves.toBeUndefined() + await expect(client.post('apps/app-1/tasks/t-1:stop', { json: {} })).resolves.toBeUndefined() } finally { await stub.stop() diff --git a/cli/src/http/error-mapper.test.ts b/cli/src/http/error-mapper.test.ts index 3244222da07..c08a3a62e48 100644 --- a/cli/src/http/error-mapper.test.ts +++ b/cli/src/http/error-mapper.test.ts @@ -74,7 +74,7 @@ describe('classifyResponse — canonical ErrorBody', () => { describe('classifyResponse 403', () => { it('maps 403 to AccessDenied (exit 4 bucket)', async () => { - const req403 = new Request('https://x/openapi/v1/apps/abc/export') + const req403 = new Request('https://x/openapi/v1/apps/abc/dsl') const res403 = new Response( JSON.stringify({ code: 'unsupported_token_type', message: 'unsupported_token_type', status: 403 }), { status: 403, headers: { 'content-type': 'application/json' } }, @@ -91,6 +91,31 @@ describe('classifyResponse 403', () => { }) }) +describe('classifyResponse 426', () => { + it('maps 426 to VersionSkew (exit 6) and surfaces the server upgrade message', async () => { + const body = { + code: 'upgrade_required', + message: 'difyctl 0.1.0 is no longer supported; upgrade to >= 0.2.0.', + status: 426, + hint: 'Upgrade difyctl: https://docs.dify.ai/en/cli/install', + } + + const err = await classified(426, body) + + expect(err.code).toBe(ErrorCode.VersionSkew) + expect(err.exit()).toBe(6) + expect(err.message).toBe('difyctl 0.1.0 is no longer supported; upgrade to >= 0.2.0.') + expect(err.serverError?.code).toBe('upgrade_required') + }) + + it('426 with no parseable ErrorBody falls back to a version message', async () => { + const err = await classified(426, 'not json') + + expect(err.code).toBe(ErrorCode.VersionSkew) + expect(err.message).toBe('client version no longer supported by the server') + }) +}) + describe('classifyResponse — non-conforming bodies (no fallback by design)', () => { it('non-JSON body yields no serverError, classification by status', async () => { const err = await classified(502, 'bad gateway') diff --git a/cli/src/http/error-mapper.ts b/cli/src/http/error-mapper.ts index 34d7637d4e0..2e29985037c 100644 --- a/cli/src/http/error-mapper.ts +++ b/cli/src/http/error-mapper.ts @@ -50,11 +50,22 @@ const ACCESS_DENIED_CLASS: StatusClass = { includeRaw: false, } +// 426 Upgrade Required: the server rejected this difyctl as too old. Give it the +// version-compat exit code so scripts can tell it apart from a generic failure. +// The server's ErrorBody.code ("upgrade_required") + message still ride along. +const VERSION_COMPAT_CLASS: StatusClass = { + code: ErrorCode.VersionSkew, + fallbackMessage: () => 'client version no longer supported by the server', + includeRaw: false, +} + function statusClass(status: number): StatusClass { if (status === 401) return AUTH_EXPIRED_CLASS if (status === 403) return ACCESS_DENIED_CLASS + if (status === 426) + return VERSION_COMPAT_CLASS if (status === 429) return RATE_LIMITED_CLASS if (status >= 500) diff --git a/cli/src/store/manager.ts b/cli/src/store/manager.ts index 37962681cec..832aba6196b 100644 --- a/cli/src/store/manager.ts +++ b/cli/src/store/manager.ts @@ -7,6 +7,7 @@ import { FileTokenStore, KeychainTokenStore } from './token-store' export const CACHE_APP_INFO = 'app-info' export const CACHE_NUDGE = 'nudge' +export const CACHE_COMPAT = 'compat' const HOSTS_FILE = 'hosts.yml' const TOKENS_FILE = 'tokens.yml' export const CONFIG_FILE_NAME = 'config.yml' diff --git a/cli/src/version/compat.test.ts b/cli/src/version/compat.test.ts index dcf3f08b259..c47b25e280f 100644 --- a/cli/src/version/compat.test.ts +++ b/cli/src/version/compat.test.ts @@ -32,14 +32,38 @@ describe('evaluateCompat', () => { expect(evaluateCompat('1.7.0', range).status).toBe('compatible') }) - it('returns unsupported when server is below minimum', () => { + it('returns too_old when server is below minimum', () => { const v = evaluateCompat('1.5.9', range) - expect(v.status).toBe('unsupported') + expect(v.status).toBe('too_old') expect(v.detail).toContain('1.5.9') }) - it('returns unsupported when server is above maximum', () => { - expect(evaluateCompat('2.0.0', range).status).toBe('unsupported') + it('returns too_new when server is above maximum', () => { + expect(evaluateCompat('2.0.0', range).status).toBe('too_new') + }) + + describe('ignores pre-release/channel suffixes (numeric-core comparison)', () => { + it('treats a pre-release of the upper bound as compatible', () => { + // 1.7.0-rc.1 has core 1.7.0 == maxDify; a suffix-sensitive range would push + // it out of [1.6.0, 1.7.0], but its numeric core is in range. + expect(evaluateCompat('1.7.0-rc.1', range).status).toBe('compatible') + }) + + it('treats a pre-release of the lower bound as compatible', () => { + expect(evaluateCompat('1.6.0-alpha', range).status).toBe('compatible') + }) + + it('still flags a pre-release whose core is below the minimum as too_old', () => { + expect(evaluateCompat('1.5.9-rc.1', range).status).toBe('too_old') + }) + + it('still flags a pre-release whose core is above the maximum as too_new', () => { + expect(evaluateCompat('2.0.0-alpha', range).status).toBe('too_new') + }) + + it('strips suffixes on the range bounds too', () => { + expect(evaluateCompat('1.6.5', { minDify: '1.6.0-alpha', maxDify: '1.7.0-rc' }).status).toBe('compatible') + }) }) it('returns unknown when server version is empty', () => { diff --git a/cli/src/version/compat.ts b/cli/src/version/compat.ts index b373ad0906b..5a1caf68cb1 100644 --- a/cli/src/version/compat.ts +++ b/cli/src/version/compat.ts @@ -1,4 +1,5 @@ -import { parseRange, satisfies, tryParse } from 'std-semver' +import type { SemVer } from 'std-semver' +import { compare, tryParse } from 'std-semver' export type DifyCompat = { readonly minDify: string @@ -14,7 +15,7 @@ export function compatString(): string { return `dify >=${difyCompat.minDify}, <=${difyCompat.maxDify}` } -export type CompatStatus = 'compatible' | 'unsupported' | 'unknown' +export type CompatStatus = 'compatible' | 'too_old' | 'too_new' | 'unknown' export type CompatVerdict = { readonly status: CompatStatus @@ -27,6 +28,12 @@ function clamp(s: string): string { return s.length > DETAIL_MAX_LEN ? `${s.slice(0, DETAIL_MAX_LEN)}…` : s } +// Numeric core (major.minor.patch) with pre-release/build stripped, so ordering +// ignores channel suffixes: a 0.2.0-rc.1 build compares equal to the 0.2.0 floor. +function core(v: SemVer): SemVer { + return { major: v.major, minor: v.minor, patch: v.patch, prerelease: [], build: [] } +} + export function evaluateCompat( serverVersion: string | undefined, range: DifyCompat = difyCompat, @@ -34,25 +41,20 @@ export function evaluateCompat( if (serverVersion === undefined || serverVersion === '') return { status: 'unknown', detail: 'server version unknown' } - const parsedServer = tryParse(serverVersion) - if (parsedServer === undefined) + const server = tryParse(serverVersion) + if (server === undefined) return { status: 'unknown', detail: `server version ${JSON.stringify(clamp(serverVersion))} is not valid semver` } - // The compat range is inclusive at both ends, exactly the format compatString prints. - const expr = `>=${range.minDify} <=${range.maxDify}` - const parsedRange = (() => { - try { - return parseRange(expr) - } - catch { - return undefined - } - })() - if (parsedRange === undefined) - return { status: 'unknown', detail: `compat range ${JSON.stringify(expr)} is not valid semver` } + const min = tryParse(range.minDify) + const max = tryParse(range.maxDify) + if (min === undefined || max === undefined) + return { status: 'unknown', detail: `compat range ${JSON.stringify(`>=${range.minDify} <=${range.maxDify}`)} is not valid semver` } - if (satisfies(parsedServer, parsedRange)) - return { status: 'compatible', detail: `server ${serverVersion} in [${range.minDify}, ${range.maxDify}]` } + if (compare(core(server), core(min)) < 0) + return { status: 'too_old', detail: `server ${serverVersion} is older than the minimum ${range.minDify}` } - return { status: 'unsupported', detail: `server ${serverVersion} outside [${range.minDify}, ${range.maxDify}]` } + if (compare(core(server), core(max)) > 0) + return { status: 'too_new', detail: `server ${serverVersion} is newer than the tested maximum ${range.maxDify}` } + + return { status: 'compatible', detail: `server ${serverVersion} in [${range.minDify}, ${range.maxDify}]` } } diff --git a/cli/src/version/enforce.test.ts b/cli/src/version/enforce.test.ts new file mode 100644 index 00000000000..4c3395530c2 --- /dev/null +++ b/cli/src/version/enforce.test.ts @@ -0,0 +1,86 @@ +import type { ServerVersionResponse } from '@dify/contracts/api/openapi/types.gen' +import type { CompatStore } from '@/cache/compat-store' +import { describe, expect, it, vi } from 'vitest' +import { ErrorCode } from '@/errors/codes' +import { enforceDifyVersion } from './enforce' + +// Injected build range in tests is __DIFYCTL_MIN_DIFY__=1.6.0 / MAX=1.7.0 (test/setup.ts): +// 1.5.0 → too_old, 1.6.4 → compatible, 99.0.0 → too_new, '' → unknown. +const HOST = 'https://cloud.dify.ai' + +function fakeStore(fresh = false): CompatStore & { readonly marked: string[] } { + const marked: string[] = [] + return { + marked, + isFreshCompatible: () => fresh, + markCompatible: async (host) => { + marked.push(host) + }, + } +} + +const server = (version: string): ServerVersionResponse => ({ version, edition: 'SELF_HOSTED' }) + +describe('enforceDifyVersion', () => { + it('throws version_skew (exit 6) when the server is too old, and never caches it', async () => { + const store = fakeStore() + const probe = vi.fn(async () => server('1.5.0')) + + await expect(enforceDifyVersion(HOST, { store, probe })).rejects.toMatchObject({ code: ErrorCode.VersionSkew }) + expect(store.marked).toHaveLength(0) + }) + + it('passes and caches when the server is compatible', async () => { + const store = fakeStore() + const probe = vi.fn(async () => server('1.6.4')) + + const res = await enforceDifyVersion(HOST, { store, probe }) + + expect(res?.version).toBe('1.6.4') + expect(store.marked).toEqual([HOST]) + }) + + it('passes (soft, no throw) and caches when the server is too new', async () => { + const store = fakeStore() + const probe = vi.fn(async () => server('99.0.0')) + + await expect(enforceDifyVersion(HOST, { store, probe })).resolves.toBeDefined() + expect(store.marked).toEqual([HOST]) + }) + + it('skips the probe entirely when the host is fresh-compatible', async () => { + const store = fakeStore(true) + const probe = vi.fn(async () => server('1.5.0')) // would throw if it ran + + await expect(enforceDifyVersion(HOST, { store, probe })).resolves.toBeUndefined() + expect(probe).not.toHaveBeenCalled() + }) + + it('re-probes despite a fresh cache when forceFresh is set', async () => { + const store = fakeStore(true) + const probe = vi.fn(async () => server('1.5.0')) + + await expect(enforceDifyVersion(HOST, { store, probe, forceFresh: true })) + .rejects + .toMatchObject({ code: ErrorCode.VersionSkew }) + expect(probe).toHaveBeenCalledOnce() + }) + + it('fails open (never blocks, never caches) when the probe errors', async () => { + const store = fakeStore() + const probe = vi.fn(async () => { + throw new Error('net down') + }) + + await expect(enforceDifyVersion(HOST, { store, probe })).resolves.toBeUndefined() + expect(store.marked).toHaveLength(0) + }) + + it('does not block or cache on an unknown server version', async () => { + const store = fakeStore() + const probe = vi.fn(async () => server('')) + + await expect(enforceDifyVersion(HOST, { store, probe })).resolves.toBeDefined() + expect(store.marked).toHaveLength(0) + }) +}) diff --git a/cli/src/version/enforce.ts b/cli/src/version/enforce.ts new file mode 100644 index 00000000000..d1b9fe87da9 --- /dev/null +++ b/cli/src/version/enforce.ts @@ -0,0 +1,69 @@ +import type { ServerVersionResponse } from '@dify/contracts/api/openapi/types.gen' +import type { CompatStore } from '@/cache/compat-store' +import { META_PROBE_TIMEOUT_MS, MetaClient } from '@/api/meta' +import { loadCompatStore } from '@/cache/compat-store' +import { newError } from '@/errors/base' +import { ErrorCode } from '@/errors/codes' +import { createHttpClient } from '@/http/client' +import { openAPIBase } from '@/util/host' +import { difyCompat, evaluateCompat } from './compat' +import { versionInfo } from './info' + +export type ServerVersionProbe = (host: string) => Promise + +const UPGRADE_HINT + = `upgrade the Dify server to >= ${difyCompat.minDify} ` + + '(https://docs.dify.ai/en/getting-started/install-self-hosted)' + +// /_version is unauthenticated; same timeout/no-retry budget as the auto-nudge probe. +const defaultProbe: ServerVersionProbe = async (host) => { + const http = createHttpClient({ baseURL: openAPIBase(host), timeoutMs: META_PROBE_TIMEOUT_MS, retryAttempts: 0 }) + return new MetaClient(http).serverVersion() +} + +export type EnforceOptions = { + readonly probe?: ServerVersionProbe + readonly store?: CompatStore + readonly forceFresh?: boolean +} + +/** + * Hard version gate for the client → server direction: refuse a Dify server older + * than this difyctl requires (its removed paths would only 404 otherwise). + * + * Cached: a host recently confirmed compatible is not re-probed for COMPAT_TTL_MS. + * Only "compatible" is cached, so a just-upgraded server clears a previous block at + * once. Fails open on any probe error — a flaky network never blocks a command. + * Returns the probed server version when it actually probed (skipped/failed → undefined), + * so the caller can reuse it. + */ +export async function enforceDifyVersion( + host: string, + opts: EnforceOptions = {}, +): Promise { + const store = opts.store ?? await loadCompatStore() + if (opts.forceFresh !== true && store.isFreshCompatible(host)) + return undefined + + const probe = opts.probe ?? defaultProbe + let server: ServerVersionResponse + try { + server = await probe(host) + } + catch { + return undefined + } + + const verdict = evaluateCompat(server.version) + if (verdict.status === 'too_old') { + throw newError( + ErrorCode.VersionSkew, + `Dify server ${server.version} is too old for difyctl ${versionInfo.version}: ${verdict.detail}`, + ).withHint(UPGRADE_HINT) + } + + if (verdict.status === 'compatible' || verdict.status === 'too_new') + await store.markCompatible(host) + + return server +} diff --git a/cli/src/version/nudge.ts b/cli/src/version/nudge.ts index a6d9fe96d4b..f8d368ee066 100644 --- a/cli/src/version/nudge.ts +++ b/cli/src/version/nudge.ts @@ -44,7 +44,9 @@ export async function maybeNudgeCompat(host: string, deps: NudgeDeps): Promise { expect(report.compat.status).toBe('compatible') }) - it('returns unsupported when server version is out of range', async () => { + it('returns too_new when server version is above range', async () => { const report = await runVersionProbe({ skipServer: false, loadActive: async () => active(), @@ -113,7 +113,7 @@ describe('runVersionProbe', () => { }) expect(report.server.reachable).toBe(true) - expect(report.compat.status).toBe('unsupported') + expect(report.compat.status).toBe('too_new') }) it('returns unknown when server returns an empty version string', async () => { diff --git a/cli/src/version/render.test.ts b/cli/src/version/render.test.ts index 2543ebf7748..0ca2f39616d 100644 --- a/cli/src/version/render.test.ts +++ b/cli/src/version/render.test.ts @@ -131,7 +131,7 @@ describe('renderVersionText', () => { compat: { minDify: '1.6.0', maxDify: '1.7.0', - status: 'unsupported', + status: 'too_new', detail: 'server 99.0.0 outside [1.6.0, 1.7.0]', }, } @@ -175,7 +175,7 @@ describe('renderVersionText', () => { compat: { minDify: '1.6.0', maxDify: '1.7.0', - status: 'unsupported', + status: 'too_new', detail: 'server 99.0.0 outside [1.6.0, 1.7.0]', }, } diff --git a/cli/src/version/render.ts b/cli/src/version/render.ts index 44dded18e28..70b725c80e3 100644 --- a/cli/src/version/render.ts +++ b/cli/src/version/render.ts @@ -15,7 +15,8 @@ export type RenderOptions = { const COMPAT_LABEL: Record = { compatible: 'ok', - unsupported: 'incompatible', + too_old: 'incompatible (server too old)', + too_new: 'incompatible (server too new)', unknown: 'unknown', } @@ -50,7 +51,8 @@ export function renderVersionText(report: VersionReport, opts: RenderOptions = { lines.push('') const verdictText = `Compatibility: ${COMPAT_LABEL[compat.status]} — ${compat.detail}` - lines.push(compat.status === 'unsupported' ? c.yellow(verdictText) : verdictText) + const incompatible = compat.status === 'too_old' || compat.status === 'too_new' + lines.push(incompatible ? c.yellow(verdictText) : verdictText) if (client.channel !== 'stable') { lines.push('') diff --git a/cli/test/e2e/suites/discovery/get-app-single.e2e.ts b/cli/test/e2e/suites/discovery/get-app-single.e2e.ts index b620eb383ef..528c08c9fa0 100644 --- a/cli/test/e2e/suites/discovery/get-app-single.e2e.ts +++ b/cli/test/e2e/suites/discovery/get-app-single.e2e.ts @@ -3,7 +3,7 @@ * * Test cases sourced from: Dify CLI Enhanced spec — Dify CLI/Discovery/Single App Query (22 cases) * - * Note: difyctl get app queries a single app via GET /apps//describe?fields=info. + * Note: difyctl get app queries a single app via GET /apps/?fields=info. * The response is returned in list-envelope format {page,limit,total,data:[...]}. */ diff --git a/cli/test/fixtures/dify-mock/server.test.ts b/cli/test/fixtures/dify-mock/server.test.ts index 7233f2a23b3..e80e20643c7 100644 --- a/cli/test/fixtures/dify-mock/server.test.ts +++ b/cli/test/fixtures/dify-mock/server.test.ts @@ -111,15 +111,15 @@ describe('dify-mock fixture server', () => { expect(body.data.map(r => r.id).sort()).toEqual(['app-3', 'app-4']) }) - it('GET /openapi/v1/apps/:id/describe returns 404 for unknown id', async () => { - const r = await fetch(`${mock.url}/openapi/v1/apps/nope/describe?workspace_id=550e8400-e29b-41d4-a716-446655440000`, { + it('GET /openapi/v1/apps/:id returns 404 for unknown id', async () => { + const r = await fetch(`${mock.url}/openapi/v1/apps/nope?workspace_id=550e8400-e29b-41d4-a716-446655440000`, { headers: { Authorization: 'Bearer dfoa_test' }, }) expect(r.status).toBe(404) }) - it('GET /openapi/v1/apps/:id/describe returns the app for known id', async () => { - const r = await fetch(`${mock.url}/openapi/v1/apps/app-1/describe?workspace_id=550e8400-e29b-41d4-a716-446655440000`, { + it('GET /openapi/v1/apps/:id returns the app for known id', async () => { + const r = await fetch(`${mock.url}/openapi/v1/apps/app-1?workspace_id=550e8400-e29b-41d4-a716-446655440000`, { headers: { Authorization: 'Bearer dfoa_test' }, }) expect(r.status).toBe(200) @@ -127,8 +127,8 @@ describe('dify-mock fixture server', () => { expect(body.info.id).toBe('app-1') }) - it('POST /openapi/v1/apps/:id/run returns SSE stream for chat app', async () => { - const r = await fetch(`${mock.url}/openapi/v1/apps/app-1/run`, { + it('POST /openapi/v1/apps/:id:run returns SSE stream for chat app', async () => { + const r = await fetch(`${mock.url}/openapi/v1/apps/app-1:run`, { method: 'POST', headers: { 'Authorization': 'Bearer dfoa_test', @@ -142,8 +142,8 @@ describe('dify-mock fixture server', () => { expect(text).toContain('"answer":"echo: "') }) - it('POST /openapi/v1/apps/:id/run returns SSE stream for workflow app', async () => { - const r = await fetch(`${mock.url}/openapi/v1/apps/app-2/run`, { + it('POST /openapi/v1/apps/:id:run returns SSE stream for workflow app', async () => { + const r = await fetch(`${mock.url}/openapi/v1/apps/app-2:run`, { method: 'POST', headers: { 'Authorization': 'Bearer dfoa_test', @@ -157,8 +157,8 @@ describe('dify-mock fixture server', () => { expect(text).toContain('"workflow_finished"') }) - it('GET /openapi/v1/apps/:id/describe?fields=info returns slim payload', async () => { - const r = await fetch(`${mock.url}/openapi/v1/apps/app-1/describe?workspace_id=550e8400-e29b-41d4-a716-446655440000&fields=info`, { + it('GET /openapi/v1/apps/:id?fields=info returns slim payload', async () => { + const r = await fetch(`${mock.url}/openapi/v1/apps/app-1?workspace_id=550e8400-e29b-41d4-a716-446655440000&fields=info`, { headers: { Authorization: 'Bearer dfoa_test' }, }) expect(r.status).toBe(200) @@ -168,8 +168,8 @@ describe('dify-mock fixture server', () => { expect(body.input_schema).toBeNull() }) - it('GET /openapi/v1/apps/:id/describe full returns parameters when present', async () => { - const r = await fetch(`${mock.url}/openapi/v1/apps/app-1/describe?workspace_id=550e8400-e29b-41d4-a716-446655440000`, { + it('GET /openapi/v1/apps/:id full returns parameters when present', async () => { + const r = await fetch(`${mock.url}/openapi/v1/apps/app-1?workspace_id=550e8400-e29b-41d4-a716-446655440000`, { headers: { Authorization: 'Bearer dfoa_test' }, }) expect(r.status).toBe(200) diff --git a/cli/test/fixtures/dify-mock/server.ts b/cli/test/fixtures/dify-mock/server.ts index 766963cd0d5..d38e723988f 100644 --- a/cli/test/fixtures/dify-mock/server.ts +++ b/cli/test/fixtures/dify-mock/server.ts @@ -15,9 +15,9 @@ export type DifyMock = { scenario: Scenario setScenario: (s: Scenario) => void stop: () => Promise - /** Body of the most recent POST to /apps/:id/run */ + /** Body of the most recent POST to /apps/:id:run */ lastRunBody: Record | null - /** Number of times POST /apps/:id/files/upload was called */ + /** Number of times POST /apps/:id/files was called */ uploadCallCount: number /** Body of the most recent POST to /workspaces/:id/apps/imports */ lastImportBody: Record | null @@ -251,7 +251,7 @@ export function buildApp(getScenario: () => Scenario, state?: MockState): Hono { }) }) - app.get('/openapi/v1/apps/:id/describe', (c) => { + app.get('/openapi/v1/apps/:id', (c) => { const id = c.req.param('id') const wsId = c.req.query('workspace_id') const fieldsRaw = c.req.query('fields') ?? '' @@ -279,7 +279,7 @@ export function buildApp(getScenario: () => Scenario, state?: MockState): Hono { }) }) - app.get('/openapi/v1/permitted-external-apps/:id/describe', (c) => { + app.get('/openapi/v1/permitted-external-apps/:id', (c) => { const id = c.req.param('id') const fieldsRaw = c.req.query('fields') ?? '' const fields = fieldsRaw === '' ? [] : fieldsRaw.split(',').map(s => s.trim()).filter(s => s !== '') @@ -307,7 +307,7 @@ export function buildApp(getScenario: () => Scenario, state?: MockState): Hono { }) }) - app.get('/openapi/v1/apps/:id/export', (c) => { + app.get('/openapi/v1/apps/:id/dsl', (c) => { const id = c.req.param('id') const found = APPS.find(a => a.id === id) if (found === undefined) @@ -315,7 +315,7 @@ export function buildApp(getScenario: () => Scenario, state?: MockState): Hono { return c.json({ data: DSL_YAML }) }) - app.get('/openapi/v1/apps/:id/check-dependencies', (c) => { + app.get('/openapi/v1/apps/:id/dependencies:check', (c) => { const id = c.req.param('id') const found = APPS.find(a => a.id === id) if (found === undefined) @@ -335,12 +335,13 @@ export function buildApp(getScenario: () => Scenario, state?: MockState): Hono { return c.json({ id: 'imp-1', status: 'completed', app_id: 'app-1', app_mode: 'chat' }, { status: 200 }) }) - app.post('/openapi/v1/workspaces/:wsId/apps/imports/:importId/confirm', (c) => { + app.post('/openapi/v1/workspaces/:wsId/apps/imports/:importId:confirm', (c) => { return c.json({ id: 'imp-1', status: 'completed', app_id: 'app-1', app_mode: 'chat' }, { status: 200 }) }) - app.post('/openapi/v1/apps/:id/run', async (c) => { - const id = c.req.param('id') + app.post('/openapi/v1/apps/:id:run', async (c) => { + // Hono drops the param adjacent to the `:run` literal; recover the app id from the path. + const id = c.req.path.replace(/^.*\/apps\//, '').replace(/:run$/, '') const body = await c.req.json() as { query?: string, inputs?: unknown } if (state !== undefined) state.lastRunBody = body as Record @@ -400,7 +401,7 @@ export function buildApp(getScenario: () => Scenario, state?: MockState): Hono { return new Response(sse, { status: 200, headers: { 'content-type': 'text/event-stream' } }) }) - app.post('/openapi/v1/apps/:id/files/upload', async (c) => { + app.post('/openapi/v1/apps/:id/files', async (c) => { if (state !== undefined) state.uploadCallCount++ const form = await c.req.formData() @@ -421,11 +422,11 @@ export function buildApp(getScenario: () => Scenario, state?: MockState): Hono { ) }) - app.post('/openapi/v1/apps/:id/tasks/:taskId/stop', (c) => { + app.post('/openapi/v1/apps/:id/tasks/:taskId:stop', (c) => { return c.json({ result: 'success' }) }) - app.post('/openapi/v1/apps/:id/form/human_input/:formToken', (c) => { + app.post('/openapi/v1/apps/:id/human-input-forms/:formToken:submit', (c) => { return c.json({}) }) diff --git a/packages/contracts/generated/api/openapi/orpc.gen.ts b/packages/contracts/generated/api/openapi/orpc.gen.ts index 47aa1b90d6a..13ac22da962 100644 --- a/packages/contracts/generated/api/openapi/orpc.gen.ts +++ b/packages/contracts/generated/api/openapi/orpc.gen.ts @@ -12,16 +12,16 @@ import { zGetAccountResponse, zGetAccountSessionsQuery, zGetAccountSessionsResponse, - zGetAppsByAppIdCheckDependenciesPath, - zGetAppsByAppIdCheckDependenciesResponse, - zGetAppsByAppIdDescribePath, - zGetAppsByAppIdDescribeQuery, - zGetAppsByAppIdDescribeResponse, - zGetAppsByAppIdExportPath, - zGetAppsByAppIdExportQuery, - zGetAppsByAppIdExportResponse, - zGetAppsByAppIdFormHumanInputByFormTokenPath, - zGetAppsByAppIdFormHumanInputByFormTokenResponse, + zGetAppsByAppIdDependenciesCheckPath, + zGetAppsByAppIdDependenciesCheckResponse, + zGetAppsByAppIdDslPath, + zGetAppsByAppIdDslQuery, + zGetAppsByAppIdDslResponse, + zGetAppsByAppIdHumanInputFormsByFormTokenPath, + zGetAppsByAppIdHumanInputFormsByFormTokenResponse, + zGetAppsByAppIdPath, + zGetAppsByAppIdQuery, + zGetAppsByAppIdResponse, zGetAppsByAppIdTasksByTaskIdEventsPath, zGetAppsByAppIdTasksByTaskIdEventsQuery, zGetAppsByAppIdTasksByTaskIdEventsResponse, @@ -30,9 +30,9 @@ import { zGetHealthResponse, zGetOauthDeviceLookupQuery, zGetOauthDeviceLookupResponse, - zGetPermittedExternalAppsByAppIdDescribePath, - zGetPermittedExternalAppsByAppIdDescribeQuery, - zGetPermittedExternalAppsByAppIdDescribeResponse, + zGetPermittedExternalAppsByAppIdPath, + zGetPermittedExternalAppsByAppIdQuery, + zGetPermittedExternalAppsByAppIdResponse, zGetPermittedExternalAppsQuery, zGetPermittedExternalAppsResponse, zGetVersionResponse, @@ -42,11 +42,14 @@ import { zGetWorkspacesByWorkspaceIdPath, zGetWorkspacesByWorkspaceIdResponse, zGetWorkspacesResponse, - zPostAppsByAppIdFilesUploadPath, - zPostAppsByAppIdFilesUploadResponse, - zPostAppsByAppIdFormHumanInputByFormTokenBody, - zPostAppsByAppIdFormHumanInputByFormTokenPath, - zPostAppsByAppIdFormHumanInputByFormTokenResponse, + zPatchWorkspacesByWorkspaceIdMembersByMemberIdBody, + zPatchWorkspacesByWorkspaceIdMembersByMemberIdPath, + zPatchWorkspacesByWorkspaceIdMembersByMemberIdResponse, + zPostAppsByAppIdFilesPath, + zPostAppsByAppIdFilesResponse, + zPostAppsByAppIdHumanInputFormsByFormTokenSubmitBody, + zPostAppsByAppIdHumanInputFormsByFormTokenSubmitPath, + zPostAppsByAppIdHumanInputFormsByFormTokenSubmitResponse, zPostAppsByAppIdRunBody, zPostAppsByAppIdRunPath, zPostAppsByAppIdRunResponse, @@ -70,9 +73,6 @@ import { zPostWorkspacesByWorkspaceIdMembersResponse, zPostWorkspacesByWorkspaceIdSwitchPath, zPostWorkspacesByWorkspaceIdSwitchResponse, - zPutWorkspacesByWorkspaceIdMembersByMemberIdRoleBody, - zPutWorkspacesByWorkspaceIdMembersByMemberIdRolePath, - zPutWorkspacesByWorkspaceIdMembersByMemberIdRoleResponse, } from './zod.gen' export const get = oc @@ -168,54 +168,36 @@ export const get5 = oc .route({ inputStructure: 'detailed', method: 'GET', - operationId: 'getAppsByAppIdCheckDependencies', - path: '/apps/{app_id}/check-dependencies', + operationId: 'getAppsByAppIdDependenciesCheck', + path: '/apps/{app_id}/dependencies:check', tags: ['openapi'], }) - .input(z.object({ params: zGetAppsByAppIdCheckDependenciesPath })) - .output(zGetAppsByAppIdCheckDependenciesResponse) + .input(z.object({ params: zGetAppsByAppIdDependenciesCheckPath })) + .output(zGetAppsByAppIdDependenciesCheckResponse) -export const checkDependencies = { +export const check = { get: get5, } +export const dependencies = { + check, +} + export const get6 = oc .route({ inputStructure: 'detailed', method: 'GET', - operationId: 'getAppsByAppIdDescribe', - path: '/apps/{app_id}/describe', + operationId: 'getAppsByAppIdDsl', + path: '/apps/{app_id}/dsl', tags: ['openapi'], }) - .input( - z.object({ - params: zGetAppsByAppIdDescribePath, - query: zGetAppsByAppIdDescribeQuery.optional(), - }), - ) - .output(zGetAppsByAppIdDescribeResponse) + .input(z.object({ params: zGetAppsByAppIdDslPath, query: zGetAppsByAppIdDslQuery.optional() })) + .output(zGetAppsByAppIdDslResponse) -export const describe = { +export const dsl = { get: get6, } -export const get7 = oc - .route({ - inputStructure: 'detailed', - method: 'GET', - operationId: 'getAppsByAppIdExport', - path: '/apps/{app_id}/export', - tags: ['openapi'], - }) - .input( - z.object({ params: zGetAppsByAppIdExportPath, query: zGetAppsByAppIdExportQuery.optional() }), - ) - .output(zGetAppsByAppIdExportResponse) - -export const export_ = { - get: get7, -} - /** * Upload a file to use as an input variable when running the app */ @@ -224,78 +206,59 @@ export const post = oc description: 'Upload a file to use as an input variable when running the app', inputStructure: 'detailed', method: 'POST', - operationId: 'postAppsByAppIdFilesUpload', - path: '/apps/{app_id}/files/upload', + operationId: 'postAppsByAppIdFiles', + path: '/apps/{app_id}/files', successStatus: 201, tags: ['openapi'], }) - .input(z.object({ params: zPostAppsByAppIdFilesUploadPath })) - .output(zPostAppsByAppIdFilesUploadResponse) - -export const upload = { - post, -} + .input(z.object({ params: zPostAppsByAppIdFilesPath })) + .output(zPostAppsByAppIdFilesResponse) export const files = { - upload, + post, } -export const get8 = oc - .route({ - inputStructure: 'detailed', - method: 'GET', - operationId: 'getAppsByAppIdFormHumanInputByFormToken', - path: '/apps/{app_id}/form/human_input/{form_token}', - tags: ['openapi'], - }) - .input(z.object({ params: zGetAppsByAppIdFormHumanInputByFormTokenPath })) - .output(zGetAppsByAppIdFormHumanInputByFormTokenResponse) - export const post2 = oc .route({ inputStructure: 'detailed', method: 'POST', - operationId: 'postAppsByAppIdFormHumanInputByFormToken', - path: '/apps/{app_id}/form/human_input/{form_token}', + operationId: 'postAppsByAppIdHumanInputFormsByFormTokenSubmit', + path: '/apps/{app_id}/human-input-forms/{form_token}:submit', tags: ['openapi'], }) .input( z.object({ - body: zPostAppsByAppIdFormHumanInputByFormTokenBody, - params: zPostAppsByAppIdFormHumanInputByFormTokenPath, + body: zPostAppsByAppIdHumanInputFormsByFormTokenSubmitBody, + params: zPostAppsByAppIdHumanInputFormsByFormTokenSubmitPath, }), ) - .output(zPostAppsByAppIdFormHumanInputByFormTokenResponse) + .output(zPostAppsByAppIdHumanInputFormsByFormTokenSubmitResponse) -export const byFormToken = { - get: get8, +export const submit = { post: post2, } -export const humanInput = { +export const get7 = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'getAppsByAppIdHumanInputFormsByFormToken', + path: '/apps/{app_id}/human-input-forms/{form_token}', + tags: ['openapi'], + }) + .input(z.object({ params: zGetAppsByAppIdHumanInputFormsByFormTokenPath })) + .output(zGetAppsByAppIdHumanInputFormsByFormTokenResponse) + +export const byFormToken = { + get: get7, + submit, +} + +export const humanInputForms = { byFormToken, } -export const form = { - humanInput, -} - -export const post3 = oc - .route({ - inputStructure: 'detailed', - method: 'POST', - operationId: 'postAppsByAppIdRun', - path: '/apps/{app_id}/run', - tags: ['openapi'], - }) - .input(z.object({ body: zPostAppsByAppIdRunBody, params: zPostAppsByAppIdRunPath })) - .output(zPostAppsByAppIdRunResponse) - -export const run = { - post: post3, -} - -export const get9 = oc +export const get8 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -312,22 +275,22 @@ export const get9 = oc .output(zGetAppsByAppIdTasksByTaskIdEventsResponse) export const events = { - get: get9, + get: get8, } -export const post4 = oc +export const post3 = oc .route({ inputStructure: 'detailed', method: 'POST', operationId: 'postAppsByAppIdTasksByTaskIdStop', - path: '/apps/{app_id}/tasks/{task_id}/stop', + path: '/apps/{app_id}/tasks/{task_id}:stop', tags: ['openapi'], }) .input(z.object({ params: zPostAppsByAppIdTasksByTaskIdStopPath })) .output(zPostAppsByAppIdTasksByTaskIdStopResponse) export const stop = { - post: post4, + post: post3, } export const byTaskId = { @@ -339,14 +302,40 @@ export const tasks = { byTaskId, } +export const post4 = oc + .route({ + inputStructure: 'detailed', + method: 'POST', + operationId: 'postAppsByAppIdRun', + path: '/apps/{app_id}:run', + tags: ['openapi'], + }) + .input(z.object({ body: zPostAppsByAppIdRunBody, params: zPostAppsByAppIdRunPath })) + .output(zPostAppsByAppIdRunResponse) + +export const run = { + post: post4, +} + +export const get9 = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'getAppsByAppId', + path: '/apps/{app_id}', + tags: ['openapi'], + }) + .input(z.object({ params: zGetAppsByAppIdPath, query: zGetAppsByAppIdQuery.optional() })) + .output(zGetAppsByAppIdResponse) + export const byAppId = { - checkDependencies, - describe, - export: export_, + get: get9, + dependencies, + dsl, files, - form, - run, + humanInputForms, tasks, + run, } export const get10 = oc @@ -456,24 +445,20 @@ export const get12 = oc .route({ inputStructure: 'detailed', method: 'GET', - operationId: 'getPermittedExternalAppsByAppIdDescribe', - path: '/permitted-external-apps/{app_id}/describe', + operationId: 'getPermittedExternalAppsByAppId', + path: '/permitted-external-apps/{app_id}', tags: ['openapi'], }) .input( z.object({ - params: zGetPermittedExternalAppsByAppIdDescribePath, - query: zGetPermittedExternalAppsByAppIdDescribeQuery.optional(), + params: zGetPermittedExternalAppsByAppIdPath, + query: zGetPermittedExternalAppsByAppIdQuery.optional(), }), ) - .output(zGetPermittedExternalAppsByAppIdDescribeResponse) - -export const describe2 = { - get: get12, -} + .output(zGetPermittedExternalAppsByAppIdResponse) export const byAppId2 = { - describe: describe2, + get: get12, } export const get13 = oc @@ -497,7 +482,7 @@ export const post9 = oc inputStructure: 'detailed', method: 'POST', operationId: 'postWorkspacesByWorkspaceIdAppsImportsByImportIdConfirm', - path: '/workspaces/{workspace_id}/apps/imports/{import_id}/confirm', + path: '/workspaces/{workspace_id}/apps/imports/{import_id}:confirm', tags: ['openapi'], }) .input(z.object({ params: zPostWorkspacesByWorkspaceIdAppsImportsByImportIdConfirmPath })) @@ -536,26 +521,6 @@ export const apps2 = { imports, } -export const put = oc - .route({ - inputStructure: 'detailed', - method: 'PUT', - operationId: 'putWorkspacesByWorkspaceIdMembersByMemberIdRole', - path: '/workspaces/{workspace_id}/members/{member_id}/role', - tags: ['openapi'], - }) - .input( - z.object({ - body: zPutWorkspacesByWorkspaceIdMembersByMemberIdRoleBody, - params: zPutWorkspacesByWorkspaceIdMembersByMemberIdRolePath, - }), - ) - .output(zPutWorkspacesByWorkspaceIdMembersByMemberIdRoleResponse) - -export const role = { - put, -} - export const delete3 = oc .route({ inputStructure: 'detailed', @@ -567,9 +532,25 @@ export const delete3 = oc .input(z.object({ params: zDeleteWorkspacesByWorkspaceIdMembersByMemberIdPath })) .output(zDeleteWorkspacesByWorkspaceIdMembersByMemberIdResponse) +export const patch = oc + .route({ + inputStructure: 'detailed', + method: 'PATCH', + operationId: 'patchWorkspacesByWorkspaceIdMembersByMemberId', + path: '/workspaces/{workspace_id}/members/{member_id}', + tags: ['openapi'], + }) + .input( + z.object({ + body: zPatchWorkspacesByWorkspaceIdMembersByMemberIdBody, + params: zPatchWorkspacesByWorkspaceIdMembersByMemberIdPath, + }), + ) + .output(zPatchWorkspacesByWorkspaceIdMembersByMemberIdResponse) + export const byMemberId = { delete: delete3, - role, + patch, } export const get14 = oc @@ -616,7 +597,7 @@ export const post12 = oc inputStructure: 'detailed', method: 'POST', operationId: 'postWorkspacesByWorkspaceIdSwitch', - path: '/workspaces/{workspace_id}/switch', + path: '/workspaces/{workspace_id}:switch', tags: ['openapi'], }) .input(z.object({ params: zPostWorkspacesByWorkspaceIdSwitchPath })) diff --git a/packages/contracts/generated/api/openapi/types.gen.ts b/packages/contracts/generated/api/openapi/types.gen.ts index cd422cd1f92..677f1101f0b 100644 --- a/packages/contracts/generated/api/openapi/types.gen.ts +++ b/packages/contracts/generated/api/openapi/types.gen.ts @@ -347,6 +347,7 @@ export type OpenApiErrorCode | 'unknown' | 'unsupported_file_type' | 'unsupported_media_type' + | 'upgrade_required' export type Package = { plugin_unique_identifier: string @@ -617,30 +618,7 @@ export type GetAppsResponses = { export type GetAppsResponse = GetAppsResponses[keyof GetAppsResponses] -export type GetAppsByAppIdCheckDependenciesData = { - body?: never - path: { - app_id: string - } - query?: never - url: '/apps/{app_id}/check-dependencies' -} - -export type GetAppsByAppIdCheckDependenciesErrors = { - default: ErrorBody -} - -export type GetAppsByAppIdCheckDependenciesError - = GetAppsByAppIdCheckDependenciesErrors[keyof GetAppsByAppIdCheckDependenciesErrors] - -export type GetAppsByAppIdCheckDependenciesResponses = { - 200: CheckDependenciesResult -} - -export type GetAppsByAppIdCheckDependenciesResponse - = GetAppsByAppIdCheckDependenciesResponses[keyof GetAppsByAppIdCheckDependenciesResponses] - -export type GetAppsByAppIdDescribeData = { +export type GetAppsByAppIdData = { body?: never path: { app_id: string @@ -648,25 +626,46 @@ export type GetAppsByAppIdDescribeData = { query?: { fields?: string } - url: '/apps/{app_id}/describe' + url: '/apps/{app_id}' } -export type GetAppsByAppIdDescribeErrors = { +export type GetAppsByAppIdErrors = { 422: ErrorBody default: ErrorBody } -export type GetAppsByAppIdDescribeError - = GetAppsByAppIdDescribeErrors[keyof GetAppsByAppIdDescribeErrors] +export type GetAppsByAppIdError = GetAppsByAppIdErrors[keyof GetAppsByAppIdErrors] -export type GetAppsByAppIdDescribeResponses = { +export type GetAppsByAppIdResponses = { 200: AppDescribeResponse } -export type GetAppsByAppIdDescribeResponse - = GetAppsByAppIdDescribeResponses[keyof GetAppsByAppIdDescribeResponses] +export type GetAppsByAppIdResponse = GetAppsByAppIdResponses[keyof GetAppsByAppIdResponses] -export type GetAppsByAppIdExportData = { +export type GetAppsByAppIdDependenciesCheckData = { + body?: never + path: { + app_id: string + } + query?: never + url: '/apps/{app_id}/dependencies:check' +} + +export type GetAppsByAppIdDependenciesCheckErrors = { + default: ErrorBody +} + +export type GetAppsByAppIdDependenciesCheckError + = GetAppsByAppIdDependenciesCheckErrors[keyof GetAppsByAppIdDependenciesCheckErrors] + +export type GetAppsByAppIdDependenciesCheckResponses = { + 200: CheckDependenciesResult +} + +export type GetAppsByAppIdDependenciesCheckResponse + = GetAppsByAppIdDependenciesCheckResponses[keyof GetAppsByAppIdDependenciesCheckResponses] + +export type GetAppsByAppIdDslData = { body?: never path: { app_id: string @@ -675,33 +674,32 @@ export type GetAppsByAppIdExportData = { include_secret?: boolean workflow_id?: string } - url: '/apps/{app_id}/export' + url: '/apps/{app_id}/dsl' } -export type GetAppsByAppIdExportErrors = { +export type GetAppsByAppIdDslErrors = { 422: ErrorBody default: ErrorBody } -export type GetAppsByAppIdExportError = GetAppsByAppIdExportErrors[keyof GetAppsByAppIdExportErrors] +export type GetAppsByAppIdDslError = GetAppsByAppIdDslErrors[keyof GetAppsByAppIdDslErrors] -export type GetAppsByAppIdExportResponses = { +export type GetAppsByAppIdDslResponses = { 200: AppDslExportResponse } -export type GetAppsByAppIdExportResponse - = GetAppsByAppIdExportResponses[keyof GetAppsByAppIdExportResponses] +export type GetAppsByAppIdDslResponse = GetAppsByAppIdDslResponses[keyof GetAppsByAppIdDslResponses] -export type PostAppsByAppIdFilesUploadData = { +export type PostAppsByAppIdFilesData = { body?: never path: { app_id: string } query?: never - url: '/apps/{app_id}/files/upload' + url: '/apps/{app_id}/files' } -export type PostAppsByAppIdFilesUploadErrors = { +export type PostAppsByAppIdFilesErrors = { 400: unknown 401: unknown 413: unknown @@ -709,79 +707,56 @@ export type PostAppsByAppIdFilesUploadErrors = { default: ErrorBody } -export type PostAppsByAppIdFilesUploadError - = PostAppsByAppIdFilesUploadErrors[keyof PostAppsByAppIdFilesUploadErrors] +export type PostAppsByAppIdFilesError = PostAppsByAppIdFilesErrors[keyof PostAppsByAppIdFilesErrors] -export type PostAppsByAppIdFilesUploadResponses = { +export type PostAppsByAppIdFilesResponses = { 201: FileResponse } -export type PostAppsByAppIdFilesUploadResponse - = PostAppsByAppIdFilesUploadResponses[keyof PostAppsByAppIdFilesUploadResponses] +export type PostAppsByAppIdFilesResponse + = PostAppsByAppIdFilesResponses[keyof PostAppsByAppIdFilesResponses] -export type GetAppsByAppIdFormHumanInputByFormTokenData = { +export type GetAppsByAppIdHumanInputFormsByFormTokenData = { body?: never path: { app_id: string form_token: string } query?: never - url: '/apps/{app_id}/form/human_input/{form_token}' + url: '/apps/{app_id}/human-input-forms/{form_token}' } -export type GetAppsByAppIdFormHumanInputByFormTokenResponses = { +export type GetAppsByAppIdHumanInputFormsByFormTokenResponses = { 200: HumanInputFormDefinitionResponse } -export type GetAppsByAppIdFormHumanInputByFormTokenResponse - = GetAppsByAppIdFormHumanInputByFormTokenResponses[keyof GetAppsByAppIdFormHumanInputByFormTokenResponses] +export type GetAppsByAppIdHumanInputFormsByFormTokenResponse + = GetAppsByAppIdHumanInputFormsByFormTokenResponses[keyof GetAppsByAppIdHumanInputFormsByFormTokenResponses] -export type PostAppsByAppIdFormHumanInputByFormTokenData = { +export type PostAppsByAppIdHumanInputFormsByFormTokenSubmitData = { body: HumanInputFormSubmitPayload path: { app_id: string form_token: string } query?: never - url: '/apps/{app_id}/form/human_input/{form_token}' + url: '/apps/{app_id}/human-input-forms/{form_token}:submit' } -export type PostAppsByAppIdFormHumanInputByFormTokenErrors = { +export type PostAppsByAppIdHumanInputFormsByFormTokenSubmitErrors = { 422: ErrorBody default: ErrorBody } -export type PostAppsByAppIdFormHumanInputByFormTokenError - = PostAppsByAppIdFormHumanInputByFormTokenErrors[keyof PostAppsByAppIdFormHumanInputByFormTokenErrors] +export type PostAppsByAppIdHumanInputFormsByFormTokenSubmitError + = PostAppsByAppIdHumanInputFormsByFormTokenSubmitErrors[keyof PostAppsByAppIdHumanInputFormsByFormTokenSubmitErrors] -export type PostAppsByAppIdFormHumanInputByFormTokenResponses = { +export type PostAppsByAppIdHumanInputFormsByFormTokenSubmitResponses = { 200: FormSubmitResponse } -export type PostAppsByAppIdFormHumanInputByFormTokenResponse - = PostAppsByAppIdFormHumanInputByFormTokenResponses[keyof PostAppsByAppIdFormHumanInputByFormTokenResponses] - -export type PostAppsByAppIdRunData = { - body: AppRunRequest - path: { - app_id: string - } - query?: never - url: '/apps/{app_id}/run' -} - -export type PostAppsByAppIdRunErrors = { - 422: ErrorBody -} - -export type PostAppsByAppIdRunError = PostAppsByAppIdRunErrors[keyof PostAppsByAppIdRunErrors] - -export type PostAppsByAppIdRunResponses = { - 200: EventStreamResponse -} - -export type PostAppsByAppIdRunResponse - = PostAppsByAppIdRunResponses[keyof PostAppsByAppIdRunResponses] +export type PostAppsByAppIdHumanInputFormsByFormTokenSubmitResponse + = PostAppsByAppIdHumanInputFormsByFormTokenSubmitResponses[keyof PostAppsByAppIdHumanInputFormsByFormTokenSubmitResponses] export type GetAppsByAppIdTasksByTaskIdEventsData = { body?: never @@ -810,7 +785,7 @@ export type PostAppsByAppIdTasksByTaskIdStopData = { task_id: string } query?: never - url: '/apps/{app_id}/tasks/{task_id}/stop' + url: '/apps/{app_id}/tasks/{task_id}:stop' } export type PostAppsByAppIdTasksByTaskIdStopErrors = { @@ -827,6 +802,28 @@ export type PostAppsByAppIdTasksByTaskIdStopResponses = { export type PostAppsByAppIdTasksByTaskIdStopResponse = PostAppsByAppIdTasksByTaskIdStopResponses[keyof PostAppsByAppIdTasksByTaskIdStopResponses] +export type PostAppsByAppIdRunData = { + body: AppRunRequest + path: { + app_id: string + } + query?: never + url: '/apps/{app_id}:run' +} + +export type PostAppsByAppIdRunErrors = { + 422: ErrorBody +} + +export type PostAppsByAppIdRunError = PostAppsByAppIdRunErrors[keyof PostAppsByAppIdRunErrors] + +export type PostAppsByAppIdRunResponses = { + 200: EventStreamResponse +} + +export type PostAppsByAppIdRunResponse + = PostAppsByAppIdRunResponses[keyof PostAppsByAppIdRunResponses] + export type PostOauthDeviceApproveData = { body: DeviceMutateRequest path?: never @@ -926,7 +923,7 @@ export type GetPermittedExternalAppsResponses = { export type GetPermittedExternalAppsResponse = GetPermittedExternalAppsResponses[keyof GetPermittedExternalAppsResponses] -export type GetPermittedExternalAppsByAppIdDescribeData = { +export type GetPermittedExternalAppsByAppIdData = { body?: never path: { app_id: string @@ -934,23 +931,23 @@ export type GetPermittedExternalAppsByAppIdDescribeData = { query?: { fields?: string } - url: '/permitted-external-apps/{app_id}/describe' + url: '/permitted-external-apps/{app_id}' } -export type GetPermittedExternalAppsByAppIdDescribeErrors = { +export type GetPermittedExternalAppsByAppIdErrors = { 422: ErrorBody default: ErrorBody } -export type GetPermittedExternalAppsByAppIdDescribeError - = GetPermittedExternalAppsByAppIdDescribeErrors[keyof GetPermittedExternalAppsByAppIdDescribeErrors] +export type GetPermittedExternalAppsByAppIdError + = GetPermittedExternalAppsByAppIdErrors[keyof GetPermittedExternalAppsByAppIdErrors] -export type GetPermittedExternalAppsByAppIdDescribeResponses = { +export type GetPermittedExternalAppsByAppIdResponses = { 200: AppDescribeResponse } -export type GetPermittedExternalAppsByAppIdDescribeResponse - = GetPermittedExternalAppsByAppIdDescribeResponses[keyof GetPermittedExternalAppsByAppIdDescribeResponses] +export type GetPermittedExternalAppsByAppIdResponse + = GetPermittedExternalAppsByAppIdResponses[keyof GetPermittedExternalAppsByAppIdResponses] export type GetWorkspacesData = { body?: never @@ -1027,7 +1024,7 @@ export type PostWorkspacesByWorkspaceIdAppsImportsByImportIdConfirmData = { workspace_id: string } query?: never - url: '/workspaces/{workspace_id}/apps/imports/{import_id}/confirm' + url: '/workspaces/{workspace_id}/apps/imports/{import_id}:confirm' } export type PostWorkspacesByWorkspaceIdAppsImportsByImportIdConfirmErrors = { @@ -1120,30 +1117,30 @@ export type DeleteWorkspacesByWorkspaceIdMembersByMemberIdResponses = { export type DeleteWorkspacesByWorkspaceIdMembersByMemberIdResponse = DeleteWorkspacesByWorkspaceIdMembersByMemberIdResponses[keyof DeleteWorkspacesByWorkspaceIdMembersByMemberIdResponses] -export type PutWorkspacesByWorkspaceIdMembersByMemberIdRoleData = { +export type PatchWorkspacesByWorkspaceIdMembersByMemberIdData = { body: MemberRoleUpdatePayload path: { member_id: string workspace_id: string } query?: never - url: '/workspaces/{workspace_id}/members/{member_id}/role' + url: '/workspaces/{workspace_id}/members/{member_id}' } -export type PutWorkspacesByWorkspaceIdMembersByMemberIdRoleErrors = { +export type PatchWorkspacesByWorkspaceIdMembersByMemberIdErrors = { 422: ErrorBody default: ErrorBody } -export type PutWorkspacesByWorkspaceIdMembersByMemberIdRoleError - = PutWorkspacesByWorkspaceIdMembersByMemberIdRoleErrors[keyof PutWorkspacesByWorkspaceIdMembersByMemberIdRoleErrors] +export type PatchWorkspacesByWorkspaceIdMembersByMemberIdError + = PatchWorkspacesByWorkspaceIdMembersByMemberIdErrors[keyof PatchWorkspacesByWorkspaceIdMembersByMemberIdErrors] -export type PutWorkspacesByWorkspaceIdMembersByMemberIdRoleResponses = { +export type PatchWorkspacesByWorkspaceIdMembersByMemberIdResponses = { 200: MemberActionResponse } -export type PutWorkspacesByWorkspaceIdMembersByMemberIdRoleResponse - = PutWorkspacesByWorkspaceIdMembersByMemberIdRoleResponses[keyof PutWorkspacesByWorkspaceIdMembersByMemberIdRoleResponses] +export type PatchWorkspacesByWorkspaceIdMembersByMemberIdResponse + = PatchWorkspacesByWorkspaceIdMembersByMemberIdResponses[keyof PatchWorkspacesByWorkspaceIdMembersByMemberIdResponses] export type PostWorkspacesByWorkspaceIdSwitchData = { body?: never @@ -1151,7 +1148,7 @@ export type PostWorkspacesByWorkspaceIdSwitchData = { workspace_id: string } query?: never - url: '/workspaces/{workspace_id}/switch' + url: '/workspaces/{workspace_id}:switch' } export type PostWorkspacesByWorkspaceIdSwitchErrors = { diff --git a/packages/contracts/generated/api/openapi/zod.gen.ts b/packages/contracts/generated/api/openapi/zod.gen.ts index 70ece880a68..b271648f8b2 100644 --- a/packages/contracts/generated/api/openapi/zod.gen.ts +++ b/packages/contracts/generated/api/openapi/zod.gen.ts @@ -27,7 +27,7 @@ export const zAppDescribeInfo = z.object({ /** * AppDescribeQuery * - * `?fields=` allow-list for GET /apps//describe. + * `?fields=` allow-list for GET /apps/. * * Empty / omitted → all blocks. Unknown member → ValidationError → 422. */ @@ -47,7 +47,7 @@ export const zAppDescribeResponse = z.object({ /** * AppDslExportQuery * - * Query parameters for GET /apps//export. + * Query parameters for GET /apps//dsl. */ export const zAppDslExportQuery = z.object({ include_secret: z.boolean().optional().default(false), @@ -254,7 +254,7 @@ export const zFileResponse = z.object({ /** * FormSubmitResponse * - * Empty 200 body for POST /apps//form/human_input/. `extra='forbid'` + * Empty 200 body for POST /apps//human-input-forms/:submit. `extra='forbid'` * pins `additionalProperties: false` so the generated contract is an exact `{}` rather * than an under-annotated open object. */ @@ -430,6 +430,7 @@ export const zOpenApiErrorCode = z.enum([ 'unknown', 'unsupported_file_type', 'unsupported_media_type', + 'upgrade_required', ]) /** @@ -581,7 +582,7 @@ export const zPermittedExternalAppsListQuery = z.object({ /** * TaskStopResponse * - * 200 body for POST /apps//tasks//stop. The handler always returns + * 200 body for POST /apps//tasks/:stop. The handler always returns * {"result": "success"}, so `result` is required (no default) — the generated contract * types it as a required `'success'` rather than an optional field. */ @@ -740,33 +741,33 @@ export const zGetAppsQuery = z.object({ */ export const zGetAppsResponse = zAppListResponse -export const zGetAppsByAppIdCheckDependenciesPath = z.object({ +export const zGetAppsByAppIdPath = z.object({ app_id: z.string(), }) -/** - * Dependencies checked - */ -export const zGetAppsByAppIdCheckDependenciesResponse = zCheckDependenciesResult - -export const zGetAppsByAppIdDescribePath = z.object({ - app_id: z.string(), -}) - -export const zGetAppsByAppIdDescribeQuery = z.object({ +export const zGetAppsByAppIdQuery = z.object({ fields: z.string().optional(), }) /** * App description */ -export const zGetAppsByAppIdDescribeResponse = zAppDescribeResponse +export const zGetAppsByAppIdResponse = zAppDescribeResponse -export const zGetAppsByAppIdExportPath = z.object({ +export const zGetAppsByAppIdDependenciesCheckPath = z.object({ app_id: z.string(), }) -export const zGetAppsByAppIdExportQuery = z.object({ +/** + * Dependencies checked + */ +export const zGetAppsByAppIdDependenciesCheckResponse = zCheckDependenciesResult + +export const zGetAppsByAppIdDslPath = z.object({ + app_id: z.string(), +}) + +export const zGetAppsByAppIdDslQuery = z.object({ include_secret: z.boolean().optional().default(false), workflow_id: z.string().optional(), }) @@ -774,18 +775,18 @@ export const zGetAppsByAppIdExportQuery = z.object({ /** * Export successful */ -export const zGetAppsByAppIdExportResponse = zAppDslExportResponse +export const zGetAppsByAppIdDslResponse = zAppDslExportResponse -export const zPostAppsByAppIdFilesUploadPath = z.object({ +export const zPostAppsByAppIdFilesPath = z.object({ app_id: z.string(), }) /** * File uploaded successfully */ -export const zPostAppsByAppIdFilesUploadResponse = zFileResponse +export const zPostAppsByAppIdFilesResponse = zFileResponse -export const zGetAppsByAppIdFormHumanInputByFormTokenPath = z.object({ +export const zGetAppsByAppIdHumanInputFormsByFormTokenPath = z.object({ app_id: z.string(), form_token: z.string(), }) @@ -793,11 +794,11 @@ export const zGetAppsByAppIdFormHumanInputByFormTokenPath = z.object({ /** * Form definition */ -export const zGetAppsByAppIdFormHumanInputByFormTokenResponse = zHumanInputFormDefinitionResponse +export const zGetAppsByAppIdHumanInputFormsByFormTokenResponse = zHumanInputFormDefinitionResponse -export const zPostAppsByAppIdFormHumanInputByFormTokenBody = zHumanInputFormSubmitPayload +export const zPostAppsByAppIdHumanInputFormsByFormTokenSubmitBody = zHumanInputFormSubmitPayload -export const zPostAppsByAppIdFormHumanInputByFormTokenPath = z.object({ +export const zPostAppsByAppIdHumanInputFormsByFormTokenSubmitPath = z.object({ app_id: z.string(), form_token: z.string(), }) @@ -805,18 +806,7 @@ export const zPostAppsByAppIdFormHumanInputByFormTokenPath = z.object({ /** * Form submitted */ -export const zPostAppsByAppIdFormHumanInputByFormTokenResponse = zFormSubmitResponse - -export const zPostAppsByAppIdRunBody = zAppRunRequest - -export const zPostAppsByAppIdRunPath = z.object({ - app_id: z.string(), -}) - -/** - * Run result (SSE stream) - */ -export const zPostAppsByAppIdRunResponse = zEventStreamResponse +export const zPostAppsByAppIdHumanInputFormsByFormTokenSubmitResponse = zFormSubmitResponse export const zGetAppsByAppIdTasksByTaskIdEventsPath = z.object({ app_id: z.string(), @@ -843,6 +833,17 @@ export const zPostAppsByAppIdTasksByTaskIdStopPath = z.object({ */ export const zPostAppsByAppIdTasksByTaskIdStopResponse = zTaskStopResponse +export const zPostAppsByAppIdRunBody = zAppRunRequest + +export const zPostAppsByAppIdRunPath = z.object({ + app_id: z.string(), +}) + +/** + * Run result (SSE stream) + */ +export const zPostAppsByAppIdRunResponse = zEventStreamResponse + export const zPostOauthDeviceApproveBody = zDeviceMutateRequest /** @@ -892,18 +893,18 @@ export const zGetPermittedExternalAppsQuery = z.object({ */ export const zGetPermittedExternalAppsResponse = zPermittedExternalAppsListResponse -export const zGetPermittedExternalAppsByAppIdDescribePath = z.object({ +export const zGetPermittedExternalAppsByAppIdPath = z.object({ app_id: z.string(), }) -export const zGetPermittedExternalAppsByAppIdDescribeQuery = z.object({ +export const zGetPermittedExternalAppsByAppIdQuery = z.object({ fields: z.string().optional(), }) /** * Permitted external app description */ -export const zGetPermittedExternalAppsByAppIdDescribeResponse = zAppDescribeResponse +export const zGetPermittedExternalAppsByAppIdResponse = zAppDescribeResponse /** * Workspace list @@ -975,9 +976,9 @@ export const zDeleteWorkspacesByWorkspaceIdMembersByMemberIdPath = z.object({ */ export const zDeleteWorkspacesByWorkspaceIdMembersByMemberIdResponse = zMemberActionResponse -export const zPutWorkspacesByWorkspaceIdMembersByMemberIdRoleBody = zMemberRoleUpdatePayload +export const zPatchWorkspacesByWorkspaceIdMembersByMemberIdBody = zMemberRoleUpdatePayload -export const zPutWorkspacesByWorkspaceIdMembersByMemberIdRolePath = z.object({ +export const zPatchWorkspacesByWorkspaceIdMembersByMemberIdPath = z.object({ member_id: z.string(), workspace_id: z.string(), }) @@ -985,7 +986,7 @@ export const zPutWorkspacesByWorkspaceIdMembersByMemberIdRolePath = z.object({ /** * Role updated */ -export const zPutWorkspacesByWorkspaceIdMembersByMemberIdRoleResponse = zMemberActionResponse +export const zPatchWorkspacesByWorkspaceIdMembersByMemberIdResponse = zMemberActionResponse export const zPostWorkspacesByWorkspaceIdSwitchPath = z.object({ workspace_id: z.string(), diff --git a/packages/contracts/openapi-ts.api.config.ts b/packages/contracts/openapi-ts.api.config.ts index efb1a226f1e..0fb424bcfb5 100644 --- a/packages/contracts/openapi-ts.api.config.ts +++ b/packages/contracts/openapi-ts.api.config.ts @@ -102,11 +102,11 @@ const segmentWords = (segment: string) => { return toWords(segment) } +// Split on `:` too so custom methods nest as their own node (apps.byAppId.run), not apps.appIdRun. +const routeNamingSegments = (routePath: string) => routePath.split(/[/:]/).filter(Boolean) + const routeWords = (routePath: string) => { - return routePath - .split('/') - .filter(Boolean) - .flatMap(segmentWords) + return routeNamingSegments(routePath).flatMap(segmentWords) } const operationId = (method: string, routePath: string) => { @@ -114,10 +114,7 @@ const operationId = (method: string, routePath: string) => { } const contractPathSegments = (operation: ApiContractOperation) => { - const segments = operation.path - .split('/') - .filter(Boolean) - .map(segment => toCamelCase(segmentWords(segment))) + const segments = routeNamingSegments(operation.path).map(segment => toCamelCase(segmentWords(segment))) return [...(segments.length > 0 ? segments : ['root']), operation.method.toLowerCase()] } From ee0068eed423a414dd2a43913624a979dffcbecf Mon Sep 17 00:00:00 2001 From: Stephen Zhou Date: Wed, 8 Jul 2026 10:35:54 +0800 Subject: [PATCH 35/70] refactor(web): migrate dataset access context (#38523) Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- .../__tests__/layout-main.spec.tsx | 30 ++++ .../[datasetId]/layout-main.tsx | 28 ++-- .../(commonLayout)/datasets/layout.spec.tsx | 12 ++ web/app/(commonLayout)/datasets/layout.tsx | 16 +- .../datasets/__tests__/mock-dataset-access.ts | 142 ++++++++++++++++++ .../access-config/__tests__/index.spec.tsx | 14 ++ .../datasets/access-config/index.tsx | 20 +-- .../datasets/create/__tests__/index.spec.tsx | 16 ++ web/app/components/datasets/create/index.tsx | 13 +- .../documents/__tests__/index.spec.tsx | 15 ++ .../document-list/__tests__/index.spec.tsx | 15 ++ .../components/document-table-row.tsx | 10 +- .../datasets/documents/components/list.tsx | 10 +- .../__tests__/index.spec.tsx | 16 ++ .../documents/create-from-pipeline/index.tsx | 13 +- .../datasets/documents/detail/index.tsx | 10 +- .../components/datasets/documents/index.tsx | 10 +- .../datasets/extra-info/api-access/card.tsx | 10 +- .../service-api/__tests__/index.spec.tsx | 14 ++ .../datasets/extra-info/service-api/index.tsx | 5 +- .../hit-testing/__tests__/index.spec.tsx | 12 ++ .../components/datasets/hit-testing/index.tsx | 14 +- .../datasets/list/__tests__/index.spec.tsx | 33 ++++ .../dataset-card/__tests__/index.spec.tsx | 12 ++ .../__tests__/operations-dropdown.spec.tsx | 14 ++ .../components/operations-dropdown.tsx | 16 +- .../datasets/list/dataset-card/index.tsx | 10 +- web/app/components/datasets/list/index.tsx | 10 +- .../settings/form/__tests__/index.spec.tsx | 17 +++ .../__tests__/basic-info-section.spec.tsx | 26 +++- .../hooks/__tests__/use-form-state.spec.ts | 15 ++ .../settings/form/hooks/use-form-state.ts | 10 +- .../__tests__/index.spec.tsx | 26 +++- .../settings/permission-selector/index.tsx | 14 +- web/context/app-context-defaults.ts | 34 +++++ web/context/app-context-state.ts | 29 +++- web/context/app-context.ts | 37 +---- 37 files changed, 629 insertions(+), 119 deletions(-) create mode 100644 web/app/components/datasets/__tests__/mock-dataset-access.ts create mode 100644 web/context/app-context-defaults.ts diff --git a/web/app/(commonLayout)/datasets/(datasetDetailLayout)/[datasetId]/__tests__/layout-main.spec.tsx b/web/app/(commonLayout)/datasets/(datasetDetailLayout)/[datasetId]/__tests__/layout-main.spec.tsx index 1a566074505..7bb628997df 100644 --- a/web/app/(commonLayout)/datasets/(datasetDetailLayout)/[datasetId]/__tests__/layout-main.spec.tsx +++ b/web/app/(commonLayout)/datasets/(datasetDetailLayout)/[datasetId]/__tests__/layout-main.spec.tsx @@ -31,14 +31,44 @@ vi.mock('@/context/app-context', () => ({ userProfile: { id: 'user-1' }, workspacePermissionKeys: [], }), + useSelector: (selector: (state: { + isCurrentWorkspaceDatasetOperator: boolean + isLoadingCurrentWorkspace: boolean + isLoadingWorkspacePermissionKeys: boolean + userProfile: { id: string } + workspacePermissionKeys: string[] + }) => unknown) => selector({ + isCurrentWorkspaceDatasetOperator: false, + isLoadingCurrentWorkspace: false, + isLoadingWorkspacePermissionKeys: false, + userProfile: { id: 'user-1' }, + workspacePermissionKeys: [], + }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createDatasetAccessAtomMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessAtomMock(importOriginal, () => ({ + userProfile: { id: 'user-1' }, + workspacePermissionKeys: [], + }), () => ({ + isRbacEnabled: mockIsRbacEnabled, + })) +}) + vi.mock('@/context/event-emitter', () => ({ useEventEmitterContextContext: () => ({ eventEmitter: undefined, }), })) +vi.mock('jotai', async (importOriginal) => { + const { createDatasetAccessJotaiMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessJotaiMock(importOriginal) +}) + vi.mock('@/hooks/use-document-title', () => ({ default: vi.fn(), })) diff --git a/web/app/(commonLayout)/datasets/(datasetDetailLayout)/[datasetId]/layout-main.tsx b/web/app/(commonLayout)/datasets/(datasetDetailLayout)/[datasetId]/layout-main.tsx index 8a505ee3077..36a9a792ea5 100644 --- a/web/app/(commonLayout)/datasets/(datasetDetailLayout)/[datasetId]/layout-main.tsx +++ b/web/app/(commonLayout)/datasets/(datasetDetailLayout)/[datasetId]/layout-main.tsx @@ -2,14 +2,19 @@ import type { FC } from 'react' import type { DataSet } from '@/models/datasets' import { cn } from '@langgenius/dify-ui/cn' -import { useSuspenseQuery } from '@tanstack/react-query' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useEffect } from 'react' import { useTranslation } from 'react-i18next' import Loading from '@/app/components/base/loading' -import { useAppContext } from '@/context/app-context' +import { + currentWorkspaceLoadingAtom, + datasetRbacEnabledAtom, + userProfileIdAtom, + workspacePermissionKeysAtom, + workspacePermissionKeysLoadingAtom, +} from '@/context/app-context-state' import DatasetDetailContext from '@/context/dataset-detail' -import { systemFeaturesQueryOptions } from '@/features/system-features/client' import useDocumentTitle from '@/hooks/use-document-title' import { usePathname, useRouter } from '@/next/navigation' import { useDatasetDetail } from '@/service/knowledge/use-dataset' @@ -58,23 +63,20 @@ const DatasetDetailLayout: FC = (props) => { const { t } = useTranslation() const router = useRouter() const pathname = usePathname() - const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) - const { - isLoadingCurrentWorkspace, - isLoadingWorkspacePermissionKeys, - userProfile, - workspacePermissionKeys, - } = useAppContext() - const isRbacEnabled = systemFeatures.rbac_enabled + const isLoadingCurrentWorkspace = useAtomValue(currentWorkspaceLoadingAtom) + const isLoadingWorkspacePermissionKeys = useAtomValue(workspacePermissionKeysLoadingAtom) + const isRbacEnabled = useAtomValue(datasetRbacEnabledAtom) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const { data: datasetRes, error, refetch: mutateDatasetRes } = useDatasetDetail(datasetId) const shouldRedirect = shouldRedirectToDatasetList(error) const datasetACLCapabilities = React.useMemo(() => getDatasetACLCapabilities(datasetRes?.permission_keys, { - currentUserId: userProfile?.id, + currentUserId, resourceMaintainer: datasetRes?.maintainer, workspacePermissionKeys, isRbacEnabled, - }), [datasetRes?.maintainer, datasetRes?.permission_keys, isRbacEnabled, userProfile?.id, workspacePermissionKeys]) + }), [datasetRes?.maintainer, datasetRes?.permission_keys, isRbacEnabled, currentUserId, workspacePermissionKeys]) const isAccessConfigPath = pathname.endsWith('/access-config') const isHitTestingPath = pathname.endsWith('/hitTesting') const isPermissionControlledPath = isAccessConfigPath || isHitTestingPath diff --git a/web/app/(commonLayout)/datasets/layout.spec.tsx b/web/app/(commonLayout)/datasets/layout.spec.tsx index 4db962ff301..98ac56c6fce 100644 --- a/web/app/(commonLayout)/datasets/layout.spec.tsx +++ b/web/app/(commonLayout)/datasets/layout.spec.tsx @@ -20,10 +20,22 @@ vi.mock('@/context/app-context', () => ({ useSelector: (selector: (state: AppContextMock) => unknown) => selector(mockUseAppContext()), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createDatasetAccessAtomMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessAtomMock(importOriginal, () => mockUseAppContext()) +}) + vi.mock('@/context/external-api-panel-context', () => ({ ExternalApiPanelProvider: ({ children }: { children: ReactNode }) => <>{children}, })) +vi.mock('jotai', async (importOriginal) => { + const { createDatasetAccessJotaiMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessJotaiMock(importOriginal) +}) + vi.mock('@/context/external-knowledge-api-context', () => ({ ExternalKnowledgeApiProvider: ({ children, enabled }: { children: ReactNode, enabled?: boolean }) => { mockExternalKnowledgeApiProviderEnabled = enabled diff --git a/web/app/(commonLayout)/datasets/layout.tsx b/web/app/(commonLayout)/datasets/layout.tsx index 8f6777dedfb..fa6d3d0569a 100644 --- a/web/app/(commonLayout)/datasets/layout.tsx +++ b/web/app/(commonLayout)/datasets/layout.tsx @@ -1,8 +1,14 @@ 'use client' +import { useAtomValue } from 'jotai' import { useEffect } from 'react' import Loading from '@/app/components/base/loading' -import { useSelector as useAppContextSelector } from '@/context/app-context' +import { + currentWorkspaceIdAtom, + currentWorkspaceLoadingAtom, + workspacePermissionKeysAtom, + workspacePermissionKeysLoadingAtom, +} from '@/context/app-context-state' import { ExternalApiPanelProvider } from '@/context/external-api-panel-context' import { ExternalKnowledgeApiProvider } from '@/context/external-knowledge-api-context' import { usePathname, useRouter } from '@/next/navigation' @@ -21,10 +27,10 @@ const isDatasetExternalConnectPath = (pathname: string) => { } export default function DatasetsLayout({ children }: { children: React.ReactNode }) { - const currentWorkspaceId = useAppContextSelector(state => state.currentWorkspace.id) - const isLoadingCurrentWorkspace = useAppContextSelector(state => state.isLoadingCurrentWorkspace) - const isLoadingWorkspacePermissionKeys = useAppContextSelector(state => state.isLoadingWorkspacePermissionKeys) - const workspacePermissionKeys = useAppContextSelector(state => state.workspacePermissionKeys) + const currentWorkspaceId = useAtomValue(currentWorkspaceIdAtom) + const isLoadingCurrentWorkspace = useAtomValue(currentWorkspaceLoadingAtom) + const isLoadingWorkspacePermissionKeys = useAtomValue(workspacePermissionKeysLoadingAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const router = useRouter() const pathname = usePathname() const isLoadingAccess = isLoadingCurrentWorkspace || !!isLoadingWorkspacePermissionKeys diff --git a/web/app/components/datasets/__tests__/mock-dataset-access.ts b/web/app/components/datasets/__tests__/mock-dataset-access.ts new file mode 100644 index 00000000000..fb6aff4650d --- /dev/null +++ b/web/app/components/datasets/__tests__/mock-dataset-access.ts @@ -0,0 +1,142 @@ +const DATASET_ACCESS_ATOM_KIND = Symbol('dataset-access-atom-kind') + +type DatasetAccessMockState = { + userProfile?: { + id?: string + name?: string + email?: string + avatar?: string + avatar_url?: string + is_password_set?: boolean + } | null + currentWorkspace?: { + id?: string + } | null + isCurrentWorkspaceOwner?: boolean + isLoadingCurrentWorkspace?: boolean + isLoadingWorkspacePermissionKeys?: boolean + workspacePermissionKeys?: string[] +} + +type DatasetAccessMockOptions = { + isRbacEnabled?: boolean +} + +type DatasetAccessAtomKind + = | 'userProfile' + | 'userProfileId' + | 'currentWorkspaceId' + | 'isCurrentWorkspaceOwner' + | 'workspacePermissionKeys' + | 'currentWorkspaceLoading' + | 'workspacePermissionKeysLoading' + | 'datasetRbacEnabled' + +type DatasetAccessMockAtom = { + [DATASET_ACCESS_ATOM_KIND]: DatasetAccessAtomKind +} + +type DatasetAccessMockRegistry = { + getState: () => DatasetAccessMockState + getOptions: () => DatasetAccessMockOptions +} + +const defaultUserProfile = { + id: 'user-1', + name: 'User', + email: 'user@example.com', + avatar: '', + avatar_url: '', + is_password_set: true, +} + +let datasetAccessMockRegistry: DatasetAccessMockRegistry | undefined + +const createMockAtom = ( + kind: DatasetAccessAtomKind, +): DatasetAccessMockAtom => ({ + [DATASET_ACCESS_ATOM_KIND]: kind, +}) + +const isDatasetAccessMockAtom = (atom: unknown): atom is DatasetAccessMockAtom => { + return typeof atom === 'object' && atom !== null && DATASET_ACCESS_ATOM_KIND in atom +} + +const getUserProfile = (state: DatasetAccessMockState) => ({ + ...defaultUserProfile, + ...state.userProfile, +}) + +const getWorkspacePermissionKeys = (state: DatasetAccessMockState) => state.workspacePermissionKeys ?? [] + +export const createDatasetAccessAtomMock = async ( + importOriginal: () => Promise, + getState: () => DatasetAccessMockState, + getOptions: () => DatasetAccessMockOptions = () => ({}), +) => { + const actual = await importOriginal() + datasetAccessMockRegistry = { + getState, + getOptions, + } + + return { + ...actual, + userProfileAtom: createMockAtom('userProfile'), + userProfileIdAtom: createMockAtom('userProfileId'), + currentWorkspaceIdAtom: createMockAtom('currentWorkspaceId'), + isCurrentWorkspaceOwnerAtom: createMockAtom('isCurrentWorkspaceOwner'), + workspacePermissionKeysAtom: createMockAtom('workspacePermissionKeys'), + currentWorkspaceLoadingAtom: createMockAtom('currentWorkspaceLoading'), + workspacePermissionKeysLoadingAtom: createMockAtom('workspacePermissionKeysLoading'), + datasetRbacEnabledAtom: createMockAtom('datasetRbacEnabled'), + } +} + +export const createDatasetAccessJotaiMock = async ( + importOriginal: () => Promise, +) => { + const actual = await importOriginal() + + return { + ...actual, + useAtomValue: (atom: unknown) => { + if (!isDatasetAccessMockAtom(atom)) + return actual.useAtomValue(atom as Parameters[0]) + + if (!datasetAccessMockRegistry) + throw new Error('Dataset access atom mock is not initialized') + + const state = datasetAccessMockRegistry.getState() + const options = datasetAccessMockRegistry.getOptions() + const userProfile = getUserProfile(state) + const workspacePermissionKeys = getWorkspacePermissionKeys(state) + + if (atom[DATASET_ACCESS_ATOM_KIND] === 'userProfile') + return userProfile + + if (atom[DATASET_ACCESS_ATOM_KIND] === 'userProfileId') + return userProfile.id + + if (atom[DATASET_ACCESS_ATOM_KIND] === 'currentWorkspaceId') + return state.currentWorkspace?.id ?? 'workspace-1' + + if (atom[DATASET_ACCESS_ATOM_KIND] === 'isCurrentWorkspaceOwner') + return state.isCurrentWorkspaceOwner ?? false + + if (atom[DATASET_ACCESS_ATOM_KIND] === 'workspacePermissionKeys') + return workspacePermissionKeys + + if (atom[DATASET_ACCESS_ATOM_KIND] === 'currentWorkspaceLoading') + return state.isLoadingCurrentWorkspace ?? false + + if (atom[DATASET_ACCESS_ATOM_KIND] === 'workspacePermissionKeysLoading') + return state.isLoadingWorkspacePermissionKeys ?? false + + if (atom[DATASET_ACCESS_ATOM_KIND] === 'datasetRbacEnabled') + return options.isRbacEnabled ?? true + + throw new Error(`Unsupported dataset access atom: ${atom[DATASET_ACCESS_ATOM_KIND]}`) + }, + } +} diff --git a/web/app/components/datasets/access-config/__tests__/index.spec.tsx b/web/app/components/datasets/access-config/__tests__/index.spec.tsx index eeeec0dcf88..df6d3d7c1fa 100644 --- a/web/app/components/datasets/access-config/__tests__/index.spec.tsx +++ b/web/app/components/datasets/access-config/__tests__/index.spec.tsx @@ -79,6 +79,14 @@ vi.mock('@/context/app-context', () => ({ useSelector: vi.fn((selector: (state: typeof mockAppContextState) => unknown) => selector(mockAppContextState)), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createDatasetAccessAtomMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessAtomMock(importOriginal, () => mockAppContextState, () => ({ + isRbacEnabled: mockIsRbacEnabled, + })) +}) + vi.mock('@/app/components/access-rules-editor', () => ({ default: (props: AccessRulesEditorProps) => { mockAccessRulesEditor.props = props @@ -88,6 +96,12 @@ vi.mock('@/app/components/access-rules-editor', () => ({ }, })) +vi.mock('jotai', async (importOriginal) => { + const { createDatasetAccessJotaiMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessJotaiMock(importOriginal) +}) + describe('DatasetAccessConfigPage', () => { beforeEach(() => { vi.clearAllMocks() diff --git a/web/app/components/datasets/access-config/index.tsx b/web/app/components/datasets/access-config/index.tsx index b1df8e3eea7..86b2aa061f1 100644 --- a/web/app/components/datasets/access-config/index.tsx +++ b/web/app/components/datasets/access-config/index.tsx @@ -2,15 +2,18 @@ import type { ResourceOpenScope } from '@/models/access-control' import { ScrollArea } from '@langgenius/dify-ui/scroll-area' -import { useSuspenseQuery } from '@tanstack/react-query' +import { useAtomValue } from 'jotai' import { useCallback, useMemo, useState } from 'react' import { useTranslation } from 'react-i18next' import AccessRulesEditor from '@/app/components/access-rules-editor' import Loading from '@/app/components/base/loading' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { + datasetRbacEnabledAtom, + userProfileIdAtom, + workspacePermissionKeysAtom, +} from '@/context/app-context-state' import { useDatasetDetailContextWithSelector } from '@/context/dataset-detail' import { useLocale } from '@/context/i18n' -import { systemFeaturesQueryOptions } from '@/features/system-features/client' import { getAccessControlTemplateLanguage } from '@/i18n-config/language' import { useDatasetAccessRules, @@ -30,16 +33,15 @@ const DatasetAccessConfigPage = ({ datasetId }: DatasetAccessConfigPageProps) => const locale = useLocale() const language = useMemo(() => getAccessControlTemplateLanguage(locale), [locale]) const dataset = useDatasetDetailContextWithSelector(state => state.dataset) - const currentUserId = useAppContextWithSelector(state => state.userProfile?.id) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) - const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) - const isRbacEnabled = systemFeatures.rbac_enabled - const canAccessConfig = useMemo(() => getDatasetACLCapabilities(dataset?.permission_keys, { + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) + const isRbacEnabled = useAtomValue(datasetRbacEnabledAtom) + const canAccessConfig = getDatasetACLCapabilities(dataset?.permission_keys, { currentUserId, resourceMaintainer: dataset?.maintainer, workspacePermissionKeys, isRbacEnabled, - }).canAccessConfig, [currentUserId, dataset?.maintainer, dataset?.permission_keys, isRbacEnabled, workspacePermissionKeys]) + }).canAccessConfig const { data: datasetAccessRulesResponse, isLoading: isLoadingDatasetAccessRules } = useDatasetAccessRules(datasetId, language, { enabled: canAccessConfig }) const { data: datasetUserAccessSettingsResponse, isLoading: isLoadingDatasetUserAccessSettings } = useDatasetUserAccessSettings(datasetId, language, { enabled: canAccessConfig }) const { mutate: updateDatasetOpenScope, isPending: isUpdatingDatasetOpenScope } = useUpdateDatasetOpenScope(datasetId) diff --git a/web/app/components/datasets/create/__tests__/index.spec.tsx b/web/app/components/datasets/create/__tests__/index.spec.tsx index 7668e778d94..04a8c68f9c9 100644 --- a/web/app/components/datasets/create/__tests__/index.spec.tsx +++ b/web/app/components/datasets/create/__tests__/index.spec.tsx @@ -54,6 +54,16 @@ vi.mock('@/context/app-context', () => ({ }, })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createDatasetAccessAtomMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessAtomMock(importOriginal, () => ({ + userProfile: { id: mockCurrentUserId }, + workspacePermissionKeys: mockWorkspacePermissionKeys, + isLoadingWorkspacePermissionKeys: mockIsLoadingWorkspacePermissionKeys, + })) +}) + // Mock modal context const mockSetShowAccountSettingModal = vi.fn() vi.mock('@/context/modal-context', () => ({ @@ -68,6 +78,12 @@ vi.mock('@/context/modal-context', () => ({ }, })) +vi.mock('jotai', async (importOriginal) => { + const { createDatasetAccessJotaiMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessJotaiMock(importOriginal) +}) + // Mock dataset detail context let mockDatasetDetail: DataSet | undefined vi.mock('@/context/dataset-detail', () => ({ diff --git a/web/app/components/datasets/create/index.tsx b/web/app/components/datasets/create/index.tsx index 1f41f7b456c..ac9653cdff7 100644 --- a/web/app/components/datasets/create/index.tsx +++ b/web/app/components/datasets/create/index.tsx @@ -3,6 +3,7 @@ import type { NotionPage } from '@/models/common' import type { CrawlOptions, CrawlResultItem, createDocumentResponse, FileItem } from '@/models/datasets' import type { RETRIEVE_METHOD } from '@/types/app' import { produce } from 'immer' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useCallback, useEffect, useState } from 'react' import { useTranslation } from 'react-i18next' @@ -10,7 +11,11 @@ import Loading from '@/app/components/base/loading' import { ACCOUNT_SETTING_TAB } from '@/app/components/header/account-setting/constants' import { useDefaultModel } from '@/app/components/header/account-setting/model-provider-page/hooks' import { useIntegrationsSetting } from '@/app/components/header/account-setting/use-integrations-setting' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { + userProfileIdAtom, + workspacePermissionKeysAtom, + workspacePermissionKeysLoadingAtom, +} from '@/context/app-context-state' import { useDatasetDetailContextWithSelector } from '@/context/dataset-detail' import { DataSourceProvider } from '@/models/common' import { DataSourceType } from '@/models/datasets' @@ -43,9 +48,9 @@ const DatasetUpdateForm = ({ datasetId }: DatasetUpdateFormProps) => { const router = useRouter() const openIntegrationsSetting = useIntegrationsSetting() const datasetDetail = useDatasetDetailContextWithSelector(state => state.dataset) - const currentUserId = useAppContextWithSelector(state => state.userProfile?.id) - const isLoadingWorkspacePermissionKeys = useAppContextWithSelector(state => state.isLoadingWorkspacePermissionKeys) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const currentUserId = useAtomValue(userProfileIdAtom) + const isLoadingWorkspacePermissionKeys = useAtomValue(workspacePermissionKeysLoadingAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const { data: embeddingsDefaultModel } = useDefaultModel(ModelTypeEnum.textEmbedding) const canAddDocumentsToDataset = !datasetId || getDatasetACLCapabilities(datasetDetail?.permission_keys, { currentUserId, diff --git a/web/app/components/datasets/documents/__tests__/index.spec.tsx b/web/app/components/datasets/documents/__tests__/index.spec.tsx index fa09e9c0766..de6e84e2ee5 100644 --- a/web/app/components/datasets/documents/__tests__/index.spec.tsx +++ b/web/app/components/datasets/documents/__tests__/index.spec.tsx @@ -55,6 +55,15 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createDatasetAccessAtomMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessAtomMock(importOriginal, () => ({ + userProfile: { id: 'test-user' }, + workspacePermissionKeys: ['dataset.create_and_management'], + })) +}) + // Mock document service hooks const mockInvalidDocumentList = vi.fn() const mockInvalidDocumentDetail = vi.fn() @@ -92,6 +101,12 @@ vi.mock('@/service/knowledge/use-document', () => ({ useInvalidDocumentDetail: vi.fn(() => mockInvalidDocumentDetail), })) +vi.mock('jotai', async (importOriginal) => { + const { createDatasetAccessJotaiMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessJotaiMock(importOriginal) +}) + // Mock segment service hooks vi.mock('@/service/knowledge/use-segment', () => ({ useSegmentListKey: 'segment-list-key', diff --git a/web/app/components/datasets/documents/components/document-list/__tests__/index.spec.tsx b/web/app/components/datasets/documents/components/document-list/__tests__/index.spec.tsx index 3840223e2db..632e9b16fcb 100644 --- a/web/app/components/datasets/documents/components/document-list/__tests__/index.spec.tsx +++ b/web/app/components/datasets/documents/components/document-list/__tests__/index.spec.tsx @@ -42,6 +42,15 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createDatasetAccessAtomMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessAtomMock(importOriginal, () => ({ + userProfile: { id: 'user-1' }, + workspacePermissionKeys: ['dataset.create_and_management'], + })) +}) + vi.mock('@/app/components/datasets/metadata/hooks/use-batch-edit-document-metadata', () => ({ default: () => ({ isShowEditModal: false, @@ -52,6 +61,12 @@ vi.mock('@/app/components/datasets/metadata/hooks/use-batch-edit-document-metada }), })) +vi.mock('jotai', async (importOriginal) => { + const { createDatasetAccessJotaiMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessJotaiMock(importOriginal) +}) + const createTestQueryClient = () => new QueryClient({ defaultOptions: { queries: { retry: false, gcTime: 0 }, diff --git a/web/app/components/datasets/documents/components/document-list/components/document-table-row.tsx b/web/app/components/datasets/documents/components/document-list/components/document-table-row.tsx index 7fedde10168..0961eb2bcd9 100644 --- a/web/app/components/datasets/documents/components/document-list/components/document-table-row.tsx +++ b/web/app/components/datasets/documents/components/document-list/components/document-table-row.tsx @@ -2,6 +2,7 @@ import type { SimpleDocumentDetail } from '@/models/datasets' import { Checkbox } from '@langgenius/dify-ui/checkbox' import { Tooltip, TooltipContent, TooltipTrigger } from '@langgenius/dify-ui/tooltip' import { pick } from 'es-toolkit/object' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useCallback } from 'react' import { useTranslation } from 'react-i18next' @@ -9,7 +10,10 @@ import ChunkingModeLabel from '@/app/components/datasets/common/chunking-mode-la import Operations from '@/app/components/datasets/documents/components/operations' import SummaryStatus from '@/app/components/datasets/documents/detail/completed/common/summary-status' import StatusItem from '@/app/components/datasets/documents/status-item' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { + userProfileIdAtom, + workspacePermissionKeysAtom, +} from '@/context/app-context-state' import { useDatasetDetailContextWithSelector } from '@/context/dataset-detail' import useTimestamp from '@/hooks/use-timestamp' import { DataSourceType } from '@/models/datasets' @@ -62,8 +66,8 @@ const DocumentTableRow = React.memo(({ const searchParams = useSearchParams() const documentNameId = React.useId() const dataset = useDatasetDetailContextWithSelector(s => s.dataset) - const currentUserId = useAppContextWithSelector(state => state.userProfile?.id) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const datasetACLCapabilities = React.useMemo(() => getDatasetACLCapabilities(dataset?.permission_keys, { currentUserId, resourceMaintainer: dataset?.maintainer, diff --git a/web/app/components/datasets/documents/components/list.tsx b/web/app/components/datasets/documents/components/list.tsx index ef965554b9c..d2d56f57cba 100644 --- a/web/app/components/datasets/documents/components/list.tsx +++ b/web/app/components/datasets/documents/components/list.tsx @@ -4,11 +4,15 @@ import { Checkbox } from '@langgenius/dify-ui/checkbox' import { CheckboxGroup } from '@langgenius/dify-ui/checkbox-group' import { Pagination } from '@langgenius/dify-ui/pagination' import { useBoolean } from 'ahooks' +import { useAtomValue } from 'jotai' import { useCallback, useMemo, useState } from 'react' import { useTranslation } from 'react-i18next' import EditMetadataBatchModal from '@/app/components/datasets/metadata/edit-metadata-batch/modal' import useBatchEditDocumentMetadata from '@/app/components/datasets/metadata/hooks/use-batch-edit-document-metadata' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { + userProfileIdAtom, + workspacePermissionKeysAtom, +} from '@/context/app-context-state' import { useDatasetDetailContextWithSelector as useDatasetDetailContext } from '@/context/dataset-detail' import { ChunkingMode, DocumentActionType } from '@/models/datasets' import { getDatasetACLCapabilities } from '@/utils/permission' @@ -61,8 +65,8 @@ const DocumentList = ({ const pageSize = pagination.limit ?? 10 const totalPages = Math.max(Math.ceil(pagination.total / pageSize), 1) const datasetConfig = useDatasetDetailContext(s => s.dataset) - const currentUserId = useAppContextWithSelector(state => state.userProfile?.id) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const datasetACLCapabilities = useMemo(() => getDatasetACLCapabilities(datasetConfig?.permission_keys, { currentUserId, resourceMaintainer: datasetConfig?.maintainer, diff --git a/web/app/components/datasets/documents/create-from-pipeline/__tests__/index.spec.tsx b/web/app/components/datasets/documents/create-from-pipeline/__tests__/index.spec.tsx index b6ec26b923c..8e699008b3b 100644 --- a/web/app/components/datasets/documents/create-from-pipeline/__tests__/index.spec.tsx +++ b/web/app/components/datasets/documents/create-from-pipeline/__tests__/index.spec.tsx @@ -63,6 +63,16 @@ vi.mock('@/context/app-context', () => ({ }, })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createDatasetAccessAtomMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessAtomMock(importOriginal, () => ({ + userProfile: { id: mockCurrentUserId }, + workspacePermissionKeys: mockWorkspacePermissionKeys, + isLoadingWorkspacePermissionKeys: mockIsLoadingWorkspacePermissionKeys, + })) +}) + vi.mock('@/service/use-billing', () => ({ useCurrentPlanVectorSpace: () => ({ data: { @@ -73,6 +83,12 @@ vi.mock('@/service/use-billing', () => ({ }), })) +vi.mock('jotai', async (importOriginal) => { + const { createDatasetAccessJotaiMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessJotaiMock(importOriginal) +}) + vi.mock('@/context/dataset-detail', () => ({ useDatasetDetailContextWithSelector: (selector: (state: { dataset: { id: string, pipeline_id: string, permission_keys: string[] } }) => unknown) => selector({ dataset: { id: 'test-dataset-id', pipeline_id: 'test-pipeline-id', permission_keys: mockDatasetPermissionKeys } }), diff --git a/web/app/components/datasets/documents/create-from-pipeline/index.tsx b/web/app/components/datasets/documents/create-from-pipeline/index.tsx index d83e9ae501d..6ee21ae9d7d 100644 --- a/web/app/components/datasets/documents/create-from-pipeline/index.tsx +++ b/web/app/components/datasets/documents/create-from-pipeline/index.tsx @@ -5,11 +5,16 @@ import type { Node } from '@/app/components/workflow/types' import type { FileIndexingEstimateResponse } from '@/models/datasets' import type { InitialDocumentDetail } from '@/models/pipeline' import { useBoolean } from 'ahooks' +import { useAtomValue } from 'jotai' import { useCallback, useEffect, useMemo, useState } from 'react' import { useTranslation } from 'react-i18next' import Loading from '@/app/components/base/loading' import { PlanUpgradeModal } from '@/app/components/billing/plan-upgrade-modal' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { + userProfileIdAtom, + workspacePermissionKeysAtom, + workspacePermissionKeysLoadingAtom, +} from '@/context/app-context-state' import { useDatasetDetailContextWithSelector } from '@/context/dataset-detail' import { useProviderContextSelector } from '@/context/provider-context' import { DatasourceType } from '@/models/pipeline' @@ -40,9 +45,9 @@ const CreateFormPipeline = () => { const enableBilling = useProviderContextSelector(state => state.enableBilling) const dataset = useDatasetDetailContextWithSelector(s => s.dataset) const pipelineId = dataset?.pipeline_id - const currentUserId = useAppContextWithSelector(state => state.userProfile?.id) - const isLoadingWorkspacePermissionKeys = useAppContextWithSelector(state => state.isLoadingWorkspacePermissionKeys) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const currentUserId = useAtomValue(userProfileIdAtom) + const isLoadingWorkspacePermissionKeys = useAtomValue(workspacePermissionKeysLoadingAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const dataSourceStore = useDataSourceStore() const canAddDocumentsToDataset = getDatasetACLCapabilities(dataset?.permission_keys, { currentUserId, diff --git a/web/app/components/datasets/documents/detail/index.tsx b/web/app/components/datasets/documents/detail/index.tsx index 1d7ad17d07e..31de23ccef1 100644 --- a/web/app/components/datasets/documents/detail/index.tsx +++ b/web/app/components/datasets/documents/detail/index.tsx @@ -4,6 +4,7 @@ import type { DocumentDisplayStatus, FileItem, FullDocumentDetail } from '@/mode import type { SegmentImportStatus } from '@/types/dataset' import { cn } from '@langgenius/dify-ui/cn' import { toast } from '@langgenius/dify-ui/toast' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useCallback, useMemo, useState } from 'react' import { useTranslation } from 'react-i18next' @@ -11,7 +12,10 @@ import Divider from '@/app/components/base/divider' import FloatRightContainer from '@/app/components/base/float-right-container' import Loading from '@/app/components/base/loading' import Metadata from '@/app/components/datasets/metadata/metadata-document' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { + userProfileIdAtom, + workspacePermissionKeysAtom, +} from '@/context/app-context-state' import { useDatasetDetailContextWithSelector } from '@/context/dataset-detail' import useBreakpoints, { MediaType } from '@/hooks/use-breakpoints' import { ChunkingMode, DisplayStatusList } from '@/models/datasets' @@ -49,8 +53,8 @@ const DocumentDetail: FC = ({ datasetId, documentId }) => { const isMobile = media === MediaType.mobile const dataset = useDatasetDetailContextWithSelector(s => s.dataset) - const currentUserId = useAppContextWithSelector(state => state.userProfile?.id) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const embeddingAvailable = !!dataset?.embedding_available const datasetACLCapabilities = useMemo( () => getDatasetACLCapabilities(dataset?.permission_keys, { diff --git a/web/app/components/datasets/documents/index.tsx b/web/app/components/datasets/documents/index.tsx index 33b27c1c46e..d64fefe1325 100644 --- a/web/app/components/datasets/documents/index.tsx +++ b/web/app/components/datasets/documents/index.tsx @@ -1,8 +1,12 @@ 'use client' import type { FC } from 'react' +import { useAtomValue } from 'jotai' import { useCallback } from 'react' import Loading from '@/app/components/base/loading' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { + userProfileIdAtom, + workspacePermissionKeysAtom, +} from '@/context/app-context-state' import { useDatasetDetailContextWithSelector } from '@/context/dataset-detail' import { useProviderContext } from '@/context/provider-context' import { DataSourceType } from '@/models/datasets' @@ -31,8 +35,8 @@ const Documents: FC = ({ datasetId }) => { const isFreePlan = plan.type === 'sandbox' const dataset = useDatasetDetailContextWithSelector(s => s.dataset) - const currentUserId = useAppContextWithSelector(state => state.userProfile?.id) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const embeddingAvailable = !!dataset?.embedding_available const datasetACLCapabilities = getDatasetACLCapabilities(dataset?.permission_keys, { currentUserId, diff --git a/web/app/components/datasets/extra-info/api-access/card.tsx b/web/app/components/datasets/extra-info/api-access/card.tsx index 3a5685e8fbb..13b65ca612d 100644 --- a/web/app/components/datasets/extra-info/api-access/card.tsx +++ b/web/app/components/datasets/extra-info/api-access/card.tsx @@ -1,10 +1,14 @@ import { cn } from '@langgenius/dify-ui/cn' import { StatusDot } from '@langgenius/dify-ui/status-dot' import { Switch } from '@langgenius/dify-ui/switch' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useCallback } from 'react' import { useTranslation } from 'react-i18next' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { + userProfileIdAtom, + workspacePermissionKeysAtom, +} from '@/context/app-context-state' import { useDatasetDetailContextWithSelector } from '@/context/dataset-detail' import { useDatasetApiAccessUrl } from '@/hooks/use-api-access-url' import Link from '@/next/link' @@ -22,8 +26,8 @@ const Card = ({ const datasetId = useDatasetDetailContextWithSelector(state => state.dataset?.id) const dataset = useDatasetDetailContextWithSelector(state => state.dataset) const mutateDatasetRes = useDatasetDetailContextWithSelector(state => state.mutateDatasetRes) - const currentUserId = useAppContextWithSelector(state => state.userProfile?.id) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const { mutateAsync: enableDatasetServiceApi } = useEnableDatasetServiceApi() const { mutateAsync: disableDatasetServiceApi } = useDisableDatasetServiceApi() diff --git a/web/app/components/datasets/extra-info/service-api/__tests__/index.spec.tsx b/web/app/components/datasets/extra-info/service-api/__tests__/index.spec.tsx index 06bb88e91ef..1a5ac70964a 100644 --- a/web/app/components/datasets/extra-info/service-api/__tests__/index.spec.tsx +++ b/web/app/components/datasets/extra-info/service-api/__tests__/index.spec.tsx @@ -15,6 +15,14 @@ vi.mock('@/context/app-context', () => ({ selector({ workspacePermissionKeys: mockWorkspacePermissionKeys }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createDatasetAccessAtomMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockWorkspacePermissionKeys, + })) +}) + vi.mock('@/next/navigation', () => ({ useRouter: () => ({ push: vi.fn(), @@ -24,6 +32,12 @@ vi.mock('@/next/navigation', () => ({ useSearchParams: () => new URLSearchParams(), })) +vi.mock('jotai', async (importOriginal) => { + const { createDatasetAccessJotaiMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessJotaiMock(importOriginal) +}) + // Mock next/link vi.mock('@/next/link', () => ({ default: ({ children, href, ...props }: { children: React.ReactNode, href: string, [key: string]: unknown }) => ( diff --git a/web/app/components/datasets/extra-info/service-api/index.tsx b/web/app/components/datasets/extra-info/service-api/index.tsx index 679e80c5df3..e48b3c2eb7a 100644 --- a/web/app/components/datasets/extra-info/service-api/index.tsx +++ b/web/app/components/datasets/extra-info/service-api/index.tsx @@ -1,11 +1,12 @@ import { cn } from '@langgenius/dify-ui/cn' import { Popover, PopoverContent, PopoverTrigger } from '@langgenius/dify-ui/popover' import { StatusDot } from '@langgenius/dify-ui/status-dot' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useCallback, useState } from 'react' import { useTranslation } from 'react-i18next' import SecretKeyModal from '@/app/components/develop/secret-key/secret-key-modal' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { workspacePermissionKeysAtom } from '@/context/app-context-state' import { hasPermission } from '@/utils/permission' import Card from './card' @@ -19,7 +20,7 @@ const ServiceApi = ({ const { t } = useTranslation() const [open, setOpen] = useState(false) const [isSecretKeyModalVisible, setIsSecretKeyModalVisible] = useState(false) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const canManageSecretKey = hasPermission(workspacePermissionKeys, 'dataset.api_key.manage') const handleOpenSecretKeyModal = useCallback(() => { diff --git a/web/app/components/datasets/hit-testing/__tests__/index.spec.tsx b/web/app/components/datasets/hit-testing/__tests__/index.spec.tsx index 50ba6fc7d6e..fa20d559603 100644 --- a/web/app/components/datasets/hit-testing/__tests__/index.spec.tsx +++ b/web/app/components/datasets/hit-testing/__tests__/index.spec.tsx @@ -83,6 +83,12 @@ vi.mock('@/context/app-context', () => ({ useSelector: (selector: (state: typeof mockAppContextState) => unknown) => selector(mockAppContextState), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createDatasetAccessAtomMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessAtomMock(importOriginal, () => mockAppContextState) +}) + const mockRecordsRefetch = vi.fn() const mockHitTestingMutateAsync = vi.fn() const mockExternalHitTestingMutateAsync = vi.fn() @@ -101,6 +107,12 @@ vi.mock('@/service/knowledge/use-dataset', () => ({ })), })) +vi.mock('jotai', async (importOriginal) => { + const { createDatasetAccessJotaiMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessJotaiMock(importOriginal) +}) + vi.mock('@/service/knowledge/use-hit-testing', () => ({ useHitTesting: vi.fn(() => ({ mutateAsync: mockHitTestingMutateAsync, diff --git a/web/app/components/datasets/hit-testing/index.tsx b/web/app/components/datasets/hit-testing/index.tsx index 3324a497838..92332e97e95 100644 --- a/web/app/components/datasets/hit-testing/index.tsx +++ b/web/app/components/datasets/hit-testing/index.tsx @@ -20,6 +20,7 @@ import { } from '@langgenius/dify-ui/drawer' import { Pagination } from '@langgenius/dify-ui/pagination' import { useBoolean } from 'ahooks' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useCallback, useEffect, useState } from 'react' import { useTranslation } from 'react-i18next' @@ -27,7 +28,10 @@ import { useContext } from 'use-context-selector' import FloatRightContainer from '@/app/components/base/float-right-container' import Loading from '@/app/components/base/loading' import docStyle from '@/app/components/datasets/documents/detail/completed/style.module.css' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { + userProfileIdAtom, + workspacePermissionKeysAtom, +} from '@/context/app-context-state' import DatasetDetailContext from '@/context/dataset-detail' import useBreakpoints, { MediaType } from '@/hooks/use-breakpoints' import { useDatasetTestingRecords } from '@/service/knowledge/use-dataset' @@ -63,13 +67,13 @@ const HitTestingPage: FC = ({ datasetId }: Props) => { const [currPage, setCurrPage] = useState(0) const { dataset: currentDataset } = useContext(DatasetDetailContext) - const currentUserId = useAppContextWithSelector(state => state.userProfile?.id) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) - const canRunRetrievalRecall = React.useMemo(() => getDatasetACLCapabilities(currentDataset?.permission_keys, { + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) + const canRunRetrievalRecall = getDatasetACLCapabilities(currentDataset?.permission_keys, { currentUserId, resourceMaintainer: currentDataset?.maintainer, workspacePermissionKeys, - }).canRetrievalRecall, [currentDataset?.maintainer, currentDataset?.permission_keys, currentUserId, workspacePermissionKeys]) + }).canRetrievalRecall const { data: recordsRes, refetch: recordsRefetch, isLoading: isRecordsLoading } = useDatasetTestingRecords(datasetId, { limit, page: currPage + 1 }, { enabled: canRunRetrievalRecall }) const total = recordsRes?.total || 0 diff --git a/web/app/components/datasets/list/__tests__/index.spec.tsx b/web/app/components/datasets/list/__tests__/index.spec.tsx index 76dfb03a9ed..9b6a71a0407 100644 --- a/web/app/components/datasets/list/__tests__/index.spec.tsx +++ b/web/app/components/datasets/list/__tests__/index.spec.tsx @@ -8,6 +8,7 @@ const mockReplace = vi.fn() let mockAppContextState = { isCurrentWorkspaceEditor: true, isCurrentWorkspaceManager: true, + isCurrentWorkspaceOwner: true, workspacePermissionKeys: ['dataset.create_and_management', 'dataset.external.connect'], } let mockIsCurrentWorkspaceOwner = true @@ -27,6 +28,12 @@ vi.mock('@/context/app-context', () => ({ useSelector: (selector: (state: typeof mockAppContextState) => unknown) => selector(mockAppContextState), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createDatasetAccessAtomMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessAtomMock(importOriginal, () => mockAppContextState) +}) + // Mock external api panel context const mockSetShowExternalApiPanel = vi.fn() vi.mock('@/context/external-api-panel-context', () => ({ @@ -36,6 +43,12 @@ vi.mock('@/context/external-api-panel-context', () => ({ }), })) +vi.mock('jotai', async (importOriginal) => { + const { createDatasetAccessJotaiMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessJotaiMock(importOriginal) +}) + // Mock useDocumentTitle hook vi.mock('@/hooks/use-document-title', () => ({ default: vi.fn(), @@ -132,6 +145,7 @@ describe('List', () => { mockAppContextState = { isCurrentWorkspaceEditor: true, isCurrentWorkspaceManager: true, + isCurrentWorkspaceOwner: true, workspacePermissionKeys: ['dataset.create_and_management', 'dataset.external.connect'], } mockIsCurrentWorkspaceOwner = true @@ -170,6 +184,7 @@ describe('List', () => { mockAppContextState = { isCurrentWorkspaceEditor: true, isCurrentWorkspaceManager: true, + isCurrentWorkspaceOwner: true, workspacePermissionKeys: ['dataset.create_and_management'], } @@ -282,6 +297,7 @@ describe('List', () => { mockAppContextState = { isCurrentWorkspaceEditor: false, isCurrentWorkspaceManager: true, + isCurrentWorkspaceOwner: true, workspacePermissionKeys: ['dataset.create_and_management'], } const { useDatasetList } = await import('@/service/knowledge/use-dataset') @@ -303,6 +319,7 @@ describe('List', () => { mockAppContextState = { isCurrentWorkspaceEditor: true, isCurrentWorkspaceManager: true, + isCurrentWorkspaceOwner: true, workspacePermissionKeys: [], } const { useDatasetList } = await import('@/service/knowledge/use-dataset') @@ -367,6 +384,7 @@ describe('List', () => { useSelector: (selector: (state: typeof mockAppContextState) => unknown) => selector({ isCurrentWorkspaceEditor: false, isCurrentWorkspaceManager: false, + isCurrentWorkspaceOwner: false, workspacePermissionKeys: ['dataset.create_and_management', 'dataset.external.connect'], }), })) @@ -417,6 +435,12 @@ describe('List', () => { }) it('should not show ExternalAPIPanel without dataset.external.connect even when panel state is open', async () => { + mockAppContextState = { + isCurrentWorkspaceEditor: true, + isCurrentWorkspaceManager: true, + isCurrentWorkspaceOwner: true, + workspacePermissionKeys: ['dataset.create_and_management'], + } vi.doMock('@/context/app-context', () => ({ useAppContext: () => ({ currentWorkspace: { role: 'admin' }, @@ -425,6 +449,7 @@ describe('List', () => { useSelector: (selector: (state: typeof mockAppContextState) => unknown) => selector({ isCurrentWorkspaceEditor: true, isCurrentWorkspaceManager: true, + isCurrentWorkspaceOwner: true, workspacePermissionKeys: ['dataset.create_and_management'], }), })) @@ -452,6 +477,7 @@ describe('List', () => { useSelector: (selector: (state: typeof mockAppContextState) => unknown) => selector({ isCurrentWorkspaceEditor: true, isCurrentWorkspaceManager: true, + isCurrentWorkspaceOwner: true, workspacePermissionKeys: ['dataset.create_and_management', 'dataset.external.connect'], }), })) @@ -481,6 +507,12 @@ describe('List', () => { }) it('should not show include all checkbox when not workspace owner', async () => { + mockAppContextState = { + isCurrentWorkspaceEditor: true, + isCurrentWorkspaceManager: true, + isCurrentWorkspaceOwner: false, + workspacePermissionKeys: ['dataset.create_and_management', 'dataset.external.connect'], + } vi.doMock('@/context/app-context', () => ({ useAppContext: () => ({ currentWorkspace: { role: 'editor' }, @@ -489,6 +521,7 @@ describe('List', () => { useSelector: (selector: (state: typeof mockAppContextState) => unknown) => selector({ isCurrentWorkspaceEditor: true, isCurrentWorkspaceManager: true, + isCurrentWorkspaceOwner: false, workspacePermissionKeys: ['dataset.create_and_management', 'dataset.external.connect'], }), })) diff --git a/web/app/components/datasets/list/dataset-card/__tests__/index.spec.tsx b/web/app/components/datasets/list/dataset-card/__tests__/index.spec.tsx index bbe58038d6e..d1735b299e0 100644 --- a/web/app/components/datasets/list/dataset-card/__tests__/index.spec.tsx +++ b/web/app/components/datasets/list/dataset-card/__tests__/index.spec.tsx @@ -53,6 +53,12 @@ vi.mock('@/context/app-context', () => ({ useSelector: (selector: (state: typeof mockAppContextState) => unknown) => selector(mockAppContextState), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createDatasetAccessAtomMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessAtomMock(importOriginal, () => mockAppContextState) +}) + vi.mock('../hooks/use-dataset-card-state', () => ({ useDatasetCardState: () => ({ modalState: { @@ -72,6 +78,12 @@ vi.mock('../hooks/use-dataset-card-state', () => ({ }), })) +vi.mock('jotai', async (importOriginal) => { + const { createDatasetAccessJotaiMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessJotaiMock(importOriginal) +}) + vi.mock('../components/corner-labels', () => ({ default: () =>
    , })) diff --git a/web/app/components/datasets/list/dataset-card/components/__tests__/operations-dropdown.spec.tsx b/web/app/components/datasets/list/dataset-card/components/__tests__/operations-dropdown.spec.tsx index 6e087158bc6..9620202a47a 100644 --- a/web/app/components/datasets/list/dataset-card/components/__tests__/operations-dropdown.spec.tsx +++ b/web/app/components/datasets/list/dataset-card/components/__tests__/operations-dropdown.spec.tsx @@ -24,6 +24,20 @@ vi.mock('@/context/app-context', () => ({ useSelector: vi.fn((selector: (state: typeof mockAppContextState) => unknown) => selector(mockAppContextState)), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createDatasetAccessAtomMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessAtomMock(importOriginal, () => mockAppContextState, () => ({ + isRbacEnabled: mockIsRbacEnabled, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createDatasetAccessJotaiMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessJotaiMock(importOriginal) +}) + describe('OperationsDropdown', () => { const createMockDataset = (overrides: Partial = {}): DataSet => ({ id: 'dataset-1', diff --git a/web/app/components/datasets/list/dataset-card/components/operations-dropdown.tsx b/web/app/components/datasets/list/dataset-card/components/operations-dropdown.tsx index d5ec29b233e..d340b2e974f 100644 --- a/web/app/components/datasets/list/dataset-card/components/operations-dropdown.tsx +++ b/web/app/components/datasets/list/dataset-card/components/operations-dropdown.tsx @@ -5,10 +5,13 @@ import { DropdownMenuContent, DropdownMenuTrigger, } from '@langgenius/dify-ui/dropdown-menu' -import { useSuspenseQuery } from '@tanstack/react-query' +import { useAtomValue } from 'jotai' import * as React from 'react' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' -import { systemFeaturesQueryOptions } from '@/features/system-features/client' +import { + datasetRbacEnabledAtom, + userProfileIdAtom, + workspacePermissionKeysAtom, +} from '@/context/app-context-state' import { getDatasetACLCapabilities } from '@/utils/permission' import Operations from '../operations' @@ -28,10 +31,9 @@ const OperationsDropdown = ({ openAccessConfig, }: OperationsDropdownProps) => { const [open, setOpen] = React.useState(false) - const currentUserId = useAppContextWithSelector(state => state.userProfile?.id) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) - const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) - const isRbacEnabled = systemFeatures.rbac_enabled + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) + const isRbacEnabled = useAtomValue(datasetRbacEnabledAtom) const datasetACLCapabilities = React.useMemo(() => getDatasetACLCapabilities(dataset.permission_keys, { currentUserId, resourceMaintainer: dataset.maintainer, diff --git a/web/app/components/datasets/list/dataset-card/index.tsx b/web/app/components/datasets/list/dataset-card/index.tsx index 0e03449aaf5..e568247a486 100644 --- a/web/app/components/datasets/list/dataset-card/index.tsx +++ b/web/app/components/datasets/list/dataset-card/index.tsx @@ -3,9 +3,13 @@ import type { KeyboardEvent, MouseEvent } from 'react' import type { DataSet } from '@/models/datasets' import { cn } from '@langgenius/dify-ui/cn' import { toast } from '@langgenius/dify-ui/toast' +import { useAtomValue } from 'jotai' import { useMemo } from 'react' import { useTranslation } from 'react-i18next' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { + userProfileIdAtom, + workspacePermissionKeysAtom, +} from '@/context/app-context-state' import { DatasetCardTags } from '@/features/tag-management/components/dataset-card-tags' import { useRouter } from '@/next/navigation' import { getDatasetACLCapabilities, hasOnlyDatasetPreviewPermission, hasPermission } from '@/utils/permission' @@ -32,8 +36,8 @@ const DatasetCard = ({ }: DatasetCardProps) => { const { t } = useTranslation() const { push } = useRouter() - const currentUserId = useAppContextWithSelector(state => state.userProfile?.id) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const datasetCard = useDatasetCardController({ dataset, onSuccess }) const { diff --git a/web/app/components/datasets/list/index.tsx b/web/app/components/datasets/list/index.tsx index ba1d7cb5565..f079aa3414b 100644 --- a/web/app/components/datasets/list/index.tsx +++ b/web/app/components/datasets/list/index.tsx @@ -1,11 +1,15 @@ 'use client' import { useBoolean, useDebounceFn } from 'ahooks' +import { useAtomValue } from 'jotai' // Libraries import { useState } from 'react' import { useTranslation } from 'react-i18next' -import { useAppContext, useSelector as useAppContextSelector } from '@/context/app-context' +import { + isCurrentWorkspaceOwnerAtom, + workspacePermissionKeysAtom, +} from '@/context/app-context-state' import { useExternalApiPanel } from '@/context/external-api-panel-context' import { TagManagementModal } from '@/features/tag-management/components/tag-management-modal' import useDocumentTitle from '@/hooks/use-document-title' @@ -22,7 +26,7 @@ import DatasetListHeader from './header' const List = () => { const { t } = useTranslation() const { push } = useRouter() - const { isCurrentWorkspaceOwner } = useAppContext() + const isCurrentWorkspaceOwner = useAtomValue(isCurrentWorkspaceOwnerAtom) const [showTagManagementModal, setShowTagManagementModal] = useState(false) const { showExternalApiPanel, setShowExternalApiPanel } = useExternalApiPanel() const [includeAll, { toggle: toggleIncludeAll }] = useBoolean(false) @@ -48,7 +52,7 @@ const List = () => { handleTagsUpdate() } - const workspacePermissionKeys = useAppContextSelector(state => state.workspacePermissionKeys) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const canCreateDataset = hasPermission(workspacePermissionKeys, 'dataset.create_and_management') const canConnectExternalDataset = hasPermission(workspacePermissionKeys, 'dataset.external.connect') const { data: apiBaseInfo } = useDatasetApiBaseUrl() diff --git a/web/app/components/datasets/settings/form/__tests__/index.spec.tsx b/web/app/components/datasets/settings/form/__tests__/index.spec.tsx index b39f4e54706..64795665890 100644 --- a/web/app/components/datasets/settings/form/__tests__/index.spec.tsx +++ b/web/app/components/datasets/settings/form/__tests__/index.spec.tsx @@ -47,6 +47,17 @@ vi.mock('@/context/app-context', () => ({ }, })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createDatasetAccessAtomMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessAtomMock(importOriginal, () => ({ + userProfile: mockUserProfile, + workspacePermissionKeys: mockWorkspacePermissionKeys, + }), () => ({ + isRbacEnabled: false, + })) +}) + const createMockDataset = (overrides: Partial = {}): DataSet => ({ id: 'dataset-1', name: 'Test Dataset', @@ -129,6 +140,12 @@ vi.mock('@/context/dataset-detail', () => ({ }, })) +vi.mock('jotai', async (importOriginal) => { + const { createDatasetAccessJotaiMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessJotaiMock(importOriginal) +}) + // Mock services vi.mock('@/service/datasets', () => ({ updateDatasetSetting: vi.fn().mockResolvedValue({}), diff --git a/web/app/components/datasets/settings/form/components/__tests__/basic-info-section.spec.tsx b/web/app/components/datasets/settings/form/components/__tests__/basic-info-section.spec.tsx index 3c0e997ad57..2466445435a 100644 --- a/web/app/components/datasets/settings/form/components/__tests__/basic-info-section.spec.tsx +++ b/web/app/components/datasets/settings/form/components/__tests__/basic-info-section.spec.tsx @@ -19,17 +19,29 @@ vi.mock('@tanstack/react-query', async (importOriginal) => { } }) -// Mock app-context -vi.mock('@/context/app-context', () => ({ - useSelector: () => ({ +const mockAppContextState = vi.hoisted(() => ({ + userProfile: { id: 'user-1', name: 'Current User', email: 'current@example.com', avatar_url: '', role: 'owner', - }), + }, })) +// Mock app-context +vi.mock('@/context/app-context', () => ({ + useSelector: () => mockAppContextState.userProfile, +})) + +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createDatasetAccessAtomMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessAtomMock(importOriginal, () => mockAppContextState, () => ({ + isRbacEnabled: false, + })) +}) + // Mock image uploader hooks for AppIconPicker vi.mock('@/app/components/base/image-uploader/hooks', () => ({ useLocalFileUploader: () => ({ @@ -47,6 +59,12 @@ vi.mock('@/app/components/base/image-uploader/hooks', () => ({ }), })) +vi.mock('jotai', async (importOriginal) => { + const { createDatasetAccessJotaiMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessJotaiMock(importOriginal) +}) + describe('BasicInfoSection', () => { const mockDataset: DataSet = { id: 'dataset-1', diff --git a/web/app/components/datasets/settings/form/hooks/__tests__/use-form-state.spec.ts b/web/app/components/datasets/settings/form/hooks/__tests__/use-form-state.spec.ts index 78db31af095..bd047dd469e 100644 --- a/web/app/components/datasets/settings/form/hooks/__tests__/use-form-state.spec.ts +++ b/web/app/components/datasets/settings/form/hooks/__tests__/use-form-state.spec.ts @@ -23,6 +23,15 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createDatasetAccessAtomMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessAtomMock(importOriginal, () => ({ + userProfile: { id: 'user-1' }, + workspacePermissionKeys: [], + })) +}) + const createDefaultMockDataset = (): DataSet => ({ id: 'dataset-1', name: 'Test Dataset', @@ -104,6 +113,12 @@ vi.mock('@/context/dataset-detail', () => ({ }, })) +vi.mock('jotai', async (importOriginal) => { + const { createDatasetAccessJotaiMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessJotaiMock(importOriginal) +}) + // Mock services vi.mock('@/service/datasets', () => ({ updateDatasetSetting: vi.fn().mockResolvedValue({}), diff --git a/web/app/components/datasets/settings/form/hooks/use-form-state.ts b/web/app/components/datasets/settings/form/hooks/use-form-state.ts index 8892d982636..fb3541fb573 100644 --- a/web/app/components/datasets/settings/form/hooks/use-form-state.ts +++ b/web/app/components/datasets/settings/form/hooks/use-form-state.ts @@ -5,12 +5,16 @@ import type { Member } from '@/models/common' import type { IconInfo, SummaryIndexSetting as SummaryIndexSettingType } from '@/models/datasets' import type { RetrievalConfig } from '@/types/app' import { toast } from '@langgenius/dify-ui/toast' +import { useAtomValue } from 'jotai' import { useCallback, useMemo, useState } from 'react' import { useTranslation } from 'react-i18next' import { isReRankModelSelected } from '@/app/components/datasets/common/check-rerank-model' import { ModelTypeEnum } from '@/app/components/header/account-setting/model-provider-page/declarations' import { useModelList } from '@/app/components/header/account-setting/model-provider-page/hooks' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { + userProfileIdAtom, + workspacePermissionKeysAtom, +} from '@/context/app-context-state' import { useDatasetDetailContextWithSelector } from '@/context/dataset-detail' import { DatasetPermission } from '@/models/datasets' import { updateDatasetSetting } from '@/service/datasets' @@ -30,8 +34,8 @@ export const useFormState = () => { const { t } = useTranslation() const currentDataset = useDatasetDetailContextWithSelector(state => state.dataset) const mutateDatasets = useDatasetDetailContextWithSelector(state => state.mutateDatasetRes) - const currentUserId = useAppContextWithSelector(state => state.userProfile?.id) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const datasetACLCapabilities = useMemo( () => getDatasetACLCapabilities(currentDataset?.permission_keys, { currentUserId, diff --git a/web/app/components/datasets/settings/permission-selector/__tests__/index.spec.tsx b/web/app/components/datasets/settings/permission-selector/__tests__/index.spec.tsx index 7793c50acc2..f37728f2724 100644 --- a/web/app/components/datasets/settings/permission-selector/__tests__/index.spec.tsx +++ b/web/app/components/datasets/settings/permission-selector/__tests__/index.spec.tsx @@ -4,17 +4,32 @@ import { renderWithSystemFeatures } from '@/__tests__/utils/mock-system-features import { DatasetPermission } from '@/models/datasets' import PermissionSelector from '../index' -// Mock app-context -vi.mock('@/context/app-context', () => ({ - useSelector: () => ({ +const mockAppContextState = vi.hoisted(() => ({ + userProfile: { id: 'user-1', name: 'Current User', email: 'current@example.com', avatar_url: '', role: 'owner', - }), + }, })) +let mockIsRbacEnabled = false + +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createDatasetAccessAtomMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessAtomMock(importOriginal, () => mockAppContextState, () => ({ + isRbacEnabled: mockIsRbacEnabled, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createDatasetAccessJotaiMock } = await import('@/app/components/datasets/__tests__/mock-dataset-access') + + return createDatasetAccessJotaiMock(importOriginal) +}) + describe('PermissionSelector', () => { const mockMemberList: Member[] = [ { id: 'user-1', name: 'Current User', email: 'current@example.com', avatar: '', avatar_url: '', role: 'owner', roles: [], last_login_at: '', created_at: '', status: 'active' }!, @@ -33,6 +48,7 @@ describe('PermissionSelector', () => { beforeEach(() => { vi.clearAllMocks() + mockIsRbacEnabled = false }) describe('Rendering', () => { @@ -409,6 +425,8 @@ describe('PermissionSelector', () => { }) it('should show access config hint and remain closed when RBAC is enabled', () => { + mockIsRbacEnabled = true + renderWithSystemFeatures(, { systemFeatures: { rbac_enabled: true, diff --git a/web/app/components/datasets/settings/permission-selector/index.tsx b/web/app/components/datasets/settings/permission-selector/index.tsx index afff19320e5..8289c9698a5 100644 --- a/web/app/components/datasets/settings/permission-selector/index.tsx +++ b/web/app/components/datasets/settings/permission-selector/index.tsx @@ -7,12 +7,14 @@ import { PopoverContent, PopoverTrigger, } from '@langgenius/dify-ui/popover' -import { useSuspenseQuery } from '@tanstack/react-query' import { useDebounceFn } from 'ahooks' +import { useAtomValue } from 'jotai' import { useCallback, useMemo, useState } from 'react' import { useTranslation } from 'react-i18next' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' -import { systemFeaturesQueryOptions } from '@/features/system-features/client' +import { + datasetRbacEnabledAtom, + userProfileAtom, +} from '@/context/app-context-state' import { DatasetPermission } from '@/models/datasets' import MemberItem from './member-item' import Item from './permission-item' @@ -35,8 +37,8 @@ const PermissionSelector = ({ onMemberSelect, }: RoleSelectorProps) => { const { t } = useTranslation() - const userProfile = useAppContextWithSelector(state => state.userProfile) - const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) + const userProfile = useAtomValue(userProfileAtom) + const isRbacEnabled = useAtomValue(datasetRbacEnabledAtom) const [open, setOpen] = useState(false) const [keywords, setKeywords] = useState('') @@ -89,7 +91,7 @@ const PermissionSelector = ({ const isAllTeamMembers = permission === DatasetPermission.allTeamMembers const isPartialMembers = permission === DatasetPermission.partialMembers const selectedMemberNames = selectedMembers.map(member => member.name).join(', ') - const isDisabledByRBAC = systemFeatures.rbac_enabled + const isDisabledByRBAC = isRbacEnabled const isDisabled = disabled || isDisabledByRBAC return ( diff --git a/web/context/app-context-defaults.ts b/web/context/app-context-defaults.ts new file mode 100644 index 00000000000..475d7e474f5 --- /dev/null +++ b/web/context/app-context-defaults.ts @@ -0,0 +1,34 @@ +import type { GetAccountProfileResponse } from '@dify/contracts/api/console/account/types.gen' +import type { ICurrentWorkspace, LangGeniusVersionResponse } from '@/models/common' + +export const userProfilePlaceholder: GetAccountProfileResponse = { + id: '', + name: '', + email: '', + avatar: '', + avatar_url: '', + is_password_set: false, +} + +export const initialLangGeniusVersionInfo: LangGeniusVersionResponse = { + current_env: '', + current_version: '', + latest_version: '', + release_date: '', + release_notes: '', + version: '', + can_auto_update: false, +} + +export const initialWorkspaceInfo: ICurrentWorkspace = { + id: '', + name: '', + plan: '', + status: '', + created_at: 0, + role: 'normal', + providers: [], + trial_credits: 200, + trial_credits_used: 0, + next_credit_reset_date: 0, +} diff --git a/web/context/app-context-state.ts b/web/context/app-context-state.ts index 780138f4d0f..04676f51099 100644 --- a/web/context/app-context-state.ts +++ b/web/context/app-context-state.ts @@ -8,6 +8,7 @@ import { atom } from 'jotai' import { atomWithQuery, atomWithSuspenseQuery, queryClientAtom } from 'jotai-tanstack-query' import { userProfileQueryOptions } from '@/features/account-profile/client' import { systemFeaturesQueryOptions } from '@/features/system-features/client' +import { defaultSystemFeatures } from '@/features/system-features/config' import { workspacePermissionKeysQueryOptions } from '@/service/access-control/use-permission-keys' import { consoleQuery } from '@/service/client' import { langGeniusVersionQueryOptions } from '@/service/lang-genius-version' @@ -15,7 +16,7 @@ import { initialLangGeniusVersionInfo, initialWorkspaceInfo, userProfilePlaceholder, -} from './app-context' +} from './app-context-defaults' import { emptyWorkspacePermissionKeys, getLangGeniusVersionInfo, @@ -29,12 +30,22 @@ const accountProfileQueryAtom = atomWithSuspenseQuery(() => userProfileQueryOpti const systemFeaturesQueryAtom = atomWithSuspenseQuery(() => systemFeaturesQueryOptions()) +const systemFeaturesAtom = atom((get): GetSystemFeaturesResponse => { + const systemFeaturesQuery = get(systemFeaturesQueryAtom) as SuspenseQueryResult + + return systemFeaturesQuery.data ?? defaultSystemFeatures +}) + export const userProfileAtom = atom((get): GetAccountProfileResponse => { const accountProfileQuery = get(accountProfileQueryAtom) as SuspenseQueryResult return accountProfileQuery.data?.profile || userProfilePlaceholder }) +export const userProfileIdAtom = atom((get) => { + return get(userProfileAtom).id +}) + const profileMetaAtom = atom((get) => { const accountProfileQuery = get(accountProfileQueryAtom) as SuspenseQueryResult @@ -58,18 +69,26 @@ export const currentWorkspaceAtom = atom((get) => { return get(normalizedCurrentWorkspaceAtom) }) +export const currentWorkspaceIdAtom = atom((get) => { + return get(currentWorkspaceAtom).id +}) + export const workspaceRoleFlagsAtom = atom((get) => { return getWorkspaceRoleFlags(get(currentWorkspaceAtom)) }) +export const isCurrentWorkspaceOwnerAtom = atom((get) => { + return get(workspaceRoleFlagsAtom).isCurrentWorkspaceOwner +}) + const workspacePermissionKeysQueryAtom = atomWithQuery((get) => { - const workspaceId = get(currentWorkspaceAtom).id + const workspaceId = get(currentWorkspaceIdAtom) return workspacePermissionKeysQueryOptions(workspaceId) }) export const workspacePermissionKeysAtom = atom((get) => { - return get(workspacePermissionKeysQueryAtom).data?.workspace.permission_keys ?? emptyWorkspacePermissionKeys + return get(workspacePermissionKeysQueryAtom).data?.workspace?.permission_keys ?? emptyWorkspacePermissionKeys }) export const workspacePermissionKeysLoadingAtom = atom((get) => { @@ -80,6 +99,10 @@ export const currentWorkspaceLoadingAtom = atom((get) => { return get(currentWorkspaceQueryAtom).isPending }) +export const datasetRbacEnabledAtom = atom((get) => { + return get(systemFeaturesAtom).rbac_enabled +}) + export const currentWorkspaceValidatingAtom = atom((get) => { return get(currentWorkspaceQueryAtom).isFetching }) diff --git a/web/context/app-context.ts b/web/context/app-context.ts index 5b8fcf22e69..e89546a1e1d 100644 --- a/web/context/app-context.ts +++ b/web/context/app-context.ts @@ -4,6 +4,11 @@ import type { GetAccountProfileResponse } from '@dify/contracts/api/console/acco import type { ICurrentWorkspace, LangGeniusVersionResponse } from '@/models/common' import { noop } from 'es-toolkit/function' import { createContext, useContext, useContextSelector } from 'use-context-selector' +import { + initialLangGeniusVersionInfo as defaultLangGeniusVersionInfo, + userProfilePlaceholder as defaultUserProfilePlaceholder, + initialWorkspaceInfo as defaultWorkspaceInfo, +} from './app-context-defaults' export type AppContextValue = { userProfile: GetAccountProfileResponse @@ -22,37 +27,11 @@ export type AppContextValue = { workspacePermissionKeys: string[] } -export const userProfilePlaceholder = { - id: '', - name: '', - email: '', - avatar: '', - avatar_url: '', - is_password_set: false, -} +export const userProfilePlaceholder = defaultUserProfilePlaceholder -export const initialLangGeniusVersionInfo = { - current_env: '', - current_version: '', - latest_version: '', - release_date: '', - release_notes: '', - version: '', - can_auto_update: false, -} +export const initialLangGeniusVersionInfo = defaultLangGeniusVersionInfo -export const initialWorkspaceInfo: ICurrentWorkspace = { - id: '', - name: '', - plan: '', - status: '', - created_at: 0, - role: 'normal', - providers: [], - trial_credits: 200, - trial_credits_used: 0, - next_credit_reset_date: 0, -} +export const initialWorkspaceInfo = defaultWorkspaceInfo export const AppContext = createContext({ userProfile: userProfilePlaceholder, From 68d8328b9c39e8aea5e257368e9d2437b0c02a08 Mon Sep 17 00:00:00 2001 From: Asuka Minato Date: Wed, 8 Jul 2026 12:07:27 +0900 Subject: [PATCH 36/70] chore: clean Db session from service (#38227) Co-authored-by: chariri Co-authored-by: WH-2099 --- api/commands/account.py | 8 +- api/commands/data_migration.py | 15 +- api/commands/plugin.py | 1 + api/commands/rbac.py | 35 +- api/controllers/common/app_access.py | 3 +- api/controllers/console/agent/composer.py | 20 +- api/controllers/console/agent/roster.py | 16 +- api/controllers/console/app/agent.py | 12 +- .../console/app/agent_app_feature.py | 2 +- .../console/app/agent_app_sandbox.py | 4 + .../console/app/agent_config_inspector.py | 6 +- .../console/app/agent_drive_inspector.py | 35 +- api/controllers/console/app/annotation.py | 25 +- api/controllers/console/app/app.py | 30 +- api/controllers/console/app/audio.py | 2 +- api/controllers/console/app/conversation.py | 4 +- api/controllers/console/app/message.py | 7 +- api/controllers/console/app/ops_trace.py | 17 +- .../console/app/permission_keys.py | 5 +- api/controllers/console/app/workflow.py | 63 +- .../console/app/workflow_comment.py | 2 +- .../console/app/workflow_draft_variable.py | 10 +- .../app/workflow_node_output_inspector.py | 10 +- api/controllers/console/auth/activate.py | 8 +- .../console/auth/data_source_bearer_auth.py | 6 +- .../console/auth/email_register.py | 8 +- .../console/auth/forgot_password.py | 10 +- api/controllers/console/auth/login.py | 26 +- api/controllers/console/auth/oauth.py | 16 +- api/controllers/console/auth/oauth_server.py | 5 +- api/controllers/console/billing/billing.py | 4 +- .../console/datasets/data_source.py | 10 +- api/controllers/console/datasets/datasets.py | 56 +- .../console/datasets/datasets_document.py | 101 +-- .../console/datasets/datasets_segments.py | 105 +-- api/controllers/console/datasets/external.py | 13 +- .../console/datasets/hit_testing_base.py | 4 +- api/controllers/console/datasets/metadata.py | 36 +- .../datasets/rag_pipeline/datasource_auth.py | 11 +- .../datasource_content_preview.py | 3 +- .../datasets/rag_pipeline/rag_pipeline.py | 20 +- .../rag_pipeline/rag_pipeline_datasets.py | 6 +- .../rag_pipeline_draft_variable.py | 6 +- .../rag_pipeline/rag_pipeline_workflow.py | 54 +- api/controllers/console/datasets/wraps.py | 8 +- api/controllers/console/explore/audio.py | 2 +- .../console/explore/conversation.py | 8 +- .../console/explore/installed_app.py | 2 +- api/controllers/console/explore/message.py | 9 +- api/controllers/console/explore/parameter.py | 3 +- .../console/explore/recommended_app.py | 6 +- .../console/explore/saved_message.py | 10 +- api/controllers/console/explore/trial.py | 18 +- api/controllers/console/extension.py | 16 +- api/controllers/console/init_validate.py | 2 +- api/controllers/console/setup.py | 4 +- api/controllers/console/socketio/workflow.py | 4 +- api/controllers/console/tag/tags.py | 14 +- api/controllers/console/workspace/account.py | 20 +- .../workspace/load_balancing_config.py | 3 + api/controllers/console/workspace/members.py | 22 +- .../console/workspace/model_providers.py | 2 +- api/controllers/console/workspace/models.py | 3 + api/controllers/console/workspace/plugin.py | 14 +- api/controllers/console/workspace/rbac.py | 5 +- api/controllers/console/workspace/snippets.py | 2 +- .../console/workspace/tool_providers.py | 2 + .../console/workspace/workspace.py | 19 +- api/controllers/files/agent_drive_archive.py | 2 + api/controllers/inner_api/app/dsl.py | 1 + .../inner_api/plugin/agent_drive.py | 5 +- .../inner_api/workspace/workspace.py | 6 +- api/controllers/openapi/account.py | 14 +- api/controllers/openapi/app_dsl.py | 1 + api/controllers/openapi/apps.py | 12 +- .../openapi/apps_permitted_external.py | 4 +- api/controllers/openapi/auth/prepare.py | 10 +- api/controllers/openapi/auth/verify.py | 4 +- api/controllers/openapi/oauth_device.py | 4 +- api/controllers/openapi/oauth_device_sso.py | 6 +- api/controllers/openapi/workspaces.py | 26 +- api/controllers/service_api/app/annotation.py | 10 +- api/controllers/service_api/app/app.py | 3 +- api/controllers/service_api/app/audio.py | 2 +- .../service_api/app/conversation.py | 14 +- api/controllers/service_api/app/message.py | 14 +- .../service_api/dataset/dataset.py | 46 +- .../service_api/dataset/document.py | 28 +- .../service_api/dataset/metadata.py | 38 +- .../rag_pipeline/rag_pipeline_workflow.py | 10 +- .../service_api/dataset/segment.py | 48 +- api/controllers/web/app.py | 11 +- api/controllers/web/audio.py | 2 +- api/controllers/web/completion.py | 6 +- api/controllers/web/conversation.py | 8 +- api/controllers/web/forgot_password.py | 4 +- api/controllers/web/login.py | 11 +- api/controllers/web/message.py | 10 +- api/controllers/web/passport.py | 2 +- api/controllers/web/saved_message.py | 6 +- api/controllers/web/wraps.py | 10 +- .../easy_ui_based_app/dataset/manager.py | 2 +- .../app/apps/advanced_chat/app_generator.py | 2 +- api/core/app/apps/agent_app/app_generator.py | 4 +- api/core/app/apps/agent_chat/app_generator.py | 2 +- api/core/app/apps/chat/app_generator.py | 2 +- .../annotation_reply/annotation_reply.py | 5 +- api/core/app/llm/quota.py | 2 + .../task_pipeline/message_cycle_manager.py | 2 +- .../index_tool_callback_handler.py | 6 +- api/core/llm_generator/llm_generator.py | 7 +- api/core/mcp/server/streamable_http.py | 10 +- api/core/provider_manager.py | 2 + api/core/rag/datasource/retrieval_service.py | 12 +- .../processor/paragraph_index_processor.py | 18 +- .../processor/parent_child_index_processor.py | 6 +- .../processor/qa_index_processor.py | 4 +- api/core/rag/summary_index/summary_index.py | 7 +- .../dataset_multi_retriever_tool.py | 4 +- .../dataset_retriever_tool.py | 4 +- .../nodes/agent_v2/dify_tools_builder.py | 3 +- .../update_provider_when_message_created.py | 1 + api/extensions/ext_login.py | 4 +- api/services/account_service.py | 160 ++-- api/services/agent/composer_service.py | 428 +++++++---- api/services/agent/roster_service.py | 1 + .../agent/skill_standardize_service.py | 4 + .../agent/skill_tool_inference_service.py | 11 +- api/services/agent_app_feature_service.py | 4 +- api/services/agent_app_sandbox_service.py | 12 +- api/services/agent_drive_service.py | 337 ++++----- api/services/agent_service.py | 12 +- api/services/agent_tool_inner_service.py | 2 +- api/services/annotation_service.py | 114 +-- api/services/api_based_extension_service.py | 8 +- api/services/app_dsl_service.py | 28 +- api/services/app_generate_service.py | 59 +- api/services/app_service.py | 131 ++-- api/services/async_workflow_service.py | 13 +- api/services/audio_service.py | 6 +- api/services/auth/api_key_auth_service.py | 8 +- api/services/billing_service.py | 4 +- api/services/conversation_service.py | 131 ++-- api/services/credential_permission_service.py | 6 +- api/services/credit_pool_service.py | 75 +- api/services/data_migration/export_service.py | 30 +- api/services/data_migration/import_service.py | 114 ++- api/services/dataset_service.py | 265 +++---- api/services/datasource_provider_service.py | 61 +- .../enterprise/account_deletion_sync.py | 9 +- api/services/enterprise/rbac_service.py | 96 +-- api/services/external_knowledge_service.py | 37 +- api/services/file_service.py | 2 +- api/services/hit_testing_service.py | 18 +- api/services/message_service.py | 57 +- api/services/metadata_service.py | 17 +- api/services/model_load_balancing_service.py | 63 +- api/services/oauth_device_flow.py | 26 +- api/services/oauth_server.py | 6 +- api/services/ops_service.py | 38 +- .../plugin/plugin_auto_upgrade_service.py | 221 +++--- .../plugin/plugin_permission_service.py | 39 +- .../rag_pipeline/pipeline_generate_service.py | 23 +- .../built_in/built_in_retrieval.py | 7 +- .../customized/customized_retrieval.py | 12 +- .../database/database_retrieval.py | 12 +- .../pipeline_template_base.py | 4 +- .../remote/remote_retrieval.py | 10 +- api/services/rag_pipeline/rag_pipeline.py | 689 ++++++++---------- .../rag_pipeline/rag_pipeline_dsl_service.py | 4 +- .../rag_pipeline_transform_service.py | 7 +- .../buildin/buildin_retrieval.py | 11 +- .../database/database_retrieval.py | 38 +- .../recommend_app/recommend_app_base.py | 8 +- .../recommend_app/remote/remote_retrieval.py | 11 +- api/services/recommended_app_service.py | 18 +- api/services/saved_message_service.py | 15 +- api/services/snippet_service.py | 10 +- api/services/summary_index_service.py | 602 +++++++-------- api/services/tag_service.py | 25 +- .../tools/builtin_tools_manage_service.py | 13 +- api/services/trigger/schedule_service.py | 23 +- .../trigger/trigger_provider_service.py | 2 +- .../trigger_subscription_operator_service.py | 7 +- api/services/trigger/webhook_service.py | 6 +- api/services/vector_service.py | 49 +- api/services/web_conversation_service.py | 19 +- api/services/webapp_auth_service.py | 32 +- .../workflow/node_output_inspector_service.py | 58 +- api/services/workflow/workflow_converter.py | 40 +- .../workflow_collaboration_service.py | 23 +- api/services/workflow_service.py | 131 ++-- api/services/workspace_service.py | 12 +- .../batch_create_segment_to_index_task.py | 2 +- api/tasks/regenerate_summary_index_task.py | 3 +- api/tasks/retry_document_indexing_task.py | 5 +- api/tasks/workflow_schedule_tasks.py | 2 +- api/tests/integration_tests/conftest.py | 2 +- .../services/plugin/test_plugin_lifecycle.py | 43 +- .../test_node_output_inspector_service.py | 51 +- .../controllers/console/app/test_app_apis.py | 6 +- .../auth/test_data_source_bearer_auth.py | 2 +- .../console/auth/test_email_register.py | 2 +- .../console/auth/test_forgot_password.py | 2 +- .../controllers/console/auth/test_oauth.py | 2 +- .../rag_pipeline/test_rag_pipeline.py | 36 +- .../console/test_api_based_extension.py | 2 +- .../openapi/test_account_sessions.py | 2 +- .../controllers/openapi/test_app_dsl.py | 2 +- .../controllers/openapi/test_app_run.py | 2 +- .../controllers/openapi/test_apps.py | 2 +- .../controllers/openapi/test_files.py | 2 +- .../service_api/dataset/test_dataset.py | 5 +- .../web/test_web_forgot_password.py | 4 +- .../controllers/web/test_wraps.py | 2 + .../auth/test_api_key_auth_service.py | 34 +- .../services/auth/test_auth_integration.py | 20 +- .../enterprise/test_account_deletion_sync.py | 20 +- .../plugin/test_plugin_permission_service.py | 10 +- .../test_rag_pipeline_service_db.py | 24 +- .../recommend_app/test_database_retrieval.py | 54 +- .../services/test_account_service.py | 6 +- .../services/test_agent_service.py | 30 +- .../services/test_annotation_service.py | 169 +++-- .../test_api_based_extension_service.py | 72 +- .../services/test_app_dsl_service.py | 40 +- .../services/test_app_generate_service.py | 82 ++- .../services/test_app_service.py | 107 +-- .../services/test_billing_service.py | 4 +- .../services/test_conversation_service.py | 80 +- .../test_conversation_service_variables.py | 18 +- .../services/test_credit_pool_service.py | 132 ++-- .../services/test_dataset_service.py | 7 +- .../test_dataset_service_permissions.py | 22 +- .../test_dataset_service_update_dataset.py | 26 +- .../test_file_service_zip_and_lookup.py | 8 +- .../services/test_hit_testing_service.py | 16 +- .../test_human_input_delivery_test.py | 3 + .../services/test_message_service.py | 136 +++- ...message_service_execution_extra_content.py | 2 + .../services/test_metadata_partial_update.py | 14 +- .../services/test_metadata_service.py | 82 ++- .../test_model_load_balancing_service.py | 18 +- .../services/test_oauth_server_service.py | 10 +- .../services/test_ops_service.py | 63 +- .../services/test_recommended_app_service.py | 48 +- .../services/test_saved_message_service.py | 44 +- .../services/test_tag_service.py | 18 +- .../services/test_web_conversation_service.py | 22 +- .../services/test_webapp_auth_service.py | 52 +- .../test_webhook_service_relationships.py | 7 +- .../services/test_workflow_app_service.py | 4 +- .../services/test_workflow_run_service.py | 8 +- .../services/test_workflow_service.py | 39 +- .../services/test_workspace_service.py | 40 +- .../test_workflow_tools_manage_service.py | 2 +- .../workflow/test_workflow_converter.py | 15 +- .../trigger/test_trigger_e2e.py | 4 +- .../commands/test_data_migration_commands.py | 8 +- .../controllers/common/test_app_access.py | 2 +- .../console/agent/test_agent_controllers.py | 68 +- .../console/app/test_agent_app_sandbox.py | 3 + .../console/app/test_annotation_api.py | 6 +- .../console/app/test_annotation_security.py | 14 +- .../console/app/test_app_response_models.py | 25 +- .../controllers/console/app/test_workflow.py | 4 +- .../test_workflow_human_input_debug_api.py | 5 +- .../test_workflow_node_output_inspector.py | 9 +- .../console/auth/test_account_activation.py | 2 +- .../auth/test_data_source_bearer_auth.py | 6 +- .../rag_pipeline/test_datasource_auth.py | 3 +- .../rag_pipeline/test_rag_pipeline.py | 84 ++- .../test_rag_pipeline_workflow.py | 18 +- .../console/datasets/test_datasets.py | 18 +- .../console/datasets/test_external.py | 7 +- .../console/explore/test_recommended_app.py | 12 +- .../console/explore/test_saved_message.py | 6 +- .../console/snippets/test_snippet_workflow.py | 5 +- .../controllers/console/tag/test_tags.py | 12 +- .../controllers/console/test_extension.py | 12 +- .../console/test_workspace_account.py | 2 +- .../workspace/test_load_balancing_config.py | 4 +- .../console/workspace/test_tool_providers.py | 4 +- .../controllers/inner_api/app/test_dsl.py | 12 +- .../inner_api/plugin/test_agent_drive.py | 10 +- .../inner_api/workspace/test_workspace.py | 2 +- .../openapi/test_workspaces_members.py | 46 +- .../service_api/app/test_annotation.py | 6 +- .../controllers/service_api/app/test_app.py | 4 +- .../service_api/app/test_completion.py | 28 +- .../service_api/app/test_conversation.py | 1 + .../service_api/app/test_message.py | 31 +- .../service_api/app/test_workflow.py | 34 +- .../test_rag_pipeline_workflow.py | 37 +- .../dataset/test_dataset_segment.py | 2 +- .../service_api/dataset/test_document.py | 2 +- .../service_api/dataset/test_metadata.py | 4 +- .../unit_tests/controllers/web/test_app.py | 4 +- .../controllers/web/test_message_list.py | 4 +- .../controllers/web/test_web_login.py | 6 +- .../apps/advanced_chat/test_app_generator.py | 2 +- .../apps/test_advanced_chat_app_generator.py | 2 +- .../unit_tests/core/app/test_llm_quota.py | 21 +- .../test_llm_generator_missing.py | 6 +- .../datasource/test_datasource_retrieval.py | 23 +- .../test_paragraph_index_processor.py | 4 +- .../test_parent_child_index_processor.py | 4 +- .../processor/test_qa_index_processor.py | 4 +- .../events/test_app_event_signals.py | 33 +- ...st_update_provider_when_message_created.py | 18 +- .../services/agent/test_agent_services.py | 275 +++++-- .../agent/test_skill_standardize_service.py | 1 + .../test_skill_tool_inference_service.py | 21 +- .../data_migration/test_export_service.py | 14 +- .../data_migration/test_import_service.py | 142 ++-- .../services/enterprise/test_rbac_service.py | 85 +-- api/tests/unit_tests/services/hit_service.py | 37 +- .../test_plugin_auto_upgrade_service.py | 103 +-- .../test_built_in_retrieval.py | 12 +- .../test_customized_retrieval.py | 6 +- .../test_database_retrieval.py | 6 +- .../test_pipeline_template_base.py | 10 +- .../test_remote_retrieval.py | 8 +- .../test_pipeline_generate_service.py | 74 +- .../rag_pipeline/test_rag_pipeline_service.py | 84 ++- .../test_rag_pipeline_transform_service.py | 71 +- .../recommend_app/test_buildin_retrieval.py | 9 +- .../recommend_app/test_remote_retrieval.py | 16 +- .../services/test_account_service.py | 74 +- .../test_agent_app_sandbox_service.py | 5 + .../services/test_agent_drive_service.py | 130 +++- .../services/test_agent_tool_inner_service.py | 12 +- .../services/test_annotation_service.py | 104 +-- .../services/test_app_generate_service.py | 91 ++- .../unit_tests/services/test_app_service.py | 26 +- .../services/test_async_workflow_service.py | 22 +- .../services/test_billing_service.py | 8 +- .../services/test_conversation_service.py | 7 +- .../test_credential_permission_service.py | 4 +- .../services/test_credit_pool_service.py | 107 ++- .../services/test_dataset_service_dataset.py | 73 +- .../services/test_dataset_service_document.py | 2 +- .../services/test_dataset_service_segment.py | 42 +- .../test_datasource_provider_service.py | 28 +- .../services/test_external_dataset_service.py | 480 +++++++----- .../unit_tests/services/test_file_service.py | 6 +- .../services/test_message_service.py | 69 +- .../services/test_metadata_bug_complete.py | 6 +- .../services/test_metadata_nullable_bug.py | 4 +- .../test_model_load_balancing_service.py | 41 +- .../services/test_oauth_device_flow.py | 10 +- .../services/test_summary_index_service.py | 182 ++--- .../services/test_trigger_provider_service.py | 6 +- .../services/test_vector_service.py | 128 ++-- .../test_workflow_collaboration_service.py | 21 +- .../services/test_workflow_service.py | 101 ++- .../test_builtin_tools_manage_service.py | 6 +- .../test_node_output_inspector_service.py | 210 +++--- .../test_workflow_converter_additional.py | 18 +- .../test_workflow_human_input_delivery.py | 4 + 360 files changed, 6607 insertions(+), 4975 deletions(-) diff --git a/api/commands/account.py b/api/commands/account.py index dfd57d43142..9ea52dfd248 100644 --- a/api/commands/account.py +++ b/api/commands/account.py @@ -25,7 +25,7 @@ def reset_password(email, new_password, password_confirm): return normalized_email = email.strip().lower() - account = AccountService.get_account_by_email_with_case_fallback(db.session, email.strip()) + account = AccountService.get_account_by_email_with_case_fallback(email.strip(), session=db.session()) if not account: click.echo(click.style(f"Account not found for email: {email}", fg="red")) @@ -67,7 +67,7 @@ def reset_email(email, new_email, email_confirm): return normalized_new_email = new_email.strip().lower() - account = AccountService.get_account_by_email_with_case_fallback(db.session, email.strip()) + account = AccountService.get_account_by_email_with_case_fallback(email.strip(), session=db.session()) if not account: click.echo(click.style(f"Account not found for email: {email}", fg="red")) @@ -133,9 +133,9 @@ def create_tenant(email: str, language: str | None = None, name: str | None = No password=new_password, language=language, create_workspace_required=False, - session=db.session, + session=db.session(), ) - TenantService.create_owner_tenant_if_not_exist(account, name, session=db.session) + TenantService.create_owner_tenant_if_not_exist(account, name, session=db.session()) click.echo( click.style( diff --git a/api/commands/data_migration.py b/api/commands/data_migration.py index bd56c41ea44..8c2627601a6 100644 --- a/api/commands/data_migration.py +++ b/api/commands/data_migration.py @@ -9,6 +9,7 @@ from uuid import UUID import click import sqlalchemy as sa import yaml +from sqlalchemy.orm import Session from core.db.session_factory import session_factory from extensions.ext_database import db @@ -108,7 +109,7 @@ def export_migration_data(input_file: str | None, output_file: str | None, overw raw_config = _load_json_object(input_file, "Export config") selection = ExportConfigParser().parse(raw_config) with session_factory.create_session() as session: - result = MigrationExportService().export(session, selection) + result = MigrationExportService().export(selection, session=session) MigrationPackageService().save_package(result.package, output_file, overwrite=overwrite) click.echo(click.style(f"Output written to {output_file}", fg="green")) _render_report(result.report_items, context=_with_output_path(result.report_context, output_file)) @@ -157,7 +158,6 @@ def import_migration_data( package = MigrationPackageService().load_package(input_file) with session_factory.create_session() as session: result = MigrationImportService().import_package( - session, ImportRequest( package=package, cli_target_tenant=target_tenant, @@ -169,6 +169,7 @@ def import_migration_data( create_app_api_token_on_import=create_app_api_token_on_import, ), ), + session=session, ) _render_report(result.report_items, context=result.report_context) except MigrationDataError as exc: @@ -217,7 +218,9 @@ def migration_data_wizard() -> None: default=True, show_default=False, ) - auto_tools = _discover_auto_tools([app for app in apps if app.id in set(app_ids)], include_referenced_tools) + auto_tools = _discover_auto_tools( + [app for app in apps if app.id in set(app_ids)], include_referenced_tools, session=db.session() + ) auto_tools = _resolve_auto_tool_names(tenant.id, auto_tools) _print_auto_tools(auto_tools) additional_tools = _prompt_additional_tools(tenant.id, auto_tools) @@ -253,7 +256,7 @@ def migration_data_wizard() -> None: output_file=output_file, ) with session_factory.create_session() as session: - result = MigrationExportService().export(session, selection) + result = MigrationExportService().export(selection, session=session) MigrationPackageService().save_package(result.package, output_file, overwrite=overwrite) click.echo(click.style(f"Output written to {output_file}", fg="green")) _print_wizard_step("Report") @@ -394,13 +397,13 @@ def _prompt_import_options() -> tuple[bool, bool, str, str]: return include_secrets, create_tokens, id_strategy, conflict_strategy -def _discover_auto_tools(apps: list[App], include_referenced_tools: bool) -> WizardToolMap: +def _discover_auto_tools(apps: list[App], include_referenced_tools: bool, *, session: Session) -> WizardToolMap: auto_tools: WizardToolMap = {"api_tools": {}, "workflow_tools": {}, "mcp_tools": {}} if not include_referenced_tools: return auto_tools discovery_service = DependencyDiscoveryService() for app in apps: - dsl_content = AppDslService.export_dsl(app_model=app, include_secret=False) + dsl_content = AppDslService.export_dsl(app_model=app, session=session, include_secret=False) raw_dsl = yaml.safe_load(dsl_content) if dsl_content else {} dsl = raw_dsl if isinstance(raw_dsl, dict) else {} for dependency in discovery_service.discover_from_dsl(dsl): diff --git a/api/commands/plugin.py b/api/commands/plugin.py index 718fa60761c..3695c742921 100644 --- a/api/commands/plugin.py +++ b/api/commands/plugin.py @@ -472,6 +472,7 @@ def backfill_plugin_auto_upgrade( try: result = PluginAutoUpgradeService.backfill_strategy_categories( current_tenant_id, + session=db.session(), ) except Exception as e: failed_count += 1 diff --git a/api/commands/rbac.py b/api/commands/rbac.py index 0793d11cbb2..be4993920ad 100644 --- a/api/commands/rbac.py +++ b/api/commands/rbac.py @@ -6,6 +6,7 @@ from concurrent.futures import ThreadPoolExecutor, as_completed import click from sqlalchemy import select +from sqlalchemy.orm import Session from configs import dify_config from core.db.session_factory import session_factory @@ -131,16 +132,35 @@ def _replace_member_role( operator_account_id: str, member_account_id: str, role_id: str, + *, + session: Session, ) -> str: RBACService.MemberRoles.replace( tenant_id=tenant_id, account_id=operator_account_id, member_account_id=member_account_id, role_ids=[role_id], + session=session, ) return member_account_id +def _replace_member_role_with_new_session( + tenant_id: str, + operator_account_id: str, + member_account_id: str, + role_id: str, +) -> str: + with session_factory.create_session() as session: + return _replace_member_role( + tenant_id=tenant_id, + operator_account_id=operator_account_id, + member_account_id=member_account_id, + role_id=role_id, + session=session, + ) + + @click.command( "rbac-migrate-member-roles", help="Migrate legacy workspace member roles into RBAC member-role bindings." ) @@ -217,14 +237,21 @@ def migrate_member_roles_to_rbac( if replace_jobs: if workers == 1: - for member_account_id, resolved_role_id in replace_jobs: - _replace_member_role(workspace_id, owner_account_id, member_account_id, resolved_role_id) - migrated_count += 1 + with session_factory.create_session() as session: + for member_account_id, resolved_role_id in replace_jobs: + _replace_member_role( + workspace_id, + owner_account_id, + member_account_id, + resolved_role_id, + session=session, + ) + migrated_count += 1 else: with ThreadPoolExecutor(max_workers=workers) as executor: futures = [ executor.submit( - _replace_member_role, + _replace_member_role_with_new_session, workspace_id, owner_account_id, member_account_id, diff --git a/api/controllers/common/app_access.py b/api/controllers/common/app_access.py index 863b69d2339..214d2de71b4 100644 --- a/api/controllers/common/app_access.py +++ b/api/controllers/common/app_access.py @@ -4,6 +4,7 @@ from collections.abc import Sequence from dataclasses import dataclass from typing import TYPE_CHECKING +from extensions.ext_database import db from services.enterprise import rbac_service as enterprise_rbac_service if TYPE_CHECKING: @@ -76,7 +77,7 @@ def resolve_app_access_filter( inner-API round trip; otherwise it is fetched here. """ if permissions is None: - permissions = enterprise_rbac_service.RBACService.MyPermissions.get(tenant_id, account_id) + permissions = enterprise_rbac_service.RBACService.MyPermissions.get(tenant_id, account_id, session=db.session()) whitelist_scope = enterprise_rbac_service.RBACService.AppAccess.whitelist_resources(tenant_id, account_id) can_manage_own_apps = _MANAGE_OWN_APPS_PERMISSION_KEY in permissions.workspace.permission_keys diff --git a/api/controllers/console/agent/composer.py b/api/controllers/console/agent/composer.py index d089772e3ab..f5d71990ade 100644 --- a/api/controllers/console/agent/composer.py +++ b/api/controllers/console/agent/composer.py @@ -16,6 +16,7 @@ from controllers.console.wraps import ( with_current_tenant_id, with_current_user_id, ) +from extensions.ext_database import db from fields.agent_fields import ( AgentAppComposerResponse, AgentComposerCandidatesResponse, @@ -69,6 +70,7 @@ class WorkflowAgentComposerApi(Resource): node_id=node_id, account_id=account_id, snapshot_id=query.snapshot_id, + session=db.session(), ), ) @@ -94,6 +96,7 @@ class WorkflowAgentComposerApi(Resource): node_id=node_id, account_id=account_id, payload=payload, + session=db.session(), ), ) @@ -126,6 +129,7 @@ class WorkflowAgentComposerCopyFromRosterApi(Resource): source_agent_id=payload.source_agent_id, source_snapshot_id=payload.source_snapshot_id, idempotency_key=payload.idempotency_key, + session=db.session(), ), ) @@ -149,8 +153,9 @@ class WorkflowAgentComposerValidateApi(Resource): tenant_id=tenant_id, payload=payload, agent_id=AgentComposerService.resolve_workflow_node_agent_id( - tenant_id=tenant_id, app_id=app_model.id, node_id=node_id + tenant_id=tenant_id, app_id=app_model.id, node_id=node_id, session=db.session() ), + session=db.session(), ) return dump_response(AgentComposerValidateResponse, {"result": "success", "errors": [], **findings}) @@ -174,6 +179,7 @@ class WorkflowAgentComposerCandidatesApi(Resource): app_id=app_model.id, node_id=node_id, user_id=current_user_id, + session=db.session(), ), ) @@ -196,7 +202,9 @@ class WorkflowAgentComposerImpactApi(Resource): ) return dump_response( AgentComposerImpactResponse, - AgentComposerService.calculate_impact(tenant_id=tenant_id, current_snapshot_id=current_snapshot_id), + AgentComposerService.calculate_impact( + tenant_id=tenant_id, current_snapshot_id=current_snapshot_id, session=db.session() + ), ) @@ -224,6 +232,7 @@ class WorkflowAgentComposerSaveToRosterApi(Resource): node_id=node_id, account_id=account_id, payload=payload, + session=db.session(), ), ) @@ -238,7 +247,7 @@ class AgentComposerApi(Resource): def get(self, tenant_id: str, agent_id: UUID): return dump_response( AgentAppComposerResponse, - AgentComposerService.load_agent_composer(tenant_id=tenant_id, agent_id=str(agent_id)), + AgentComposerService.load_agent_composer(tenant_id=tenant_id, agent_id=str(agent_id), session=db.session()), ) @console_ns.expect(console_ns.models[ComposerSavePayload.__name__]) @@ -259,6 +268,7 @@ class AgentComposerApi(Resource): agent_id=str(agent_id), account_id=account_id, payload=payload, + session=db.session(), ), ) @@ -274,7 +284,7 @@ class AgentComposerValidateApi(Resource): @account_initialization_required @with_current_tenant_id def post(self, tenant_id: str, agent_id: UUID): - AgentComposerService.load_agent_composer(tenant_id=tenant_id, agent_id=str(agent_id)) + AgentComposerService.load_agent_composer(tenant_id=tenant_id, agent_id=str(agent_id), session=db.session()) payload = ComposerSavePayload.model_validate(console_ns.payload or {}) ComposerConfigValidator.validate_publish_payload(payload) AgentComposerService.validate_knowledge_datasets(tenant_id=tenant_id, agent_soul=payload.agent_soul) @@ -282,6 +292,7 @@ class AgentComposerValidateApi(Resource): tenant_id=tenant_id, payload=payload, agent_id=str(agent_id), + session=db.session(), ) return dump_response(AgentComposerValidateResponse, {"result": "success", "errors": [], **findings}) @@ -303,5 +314,6 @@ class AgentComposerCandidatesApi(Resource): tenant_id=tenant_id, agent_id=str(agent_id), user_id=current_user_id, + session=db.session(), ), ) diff --git a/api/controllers/console/agent/roster.py b/api/controllers/console/agent/roster.py index 349826e54d7..1467cc0c246 100644 --- a/api/controllers/console/agent/roster.py +++ b/api/controllers/console/agent/roster.py @@ -534,7 +534,7 @@ class AgentAppListApi(Resource): status="normal", ) - app_pagination = AppService().get_paginate_apps(current_user.id, current_tenant_id, params, db.session) + app_pagination = AppService().get_paginate_apps(current_user.id, current_tenant_id, params, db.session()) if app_pagination is None: empty = AgentAppPagination(page=args.page, limit=args.limit, total=0, has_more=False, data=[]) return empty.model_dump(mode="json") @@ -567,7 +567,7 @@ class AgentAppListApi(Resource): icon_background=args.icon_background, ) - app = AppService().create_app(current_tenant_id, params, current_user) + app = AppService().create_app(current_tenant_id, params, current_user, session=db.session()) return _serialize_agent_app_detail(app, current_user=current_user), 201 @@ -607,7 +607,7 @@ class AgentAppApi(Resource): "max_active_requests": args.max_active_requests or 0, "role": args.role, } - updated = AppService().update_app(app_model, args_dict) + updated = AppService().update_app(app_model, args_dict, session=db.session()) return _serialize_agent_app_detail(updated, current_user=current_user) @console_ns.response(204, "Agent app deleted successfully") @@ -619,7 +619,7 @@ class AgentAppApi(Resource): @with_current_tenant_id def delete(self, tenant_id: str, agent_id: UUID): app_model = _resolve_agent_app_model(tenant_id=tenant_id, agent_id=agent_id) - AppService().delete_app(app_model) + AppService().delete_app(app_model, session=db.session()) return "", 204 @@ -668,6 +668,7 @@ class AgentPublishApi(Resource): agent_id=str(agent_id), account_id=current_user.id, version_note=args.version_note, + session=db.session(), ) @@ -688,6 +689,7 @@ class AgentBuildDraftCheckoutApi(Resource): agent_id=str(agent_id), account_id=current_user.id, force=args.force, + session=db.session(), ) @@ -705,6 +707,7 @@ class AgentBuildDraftApi(Resource): tenant_id=tenant_id, agent_id=str(agent_id), account_id=current_user.id, + session=db.session(), ) @console_ns.expect(console_ns.models[ComposerSavePayload.__name__]) @@ -722,6 +725,7 @@ class AgentBuildDraftApi(Resource): agent_id=str(agent_id), account_id=current_user.id, payload=payload, + session=db.session(), ) @console_ns.response(200, "Agent build draft discarded", console_ns.models[AgentSimpleResultResponse.__name__]) @@ -736,6 +740,7 @@ class AgentBuildDraftApi(Resource): tenant_id=tenant_id, agent_id=str(agent_id), account_id=current_user.id, + session=db.session(), ) @@ -753,6 +758,7 @@ class AgentBuildDraftApplyApi(Resource): tenant_id=tenant_id, agent_id=str(agent_id), account_id=current_user.id, + session=db.session(), ) @@ -810,7 +816,7 @@ class AgentApiStatusApi(Resource): def post(self, tenant_id: str, agent_id: UUID): app_model = _resolve_agent_app_model(tenant_id=tenant_id, agent_id=agent_id) args = AgentApiStatusPayload.model_validate(console_ns.payload) - app_model = AppService().update_app_api_status(app_model, args.enable_api) + app_model = AppService().update_app_api_status(app_model, args.enable_api, session=db.session()) return _serialize_agent_api_access(app_model) diff --git a/api/controllers/console/app/agent.py b/api/controllers/console/app/agent.py index 99164b4755a..81d17ace37a 100644 --- a/api/controllers/console/app/agent.py +++ b/api/controllers/console/app/agent.py @@ -172,7 +172,7 @@ register_response_schema_models( def _resolve_agent_id(app_model: App, node_id: str | None) -> str | None: if node_id and app_model.mode != AppMode.AGENT: return AgentComposerService.resolve_workflow_node_agent_id( - tenant_id=app_model.tenant_id, app_id=app_model.id, node_id=node_id + tenant_id=app_model.tenant_id, app_id=app_model.id, node_id=node_id, session=db.session() ) return app_model.bound_agent_id @@ -202,6 +202,7 @@ def _upload_skill_for_app(*, current_user: Account, app_model: App): tenant_id=app_model.tenant_id, user_id=current_user.id, agent_id=agent_id, + session=db.session(), ) except (SkillPackageError, AgentDriveError) as exc: return {"code": exc.code, "message": exc.message}, exc.status_code @@ -240,6 +241,7 @@ def _commit_drive_file_for_app(*, current_user: Account, app_model: App, allow_n value_owned_by_drive=True, ) ], + session=db.session(), ) except AgentDriveError as exc: return {"code": exc.code, "message": exc.message}, exc.status_code @@ -273,6 +275,7 @@ def _delete_drive_file_for_app(*, current_user: Account, app_model: App, allow_n user_id=current_user.id, agent_id=agent_id, items=[DriveCommitItem(key=key, file_ref=None)], + session=db.session(), ) except AgentDriveError as exc: return {"code": exc.code, "message": exc.message}, exc.status_code @@ -298,6 +301,7 @@ def _delete_skill_for_app(*, current_user: Account, app_model: App, slug: str, a DriveCommitItem(key=f"{slug}/SKILL.md", file_ref=None), DriveCommitItem(key=f"{slug}/.DIFY-SKILL-FULL.zip", file_ref=None), ], + session=db.session(), ) except AgentDriveError as exc: return {"code": exc.code, "message": exc.message}, exc.status_code @@ -313,7 +317,9 @@ def _infer_skill_tools_for_app(*, app_model: App, slug: str): if "/" in slug or not slug.strip(): return {"code": "drive_key_invalid", "message": "skill slug must be a single path segment"}, 400 try: - return SkillToolInferenceService().infer(tenant_id=app_model.tenant_id, agent_id=agent_id, slug=slug) + return SkillToolInferenceService().infer( + tenant_id=app_model.tenant_id, agent_id=agent_id, slug=slug, session=db.session() + ) except SkillToolInferenceError as exc: return {"code": exc.code, "message": exc.message}, exc.status_code @@ -335,7 +341,7 @@ class AgentLogApi(Resource): """Get agent logs""" args = AgentLogQuery.model_validate(request.args.to_dict(flat=True)) - return AgentService.get_agent_logs(app_model, args.conversation_id, args.message_id) + return AgentService.get_agent_logs(app_model, args.conversation_id, args.message_id, db.session()) @console_ns.route("/agent//skills/upload") diff --git a/api/controllers/console/app/agent_app_feature.py b/api/controllers/console/app/agent_app_feature.py index 6990886a511..edd2f31f75f 100644 --- a/api/controllers/console/app/agent_app_feature.py +++ b/api/controllers/console/app/agent_app_feature.py @@ -93,7 +93,7 @@ class AgentAppFeatureConfigResource(Resource): app_model=app_model, account=current_user, config=args.model_dump(exclude_none=True), - session=db.session, + session=db.session(), ) app_model_config_was_updated.send(app_model, app_model_config=new_app_model_config) diff --git a/api/controllers/console/app/agent_app_sandbox.py b/api/controllers/console/app/agent_app_sandbox.py index 4324f425a08..6f3811ccdc8 100644 --- a/api/controllers/console/app/agent_app_sandbox.py +++ b/api/controllers/console/app/agent_app_sandbox.py @@ -25,6 +25,7 @@ from controllers.console import console_ns from controllers.console.agent.app_helpers import resolve_agent_runtime_app_model from controllers.console.app.wraps import get_app_model from controllers.console.wraps import account_initialization_required, setup_required, with_current_tenant_id +from extensions.ext_database import db from fields.base import ResponseModel from libs.login import login_required from models.model import App, AppMode @@ -269,6 +270,7 @@ class WorkflowAgentSandboxListResource(Resource): node_id=node_id, node_execution_id=query.node_execution_id, path=query.path, + session=db.session(), ) except Exception as exc: return _handle(exc) @@ -305,6 +307,7 @@ class WorkflowAgentSandboxReadResource(Resource): node_id=node_id, node_execution_id=query.node_execution_id, path=query.path, + session=db.session(), ) except Exception as exc: return _handle(exc) @@ -334,6 +337,7 @@ class WorkflowAgentSandboxUploadResource(Resource): node_id=node_id, node_execution_id=payload.node_execution_id, path=payload.path, + session=db.session(), ) except Exception as exc: return _handle(exc) diff --git a/api/controllers/console/app/agent_config_inspector.py b/api/controllers/console/app/agent_config_inspector.py index 83824d6434f..0f7aa80ca78 100644 --- a/api/controllers/console/app/agent_config_inspector.py +++ b/api/controllers/console/app/agent_config_inspector.py @@ -253,6 +253,7 @@ def _resolve_agent_id(app_model: App, node_id: str | None) -> str | None: tenant_id=app_model.tenant_id, app_id=app_model.id, node_id=node_id, + session=db.session(), ) return app_model.bound_agent_id @@ -288,13 +289,16 @@ def _resolve_console_version( tenant_id=tenant_id, agent_id=agent_id, account_id=account_id, + session=db.session(), ) draft = state.get("draft") or {} draft_id = draft.get("id") if isinstance(draft_id, str) and draft_id: return draft_id, AgentConfigVersionKind.BUILD_DRAFT else: - state = AgentComposerService.load_agent_composer(tenant_id=tenant_id, agent_id=agent_id) + state = AgentComposerService.load_agent_composer( + tenant_id=tenant_id, agent_id=agent_id, session=db.session() + ) draft = state.get("draft") or {} draft_id = draft.get("id") if isinstance(draft_id, str) and draft_id: diff --git a/api/controllers/console/app/agent_drive_inspector.py b/api/controllers/console/app/agent_drive_inspector.py index 473e7364b3e..5166393b3d9 100644 --- a/api/controllers/console/app/agent_drive_inspector.py +++ b/api/controllers/console/app/agent_drive_inspector.py @@ -28,6 +28,7 @@ from controllers.console import console_ns from controllers.console.agent.app_helpers import resolve_agent_runtime_app_model from controllers.console.app.wraps import get_app_model from controllers.console.wraps import account_initialization_required, setup_required, with_current_tenant_id +from extensions.ext_database import db from fields.base import ResponseModel from libs.login import login_required from models.model import App, AppMode @@ -147,7 +148,7 @@ def _resolve_agent_id(app_model: App, node_id: str | None) -> str | None: """Agent identity for the drive: app-bound agent, or the workflow node binding.""" if node_id: return AgentComposerService.resolve_workflow_node_agent_id( - tenant_id=app_model.tenant_id, app_id=app_model.id, node_id=node_id + tenant_id=app_model.tenant_id, app_id=app_model.id, node_id=node_id, session=db.session() ) return app_model.bound_agent_id @@ -184,7 +185,9 @@ class AgentDriveListByAgentApi(Resource): query = query_params_from_request(AgentDriveListByAgentQuery) resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) try: - items = AgentDriveService().manifest(tenant_id=tenant_id, agent_id=str(agent_id), prefix=query.prefix) + items = AgentDriveService().manifest( + tenant_id=tenant_id, agent_id=str(agent_id), prefix=query.prefix, session=db.session() + ) except AgentDriveError as exc: return _handle(exc) return {"items": [{k: v for k, v in item.items() if k != "file_id"} for item in items]} @@ -203,7 +206,7 @@ class AgentDriveSkillListByAgentApi(Resource): def get(self, tenant_id: str, agent_id: UUID): resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) try: - items = AgentDriveService().list_skills(tenant_id=tenant_id, agent_id=str(agent_id)) + items = AgentDriveService().list_skills(tenant_id=tenant_id, agent_id=str(agent_id), session=db.session()) except AgentDriveError as exc: return _handle(exc) return {"items": items} @@ -227,6 +230,7 @@ class AgentDriveSkillInspectByAgentApi(Resource): tenant_id=tenant_id, agent_id=str(agent_id), skill_path=skill_path, + session=db.session(), ) ) except AgentDriveError as exc: @@ -247,7 +251,9 @@ class AgentDrivePreviewByAgentApi(Resource): query = query_params_from_request(AgentDriveFileByAgentQuery) resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) try: - return AgentDriveService().preview(tenant_id=tenant_id, agent_id=str(agent_id), key=query.key) + return AgentDriveService().preview( + tenant_id=tenant_id, agent_id=str(agent_id), key=query.key, session=db.session() + ) except AgentDriveError as exc: return _handle(exc) @@ -266,7 +272,9 @@ class AgentDriveDownloadByAgentApi(Resource): query = query_params_from_request(AgentDriveFileByAgentQuery) resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) try: - url = AgentDriveService().download_url(tenant_id=tenant_id, agent_id=str(agent_id), key=query.key) + url = AgentDriveService().download_url( + tenant_id=tenant_id, agent_id=str(agent_id), key=query.key, session=db.session() + ) except AgentDriveError as exc: return _handle(exc) return {"url": url} @@ -288,7 +296,9 @@ class AgentDriveListApi(Resource): if not agent_id: return _agent_not_bound() try: - items = AgentDriveService().manifest(tenant_id=app_model.tenant_id, agent_id=agent_id, prefix=query.prefix) + items = AgentDriveService().manifest( + tenant_id=app_model.tenant_id, agent_id=agent_id, prefix=query.prefix, session=db.session() + ) except AgentDriveError as exc: return _handle(exc) # the inner manifest exposes file_id for agent-side pulls; the console @@ -312,7 +322,9 @@ class AgentDriveSkillListApi(Resource): if not agent_id: return _agent_not_bound() try: - items = AgentDriveService().list_skills(tenant_id=app_model.tenant_id, agent_id=agent_id) + items = AgentDriveService().list_skills( + tenant_id=app_model.tenant_id, agent_id=agent_id, session=db.session() + ) except AgentDriveError as exc: return _handle(exc) return {"items": items} @@ -345,6 +357,7 @@ class AgentDriveSkillInspectApi(Resource): tenant_id=app_model.tenant_id, agent_id=agent_id, skill_path=skill_path, + session=db.session(), ) ) except AgentDriveError as exc: @@ -367,7 +380,9 @@ class AgentDrivePreviewApi(Resource): if not agent_id: return _agent_not_bound() try: - return AgentDriveService().preview(tenant_id=app_model.tenant_id, agent_id=agent_id, key=query.key) + return AgentDriveService().preview( + tenant_id=app_model.tenant_id, agent_id=agent_id, key=query.key, session=db.session() + ) except AgentDriveError as exc: return _handle(exc) @@ -388,7 +403,9 @@ class AgentDriveDownloadApi(Resource): if not agent_id: return _agent_not_bound() try: - url = AgentDriveService().download_url(tenant_id=app_model.tenant_id, agent_id=agent_id, key=query.key) + url = AgentDriveService().download_url( + tenant_id=app_model.tenant_id, agent_id=agent_id, key=query.key, session=db.session() + ) except AgentDriveError as exc: return _handle(exc) return {"url": url} diff --git a/api/controllers/console/app/annotation.py b/api/controllers/console/app/annotation.py index d14c7d2a7dc..961f9e2f1d8 100644 --- a/api/controllers/console/app/annotation.py +++ b/api/controllers/console/app/annotation.py @@ -211,7 +211,7 @@ class AppAnnotationSettingDetailApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) def get(self, app_id: UUID): - result = AppAnnotationService.get_app_annotation_setting_by_app_id(str(app_id)) + result = AppAnnotationService.get_app_annotation_setting_by_app_id(str(app_id), session=db.session()) return dump_response(AnnotationSettingResponse, result), 200 @@ -235,7 +235,7 @@ class AppAnnotationSettingUpdateApi(Resource): setting_args: UpdateAnnotationSettingArgs = {"score_threshold": args.score_threshold} result = AppAnnotationService.update_app_annotation_setting( - str(app_id), annotation_setting_id_str, setting_args + str(app_id), annotation_setting_id_str, setting_args, session=db.session() ) return dump_response(AnnotationSettingResponse, result), 200 @@ -292,7 +292,9 @@ class AnnotationApi(Resource): limit = args.limit keyword = args.keyword - annotation_list, total = AppAnnotationService.get_annotation_list_by_app_id(str(app_id), page, limit, keyword) + annotation_list, total = AppAnnotationService.get_annotation_list_by_app_id( + str(app_id), page, limit, keyword, session=db.session() + ) annotation_models = TypeAdapter(list[Annotation]).validate_python(annotation_list, from_attributes=True) return AnnotationList( data=annotation_models, has_more=len(annotation_list) == limit, limit=limit, total=total, page=page @@ -321,7 +323,9 @@ class AnnotationApi(Resource): upsert_args["message_id"] = args.message_id if args.question is not None: upsert_args["question"] = args.question - annotation = AppAnnotationService.up_insert_app_annotation_from_message(upsert_args, str(app_id)) + annotation = AppAnnotationService.up_insert_app_annotation_from_message( + upsert_args, str(app_id), session=db.session() + ) return dump_response(Annotation, annotation), 201 @setup_required @@ -345,11 +349,11 @@ class AnnotationApi(Resource): }, 400 app_ref = _get_app_ref(str(app_id)) - AppAnnotationService.delete_app_annotations_in_batch(app_ref, annotation_ids) + AppAnnotationService.delete_app_annotations_in_batch(app_ref, annotation_ids, session=db.session()) return "", 204 # If no annotation_ids are provided, handle clearing all annotations else: - AppAnnotationService.clear_all_annotations(str(app_id)) + AppAnnotationService.clear_all_annotations(str(app_id), session=db.session()) return "", 204 @@ -370,7 +374,7 @@ class AnnotationExportApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) def get(self, app_id: UUID): - annotation_list = AppAnnotationService.export_annotation_list_by_app_id(str(app_id)) + annotation_list = AppAnnotationService.export_annotation_list_by_app_id(str(app_id), session=db.session()) annotation_models = TypeAdapter(list[Annotation]).validate_python(annotation_list, from_attributes=True) return ( AnnotationExportList(data=annotation_models).model_dump(mode="json"), @@ -406,7 +410,7 @@ class AnnotationUpdateDeleteApi(Resource): update_args["question"] = args.question app_ref = _get_app_ref(str(app_id)) annotation_ref = AppRefService.create_annotation_ref(app_ref, str(annotation_id)) - annotation = AppAnnotationService.update_app_annotation_directly(update_args, annotation_ref, db.session) + annotation = AppAnnotationService.update_app_annotation_directly(update_args, annotation_ref, db.session()) return Annotation.model_validate(annotation, from_attributes=True).model_dump(mode="json") @setup_required @@ -418,7 +422,7 @@ class AnnotationUpdateDeleteApi(Resource): def delete(self, app_id: UUID, annotation_id: UUID): app_ref = _get_app_ref(str(app_id)) annotation_ref = AppRefService.create_annotation_ref(app_ref, str(annotation_id)) - AppAnnotationService.delete_app_annotation(annotation_ref, db.session) + AppAnnotationService.delete_app_annotation(annotation_ref, db.session()) return "", 204 @@ -477,7 +481,7 @@ class AnnotationBatchImportApi(Resource): return dump_response( AnnotationBatchImportResponse, - AppAnnotationService.batch_import_app_annotations(str(app_id), file), + AppAnnotationService.batch_import_app_annotations(str(app_id), file, session=db.session()), ) @@ -538,6 +542,7 @@ class AnnotationHitHistoryListApi(Resource): annotation_ref, page, limit, + session=db.session(), ) history_models = TypeAdapter(list[AnnotationHitHistory]).validate_python( annotation_hit_history_list, from_attributes=True diff --git a/api/controllers/console/app/app.py b/api/controllers/console/app/app.py index a1423318c72..4f0022b37b4 100644 --- a/api/controllers/console/app/app.py +++ b/api/controllers/console/app/app.py @@ -584,6 +584,7 @@ class AppListApi(Resource): permissions = enterprise_rbac_service.RBACService.MyPermissions.get( str(current_tenant_id), current_user_id, + session=db.session(), ) if dify_config.RBAC_ENABLED: access_filter = resolve_app_access_filter( @@ -595,7 +596,7 @@ class AppListApi(Resource): # get app list app_service = AppService() - app_pagination = app_service.get_paginate_apps(current_user_id, current_tenant_id, params, db.session) + app_pagination = app_service.get_paginate_apps(current_user_id, current_tenant_id, params, session) if not app_pagination: response = AppPagination(page=args.page, limit=args.limit, total=0, has_more=False, data=[]) return response.model_dump(mode="json"), 200 @@ -643,11 +644,12 @@ class AppListApi(Resource): ) app_service = AppService() - app = app_service.create_app(current_tenant_id, params, current_user) + app = app_service.create_app(current_tenant_id, params, current_user, session=db.session()) permission_keys_map = enterprise_rbac_service.RBACService.AppPermissions.batch_get( str(current_tenant_id), current_user.id, [str(app.id)], + session=db.session(), ) app_detail = AppDetailWithSite.model_validate(app, from_attributes=True).model_copy( update={"permission_keys": permission_keys_map.get(str(app.id), [])} @@ -681,7 +683,7 @@ class StarredAppListApi(Resource): is_created_by_me=args.is_created_by_me, ) - app_pagination = AppService().get_paginate_starred_apps(current_user_id, current_tenant_id, params, db.session) + app_pagination = AppService().get_paginate_starred_apps(current_user_id, current_tenant_id, params, session) if not app_pagination: empty = AppPagination(page=args.page, limit=args.limit, total=0, has_more=False, data=[]) return empty.model_dump(mode="json"), 200 @@ -705,7 +707,7 @@ class AppStarApi(Resource): @with_session @get_app_model(mode=None) def post(self, session: Session, current_user_id: str, app_model: App): - AppService.star_app(session, app=app_model, account_id=current_user_id) + AppService.star_app(app=app_model, account_id=current_user_id, session=session) return SimpleResultResponse(result="success").model_dump(mode="json") @console_ns.doc("unstar_app") @@ -721,7 +723,7 @@ class AppStarApi(Resource): @with_session @get_app_model(mode=None) def delete(self, session: Session, current_user_id: str, app_model: App): - AppService.unstar_app(session, app=app_model, account_id=current_user_id) + AppService.unstar_app(app=app_model, account_id=current_user_id, session=session) return SimpleResultResponse(result="success").model_dump(mode="json") @@ -753,6 +755,7 @@ class AppApi(Resource): str(current_tenant_id), current_user.id, app_id=str(app_model.id), + session=db.session(), ) permission_keys_map = permissions.app.permission_keys_by_resource_ids([str(app_model.id)]) @@ -789,7 +792,7 @@ class AppApi(Resource): "use_icon_as_answer_icon": args.use_icon_as_answer_icon or False, "max_active_requests": args.max_active_requests or 0, } - app_model = app_service.update_app(app_model, args_dict) + app_model = app_service.update_app(app_model, args_dict, session=db.session()) return dump_response(AppDetailWithSite, app_model) @console_ns.doc("delete_app") @@ -806,7 +809,7 @@ class AppApi(Resource): def delete(self, app_model: App): """Delete app""" app_service = AppService() - app_service.delete_app(app_model) + app_service.delete_app(app_model, session=db.session()) return "", 204 @@ -835,7 +838,7 @@ class AppCopyApi(Resource): with Session(db.engine, expire_on_commit=False) as session: import_service = AppDslService(session) - yaml_content = import_service.export_dsl(app_model=app_model, include_secret=True) + yaml_content = import_service.export_dsl(app_model=app_model, session=session, include_secret=True) result = import_service.import_app( account=current_user, import_mode=ImportMode.YAML_CONTENT, @@ -877,6 +880,7 @@ class AppCopyApi(Resource): str(current_tenant_id), current_user.id, [str(app.id)], + session=db.session(), ) response_model = AppDetailWithSite.model_validate(app, from_attributes=True).model_copy( update={"permission_keys": permission_keys_map.get(str(app.id), [])} @@ -905,6 +909,7 @@ class AppExportApi(Resource): response = AppExportResponse( data=AppDslService.export_dsl( app_model=app_model, + session=db.session(), include_secret=args.include_secret, workflow_id=args.workflow_id, ) @@ -929,7 +934,7 @@ class AppPublishToCreatorsPlatformApi(Resource): if not dify_config.CREATORS_PLATFORM_FEATURES_ENABLED: return {"error": "Creators Platform features are not enabled"}, 403 - dsl_content = AppDslService.export_dsl(app_model=app_model, include_secret=False) + dsl_content = AppDslService.export_dsl(app_model=app_model, session=db.session(), include_secret=False) dsl_bytes = dsl_content.encode("utf-8") claim_code = upload_dsl(dsl_bytes) @@ -955,7 +960,7 @@ class AppNameApi(Resource): args = AppNamePayload.model_validate(console_ns.payload) app_service = AppService() - app_model = app_service.update_app_name(app_model, args.name) + app_model = app_service.update_app_name(app_model, args.name, session=db.session()) return dump_response(AppDetail, app_model) @@ -982,6 +987,7 @@ class AppIconApi(Resource): args.icon or "", args.icon_background or "", args.icon_type, + session=db.session(), ) return dump_response(AppDetail, app_model) @@ -1004,7 +1010,7 @@ class AppSiteStatus(Resource): args = AppSiteStatusPayload.model_validate(console_ns.payload) app_service = AppService() - app_model = app_service.update_app_site_status(app_model, args.enable_site) + app_model = app_service.update_app_site_status(app_model, args.enable_site, session=db.session()) return dump_response(AppDetail, app_model) @@ -1026,7 +1032,7 @@ class AppApiStatus(Resource): args = AppApiStatusPayload.model_validate(console_ns.payload) app_service = AppService() - app_model = app_service.update_app_api_status(app_model, args.enable_api) + app_model = app_service.update_app_api_status(app_model, args.enable_api, session=db.session()) return dump_response(AppDetail, app_model) diff --git a/api/controllers/console/app/audio.py b/api/controllers/console/app/audio.py index c6cd71f30f1..0c9ed786a1e 100644 --- a/api/controllers/console/app/audio.py +++ b/api/controllers/console/app/audio.py @@ -161,7 +161,7 @@ class ChatMessageTextApi(Resource): # response-contract:ignore return AudioService.transcript_tts( app_model=app_model, - session=db.session, + session=db.session(), text=payload.text, voice=payload.voice, message_ref=message_ref, diff --git a/api/controllers/console/app/conversation.py b/api/controllers/console/app/conversation.py index a80935e5e33..b7d422d30b1 100644 --- a/api/controllers/console/app/conversation.py +++ b/api/controllers/console/app/conversation.py @@ -200,7 +200,7 @@ class CompletionConversationDetailApi(Resource): conversation_id_str = str(conversation_id) try: - ConversationService.delete(app_model, conversation_id_str, current_user) + ConversationService.delete(app_model, conversation_id_str, current_user, session=db.session()) except ConversationNotExistsError: raise NotFound("Conversation Not Exists.") @@ -354,7 +354,7 @@ class ChatConversationDetailApi(Resource): conversation_id_str = str(conversation_id) try: - ConversationService.delete(app_model, conversation_id_str, current_user) + ConversationService.delete(app_model, conversation_id_str, current_user, session=db.session()) except ConversationNotExistsError: raise NotFound("Conversation Not Exists.") diff --git a/api/controllers/console/app/message.py b/api/controllers/console/app/message.py index 6a44ca3db8a..958b356de94 100644 --- a/api/controllers/console/app/message.py +++ b/api/controllers/console/app/message.py @@ -363,6 +363,7 @@ def _list_chat_messages(*, app_model: App, current_user: Account | None = None): app_model=app_model, conversation_id=args.conversation_id, user=current_user, + session=db.session(), ) except ConversationNotExistsError: raise NotFound("Conversation Not Exists.") @@ -474,7 +475,11 @@ def _get_message_suggested_questions(*, current_user: Account, app_model: App, m try: questions = MessageService.get_suggested_questions_after_answer( - app_model=app_model, message_id=message_id_str, user=current_user, invoke_from=InvokeFrom.DEBUGGER + app_model=app_model, + message_id=message_id_str, + user=current_user, + invoke_from=InvokeFrom.DEBUGGER, + session=db.session(), ) except MessageNotExistsError: raise NotFound("Message not found") diff --git a/api/controllers/console/app/ops_trace.py b/api/controllers/console/app/ops_trace.py index 46d5ea56e20..e86f65fc035 100644 --- a/api/controllers/console/app/ops_trace.py +++ b/api/controllers/console/app/ops_trace.py @@ -17,6 +17,7 @@ from controllers.console.wraps import ( rbac_permission_required, setup_required, ) +from extensions.ext_database import db from fields.base import ResponseModel from libs.login import login_required from models import App @@ -78,7 +79,7 @@ class TraceAppConfigApi(Resource): try: trace_config = OpsService.get_tracing_app_config( - app_id=app_model.id, tracing_provider=args.tracing_provider + app_id=app_model.id, tracing_provider=args.tracing_provider, session=db.session() ) if not trace_config: return {"has_not_configured": True} @@ -109,7 +110,10 @@ class TraceAppConfigApi(Resource): try: result = OpsService.create_tracing_app_config( - app_id=app_model.id, tracing_provider=args.tracing_provider, tracing_config=args.tracing_config + app_id=app_model.id, + tracing_provider=args.tracing_provider, + tracing_config=args.tracing_config, + session=db.session(), ) if not result: raise TracingConfigIsExist() @@ -142,7 +146,10 @@ class TraceAppConfigApi(Resource): try: result = OpsService.update_tracing_app_config( - app_id=app_model.id, tracing_provider=args.tracing_provider, tracing_config=args.tracing_config + app_id=app_model.id, + tracing_provider=args.tracing_provider, + tracing_config=args.tracing_config, + session=db.session(), ) if not result: raise TracingConfigNotExist() @@ -168,7 +175,9 @@ class TraceAppConfigApi(Resource): args = TraceProviderQuery.model_validate(request.args.to_dict(flat=True)) try: - result = OpsService.delete_tracing_app_config(app_id=app_model.id, tracing_provider=args.tracing_provider) + result = OpsService.delete_tracing_app_config( + app_id=app_model.id, tracing_provider=args.tracing_provider, session=db.session() + ) if not result: raise TracingConfigNotExist() return "", 204 diff --git a/api/controllers/console/app/permission_keys.py b/api/controllers/console/app/permission_keys.py index 810ea04e377..be10f904021 100644 --- a/api/controllers/console/app/permission_keys.py +++ b/api/controllers/console/app/permission_keys.py @@ -1,6 +1,9 @@ +from extensions.ext_database import db from services.enterprise import rbac_service as enterprise_rbac_service def get_app_permission_keys(tenant_id: str, account_id: str | None, app_id: str) -> list[str]: - permission_keys_map = enterprise_rbac_service.RBACService.AppPermissions.batch_get(tenant_id, account_id, [app_id]) + permission_keys_map = enterprise_rbac_service.RBACService.AppPermissions.batch_get( + tenant_id, account_id, [app_id], session=db.session() + ) return permission_keys_map.get(app_id, []) diff --git a/api/controllers/console/app/workflow.py b/api/controllers/console/app/workflow.py index 609dbfb82c5..53c7c6ea788 100644 --- a/api/controllers/console/app/workflow.py +++ b/api/controllers/console/app/workflow.py @@ -2,7 +2,7 @@ import json import logging from collections.abc import Sequence from datetime import datetime -from typing import Any, NotRequired, TypedDict, cast +from typing import Any, NotRequired, TypedDict from flask import abort, request from flask_restx import Resource, fields @@ -522,7 +522,7 @@ class DraftWorkflowApi(Resource): """ # fetch draft workflow by app_model workflow_service = WorkflowService() - workflow = workflow_service.get_draft_workflow(app_model=app_model) + workflow = workflow_service.get_draft_workflow(app_model=app_model, session=db.session()) if not workflow: raise DraftWorkflowNotExist() @@ -533,7 +533,7 @@ class DraftWorkflowApi(Resource): # front-end can treat draft graph node data as the editing source. response = WorkflowResponse.model_validate(workflow, from_attributes=True).model_dump(mode="json") response["graph"] = WorkflowAgentPublishService.project_draft_bindings_to_graph( - session=cast(Session, db.session), + session=db.session(), draft_workflow=workflow, ) return response @@ -602,6 +602,7 @@ class DraftWorkflowApi(Resource): account=current_user, environment_variables=environment_variables, conversation_variables=conversation_variables, + session=db.session(), ) except WorkflowHashNotEqualError: raise DraftWorkflowNotSync() @@ -695,7 +696,12 @@ class AdvancedChatDraftRunIterationNodeApi(Resource): try: response = AppGenerateService.generate_single_iteration( - app_model=app_model, user=current_user, node_id=node_id, args=args, streaming=True + app_model=app_model, + user=current_user, + node_id=node_id, + args=args, + session=db.session(), + streaming=True, ) return helper.compact_generate_response(response) @@ -738,7 +744,12 @@ class WorkflowDraftRunIterationNodeApi(Resource): try: response = AppGenerateService.generate_single_iteration( - app_model=app_model, user=current_user, node_id=node_id, args=args, streaming=True + app_model=app_model, + user=current_user, + node_id=node_id, + args=args, + session=db.session(), + streaming=True, ) return helper.compact_generate_response(response) @@ -777,7 +788,12 @@ class AdvancedChatDraftRunLoopNodeApi(Resource): try: response = AppGenerateService.generate_single_loop( - app_model=app_model, user=current_user, node_id=node_id, args=args, streaming=True + app_model=app_model, + user=current_user, + node_id=node_id, + args=args, + session=db.session(), + streaming=True, ) return helper.compact_generate_response(response) @@ -820,7 +836,12 @@ class WorkflowDraftRunLoopNodeApi(Resource): try: response = AppGenerateService.generate_single_loop( - app_model=app_model, user=current_user, node_id=node_id, args=args, streaming=True + app_model=app_model, + user=current_user, + node_id=node_id, + args=args, + session=db.session(), + streaming=True, ) return helper.compact_generate_response(response) @@ -897,6 +918,7 @@ class AdvancedChatDraftHumanInputFormPreviewApi(Resource): account=current_user, node_id=node_id, inputs=inputs, + session=db.session(), ) return jsonable_encoder(preview) @@ -932,6 +954,7 @@ class AdvancedChatDraftHumanInputFormRunApi(Resource): form_inputs=args.form_inputs, inputs=args.inputs, action=args.action, + session=db.session(), ) return jsonable_encoder(result) @@ -963,6 +986,7 @@ class WorkflowDraftHumanInputFormPreviewApi(Resource): account=current_user, node_id=node_id, inputs=inputs, + session=db.session(), ) return jsonable_encoder(preview) @@ -998,6 +1022,7 @@ class WorkflowDraftHumanInputFormRunApi(Resource): form_inputs=args.form_inputs, inputs=args.inputs, action=args.action, + session=db.session(), ) return jsonable_encoder(result) @@ -1028,6 +1053,7 @@ class WorkflowDraftHumanInputDeliveryTestApi(Resource): node_id=node_id, delivery_method_id=args.delivery_method_id, inputs=args.inputs, + session=db.session(), ) return jsonable_encoder({}) @@ -1138,7 +1164,7 @@ class DraftWorkflowNodeRunApi(Resource): workflow_srv = WorkflowService() # fetch draft workflow by app_model - draft_workflow = workflow_srv.get_draft_workflow(app_model=app_model) + draft_workflow = workflow_srv.get_draft_workflow(app_model=app_model, session=db.session()) if not draft_workflow: raise ValueError("Workflow not initialized") files = _parse_file(draft_workflow, args.get("files")) @@ -1181,7 +1207,7 @@ class PublishedWorkflowApi(Resource): """ # fetch published workflow by app_model workflow_service = WorkflowService() - workflow = workflow_service.get_published_workflow(app_model=app_model) + workflow = workflow_service.get_published_workflow(app_model=app_model, session=db.session()) # return workflow, if not found, return None if workflow is None: @@ -1323,7 +1349,9 @@ class ConvertToWorkflowApi(Resource): # convert to workflow mode workflow_service = WorkflowService() - new_app_model = workflow_service.convert_to_workflow(app_model=app_model, account=current_user, args=args) + new_app_model = workflow_service.convert_to_workflow( + app_model=app_model, account=current_user, args=args, session=db.session() + ) # return app id return { @@ -1358,7 +1386,9 @@ class WorkflowFeaturesApi(Resource): features = args.features.model_dump(mode="json", exclude_unset=True) workflow_service = WorkflowService() - workflow_service.update_draft_workflow_features(app_model=app_model, features=features, account=current_user) + workflow_service.update_draft_workflow_features( + app_model=app_model, features=features, account=current_user, session=db.session() + ) return {"result": "success"} @@ -1439,6 +1469,7 @@ class DraftWorkflowRestoreApi(Resource): app_model=app_model, workflow_id=workflow_id, account=current_user, + session=db.session(), ) except IsDraftWorkflowError as exc: raise BadRequest(RESTORE_SOURCE_WORKFLOW_MUST_BE_PUBLISHED_MESSAGE) from exc @@ -1553,7 +1584,7 @@ class DraftWorkflowNodeLastRunApi(Resource): @get_app_model(mode=[AppMode.ADVANCED_CHAT, AppMode.WORKFLOW]) def get(self, app_model: App, node_id: str): srv = WorkflowService() - workflow = srv.get_draft_workflow(app_model) + workflow = srv.get_draft_workflow(app_model, session=db.session()) if not workflow: raise NotFound("Workflow not found") node_exec = srv.get_node_last_run( @@ -1606,7 +1637,7 @@ class DraftWorkflowTriggerRunApi(Resource): args = DraftWorkflowTriggerRunPayload.model_validate(console_ns.payload or {}) node_id = args.node_id workflow_service = WorkflowService() - draft_workflow = workflow_service.get_draft_workflow(app_model) + draft_workflow = workflow_service.get_draft_workflow(app_model, session=db.session()) if not draft_workflow: raise ValueError("Workflow not found") @@ -1675,7 +1706,7 @@ class DraftWorkflowTriggerNodeApi(Resource): """ workflow_service = WorkflowService() - draft_workflow = workflow_service.get_draft_workflow(app_model) + draft_workflow = workflow_service.get_draft_workflow(app_model, session=db.session()) if not draft_workflow: raise ValueError("Workflow not found") @@ -1759,7 +1790,7 @@ class DraftWorkflowTriggerRunAllApi(Resource): args = DraftWorkflowTriggerRunAllPayload.model_validate(console_ns.payload or {}) node_ids = args.node_ids workflow_service = WorkflowService() - draft_workflow = workflow_service.get_draft_workflow(app_model) + draft_workflow = workflow_service.get_draft_workflow(app_model, session=db.session()) if not draft_workflow: raise ValueError("Workflow not found") @@ -1828,7 +1859,7 @@ class WorkflowOnlineUsersApi(Resource): return {"data": []} workflow_service = WorkflowService() - accessible_app_ids = workflow_service.get_accessible_app_ids(app_ids, current_tenant_id) + accessible_app_ids = workflow_service.get_accessible_app_ids(app_ids, current_tenant_id, session=db.session()) ordered_accessible_app_ids = [app_id for app_id in app_ids if app_id in accessible_app_ids] users_json_by_app_id: dict[str, Any] = {} diff --git a/api/controllers/console/app/workflow_comment.py b/api/controllers/console/app/workflow_comment.py index 64df78f3748..9de91ac59d5 100644 --- a/api/controllers/console/app/workflow_comment.py +++ b/api/controllers/console/app/workflow_comment.py @@ -490,7 +490,7 @@ class WorkflowCommentMentionUsersApi(Resource): current_tenant = current_user.current_tenant # need the tenant object here if current_tenant is None: raise ValueError("current tenant is required") - members = TenantService.get_tenant_members(current_tenant, session=db.session) + members = TenantService.get_tenant_members(current_tenant, session=db.session()) users = TypeAdapter(list[AccountWithRole]).validate_python(members, from_attributes=True) response = WorkflowCommentMentionUsersPayload(users=users) return response.model_dump(mode="json"), 200 diff --git a/api/controllers/console/app/workflow_draft_variable.py b/api/controllers/console/app/workflow_draft_variable.py index 0ccc67f642d..1ffc01e3cef 100644 --- a/api/controllers/console/app/workflow_draft_variable.py +++ b/api/controllers/console/app/workflow_draft_variable.py @@ -337,7 +337,7 @@ class WorkflowVariableCollectionApi(Resource): # fetch draft workflow by app_model workflow_service = WorkflowService() - workflow_exist = workflow_service.is_workflow_exist(app_model=app_model) + workflow_exist = workflow_service.is_workflow_exist(app_model=app_model, session=db.session()) if not workflow_exist: raise DraftWorkflowNotExist() @@ -553,7 +553,7 @@ class VariableResetApi(Resource): ) workflow_srv = WorkflowService() - draft_workflow = workflow_srv.get_draft_workflow(app_model) + draft_workflow = workflow_srv.get_draft_workflow(app_model, session=db.session()) if draft_workflow is None: raise NotFoundError( f"Draft workflow not found, app_id={app_model.id}", @@ -606,7 +606,7 @@ class ConversationVariableCollectionApi(Resource): # NOTE(QuantumGhost): Prefill conversation variables into the draft variables table # so their IDs can be returned to the caller. workflow_srv = WorkflowService() - draft_workflow = workflow_srv.get_draft_workflow(app_model) + draft_workflow = workflow_srv.get_draft_workflow(app_model, session=db.session()) if draft_workflow is None: raise NotFoundError(description=f"draft workflow not found, id={app_model.id}") draft_var_srv = WorkflowDraftVariableService(db.session()) @@ -646,6 +646,7 @@ class ConversationVariableCollectionApi(Resource): app_model=app_model, account=current_user, conversation_variables=conversation_variables, + session=db.session(), ) return {"result": "success"} @@ -683,7 +684,7 @@ class EnvironmentVariableCollectionApi(Resource): """ # fetch draft workflow by app_model workflow_service = WorkflowService() - workflow = workflow_service.get_draft_workflow(app_model=app_model) + workflow = workflow_service.get_draft_workflow(app_model=app_model, session=db.session()) if workflow is None: raise DraftWorkflowNotExist() @@ -740,6 +741,7 @@ class EnvironmentVariableCollectionApi(Resource): app_model=app_model, account=current_user, environment_variables=environment_variables, + session=db.session(), ) return {"result": "success"} diff --git a/api/controllers/console/app/workflow_node_output_inspector.py b/api/controllers/console/app/workflow_node_output_inspector.py index 6ed59d6c566..ea45a718a02 100644 --- a/api/controllers/console/app/workflow_node_output_inspector.py +++ b/api/controllers/console/app/workflow_node_output_inspector.py @@ -41,6 +41,7 @@ from controllers.console.wraps import ( rbac_permission_required, setup_required, ) +from extensions.ext_database import db from libs.exception import BaseHTTPException from libs.login import login_required from models import App, AppMode @@ -92,7 +93,9 @@ def _serve_snapshot(app_model: App, run_id: UUID) -> dict: Flask request context. """ try: - snapshot = _service().snapshot_workflow_run(app_model=app_model, workflow_run_id=str(run_id)) + snapshot = _service().snapshot_workflow_run( + app_model=app_model, workflow_run_id=str(run_id), session=db.session() + ) except NodeOutputInspectorError as error: raise _InspectorNotFound(error) from error return snapshot.model_dump(mode="json") @@ -105,6 +108,7 @@ def _serve_node_detail(app_model: App, run_id: UUID, node_id: str) -> dict: app_model=app_model, workflow_run_id=str(run_id), node_id=node_id, + session=db.session(), ) except NodeOutputInspectorError as error: raise _InspectorNotFound(error) from error @@ -119,6 +123,7 @@ def _serve_output_preview(app_model: App, run_id: UUID, node_id: str, output_nam workflow_run_id=str(run_id), node_id=node_id, output_name=output_name, + session=db.session(), ) except NodeOutputInspectorError as error: raise _InspectorNotFound(error) from error @@ -245,7 +250,7 @@ def _stream_inspector_events(app_model: App, run_id: UUID) -> Iterator[str]: # if the run is gone (raised before yielding any bytes, so Flask turns it # into the normal HTTP 404 path). try: - snapshot = service.snapshot_workflow_run(app_model=app_model, workflow_run_id=run_id_str) + snapshot = service.snapshot_workflow_run(app_model=app_model, workflow_run_id=run_id_str, session=db.session()) except NodeOutputInspectorError as error: raise _InspectorNotFound(error) from error @@ -308,6 +313,7 @@ def _stream_inspector_events(app_model: App, run_id: UUID) -> Iterator[str]: app_model=app_model, workflow_run_id=run_id_str, node_id=message.node_id, + session=db.session(), ) except NodeOutputInspectorError: # Node may not appear in the graph yet (race with persistence); skip. diff --git a/api/controllers/console/auth/activate.py b/api/controllers/console/auth/activate.py index b6045685b55..1f58dbe910f 100644 --- a/api/controllers/console/auth/activate.py +++ b/api/controllers/console/auth/activate.py @@ -90,7 +90,7 @@ class ActivateCheckApi(Resource): token = args.token invitation = RegisterService.get_invitation_with_case_fallback( - workspaceId, args.email, token, session=db.session + workspaceId, args.email, token, session=db.session() ) if invitation: data = invitation.get("data", {}) @@ -140,7 +140,7 @@ class ActivateApi(Resource): normalized_request_email = args.email.lower() if args.email else None invitation = RegisterService.get_invitation_with_case_fallback( - args.workspace_id, args.email, args.token, session=db.session + args.workspace_id, args.email, args.token, session=db.session() ) if invitation is None: raise AlreadyActivateError() @@ -178,7 +178,7 @@ class ActivateApi(Resource): RegisterService.revoke_token(args.workspace_id, normalized_request_email, args.token) if membership_id is None: - TenantService.create_tenant_member(tenant, account, db.session, role=role) + TenantService.create_tenant_member(tenant, account, db.session(), role=role) if setup_fields: account.name = setup_fields[0] @@ -188,6 +188,6 @@ class ActivateApi(Resource): account.status = AccountStatus.ACTIVE account.initialized_at = naive_utc_now() - TenantService.switch_tenant(account, tenant.id, session=db.session) + TenantService.switch_tenant(account, tenant.id, session=db.session()) return {"result": "success"} diff --git a/api/controllers/console/auth/data_source_bearer_auth.py b/api/controllers/console/auth/data_source_bearer_auth.py index 11fab84a831..fac725e8534 100644 --- a/api/controllers/console/auth/data_source_bearer_auth.py +++ b/api/controllers/console/auth/data_source_bearer_auth.py @@ -59,7 +59,7 @@ class ApiKeyAuthDataSource(Resource): @account_initialization_required @with_current_tenant_id def get(self, current_tenant_id: str): - data_source_api_key_bindings = ApiKeyAuthService.get_provider_auth_list(db.session(), current_tenant_id) + data_source_api_key_bindings = ApiKeyAuthService.get_provider_auth_list(current_tenant_id, session=db.session()) if data_source_api_key_bindings: return { "sources": [ @@ -93,7 +93,7 @@ class ApiKeyAuthDataSourceBinding(Resource): data = payload.model_dump() ApiKeyAuthService.validate_api_key_auth_args(data) try: - ApiKeyAuthService.create_provider_auth(db.session(), current_tenant_id, data) + ApiKeyAuthService.create_provider_auth(current_tenant_id, data, session=db.session()) except Exception as e: raise ApiKeyAuthFailedError(str(e)) return {"result": "success"}, 200 @@ -110,6 +110,6 @@ class ApiKeyAuthDataSourceBindingDelete(Resource): @with_current_tenant_id def delete(self, current_tenant_id: str, binding_id: UUID): # The role of the current user in the table must be admin or owner - ApiKeyAuthService.delete_provider_auth(db.session(), current_tenant_id, str(binding_id)) + ApiKeyAuthService.delete_provider_auth(current_tenant_id, str(binding_id), session=db.session()) return "", 204 diff --git a/api/controllers/console/auth/email_register.py b/api/controllers/console/auth/email_register.py index ba4fc1275d9..d89caa9224f 100644 --- a/api/controllers/console/auth/email_register.py +++ b/api/controllers/console/auth/email_register.py @@ -101,7 +101,7 @@ class EmailRegisterSendEmailApi(Resource): if dify_config.BILLING_ENABLED and BillingService.is_email_in_freeze(normalized_email): raise AccountInFreezeError() - account = AccountService.get_account_by_email_with_case_fallback(db.session, args.email) + account = AccountService.get_account_by_email_with_case_fallback(args.email, session=db.session()) token = AccountService.send_email_register_email(email=normalized_email, account=account, language=language) return {"result": "success", "data": token} @@ -176,7 +176,7 @@ class EmailRegisterResetApi(Resource): email = register_data.get("email", "") normalized_email = email.lower() - account = AccountService.get_account_by_email_with_case_fallback(db.session, email) + account = AccountService.get_account_by_email_with_case_fallback(email, session=db.session()) if account: raise EmailAlreadyInUseError() @@ -187,7 +187,7 @@ class EmailRegisterResetApi(Resource): timezone=args.timezone, language=args.language, ) - token_pair = AccountService.login(account=account, session=db.session, ip_address=extract_remote_ip(request)) + token_pair = AccountService.login(account=account, session=db.session(), ip_address=extract_remote_ip(request)) AccountService.reset_login_error_rate_limit(normalized_email) return {"result": "success", "data": token_pair.model_dump()} @@ -206,7 +206,7 @@ class EmailRegisterResetApi(Resource): password=password, interface_language=get_valid_language(language), timezone=timezone, - session=db.session, + session=db.session(), ) except AccountRegisterError: raise AccountInFreezeError() diff --git a/api/controllers/console/auth/forgot_password.py b/api/controllers/console/auth/forgot_password.py index 8df9600070c..6456bb480f4 100644 --- a/api/controllers/console/auth/forgot_password.py +++ b/api/controllers/console/auth/forgot_password.py @@ -82,7 +82,7 @@ class ForgotPasswordSendEmailApi(Resource): else: language = "en-US" - account = AccountService.get_account_by_email_with_case_fallback(db.session, args.email) + account = AccountService.get_account_by_email_with_case_fallback(args.email, session=db.session()) token = AccountService.send_reset_password_email( account=account, @@ -180,7 +180,7 @@ class ForgotPasswordResetApi(Resource): password_hashed = hash_password(args.new_password, salt) email = reset_data.get("email", "") - account = AccountService.get_account_by_email_with_case_fallback(db.session, email) + account = AccountService.get_account_by_email_with_case_fallback(email, session=db.session()) if account: account = db.session.merge(account) @@ -198,10 +198,10 @@ class ForgotPasswordResetApi(Resource): # Create workspace if needed if ( - not TenantService.get_join_tenants(account, session=db.session) + not TenantService.get_join_tenants(account, session=db.session()) and FeatureService.get_system_features().is_allow_create_workspace ): - tenant = TenantService.create_tenant(f"{account.name}'s Workspace", session=db.session) - TenantService.create_tenant_member(tenant, account, db.session, role="owner") + tenant = TenantService.create_tenant(f"{account.name}'s Workspace", session=db.session()) + TenantService.create_tenant_member(tenant, account, db.session(), role="owner") account.current_tenant = tenant tenant_was_created.send(tenant) diff --git a/api/controllers/console/auth/login.py b/api/controllers/console/auth/login.py index 5165fc3003a..486f79bcae2 100644 --- a/api/controllers/console/auth/login.py +++ b/api/controllers/console/auth/login.py @@ -126,7 +126,7 @@ class LoginApi(Resource): invitation_data: InvitationDetailDict | None = None if invite_token: invitation_data = RegisterService.get_invitation_with_case_fallback( - None, request_email, invite_token, session=db.session + None, request_email, invite_token, session=db.session() ) if invitation_data is None: invite_token = None @@ -153,7 +153,7 @@ class LoginApi(Resource): _log_console_login_failure(email=normalized_email, reason=LoginFailureReason.INVALID_CREDENTIALS) raise AuthenticationFailedError() from exc # SELF_HOSTED only have one workspace - tenants = TenantService.get_join_tenants(account, session=db.session) + tenants = TenantService.get_join_tenants(account, session=db.session()) if len(tenants) == 0: system_features = FeatureService.get_system_features() @@ -165,7 +165,7 @@ class LoginApi(Resource): data="workspace not found, please contact system admin to invite you to join in a workspace", ).model_dump(mode="json") - token_pair = AccountService.login(account=account, session=db.session, ip_address=extract_remote_ip(request)) + token_pair = AccountService.login(account=account, session=db.session(), ip_address=extract_remote_ip(request)) AccountService.reset_login_error_rate_limit(normalized_email) # Create response with cookies instead of returning tokens in body @@ -301,7 +301,7 @@ class EmailCodeLoginApi(Resource): _log_console_login_failure(email=user_email, reason=LoginFailureReason.ACCOUNT_IN_FREEZE) raise AccountInFreezeError() if account: - tenants = TenantService.get_join_tenants(account, session=db.session) + tenants = TenantService.get_join_tenants(account, session=db.session()) if not tenants: workspaces = FeatureService.get_system_features().license.workspaces if not workspaces.is_available(): @@ -309,8 +309,8 @@ class EmailCodeLoginApi(Resource): if not FeatureService.get_system_features().is_allow_create_workspace: raise NotAllowedCreateWorkspace() else: - new_tenant = TenantService.create_tenant(f"{account.name}'s Workspace", session=db.session) - TenantService.create_tenant_member(new_tenant, account, db.session, role="owner") + new_tenant = TenantService.create_tenant(f"{account.name}'s Workspace", session=db.session()) + TenantService.create_tenant_member(new_tenant, account, db.session(), role="owner") account.current_tenant = new_tenant tenant_was_created.send(new_tenant) @@ -321,7 +321,7 @@ class EmailCodeLoginApi(Resource): name=user_email, interface_language=get_valid_language(language), timezone=args.timezone, - session=db.session, + session=db.session(), ) except WorkSpaceNotAllowedCreateError: raise NotAllowedCreateWorkspace() @@ -330,7 +330,7 @@ class EmailCodeLoginApi(Resource): raise AccountInFreezeError() except WorkspacesLimitExceededError: raise WorkspacesLimitExceeded() - token_pair = AccountService.login(account, session=db.session, ip_address=extract_remote_ip(request)) + token_pair = AccountService.login(account, session=db.session(), ip_address=extract_remote_ip(request)) AccountService.reset_login_error_rate_limit(user_email) # Create response with cookies instead of returning tokens in body @@ -358,7 +358,7 @@ class RefreshTokenApi(Resource): ), 401 try: - new_token_pair = AccountService.refresh_token(refresh_token, session=db.session) + new_token_pair = AccountService.refresh_token(refresh_token, session=db.session()) except Unauthorized as exc: return SimpleResultMessageResponse(result="fail", message=exc.description or "Unauthorized.").model_dump( mode="json" @@ -378,22 +378,22 @@ class RefreshTokenApi(Resource): def _get_account_with_case_fallback(email: str): - account = AccountService.get_user_through_email(email, session=db.session) + account = AccountService.get_user_through_email(email, session=db.session()) if account or email == email.lower(): return account - return AccountService.get_user_through_email(email.lower(), session=db.session) + return AccountService.get_user_through_email(email.lower(), session=db.session()) def _authenticate_account_with_case_fallback( original_email: str, normalized_email: str, password: str, invite_token: str | None ): try: - return AccountService.authenticate(original_email, password, invite_token, session=db.session) + return AccountService.authenticate(original_email, password, invite_token, session=db.session()) except services.errors.account.AccountPasswordError: if original_email == normalized_email: raise - return AccountService.authenticate(normalized_email, password, invite_token, session=db.session) + return AccountService.authenticate(normalized_email, password, invite_token, session=db.session()) def _log_console_login_failure(*, email: str, reason: LoginFailureReason) -> None: diff --git a/api/controllers/console/auth/oauth.py b/api/controllers/console/auth/oauth.py index 65f3a5addde..5afafd43131 100644 --- a/api/controllers/console/auth/oauth.py +++ b/api/controllers/console/auth/oauth.py @@ -195,7 +195,7 @@ class OAuthCallback(Resource): db.session.commit() try: - TenantService.create_owner_tenant_if_not_exist(account, session=db.session) + TenantService.create_owner_tenant_if_not_exist(account, session=db.session()) except Unauthorized: return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin?message=Workspace not found.") except WorkSpaceNotAllowedCreateError: @@ -206,7 +206,7 @@ class OAuthCallback(Resource): token_pair = AccountService.login( account=account, - session=db.session, + session=db.session(), ip_address=extract_remote_ip(request), ) @@ -225,7 +225,7 @@ def _get_account_by_openid_or_email(provider: str, user_info: OAuthUserInfo) -> account: Account | None = Account.get_by_openid(provider, user_info.id) if not account: - account = AccountService.get_account_by_email_with_case_fallback(db.session, user_info.email) + account = AccountService.get_account_by_email_with_case_fallback(user_info.email, session=db.session()) return account @@ -241,13 +241,13 @@ def _generate_account( oauth_new_user = False if account: - tenants = TenantService.get_join_tenants(account, session=db.session) + tenants = TenantService.get_join_tenants(account, session=db.session()) if not tenants: if not FeatureService.get_system_features().is_allow_create_workspace: raise WorkSpaceNotAllowedCreateError() else: - new_tenant = TenantService.create_tenant(f"{account.name}'s Workspace", session=db.session) - TenantService.create_tenant_member(new_tenant, account, db.session, role="owner") + new_tenant = TenantService.create_tenant(f"{account.name}'s Workspace", session=db.session()) + TenantService.create_tenant_member(new_tenant, account, db.session(), role="owner") account.current_tenant = new_tenant tenant_was_created.send(new_tenant) @@ -273,10 +273,10 @@ def _generate_account( provider=provider, language=interface_language, timezone=timezone, - session=db.session, + session=db.session(), ) # Link account - AccountService.link_account_integrate(provider, user_info.id, account, session=db.session) + AccountService.link_account_integrate(provider, user_info.id, account, session=db.session()) return account, oauth_new_user diff --git a/api/controllers/console/auth/oauth_server.py b/api/controllers/console/auth/oauth_server.py index 46e2983c12b..d068fb0785e 100644 --- a/api/controllers/console/auth/oauth_server.py +++ b/api/controllers/console/auth/oauth_server.py @@ -10,6 +10,7 @@ from werkzeug.exceptions import BadRequest, NotFound from controllers.common.schema import register_response_schema_models, register_schema_models from controllers.console.wraps import account_initialization_required, setup_required, with_current_user +from extensions.ext_database import db from graphon.model_runtime.utils.encoders import jsonable_encoder from libs.login import login_required from models import Account @@ -131,7 +132,9 @@ def oauth_server_access_token_required[T, **P, R]( response.headers["WWW-Authenticate"] = "Bearer" return response - account = OAuthServerService.validate_oauth_access_token(oauth_provider_app.client_id, access_token) + account = OAuthServerService.validate_oauth_access_token( + oauth_provider_app.client_id, access_token, db.session() + ) if not account: response = jsonify({"error": "access_token or client_id is invalid"}) response.status_code = 401 diff --git a/api/controllers/console/billing/billing.py b/api/controllers/console/billing/billing.py index d6974fe129c..3a983b50176 100644 --- a/api/controllers/console/billing/billing.py +++ b/api/controllers/console/billing/billing.py @@ -56,7 +56,7 @@ class Subscription(Resource): @with_current_tenant_id def get(self, current_tenant_id: str, current_user: Account): args = SubscriptionQuery.model_validate(request.args.to_dict(flat=True)) - BillingService.is_tenant_owner_or_admin(db.session, current_user) + BillingService.is_tenant_owner_or_admin(current_user, session=db.session()) return BillingService.get_subscription(args.plan, args.interval, current_user.email, current_tenant_id) @@ -70,7 +70,7 @@ class Invoices(Resource): @with_current_user @with_current_tenant_id def get(self, current_tenant_id: str, current_user: Account): - BillingService.is_tenant_owner_or_admin(db.session, current_user) + BillingService.is_tenant_owner_or_admin(current_user, session=db.session()) return BillingService.get_invoices(current_user.email, current_tenant_id) diff --git a/api/controllers/console/datasets/data_source.py b/api/controllers/console/datasets/data_source.py index b2c8bda0581..17f027df9b3 100644 --- a/api/controllers/console/datasets/data_source.py +++ b/api/controllers/console/datasets/data_source.py @@ -245,7 +245,7 @@ class DataSourceNotionListApi(Resource): exist_page_ids = [] # import notion in the exist dataset if query.dataset_id: - dataset = DatasetService.get_dataset(query.dataset_id, db.session) + dataset = DatasetService.get_dataset(query.dataset_id, db.session()) if not dataset: raise NotFound("Dataset not found.") if dataset.data_source_type != "notion_import": @@ -400,11 +400,11 @@ class DataSourceNotionDatasetSyncApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT) def get(self, dataset_id: UUID) -> tuple[dict[str, str], int]: dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - documents = DocumentService.get_document_by_dataset_id(dataset_id_str, db.session) + documents = DocumentService.get_document_by_dataset_id(dataset_id_str, db.session()) for document in documents: document_indexing_sync_task.delay(dataset_id_str, document.id) return {"result": "success"}, 200 @@ -420,11 +420,11 @@ class DataSourceNotionDocumentSyncApi(Resource): def get(self, dataset_id: UUID, document_id: UUID) -> tuple[dict[str, str], int]: dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if document is None: raise NotFound("Document not found.") document_indexing_sync_task.delay(dataset_id_str, document_id_str) diff --git a/api/controllers/console/datasets/datasets.py b/api/controllers/console/datasets/datasets.py index 0bef535d82b..c8ca1d621c4 100644 --- a/api/controllers/console/datasets/datasets.py +++ b/api/controllers/console/datasets/datasets.py @@ -418,6 +418,7 @@ class DatasetListApi(Resource): permissions = enterprise_rbac_service.RBACService.MyPermissions.get( str(current_tenant_id), current_user.id, + session=db.session(), ) accessible_dataset_ids: list[str] | None = None @@ -461,7 +462,7 @@ class DatasetListApi(Resource): datasets, total = DatasetService.get_datasets( query.page, query.limit, - db.session, + db.session(), current_tenant_id, current_user, query.keyword, @@ -573,6 +574,7 @@ class DatasetListApi(Resource): current_tenant_id, current_user.id, [dataset.id], + session=session, ) item = DatasetDetailWithPartialMembersResponse.model_validate(dataset, from_attributes=True).model_dump( @@ -602,17 +604,18 @@ class DatasetApi(Resource): @with_current_tenant_id def get(self, current_tenant_id: str, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) permissions = enterprise_rbac_service.RBACService.MyPermissions.get( current_tenant_id, current_user.id, dataset_id=dataset_id_str, + session=db.session(), ) permission_keys_map = permissions.dataset.permission_keys_by_resource_ids([dataset_id_str]) data = dump_response(DatasetDetailResponse, dataset) @@ -622,7 +625,7 @@ class DatasetApi(Resource): provider_id = ModelProviderID(dataset.embedding_model_provider) data["embedding_model_provider"] = str(provider_id) if data.get("permission") == "partial_members": - part_users_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session) + part_users_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session()) data.update({"partial_member_list": part_users_list}) # check embedding setting @@ -666,7 +669,7 @@ class DatasetApi(Resource): @with_session def patch(self, session: Session, current_tenant_id: str, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") @@ -685,10 +688,10 @@ class DatasetApi(Resource): # The role of the current user in the ta table must be admin, owner, editor, or dataset_operator if not dify_config.RBAC_ENABLED: DatasetPermissionService.check_permission( - session, current_user, dataset, payload.permission, payload.partial_member_list + current_user, dataset, payload.permission, payload.partial_member_list, session=session ) - dataset = DatasetService.update_dataset(session, dataset_id_str, payload_data, current_user) + dataset = DatasetService.update_dataset(dataset_id_str, payload_data, current_user, session=session) if dataset is None: raise NotFound("Dataset not found.") @@ -697,6 +700,7 @@ class DatasetApi(Resource): current_tenant_id, current_user.id, [dataset_id_str], + session=session, ) result_data = dump_response(DatasetDetailResponse, dataset) result_data["permission_keys"] = permission_keys_map.get(dataset_id_str, []) @@ -704,13 +708,13 @@ class DatasetApi(Resource): if payload.partial_member_list is not None and payload.permission == DatasetPermissionEnum.PARTIAL_TEAM: DatasetPermissionService.update_partial_member_list( - tenant_id, dataset_id_str, payload.partial_member_list, db.session + tenant_id, dataset_id_str, payload.partial_member_list, db.session() ) # clear partial member list when permission is only_me or all_team_members elif payload.permission in {DatasetPermissionEnum.ONLY_ME, DatasetPermissionEnum.ALL_TEAM}: - DatasetPermissionService.clear_partial_member_list(dataset_id_str, db.session) + DatasetPermissionService.clear_partial_member_list(dataset_id_str, db.session()) - partial_member_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session) + partial_member_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session()) result_data.update({"partial_member_list": partial_member_list}) return result_data, 200 @@ -729,8 +733,8 @@ class DatasetApi(Resource): raise Forbidden() try: - if DatasetService.delete_dataset(dataset_id_str, current_user, db.session): - DatasetPermissionService.clear_partial_member_list(dataset_id_str, db.session) + if DatasetService.delete_dataset(dataset_id_str, current_user, db.session()): + DatasetPermissionService.clear_partial_member_list(dataset_id_str, db.session()) return "", 204 else: raise NotFound("Dataset not found.") @@ -755,7 +759,7 @@ class DatasetUseCheckApi(Resource): def get(self, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset_is_using = DatasetService.dataset_use_check(dataset_id_str, db.session) + dataset_is_using = DatasetService.dataset_use_check(dataset_id_str, db.session()) return {"is_using": dataset_is_using}, 200 @@ -776,12 +780,12 @@ class DatasetQueryApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) def get(self, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -917,16 +921,16 @@ class DatasetRelatedAppListApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) def get(self, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - app_dataset_joins = DatasetService.get_related_apps(dataset.id, db.session) + app_dataset_joins = DatasetService.get_related_apps(dataset.id, db.session()) related_apps = [] for app_dataset_join in app_dataset_joins: @@ -1101,7 +1105,7 @@ class DatasetEnableApiApi(Resource): def post(self, dataset_id: UUID, status: str): dataset_id_str = str(dataset_id) - DatasetService.update_dataset_api_status(dataset_id_str, status == "enable", db.session) + DatasetService.update_dataset_api_status(dataset_id_str, status == "enable", db.session()) return {"result": "success"}, 200 @@ -1170,10 +1174,10 @@ class DatasetErrorDocs(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) def get(self, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - results = DocumentService.get_error_documents_by_dataset_id(dataset_id_str, db.session) + results = DocumentService.get_error_documents_by_dataset_id(dataset_id_str, db.session()) return dump_response(ErrorDocsResponse, {"data": results, "total": len(results)}), 200 @@ -1197,15 +1201,15 @@ class DatasetPermissionUserListApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) def get(self, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - partial_members_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session) + partial_members_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session()) return dump_response(PartialMemberListResponse, {"data": partial_members_list}), 200 @@ -1227,8 +1231,8 @@ class DatasetAutoDisableLogApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) def get(self, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - auto_disable_logs = DatasetService.get_dataset_auto_disable_logs(dataset_id_str, db.session) + auto_disable_logs = DatasetService.get_dataset_auto_disable_logs(dataset_id_str, db.session()) return dump_response(AutoDisableLogsResponse, auto_disable_logs), 200 diff --git a/api/controllers/console/datasets/datasets_document.py b/api/controllers/console/datasets/datasets_document.py index a6263c8e2e3..ee441704b20 100644 --- a/api/controllers/console/datasets/datasets_document.py +++ b/api/controllers/console/datasets/datasets_document.py @@ -183,16 +183,16 @@ class DocumentResource(Resource): def get_document( self, dataset_id: str, document_id: str, current_user: Account, current_tenant_id: str ) -> Document: - dataset = DatasetService.get_dataset(dataset_id, db.session) + dataset = DatasetService.get_dataset(dataset_id, db.session()) if not dataset: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - document = DocumentService.get_document(dataset_id, document_id, session=db.session) + document = DocumentService.get_document(dataset_id, document_id, session=db.session()) if not document: raise NotFound("Document not found.") @@ -203,16 +203,16 @@ class DocumentResource(Resource): return document def get_batch_documents(self, dataset_id: str, batch: str, current_user: Account) -> Sequence[Document]: - dataset = DatasetService.get_dataset(dataset_id, db.session) + dataset = DatasetService.get_dataset(dataset_id, db.session()) if not dataset: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - documents = DocumentService.get_batch_documents(dataset_id, batch, db.session) + documents = DocumentService.get_batch_documents(dataset_id, batch, db.session()) if not documents: raise NotFound("Documents not found.") @@ -243,13 +243,13 @@ class GetProcessRuleApi(Resource): # get the latest process rule document = db.get_or_404(Document, document_id) - dataset = DatasetService.get_dataset(document.dataset_id, db.session) + dataset = DatasetService.get_dataset(document.dataset_id, db.session()) if not dataset: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -319,12 +319,12 @@ class DatasetDocumentListApi(Resource): ) except (ArgumentTypeError, ValueError, Exception): fetch = False - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -376,6 +376,7 @@ class DatasetDocumentListApi(Resource): documents=documents, dataset=dataset, tenant_id=current_tenant_id, + session=db.session(), ) if fetch: @@ -423,7 +424,7 @@ class DatasetDocumentListApi(Resource): def post(self, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") @@ -433,7 +434,7 @@ class DatasetDocumentListApi(Resource): raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -447,9 +448,9 @@ class DatasetDocumentListApi(Resource): try: documents, batch = DocumentService.save_document_with_dataset_id( - dataset, knowledge_config, current_user, session=db.session + dataset, knowledge_config, current_user, session=db.session() ) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -468,7 +469,7 @@ class DatasetDocumentListApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) def delete(self, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") # check user's model setting @@ -477,7 +478,7 @@ class DatasetDocumentListApi(Resource): try: document_ids = request.args.getlist("document_id") dataset_ref = DatasetRefService.create_dataset_ref(dataset) - DocumentService.delete_documents(dataset_ref, document_ids, dataset.doc_form, db.session) + DocumentService.delete_documents(dataset_ref, document_ids, dataset.doc_form, db.session()) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Cannot delete document during indexing.") @@ -536,7 +537,7 @@ class DatasetInitApi(Resource): tenant_id=current_tenant_id, knowledge_config=knowledge_config, account=current_user, - session=db.session, + session=db.session(), ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -873,7 +874,7 @@ class DocumentApi(DocumentResource): if metadata == "only": response = {"id": document.id, "doc_type": document.doc_type, "doc_metadata": document.doc_metadata_details} elif metadata == "without": - dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, db.session) + dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, db.session()) document_process_rules = document.dataset_process_rule.to_dict() if document.dataset_process_rule else {} response = { "id": document.id, @@ -907,7 +908,7 @@ class DocumentApi(DocumentResource): "need_summary": document.need_summary if document.need_summary is not None else False, } else: - dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, db.session) + dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, db.session()) document_process_rules = document.dataset_process_rule.to_dict() if document.dataset_process_rule else {} response = { "id": document.id, @@ -956,7 +957,7 @@ class DocumentApi(DocumentResource): def delete(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID): dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") # check user's model setting @@ -965,7 +966,7 @@ class DocumentApi(DocumentResource): document = self.get_document(dataset_id_str, document_id_str, current_user, current_tenant_id) try: - DocumentService.delete_document(document, db.session) + DocumentService.delete_document(document, db.session()) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Cannot delete document during indexing.") @@ -989,7 +990,7 @@ class DocumentDownloadApi(DocumentResource): def get(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID) -> dict[str, Any]: # Reuse the shared permission/tenant checks implemented in DocumentResource. document = self.get_document(str(dataset_id), str(document_id), current_user, current_tenant_id) - return {"url": DocumentService.get_document_download_url(document, db.session)} + return {"url": DocumentService.get_document_download_url(document, db.session())} @console_ns.route("/datasets//documents/download-zip") @@ -1019,7 +1020,7 @@ class DocumentBatchDownloadZipApi(DocumentResource): document_ids=document_ids, tenant_id=current_tenant_id, current_user=current_user, - session=db.session, + session=db.session(), ) # Delegate ZIP packing to FileService, but keep Flask response+cleanup in the route. @@ -1168,7 +1169,7 @@ class DocumentStatusApi(DocumentResource): self, current_user: Account, dataset_id: UUID, action: Literal["enable", "disable", "archive", "un_archive"] ): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") @@ -1180,12 +1181,12 @@ class DocumentStatusApi(DocumentResource): DatasetService.check_dataset_model_setting(dataset) # check user's permission - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) document_ids = request.args.getlist("document_id") try: - DocumentService.batch_update_document_status(dataset, document_ids, action, current_user, db.session) + DocumentService.batch_update_document_status(dataset, document_ids, action, current_user, db.session()) except services.errors.document.DocumentIndexingError as e: raise InvalidActionError(str(e)) except ValueError as e: @@ -1209,11 +1210,11 @@ class DocumentPauseApi(DocumentResource): dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) # 404 if document not found if document is None: @@ -1225,7 +1226,7 @@ class DocumentPauseApi(DocumentResource): try: # pause document - DocumentService.pause_document(document, db.session) + DocumentService.pause_document(document, db.session()) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Cannot pause completed document.") @@ -1244,10 +1245,10 @@ class DocumentRecoverApi(DocumentResource): """recover document.""" dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) # 404 if document not found if document is None: @@ -1258,7 +1259,7 @@ class DocumentRecoverApi(DocumentResource): raise ArchivedDocumentImmutableError() try: # pause document - DocumentService.recover_document(document, db.session) + DocumentService.recover_document(document, db.session()) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Document is not in paused status.") @@ -1278,13 +1279,13 @@ class DocumentRetryApi(DocumentResource): """retry document.""" payload = DocumentRetryPayload.model_validate(console_ns.payload or {}) dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) retry_documents = [] if not dataset: raise NotFound("Dataset not found.") for document_id in payload.document_ids: try: - document = DocumentService.get_document(dataset.id, document_id, session=db.session) + document = DocumentService.get_document(dataset.id, document_id, session=db.session()) # 404 if document not found if document is None: @@ -1302,7 +1303,7 @@ class DocumentRetryApi(DocumentResource): logger.exception("Failed to retry document, document id: %s", document_id) continue # retry document - DocumentService.retry_document(dataset_id_str, retry_documents, db.session) + DocumentService.retry_document(dataset_id_str, retry_documents, db.session()) return "", 204 @@ -1320,14 +1321,14 @@ class DocumentRenameApi(DocumentResource): # 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() - dataset = DatasetService.get_dataset(dataset_id, db.session) + dataset = DatasetService.get_dataset(dataset_id, db.session()) if not dataset: raise NotFound("Dataset not found.") - DatasetService.check_dataset_operator_permission(current_user, dataset, session=db.session) + DatasetService.check_dataset_operator_permission(current_user, dataset, session=db.session()) payload = DocumentRenamePayload.model_validate(console_ns.payload or {}) try: - document = DocumentService.rename_document(str(dataset_id), str(document_id), payload.name, db.session) + document = DocumentService.rename_document(str(dataset_id), str(document_id), payload.name, db.session()) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Cannot delete document during indexing.") @@ -1345,11 +1346,11 @@ class WebsiteDocumentSyncApi(DocumentResource): def get(self, current_tenant_id: str, dataset_id: UUID, document_id: UUID): """sync website document.""" dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") document_id_str = str(document_id) - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") if document.tenant_id != current_tenant_id: @@ -1360,7 +1361,7 @@ class WebsiteDocumentSyncApi(DocumentResource): if DocumentService.check_archived(document): raise ArchivedDocumentImmutableError() # sync document - DocumentService.sync_website_document(dataset_id_str, document, db.session) + DocumentService.sync_website_document(dataset_id_str, document, db.session()) return {"result": "success"}, 200 @@ -1380,10 +1381,10 @@ class DocumentPipelineExecutionLogApi(DocumentResource): dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") log = db.session.scalar( @@ -1438,7 +1439,7 @@ class DocumentGenerateSummaryApi(Resource): dataset_id_str = str(dataset_id) # Get dataset - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") @@ -1447,7 +1448,7 @@ class DocumentGenerateSummaryApi(Resource): raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -1472,7 +1473,7 @@ class DocumentGenerateSummaryApi(Resource): raise ValueError("Summary index is not enabled for this dataset. Please enable it in the dataset settings.") # Verify all documents exist and belong to the dataset - documents = DocumentService.get_documents_by_ids(dataset_id_str, document_list, db.session) + documents = DocumentService.get_documents_by_ids(dataset_id_str, document_list, db.session()) if len(documents) != len(document_list): found_ids = {doc.id for doc in documents} @@ -1488,7 +1489,7 @@ class DocumentGenerateSummaryApi(Resource): DocumentService.update_documents_need_summary( dataset_id=dataset_id_str, document_ids=document_ids_to_update, - session=db.session, + session=db.session(), need_summary=True, ) @@ -1539,13 +1540,13 @@ class DocumentSummaryStatusApi(DocumentResource): document_id_str = str(document_id) # Get dataset - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # Check permissions try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -1555,7 +1556,7 @@ class DocumentSummaryStatusApi(DocumentResource): result = SummaryIndexService.get_document_summary_status_detail( document_id=document_id_str, dataset_id=dataset_id_str, - session=db.session, + session=db.session(), ) return result, 200 diff --git a/api/controllers/console/datasets/datasets_segments.py b/api/controllers/console/datasets/datasets_segments.py index 5cccd2453dc..e4f2abeb844 100644 --- a/api/controllers/console/datasets/datasets_segments.py +++ b/api/controllers/console/datasets/datasets_segments.py @@ -173,7 +173,7 @@ def _get_segment_for_document( raise NotFound("Document not found.") segment_ref = DatasetRefService.create_segment_ref(document_ref, segment_id) - segment = SegmentService.get_segment_by_ref(segment_ref) + segment = SegmentService.get_segment_by_ref(segment_ref, db.session()) if not segment: raise NotFound("Segment not found.") return segment_ref, segment @@ -193,16 +193,16 @@ class DatasetDocumentSegmentListApi(Resource): def get(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID): dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") @@ -278,7 +278,7 @@ class DatasetDocumentSegmentListApi(Resource): summaries: dict[str, str | None] = {} if segment_ids: summary_records = SummaryIndexService.get_segments_summaries( - segment_ids=segment_ids, dataset_id=dataset_id_str + segment_ids=segment_ids, dataset_id=dataset_id_str, session=db.session() ) summaries = {chunk_id: summary.summary_content for chunk_id, summary in summary_records.items()} @@ -303,14 +303,14 @@ class DatasetDocumentSegmentListApi(Resource): def delete(self, current_user: Account, dataset_id: UUID, document_id: UUID): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") segment_ids = request.args.getlist("segment_id") @@ -319,10 +319,10 @@ class DatasetDocumentSegmentListApi(Resource): if not current_user.is_dataset_editor: raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - SegmentService.delete_segments(segment_ids, document, dataset, db.session) + SegmentService.delete_segments(segment_ids, document, dataset, db.session()) return "", 204 @@ -348,11 +348,11 @@ class DatasetDocumentSegmentApi(Resource): action: Literal["enable", "disable"], ): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") # check user's model setting @@ -362,7 +362,7 @@ class DatasetDocumentSegmentApi(Resource): raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: @@ -388,7 +388,7 @@ class DatasetDocumentSegmentApi(Resource): if cache_result is not None: raise InvalidActionError("Document is being indexed, please try again later") try: - SegmentService.update_segments_status(segment_ids, action, dataset, document, db.session) + SegmentService.update_segments_status(segment_ids, action, dataset, document, db.session()) except Exception as e: raise InvalidActionError(str(e)) return dump_response(SimpleResultResponse, {"result": "success"}), 200 @@ -411,12 +411,12 @@ class DatasetDocumentSegmentAddApi(Resource): def post(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") if not current_user.is_dataset_editor: @@ -438,15 +438,20 @@ class DatasetDocumentSegmentAddApi(Resource): except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) # validate args payload = SegmentCreatePayload.model_validate(console_ns.payload or {}) payload_dict = payload.model_dump(exclude_none=True) SegmentService.segment_create_args_validate(payload_dict, document) - segment = type_cast(DocumentSegment, SegmentService.create_segment(payload_dict, document, dataset, db.session)) - summary = SummaryIndexService.get_segment_summary(segment_id=segment.id, dataset_id=dataset_id_str) + segment = type_cast( + DocumentSegment, + SegmentService.create_segment(payload_dict, document, dataset, db.session()), + ) + summary = SummaryIndexService.get_segment_summary( + segment_id=segment.id, dataset_id=dataset_id_str, session=db.session() + ) response = { "data": segment_response_with_summary(segment, summary.summary_content if summary else None), "doc_form": document.doc_form, @@ -472,21 +477,21 @@ class DatasetDocumentSegmentUpdateApi(Resource): ): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") # The role of the current user in the ta table must be admin, owner, dataset_operator, or editor if not current_user.is_dataset_editor: raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: @@ -518,9 +523,11 @@ class DatasetDocumentSegmentUpdateApi(Resource): segment, document, dataset, - db.session, + db.session(), + ) + summary = SummaryIndexService.get_segment_summary( + segment_id=segment.id, dataset_id=dataset_id_str, session=db.session() ) - summary = SummaryIndexService.get_segment_summary(segment_id=segment.id, dataset_id=dataset_id_str) response = { "data": segment_response_with_summary(segment, summary.summary_content if summary else None), "doc_form": document.doc_form, @@ -541,26 +548,26 @@ class DatasetDocumentSegmentUpdateApi(Resource): ): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") # The role of the current user in the ta table must be admin, owner, dataset_operator, or editor if not current_user.is_dataset_editor: raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) segment_id_str = str(segment_id) _, segment = _get_segment_for_document(dataset, document, segment_id_str) - SegmentService.delete_segment(segment, document, dataset, db.session) + SegmentService.delete_segment(segment, document, dataset, db.session()) return "", 204 @@ -583,12 +590,12 @@ class DatasetDocumentSegmentBatchImportApi(Resource): def post(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") @@ -658,18 +665,18 @@ class ChildChunkAddApi(Resource): ): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") if not current_user.is_dataset_editor: raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) # check embedding model setting @@ -693,7 +700,7 @@ class ChildChunkAddApi(Resource): # validate args try: payload = ChildChunkCreatePayload.model_validate(console_ns.payload or {}) - child_chunk = SegmentService.create_child_chunk(payload.content, segment, document, dataset, db.session) + child_chunk = SegmentService.create_child_chunk(payload.content, segment, document, dataset, db.session()) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) return dump_response(ChildChunkDetailResponse, {"data": child_chunk}), 200 @@ -709,14 +716,14 @@ class ChildChunkAddApi(Resource): def get(self, current_tenant_id: str, dataset_id: UUID, document_id: UUID, segment_id: UUID): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") segment_id_str = str(segment_id) @@ -759,21 +766,21 @@ class ChildChunkAddApi(Resource): ): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") # The role of the current user in the ta table must be admin, owner, dataset_operator, or editor if not current_user.is_dataset_editor: raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) segment_id_str = str(segment_id) @@ -781,7 +788,7 @@ class ChildChunkAddApi(Resource): # validate args payload = ChildChunkBatchUpdatePayload.model_validate(console_ns.payload or {}) try: - child_chunks = SegmentService.update_child_chunks(payload.chunks, segment, document, dataset, db.session) + child_chunks = SegmentService.update_child_chunks(payload.chunks, segment, document, dataset, db.session()) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) return dump_response(ChildChunkBatchUpdateResponse, {"data": child_chunks}), 200 @@ -811,31 +818,31 @@ class ChildChunkUpdateApi(Resource): ): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") # The role of the current user in the ta table must be admin, owner, dataset_operator, or editor if not current_user.is_dataset_editor: raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) segment_id_str = str(segment_id) segment_ref, _ = _get_segment_for_document(dataset, document, segment_id_str) child_chunk_id_str = str(child_chunk_id) - child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref) + child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref, db.session()) if not child_chunk: raise NotFound("Child chunk not found.") try: - SegmentService.delete_child_chunk(child_chunk, dataset, db.session) + SegmentService.delete_child_chunk(child_chunk, dataset, db.session()) except ChildChunkDeleteIndexServiceError as e: raise ChildChunkDeleteIndexError(str(e)) return "", 204 @@ -862,34 +869,34 @@ class ChildChunkUpdateApi(Resource): ): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") # The role of the current user in the ta table must be admin, owner, dataset_operator, or editor if not current_user.is_dataset_editor: raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) segment_id_str = str(segment_id) segment_ref, segment = _get_segment_for_document(dataset, document, segment_id_str) child_chunk_id_str = str(child_chunk_id) - child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref) + child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref, db.session()) if not child_chunk: raise NotFound("Child chunk not found.") # validate args try: payload = ChildChunkUpdatePayload.model_validate(console_ns.payload or {}) child_chunk = SegmentService.update_child_chunk( - payload.content, child_chunk, segment, document, dataset, db.session + payload.content, child_chunk, segment, document, dataset, db.session() ) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) diff --git a/api/controllers/console/datasets/external.py b/api/controllers/console/datasets/external.py index 5b036641d4d..9cdca96f69a 100644 --- a/api/controllers/console/datasets/external.py +++ b/api/controllers/console/datasets/external.py @@ -299,7 +299,9 @@ class ExternalApiTemplateApi(Resource): if not (current_user.has_edit_permission or current_user.is_dataset_operator): raise Forbidden() - ExternalDatasetService.delete_external_knowledge_api(session, current_tenant_id, external_knowledge_api_id_str) + ExternalDatasetService.delete_external_knowledge_api( + current_tenant_id, external_knowledge_api_id_str, session=session + ) return "", 204 @@ -318,9 +320,7 @@ class ExternalApiUseCheckApi(Resource): external_knowledge_api_id_str = str(external_knowledge_api_id) external_knowledge_api_is_using, count = ExternalDatasetService.external_knowledge_api_use_check( - session, - external_knowledge_api_id_str, - current_tenant_id, + external_knowledge_api_id_str, current_tenant_id, session=session ) return {"is_using": external_knowledge_api_is_using, "count": count}, 200 @@ -366,6 +366,7 @@ class ExternalDatasetCreateApi(Resource): str(current_tenant_id), current_user.id, [dataset_id_str], + session=session, ) item["permission_keys"] = permission_keys_map.get(dataset_id_str, []) @@ -393,12 +394,12 @@ class ExternalKnowledgeHitTestingApi(Resource): @with_session def post(self, session: Session, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) diff --git a/api/controllers/console/datasets/hit_testing_base.py b/api/controllers/console/datasets/hit_testing_base.py index cc02a990168..656a426c125 100644 --- a/api/controllers/console/datasets/hit_testing_base.py +++ b/api/controllers/console/datasets/hit_testing_base.py @@ -86,12 +86,12 @@ class DatasetsHitTestingBase: dataset_id: str, current_user: Account | None = None, current_tenant_id: str | None = None ) -> Dataset: current_user, _ = resolve_account_fallback(current_user, current_tenant_id) - dataset = DatasetService.get_dataset(dataset_id, db.session) + dataset = DatasetService.get_dataset(dataset_id, db.session()) if dataset is None: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) diff --git a/api/controllers/console/datasets/metadata.py b/api/controllers/console/datasets/metadata.py index 8802fcf2814..42ae4903673 100644 --- a/api/controllers/console/datasets/metadata.py +++ b/api/controllers/console/datasets/metadata.py @@ -61,13 +61,13 @@ class DatasetMetadataCreateApi(Resource): metadata_args = MetadataArgs.model_validate(console_ns.payload or {}) dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) metadata = MetadataService.create_metadata( - db.session(), dataset_id_str, metadata_args, current_user, current_tenant_id + dataset_id_str, metadata_args, current_user, current_tenant_id, session=db.session() ) return dump_response(DatasetMetadataResponse, metadata), 201 @@ -81,10 +81,10 @@ class DatasetMetadataCreateApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT) def get(self, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - metadata = MetadataService.get_dataset_metadatas(db.session(), dataset) + metadata = MetadataService.get_dataset_metadatas(dataset, session=db.session()) return dump_response(DatasetMetadataListResponse, metadata), 200 @@ -105,13 +105,13 @@ class DatasetMetadataApi(Resource): dataset_id_str = str(dataset_id) metadata_id_str = str(metadata_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) metadata = MetadataService.update_metadata_name( - db.session(), dataset_id_str, metadata_id_str, name, current_user, current_tenant_id + dataset_id_str, metadata_id_str, name, current_user, current_tenant_id, session=db.session() ) return dump_response(DatasetMetadataResponse, metadata), 200 @@ -125,12 +125,12 @@ class DatasetMetadataApi(Resource): def delete(self, current_user: Account, dataset_id: UUID, metadata_id: UUID): dataset_id_str = str(dataset_id) metadata_id_str = str(metadata_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) - MetadataService.delete_metadata(db.session(), dataset_id_str, metadata_id_str) + MetadataService.delete_metadata(dataset_id_str, metadata_id_str, session=db.session()) # Frontend callers only await success and invalidate metadata caches; no response body is consumed. return "", 204 @@ -162,16 +162,16 @@ class DatasetMetadataBuiltInFieldActionApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) def post(self, current_user: Account, dataset_id: UUID, action: Literal["enable", "disable"]): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) match action: case "enable": - MetadataService.enable_built_in_field(db.session(), dataset) + MetadataService.enable_built_in_field(dataset, session=db.session()) case "disable": - MetadataService.disable_built_in_field(db.session(), dataset) + MetadataService.disable_built_in_field(dataset, session=db.session()) # Frontend callers only await success and invalidate metadata caches; no response body is consumed. return "", 204 @@ -191,14 +191,14 @@ class DocumentMetadataEditApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) def post(self, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) metadata_args = MetadataOperationData.model_validate(console_ns.payload or {}) - MetadataService.update_documents_metadata(db.session(), dataset, metadata_args, current_user) + MetadataService.update_documents_metadata(dataset, metadata_args, current_user, session=db.session()) # Frontend callers only await success and invalidate caches; no response body is consumed. return "", 204 diff --git a/api/controllers/console/datasets/rag_pipeline/datasource_auth.py b/api/controllers/console/datasets/rag_pipeline/datasource_auth.py index 389515bd4f9..57d6b628d4b 100644 --- a/api/controllers/console/datasets/rag_pipeline/datasource_auth.py +++ b/api/controllers/console/datasets/rag_pipeline/datasource_auth.py @@ -23,6 +23,7 @@ from core.entities.provider_entities import ProviderConfig from core.plugin.entities.plugin_daemon import PluginOAuthAuthorizationUrlResponse from core.plugin.impl.oauth import OAuthHandler from core.tools.entities.common_entities import I18nObject +from extensions.ext_database import db from fields.base import ResponseModel from graphon.model_runtime.errors.validate import CredentialsValidateFailedError from libs.helper import dump_response @@ -309,6 +310,7 @@ class DatasourceAuth(Resource): provider=datasource_provider_id.provider_name, plugin_id=datasource_provider_id.plugin_id, user=user, + session=db.session(), ) return dump_response(DatasourceCredentialListResponse, {"result": datasources}), 200 @@ -335,6 +337,7 @@ class DatasourceAuthDeleteApi(Resource): auth_id=payload.credential_id, provider=provider_name, plugin_id=plugin_id, + session=db.session(), ) return SimpleResultResponse(result="success").model_dump(mode="json"), 200 @@ -380,7 +383,9 @@ class DatasourceAuthListApi(Resource): @with_current_tenant_id def get(self, current_tenant_id: str): datasource_provider_service = DatasourceProviderService() - datasources = datasource_provider_service.get_all_datasource_credentials(tenant_id=current_tenant_id) + datasources = datasource_provider_service.get_all_datasource_credentials( + tenant_id=current_tenant_id, session=db.session() + ) return dump_response(DatasourceProviderAuthListResponse, {"result": datasources}), 200 @@ -397,7 +402,9 @@ class DatasourceHardCodeAuthListApi(Resource): @with_current_tenant_id def get(self, current_tenant_id: str): datasource_provider_service = DatasourceProviderService() - datasources = datasource_provider_service.get_hard_code_datasource_credentials(tenant_id=current_tenant_id) + datasources = datasource_provider_service.get_hard_code_datasource_credentials( + tenant_id=current_tenant_id, session=db.session() + ) return dump_response(DatasourceProviderAuthListResponse, {"result": datasources}), 200 diff --git a/api/controllers/console/datasets/rag_pipeline/datasource_content_preview.py b/api/controllers/console/datasets/rag_pipeline/datasource_content_preview.py index 213337fedc9..873ba130064 100644 --- a/api/controllers/console/datasets/rag_pipeline/datasource_content_preview.py +++ b/api/controllers/console/datasets/rag_pipeline/datasource_content_preview.py @@ -9,6 +9,7 @@ from controllers.common.schema import register_schema_models from controllers.console import console_ns from controllers.console.datasets.wraps import get_rag_pipeline from controllers.console.wraps import account_initialization_required, setup_required, with_current_user +from extensions.ext_database import db from libs.login import login_required from models import Account from models.dataset import Pipeline @@ -41,7 +42,7 @@ class DataSourceContentPreviewApi(Resource): inputs = args.inputs datasource_type = args.datasource_type - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) preview_content = rag_pipeline_service.run_datasource_node_preview( pipeline=pipeline, node_id=node_id, diff --git a/api/controllers/console/datasets/rag_pipeline/rag_pipeline.py b/api/controllers/console/datasets/rag_pipeline/rag_pipeline.py index 4027fa487a2..2d824afb6ef 100644 --- a/api/controllers/console/datasets/rag_pipeline/rag_pipeline.py +++ b/api/controllers/console/datasets/rag_pipeline/rag_pipeline.py @@ -108,7 +108,10 @@ class PipelineTemplateListApi(Resource): query = PipelineTemplateListQuery.model_validate(request.args.to_dict(flat=True)) # get pipeline templates pipeline_templates = RagPipelineService.get_pipeline_templates( - session, query.type, query.language, current_tenant_id + type=query.type, + language=query.language, + current_tenant_id=current_tenant_id, + session=session, ) return dump_response(PipelineTemplateListResponse, pipeline_templates), 200 @@ -124,8 +127,11 @@ class PipelineTemplateDetailApi(Resource): @with_session def get(self, session: Session, template_id: str) -> JsonResponseWithStatus: query = PipelineTemplateDetailQuery.model_validate(request.args.to_dict(flat=True)) - rag_pipeline_service = RagPipelineService() - pipeline_template = rag_pipeline_service.get_pipeline_template_detail(session, template_id, query.type) + pipeline_template = RagPipelineService.get_pipeline_template_detail( + template_id, + type=query.type, + session=session, + ) if pipeline_template is None: raise NotFound("Pipeline template not found from upstream service.") return dump_response(PipelineTemplateDetailResponse, pipeline_template), 200 @@ -145,7 +151,7 @@ class CustomizedPipelineTemplateApi(Resource): payload = CustomizedPipelineTemplatePayload.model_validate(console_ns.payload or {}) pipeline_template_info = PipelineTemplateInfoEntity.model_validate(payload.model_dump()) RagPipelineService.update_customized_pipeline_template( - template_id, pipeline_template_info, current_user, current_tenant_id + template_id, pipeline_template_info, current_user, current_tenant_id, session=db.session() ) return "", 204 @@ -156,7 +162,7 @@ class CustomizedPipelineTemplateApi(Resource): @enterprise_license_required @with_current_tenant_id def delete(self, current_tenant_id: str, template_id: str) -> tuple[str, int]: - RagPipelineService.delete_customized_pipeline_template(template_id, current_tenant_id) + RagPipelineService.delete_customized_pipeline_template(template_id, current_tenant_id, session=db.session()) return "", 204 @setup_required @@ -188,8 +194,8 @@ class PublishCustomizedPipelineTemplateApi(Resource): @with_current_tenant_id def post(self, current_tenant_id: str, current_user: Account, pipeline_id: str) -> tuple[str, int]: payload = CustomizedPipelineTemplatePayload.model_validate(console_ns.payload or {}) - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) rag_pipeline_service.publish_customized_pipeline_template( - pipeline_id, payload.model_dump(), current_user, current_tenant_id + pipeline_id, payload.model_dump(), current_user, current_tenant_id, session=db.session() ) return "", 204 diff --git a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_datasets.py b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_datasets.py index a373c8b1a41..5ad764871e4 100644 --- a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_datasets.py +++ b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_datasets.py @@ -65,7 +65,7 @@ class CreateRagPipelineDatasetApi(Resource): yaml_content=payload.yaml_content, ) try: - rag_pipeline_dsl_service = RagPipelineDslService(db.session) + rag_pipeline_dsl_service = RagPipelineDslService(db.session()) import_info = rag_pipeline_dsl_service.create_rag_pipeline_dataset( tenant_id=current_tenant_id, rag_pipeline_dataset_create_entity=rag_pipeline_dataset_create_entity, @@ -75,7 +75,7 @@ class CreateRagPipelineDatasetApi(Resource): current_tenant_id, import_info["dataset_id"], rag_pipeline_dataset_create_entity.partial_member_list, - db.session, + db.session(), ) db.session.commit() except services.errors.dataset.DatasetNameDuplicateError: @@ -110,6 +110,6 @@ class CreateEmptyRagPipelineDatasetApi(Resource): permission=DatasetPermissionEnum.ONLY_ME, partial_member_list=None, ), - session=db.session, + session=db.session(), ) return dump_response(DatasetDetailResponse, dataset), 201 diff --git a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_draft_variable.py b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_draft_variable.py index af417f24dfe..25628a67177 100644 --- a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_draft_variable.py +++ b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_draft_variable.py @@ -98,7 +98,7 @@ class RagPipelineVariableCollectionApi(Resource): query = PaginationQuery.model_validate(request.args.to_dict()) # fetch draft workflow by app_model - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) workflow_exist = rag_pipeline_service.is_workflow_exist(pipeline=pipeline) if not workflow_exist: raise DraftWorkflowNotExist() @@ -290,7 +290,7 @@ class RagPipelineVariableResetApi(Resource): session=db.session(), ) - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) draft_workflow = rag_pipeline_service.get_draft_workflow(pipeline=pipeline) if draft_workflow is None: raise NotFoundError( @@ -347,7 +347,7 @@ class RagPipelineEnvironmentVariableCollectionApi(Resource): Get draft workflow """ # fetch draft workflow by app_model - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) workflow = rag_pipeline_service.get_draft_workflow(pipeline=pipeline) if workflow is None: raise DraftWorkflowNotExist() diff --git a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py index c52385f6cf2..a61fc2639db 100644 --- a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py +++ b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py @@ -197,7 +197,7 @@ class DraftRagPipelineApi(Resource): Get draft rag pipeline's workflow """ # fetch draft workflow by app_model - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) workflow = rag_pipeline_service.get_draft_workflow(pipeline=pipeline) if not workflow: @@ -231,7 +231,7 @@ class DraftRagPipelineApi(Resource): return {"message": "Invalid JSON data"}, 400 else: abort(415) - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) try: environment_variables_list = Workflow.normalize_environment_variable_mappings( @@ -283,7 +283,7 @@ class RagPipelineDraftRunIterationNodeApi(Resource): try: response = PipelineGenerateService.generate_single_iteration( - pipeline=pipeline, user=current_user, node_id=node_id, args=args, streaming=True + pipeline=pipeline, user=current_user, node_id=node_id, args=args, session=db.session(), streaming=True ) return helper.compact_generate_response(response) @@ -318,7 +318,7 @@ class RagPipelineDraftRunLoopNodeApi(Resource): try: response = PipelineGenerateService.generate_single_loop( - pipeline=pipeline, user=current_user, node_id=node_id, args=args, streaming=True + pipeline=pipeline, user=current_user, node_id=node_id, args=args, session=db.session(), streaming=True ) return helper.compact_generate_response(response) @@ -343,8 +343,8 @@ class DraftRagPipelineRunApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_current_user - @get_rag_pipeline @with_session + @get_rag_pipeline def post(self, session: Session, current_user: Account, pipeline: Pipeline): """ Run draft workflow @@ -377,8 +377,8 @@ class PublishedRagPipelineRunApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_current_user - @get_rag_pipeline @with_session + @get_rag_pipeline def post(self, session: Session, current_user: Account, pipeline: Pipeline): """ Run published workflow @@ -419,7 +419,7 @@ class RagPipelinePublishedDatasourceNodeRunApi(Resource): """ payload = DatasourceNodeRunPayload.model_validate(console_ns.payload or {}) - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) return helper.compact_generate_response( PipelineGenerator.convert_to_event_stream( rag_pipeline_service.run_datasource_workflow_node( @@ -452,7 +452,7 @@ class RagPipelineDraftDatasourceNodeRunApi(Resource): """ payload = DatasourceNodeRunPayload.model_validate(console_ns.payload or {}) - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) return helper.compact_generate_response( PipelineGenerator.convert_to_event_stream( rag_pipeline_service.run_datasource_workflow_node( @@ -490,7 +490,7 @@ class RagPipelineDraftNodeRunApi(Resource): payload = NodeRunRequiredPayload.model_validate(console_ns.payload or {}) inputs = payload.inputs - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) workflow_node_execution = rag_pipeline_service.run_draft_workflow_node( pipeline=pipeline, node_id=node_id, user_inputs=inputs, account=current_user ) @@ -543,7 +543,7 @@ class PublishedRagPipelineApi(Resource): if not pipeline.is_published: return None # fetch published workflow by pipeline - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) workflow = rag_pipeline_service.get_published_workflow(pipeline=pipeline) # return workflow, if not found, return None @@ -564,9 +564,9 @@ class PublishedRagPipelineApi(Resource): """ Publish workflow """ - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) workflow = rag_pipeline_service.publish_workflow( - session=db.session, # type: ignore[reportArgumentType,arg-type] + session=db.session(), pipeline=pipeline, account=current_user, ) @@ -599,7 +599,7 @@ class DefaultRagPipelineBlockConfigsApi(Resource): Get default block config """ # Get default block configs - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) return rag_pipeline_service.get_default_block_configs() @@ -631,7 +631,7 @@ class DefaultRagPipelineBlockConfigApi(Resource): raise ValueError("Invalid filters") # Get default block configs - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) return rag_pipeline_service.get_default_block_config(node_type=block_type, filters=filters) @@ -666,7 +666,7 @@ class PublishedAllRagPipelineApi(Resource): if user_id != current_user.id: raise Forbidden() - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) with sessionmaker(db.engine).begin() as session: workflows, has_more = rag_pipeline_service.get_all_published_workflow( session=session, @@ -698,7 +698,7 @@ class RagPipelineDraftWorkflowRestoreApi(Resource): @with_current_user @get_rag_pipeline def post(self, current_user: Account, pipeline: Pipeline, workflow_id: str): - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) try: workflow = rag_pipeline_service.restore_published_workflow_to_draft( @@ -743,7 +743,7 @@ class RagPipelineByIdApi(Resource): if not update_data: return {"message": "No valid fields to update"}, 400 - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) workflow_ref = WorkflowRefService.create_pipeline_workflow_ref(pipeline, workflow_id) # Create a session and manage the transaction @@ -809,7 +809,7 @@ class PublishedRagPipelineSecondStepApi(Resource): """ query = NodeIdQuery.model_validate(request.args.to_dict()) node_id = query.node_id - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) variables = rag_pipeline_service.get_second_step_parameters(pipeline=pipeline, node_id=node_id, is_draft=False) return { "variables": variables, @@ -832,7 +832,7 @@ class PublishedRagPipelineFirstStepApi(Resource): """ query = NodeIdQuery.model_validate(request.args.to_dict()) node_id = query.node_id - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) variables = rag_pipeline_service.get_first_step_parameters(pipeline=pipeline, node_id=node_id, is_draft=False) return { "variables": variables, @@ -855,7 +855,7 @@ class DraftRagPipelineFirstStepApi(Resource): """ query = NodeIdQuery.model_validate(request.args.to_dict()) node_id = query.node_id - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) variables = rag_pipeline_service.get_first_step_parameters(pipeline=pipeline, node_id=node_id, is_draft=True) return { "variables": variables, @@ -879,7 +879,7 @@ class DraftRagPipelineSecondStepApi(Resource): query = NodeIdQuery.model_validate(request.args.to_dict()) node_id = query.node_id - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) variables = rag_pipeline_service.get_second_step_parameters(pipeline=pipeline, node_id=node_id, is_draft=True) return { "variables": variables, @@ -913,7 +913,7 @@ class RagPipelineWorkflowRunListApi(Resource): "limit": query.limit, } - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) result = rag_pipeline_service.get_rag_pipeline_paginate_workflow_runs(pipeline=pipeline, args=args) return WorkflowRunPaginationResponse.model_validate(result, from_attributes=True).model_dump(mode="json") @@ -936,7 +936,7 @@ class RagPipelineWorkflowRunDetailApi(Resource): """ run_id_str = str(run_id) - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) workflow_run = rag_pipeline_service.get_rag_pipeline_workflow_run(pipeline=pipeline, run_id=run_id_str) if workflow_run is None: raise NotFound("Workflow run not found") @@ -962,7 +962,7 @@ class RagPipelineWorkflowRunNodeExecutionListApi(Resource): """ run_id_str = str(run_id) - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) user = cast("Account | EndUser", current_user) node_executions = rag_pipeline_service.get_rag_pipeline_workflow_run_node_executions( pipeline=pipeline, @@ -998,7 +998,7 @@ class RagPipelineWorkflowLastRunApi(Resource): @account_initialization_required @get_rag_pipeline def get(self, pipeline: Pipeline, node_id: str): - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) workflow = rag_pipeline_service.get_draft_workflow(pipeline=pipeline) if not workflow: raise NotFound("Workflow not found") @@ -1051,7 +1051,7 @@ class RagPipelineDatasourceVariableApi(Resource): """ args = DatasourceVariablesPayload.model_validate(console_ns.payload or {}).model_dump() - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) workflow_node_execution = rag_pipeline_service.set_datasource_variables( pipeline=pipeline, args=args, @@ -1074,6 +1074,6 @@ class RagPipelineRecommendedPluginApi(Resource): def get(self, current_tenant_id: str, current_user: Account): query = RagPipelineRecommendedPluginQuery.model_validate(request.args.to_dict()) - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) recommended_plugins = rag_pipeline_service.get_recommended_plugins(query.type, current_user, current_tenant_id) return recommended_plugins diff --git a/api/controllers/console/datasets/wraps.py b/api/controllers/console/datasets/wraps.py index b58a07029c8..b5a9cd753ff 100644 --- a/api/controllers/console/datasets/wraps.py +++ b/api/controllers/console/datasets/wraps.py @@ -2,6 +2,7 @@ from collections.abc import Callable from functools import wraps from sqlalchemy import select +from sqlalchemy.orm import Session from controllers.console.datasets.error import PipelineNotFoundError from extensions.ext_database import db @@ -22,9 +23,10 @@ def get_rag_pipeline[**P, R](view_func: Callable[P, R]) -> Callable[P, R]: del kwargs["pipeline_id"] - pipeline = db.session.scalar( - select(Pipeline).where(Pipeline.id == pipeline_id, Pipeline.tenant_id == current_tenant_id).limit(1) - ) + stmt = select(Pipeline).where(Pipeline.id == pipeline_id, Pipeline.tenant_id == current_tenant_id).limit(1) + # Migrated handlers pass the request Session as args[1]; legacy handlers still use db.session. + session = args[1] if len(args) > 1 and isinstance(args[1], Session) else db.session + pipeline = session.scalar(stmt) if not pipeline: raise PipelineNotFoundError() diff --git a/api/controllers/console/explore/audio.py b/api/controllers/console/explore/audio.py index c0b86c19e43..e5f98f0f655 100644 --- a/api/controllers/console/explore/audio.py +++ b/api/controllers/console/explore/audio.py @@ -113,7 +113,7 @@ class ChatTextApi(InstalledAppResource): response = AudioService.transcript_tts( app_model=app_model, - session=db.session, + session=db.session(), text=text, voice=voice, message_ref=message_ref, diff --git a/api/controllers/console/explore/conversation.py b/api/controllers/console/explore/conversation.py index 2004e648f19..25239203d8d 100644 --- a/api/controllers/console/explore/conversation.py +++ b/api/controllers/console/explore/conversation.py @@ -111,7 +111,7 @@ class ConversationApi(InstalledAppResource): conversation_id = str(c_id) try: - ConversationService.delete(app_model, conversation_id, current_user) + ConversationService.delete(app_model, conversation_id, current_user, session=db.session()) except ConversationNotExistsError: raise NotFound("Conversation Not Exists.") @@ -140,7 +140,7 @@ class ConversationRenameApi(InstalledAppResource): try: conversation = ConversationService.rename( - app_model, conversation_id, current_user, payload.name, payload.auto_generate + app_model, conversation_id, current_user, payload.name, payload.auto_generate, session=db.session() ) return ( TypeAdapter(SimpleConversation) @@ -169,7 +169,7 @@ class ConversationPinApi(InstalledAppResource): conversation_id = str(c_id) try: - WebConversationService.pin(app_model, conversation_id, current_user) + WebConversationService.pin(app_model, conversation_id, current_user, db.session()) except ConversationNotExistsError: raise NotFound("Conversation Not Exists.") @@ -192,6 +192,6 @@ class ConversationUnPinApi(InstalledAppResource): raise NotChatAppError() conversation_id = str(c_id) - WebConversationService.unpin(app_model, conversation_id, current_user) + WebConversationService.unpin(app_model, conversation_id, current_user, db.session()) return ResultResponse(result="success").model_dump(mode="json") diff --git a/api/controllers/console/explore/installed_app.py b/api/controllers/console/explore/installed_app.py index 71cb03ce6a0..1fe1201bab7 100644 --- a/api/controllers/console/explore/installed_app.py +++ b/api/controllers/console/explore/installed_app.py @@ -181,7 +181,7 @@ class InstalledAppsListApi(Resource): if current_user.current_tenant is None: raise ValueError("current_user.current_tenant must not be None") - current_user.role = TenantService.get_user_role(current_user, current_user.current_tenant, session=db.session) + current_user.role = TenantService.get_user_role(current_user, current_user.current_tenant, session=db.session()) installed_app_list: list[dict[str, Any]] = [] for installed_app, app_model in installed_apps: installed_app_list.append( diff --git a/api/controllers/console/explore/message.py b/api/controllers/console/explore/message.py index 0e27e2db25b..7b316b0382d 100644 --- a/api/controllers/console/explore/message.py +++ b/api/controllers/console/explore/message.py @@ -27,6 +27,7 @@ from controllers.console.explore.wraps import InstalledAppResource from controllers.console.wraps import with_current_user from core.app.entities.app_invoke_entities import InvokeFrom from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError +from extensions.ext_database import db from fields.conversation_fields import ResultResponse from fields.message_fields import ( ExploreMessageInfiniteScrollPagination, @@ -91,6 +92,7 @@ class MessageListApi(InstalledAppResource): args.conversation_id, args.first_id or None, args.limit, + session=db.session(), ) adapter = TypeAdapter(ExploreMessageListItem) items = [adapter.validate_python(message, from_attributes=True) for message in pagination.data] @@ -129,6 +131,7 @@ class MessageFeedbackApi(InstalledAppResource): user=current_user, rating=FeedbackRating(payload.rating) if payload.rating else None, content=payload.content, + session=db.session(), ) except MessageNotExistsError: raise NotFound("Message Not Exists.") @@ -207,7 +210,11 @@ class MessageSuggestedQuestionApi(InstalledAppResource): try: questions = MessageService.get_suggested_questions_after_answer( - app_model=app_model, user=current_user, message_id=message_id_str, invoke_from=InvokeFrom.EXPLORE + app_model=app_model, + user=current_user, + message_id=message_id_str, + invoke_from=InvokeFrom.EXPLORE, + session=db.session(), ) except MessageNotExistsError: raise NotFound("Message not found") diff --git a/api/controllers/console/explore/parameter.py b/api/controllers/console/explore/parameter.py index 0bc6e032bf0..680885f9bd3 100644 --- a/api/controllers/console/explore/parameter.py +++ b/api/controllers/console/explore/parameter.py @@ -8,6 +8,7 @@ from controllers.console import console_ns from controllers.console.app.error import AppUnavailableError from controllers.console.explore.wraps import InstalledAppResource from core.app.app_config.common.parameters_mapping import get_parameters_from_feature_dict +from extensions.ext_database import db from models.model import AppMode, InstalledApp from services.app_service import AppService @@ -64,4 +65,4 @@ class ExploreAppMetaApi(InstalledAppResource): app_model = installed_app.app if not app_model: raise ValueError("App not found") - return AppService().get_app_meta(app_model) + return AppService().get_app_meta(app_model, session=db.session()) diff --git a/api/controllers/console/explore/recommended_app.py b/api/controllers/console/explore/recommended_app.py index abe170bf90a..79eaa305d61 100644 --- a/api/controllers/console/explore/recommended_app.py +++ b/api/controllers/console/explore/recommended_app.py @@ -120,7 +120,7 @@ class RecommendedAppListApi(Resource): language_prefix = _resolve_language(args.language, current_user) return RecommendedAppListResponse.model_validate( - RecommendedAppService.get_recommended_apps_and_categories(db.session, language_prefix), + RecommendedAppService.get_recommended_apps_and_categories(language_prefix, session=db.session()), from_attributes=True, ).model_dump(mode="json") @@ -137,7 +137,7 @@ class LearnDifyAppListApi(Resource): language_prefix = _resolve_language(args.language, current_user) return LearnDifyAppListResponse.model_validate( - RecommendedAppService.get_learn_dify_apps(db.session, language_prefix), + RecommendedAppService.get_learn_dify_apps(language_prefix, session=db.session()), from_attributes=True, ).model_dump(mode="json") @@ -148,4 +148,4 @@ class RecommendedAppApi(Resource): @login_required @account_initialization_required def get(self, app_id: UUID): - return RecommendedAppService.get_recommend_app_detail(db.session, str(app_id)) + return RecommendedAppService.get_recommend_app_detail(str(app_id), session=db.session()) diff --git a/api/controllers/console/explore/saved_message.py b/api/controllers/console/explore/saved_message.py index ce43ff18c93..e3fd730a3cc 100644 --- a/api/controllers/console/explore/saved_message.py +++ b/api/controllers/console/explore/saved_message.py @@ -38,11 +38,7 @@ class SavedMessageListApi(InstalledAppResource): args = SavedMessageListQuery.model_validate(request.args.to_dict()) pagination = SavedMessageService.pagination_by_last_id( - db.session(), - app_model, - current_user, - str(args.last_id) if args.last_id else None, - args.limit, + app_model, current_user, str(args.last_id) if args.last_id else None, args.limit, session=db.session() ) adapter = TypeAdapter(SavedMessageItem) items = [adapter.validate_python(message, from_attributes=True) for message in pagination.data] @@ -65,7 +61,7 @@ class SavedMessageListApi(InstalledAppResource): payload = SavedMessageCreatePayload.model_validate(console_ns.payload or {}) try: - SavedMessageService.save(db.session(), app_model, current_user, str(payload.message_id)) + SavedMessageService.save(app_model, current_user, str(payload.message_id), session=db.session()) except MessageNotExistsError: raise NotFound("Message Not Exists.") @@ -88,6 +84,6 @@ class SavedMessageApi(InstalledAppResource): if app_model.mode != "completion": raise NotCompletionAppError() - SavedMessageService.delete(db.session(), app_model, current_user, message_id_str) + SavedMessageService.delete(app_model, current_user, message_id_str, session=db.session()) return "", 204 diff --git a/api/controllers/console/explore/trial.py b/api/controllers/console/explore/trial.py index b28116c9a2d..d01eb9c38b1 100644 --- a/api/controllers/console/explore/trial.py +++ b/api/controllers/console/explore/trial.py @@ -431,7 +431,7 @@ class TrialAppWorkflowRunApi(TrialAppResource): invoke_from=InvokeFrom.EXPLORE, streaming=True, ) - RecommendedAppService.add_trial_app_record(db.session, app_id, user_id) + RecommendedAppService.add_trial_app_record(app_id, user_id, session=session) # response-contract:ignore compact_generate_response return helper.compact_generate_response(response) except ProviderTokenNotInitError as ex: @@ -511,7 +511,7 @@ class TrialChatApi(TrialAppResource): invoke_from=InvokeFrom.EXPLORE, streaming=True, ) - RecommendedAppService.add_trial_app_record(db.session, app_id, user_id) + RecommendedAppService.add_trial_app_record(app_id, user_id, session=session) # response-contract:ignore compact_generate_response return helper.compact_generate_response(response) except services.errors.conversation.ConversationNotExistsError: @@ -551,7 +551,11 @@ class TrialMessageSuggestedQuestionApi(TrialAppResource): try: questions = MessageService.get_suggested_questions_after_answer( - app_model=app_model, user=current_user, message_id=message_id, invoke_from=InvokeFrom.EXPLORE + app_model=app_model, + user=current_user, + message_id=message_id, + invoke_from=InvokeFrom.EXPLORE, + session=db.session(), ) except MessageNotExistsError: raise NotFound("Message not found") @@ -589,7 +593,7 @@ class TrialChatAudioApi(TrialAppResource): user_id = current_user.id response = AudioService.transcript_asr(app_model=app_model, file=file, end_user=None) - RecommendedAppService.add_trial_app_record(db.session, app_id, user_id) + RecommendedAppService.add_trial_app_record(app_id, user_id, session=db.session()) return response except services.errors.app_model_config.AppModelConfigBrokenError: logger.exception("App model config broken.") @@ -645,12 +649,12 @@ class TrialChatTextApi(TrialAppResource): response = AudioService.transcript_tts( app_model=app_model, - session=db.session, + session=db.session(), text=text, voice=voice, message_ref=message_ref, ) - RecommendedAppService.add_trial_app_record(db.session, app_id, user_id) + RecommendedAppService.add_trial_app_record(app_id, user_id, session=db.session()) return response except services.errors.app_model_config.AppModelConfigBrokenError: logger.exception("App model config broken.") @@ -709,7 +713,7 @@ class TrialCompletionApi(TrialAppResource): streaming=streaming, ) - RecommendedAppService.add_trial_app_record(db.session, app_id, user_id) + RecommendedAppService.add_trial_app_record(app_id, user_id, session=session) # response-contract:ignore compact_generate_response return helper.compact_generate_response(response) except services.errors.conversation.ConversationNotExistsError: diff --git a/api/controllers/console/extension.py b/api/controllers/console/extension.py index 4b149b9c08d..cc06204a905 100644 --- a/api/controllers/console/extension.py +++ b/api/controllers/console/extension.py @@ -112,7 +112,7 @@ class APIBasedExtensionAPI(Resource): def get(self, current_tenant_id: str): return dump_response( APIBasedExtensionListResponse, - APIBasedExtensionService.get_all_by_tenant_id(db.session(), current_tenant_id), + APIBasedExtensionService.get_all_by_tenant_id(current_tenant_id, session=db.session()), ) @console_ns.doc("create_api_based_extension") @@ -133,7 +133,7 @@ class APIBasedExtensionAPI(Resource): api_key=payload.api_key, ) - extension = APIBasedExtensionService.save(db.session(), extension_data) + extension = APIBasedExtensionService.save(extension_data, session=db.session()) return APIBasedExtensionResponse( id=extension.id, name=extension.name, @@ -158,7 +158,9 @@ class APIBasedExtensionDetailAPI(Resource): return dump_response( APIBasedExtensionResponse, - APIBasedExtensionService.get_with_tenant_id(db.session(), current_tenant_id, api_based_extension_id), + APIBasedExtensionService.get_with_tenant_id( + current_tenant_id, api_based_extension_id, session=db.session() + ), ) @console_ns.doc("update_api_based_extension") @@ -174,7 +176,7 @@ class APIBasedExtensionDetailAPI(Resource): api_based_extension_id = str(id) extension_data_from_db = APIBasedExtensionService.get_with_tenant_id( - db.session(), current_tenant_id, api_based_extension_id + current_tenant_id, api_based_extension_id, session=db.session() ) payload = APIBasedExtensionPayload.model_validate(console_ns.payload or {}) @@ -187,7 +189,7 @@ class APIBasedExtensionDetailAPI(Resource): extension_data_from_db.api_key = payload.api_key api_key_for_response = payload.api_key - APIBasedExtensionService.save(db.session(), extension_data_from_db) + APIBasedExtensionService.save(extension_data_from_db, session=db.session()) return APIBasedExtensionResponse( id=extension_data_from_db.id, name=extension_data_from_db.name, @@ -208,9 +210,9 @@ class APIBasedExtensionDetailAPI(Resource): api_based_extension_id = str(id) extension_data_from_db = APIBasedExtensionService.get_with_tenant_id( - db.session(), current_tenant_id, api_based_extension_id + current_tenant_id, api_based_extension_id, session=db.session() ) - APIBasedExtensionService.delete(db.session(), extension_data_from_db) + APIBasedExtensionService.delete(extension_data_from_db, session=db.session()) return "", 204 diff --git a/api/controllers/console/init_validate.py b/api/controllers/console/init_validate.py index 27f6bcc36dc..f155f222e19 100644 --- a/api/controllers/console/init_validate.py +++ b/api/controllers/console/init_validate.py @@ -50,7 +50,7 @@ def get_init_status() -> InitStatusResponse: @only_edition_self_hosted def validate_init_password(payload: InitValidatePayload) -> InitValidateResponse: """Validate initialization password.""" - tenant_count = TenantService.get_tenant_count(session=db.session) + tenant_count = TenantService.get_tenant_count(session=db.session()) if tenant_count > 0: raise AlreadySetupError() diff --git a/api/controllers/console/setup.py b/api/controllers/console/setup.py index 2b99693a9ca..e0a0fba3329 100644 --- a/api/controllers/console/setup.py +++ b/api/controllers/console/setup.py @@ -79,7 +79,7 @@ def setup_system(payload: SetupRequestPayload) -> SetupResponse: if get_setup_status(): raise AlreadySetupError() - tenant_count = TenantService.get_tenant_count(session=db.session) + tenant_count = TenantService.get_tenant_count(session=db.session()) if tenant_count > 0: raise AlreadySetupError() @@ -94,7 +94,7 @@ def setup_system(payload: SetupRequestPayload) -> SetupResponse: password=payload.password, ip_address=extract_remote_ip(request), language=payload.language, - session=db.session, + session=db.session(), ) mark_setup_completed() diff --git a/api/controllers/console/socketio/workflow.py b/api/controllers/console/socketio/workflow.py index 99e56df3cb8..db5a4144dd3 100644 --- a/api/controllers/console/socketio/workflow.py +++ b/api/controllers/console/socketio/workflow.py @@ -44,7 +44,7 @@ def socket_connect(sid, environ, auth): return False with sio.app.app_context(): - user = AccountService.load_logged_in_account(account_id=user_id, session=db.session) + user = AccountService.load_logged_in_account(account_id=user_id, session=db.session()) if not user: logging.warning("Socket connect rejected: user not found (user_id=%s, sid=%s)", user_id, sid) return False @@ -69,7 +69,7 @@ def handle_user_connect(sid, data): if not workflow_id: return {"msg": "workflow_id is required"}, 400 - result = collaboration_service.authorize_and_join_workflow_room(workflow_id, sid) + result = collaboration_service.authorize_and_join_workflow_room(workflow_id, sid, session=db.session()) if not result: return {"msg": "unauthorized"}, 401 diff --git a/api/controllers/console/tag/tags.py b/api/controllers/console/tag/tags.py index c4ec925c9a3..86c1ad9c54c 100644 --- a/api/controllers/console/tag/tags.py +++ b/api/controllers/console/tag/tags.py @@ -137,7 +137,7 @@ class TagListApi(Resource): def get(self, current_tenant_id: str): raw_args = request.args.to_dict() param = TagListQueryParam.model_validate(raw_args) - tags = TagService.get_tags(db.session(), param.type, current_tenant_id, param.keyword) + tags = TagService.get_tags(param.type, current_tenant_id, param.keyword, session=db.session()) return dump_response(TagListResponse, tags), 200 @@ -154,7 +154,7 @@ class TagListApi(Resource): payload = TagBasePayload.model_validate(console_ns.payload or {}) _enforce_snippet_tag_rbac_if_needed(payload.type) - tag = TagService.save_tags(SaveTagPayload(name=payload.name, type=payload.type), db.session) + tag = TagService.save_tags(SaveTagPayload(name=payload.name, type=payload.type), db.session()) return dump_response(TagResponse, {"id": tag.id, "name": tag.name, "type": tag.type, "binding_count": 0}), 200 @@ -175,9 +175,9 @@ class TagUpdateDeleteApi(Resource): payload = TagUpdateRequestPayload.model_validate(console_ns.payload or {}) _enforce_snippet_tag_rbac_by_tag_id(tag_id_str) - tag = TagService.update_tags(UpdateTagPayload(name=payload.name), tag_id_str, db.session) + tag = TagService.update_tags(UpdateTagPayload(name=payload.name), tag_id_str, db.session()) - binding_count = TagService.get_tag_binding_count(tag_id_str, db.session) + binding_count = TagService.get_tag_binding_count(tag_id_str, db.session()) return ( dump_response( @@ -196,7 +196,7 @@ class TagUpdateDeleteApi(Resource): tag_id_str = str(tag_id) _enforce_snippet_tag_rbac_by_tag_id(tag_id_str) - TagService.delete_tag(tag_id_str, db.session) + TagService.delete_tag(tag_id_str, db.session()) return "", 204 @@ -223,7 +223,7 @@ def _create_tag_bindings(current_user: Account) -> tuple[dict[str, str], int]: target_id=payload.target_id, type=payload.type, ), - db.session, + db.session(), ) return {"result": "success"}, 200 @@ -239,7 +239,7 @@ def _remove_tag_bindings(current_user: Account) -> tuple[dict[str, str], int]: target_id=payload.target_id, type=payload.type, ), - db.session, + db.session(), ) return {"result": "success"}, 200 diff --git a/api/controllers/console/workspace/account.py b/api/controllers/console/workspace/account.py index 2f4ef1c5b42..6a06ed2d3e6 100644 --- a/api/controllers/console/workspace/account.py +++ b/api/controllers/console/workspace/account.py @@ -317,7 +317,7 @@ class AccountNameApi(Resource): def post(self, current_user: Account): payload = console_ns.payload or {} args = AccountNamePayload.model_validate(payload) - updated_account = AccountService.update_account(current_user, session=db.session, name=args.name) + updated_account = AccountService.update_account(current_user, session=db.session(), name=args.name) return dump_response(AccountResponse, updated_account) @@ -363,7 +363,7 @@ class AccountAvatarApi(Resource): payload = console_ns.payload or {} args = AccountAvatarPayload.model_validate(payload) - updated_account = AccountService.update_account(current_user, session=db.session, avatar=args.avatar) + updated_account = AccountService.update_account(current_user, session=db.session(), avatar=args.avatar) return dump_response(AccountResponse, updated_account) @@ -381,7 +381,7 @@ class AccountInterfaceLanguageApi(Resource): args = AccountInterfaceLanguagePayload.model_validate(payload) updated_account = AccountService.update_account( - current_user, session=db.session, interface_language=args.interface_language + current_user, session=db.session(), interface_language=args.interface_language ) return dump_response(AccountResponse, updated_account) @@ -400,7 +400,7 @@ class AccountInterfaceThemeApi(Resource): args = AccountInterfaceThemePayload.model_validate(payload) updated_account = AccountService.update_account( - current_user, session=db.session, interface_theme=args.interface_theme + current_user, session=db.session(), interface_theme=args.interface_theme ) return dump_response(AccountResponse, updated_account) @@ -418,7 +418,7 @@ class AccountTimezoneApi(Resource): payload = console_ns.payload or {} args = AccountTimezonePayload.model_validate(payload) - updated_account = AccountService.update_account(current_user, session=db.session, timezone=args.timezone) + updated_account = AccountService.update_account(current_user, session=db.session(), timezone=args.timezone) return dump_response(AccountResponse, updated_account) @@ -437,7 +437,7 @@ class AccountPasswordApi(Resource): try: assert args.password is not None - AccountService.update_account_password(current_user, args.password, args.new_password, session=db.session) + AccountService.update_account_password(current_user, args.password, args.new_password, session=db.session()) except ServiceCurrentPasswordIncorrectError: raise CurrentPasswordIncorrectError() @@ -514,7 +514,7 @@ class AccountDeleteApi(Resource): if not AccountService.verify_account_deletion_code(args.token, args.code): raise InvalidAccountDeletionCodeError() - AccountService.delete_account(account) + AccountService.delete_account(account, session=db.session()) return SimpleResultResponse(result="success").model_dump(mode="json") @@ -726,7 +726,7 @@ class ChangeEmailResetApi(Resource): if AccountService.is_account_in_freeze(normalized_new_email): raise AccountInFreezeError() - if not AccountService.check_email_unique(normalized_new_email, session=db.session): + if not AccountService.check_email_unique(normalized_new_email, session=db.session()): raise EmailAlreadyInUseError() reset_data = AccountService.get_change_email_data(args.token) @@ -751,7 +751,7 @@ class ChangeEmailResetApi(Resource): AccountService.revoke_change_email_token(args.token) updated_account = AccountService.update_account_email( - current_user, email=normalized_new_email, session=db.session + current_user, email=normalized_new_email, session=db.session() ) AccountService.send_change_email_completed_notify_email( @@ -772,6 +772,6 @@ class CheckEmailUnique(Resource): normalized_email = args.email.lower() if AccountService.is_account_in_freeze(normalized_email): raise AccountInFreezeError() - if not AccountService.check_email_unique(normalized_email, session=db.session): + if not AccountService.check_email_unique(normalized_email, session=db.session()): raise EmailAlreadyInUseError() return SimpleResultResponse(result="success").model_dump(mode="json") diff --git a/api/controllers/console/workspace/load_balancing_config.py b/api/controllers/console/workspace/load_balancing_config.py index 5983a4e10be..abeb691be03 100644 --- a/api/controllers/console/workspace/load_balancing_config.py +++ b/api/controllers/console/workspace/load_balancing_config.py @@ -10,6 +10,7 @@ from controllers.console.wraps import ( with_current_tenant_id, with_current_user, ) +from extensions.ext_database import db from fields.base import ResponseModel from graphon.model_runtime.entities.model_entities import ModelType from graphon.model_runtime.errors.validate import CredentialsValidateFailedError @@ -69,6 +70,7 @@ class LoadBalancingCredentialsValidateApi(Resource): model=payload.model, model_type=payload.model_type, credentials=payload.credentials, + session=db.session(), ) except CredentialsValidateFailedError as ex: result = False @@ -118,6 +120,7 @@ class LoadBalancingConfigCredentialsValidateApi(Resource): model=payload.model, model_type=payload.model_type, credentials=payload.credentials, + session=db.session(), config_id=config_id, ) except CredentialsValidateFailedError as ex: diff --git a/api/controllers/console/workspace/members.py b/api/controllers/console/workspace/members.py index 7e44d511bcf..72330aba5f0 100644 --- a/api/controllers/console/workspace/members.py +++ b/api/controllers/console/workspace/members.py @@ -135,7 +135,7 @@ def _normalize_enum_value(value: object) -> str: def _count_new_member_invites(tenant_id: str, emails: list[str]) -> int: new_member_count = 0 for email in emails: - account = AccountService.get_account_by_email_with_case_fallback(db.session, email) + account = AccountService.get_account_by_email_with_case_fallback(email, session=db.session()) if not account: new_member_count += 1 continue @@ -190,7 +190,7 @@ class MemberListApi(Resource): current_user, _ = current_account_with_tenant() if not current_user.current_tenant: raise ValueError("No current tenant") - members = TenantService.get_tenant_members(current_user.current_tenant, session=db.session) + members = TenantService.get_tenant_members(current_user.current_tenant, session=db.session()) if dify_config.RBAC_ENABLED: member_ids = [member.id for member in members] member_roles = enterprise_rbac_service.RBACService.MemberRoles.batch_get( @@ -275,7 +275,7 @@ class MemberInviteEmailApi(Resource): language=interface_language, role=invitee_role, inviter=inviter, - session=db.session, + session=db.session(), ) encoded_invitee_email = parse.quote(invitee_email) invitation_results.append( @@ -323,7 +323,7 @@ class MemberCancelInviteApi(Resource): else: try: TenantService.remove_member_from_tenant( - current_user.current_tenant, member, current_user, session=db.session + current_user.current_tenant, member, current_user, session=db.session() ) except services.errors.account.CannotOperateSelfError as e: return {"code": "cannot-operate-self", "message": str(e)}, HTTPStatus.BAD_REQUEST @@ -368,7 +368,7 @@ class MemberUpdateRoleApi(Resource): try: assert member is not None, "Member not found" TenantService.update_member_role( - current_user.current_tenant, member, new_role, current_user, session=db.session + current_user.current_tenant, member, new_role, current_user, session=db.session() ) except services.errors.account.CannotOperateSelfError as e: return {"code": "cannot-operate-self", "message": str(e)}, HTTPStatus.BAD_REQUEST @@ -396,7 +396,7 @@ class DatasetOperatorMemberListApi(Resource): def get(self, current_user: Account): if not current_user.current_tenant: raise ValueError("No current tenant") - members = TenantService.get_dataset_operator_members(current_user.current_tenant, session=db.session) + members = TenantService.get_dataset_operator_members(current_user.current_tenant, session=db.session()) return dump_response(AccountWithRoleListResponse, {"accounts": members}), HTTPStatus.OK @@ -420,7 +420,7 @@ class SendOwnerTransferEmailApi(Resource): # check if the current user is the owner of the workspace if not current_user.current_tenant: raise ValueError("No current tenant") - if not TenantService.is_owner(current_user, current_user.current_tenant, session=db.session): + if not TenantService.is_owner(current_user, current_user.current_tenant, session=db.session()): raise NotOwnerError() if args.language is not None and args.language == "zh-Hans": @@ -455,7 +455,7 @@ class OwnerTransferCheckApi(Resource): # check if the current user is the owner of the workspace if not current_user.current_tenant: raise ValueError("No current tenant") - if not TenantService.is_owner(current_user, current_user.current_tenant, session=db.session): + if not TenantService.is_owner(current_user, current_user.current_tenant, session=db.session()): raise NotOwnerError() user_email = current_user.email @@ -501,7 +501,7 @@ class OwnerTransfer(Resource): # check if the current user is the owner of the workspace if not current_user.current_tenant: raise ValueError("No current tenant") - if not TenantService.is_owner(current_user, current_user.current_tenant, session=db.session): + if not TenantService.is_owner(current_user, current_user.current_tenant, session=db.session()): raise NotOwnerError() if current_user.id == str(member_id): @@ -522,13 +522,13 @@ class OwnerTransfer(Resource): if not current_user.current_tenant: raise ValueError("No current tenant") - if not TenantService.is_member(member, current_user.current_tenant, session=db.session): + if not TenantService.is_member(member, current_user.current_tenant, session=db.session()): raise MemberNotInTenantError() try: assert member is not None, "Member not found" TenantService.update_member_role( - current_user.current_tenant, member, "owner", current_user, session=db.session + current_user.current_tenant, member, "owner", current_user, session=db.session() ) AccountService.send_new_owner_transfer_notify_email( diff --git a/api/controllers/console/workspace/model_providers.py b/api/controllers/console/workspace/model_providers.py index 779399f9055..3bafa2ab6a6 100644 --- a/api/controllers/console/workspace/model_providers.py +++ b/api/controllers/console/workspace/model_providers.py @@ -353,7 +353,7 @@ class ModelProviderPaymentCheckoutUrlApi(Resource): def get(self, current_tenant_id: str, current_user: Account, provider: str): if provider != "anthropic": raise ValueError(f"provider name {provider} is invalid") - BillingService.is_tenant_owner_or_admin(db.session, current_user) + BillingService.is_tenant_owner_or_admin(current_user, session=db.session()) data = BillingService.get_model_provider_payment_link( provider_name=provider, tenant_id=current_tenant_id, diff --git a/api/controllers/console/workspace/models.py b/api/controllers/console/workspace/models.py index 1da72ef4362..0f735a7479e 100644 --- a/api/controllers/console/workspace/models.py +++ b/api/controllers/console/workspace/models.py @@ -24,6 +24,7 @@ from controllers.console.wraps import ( with_current_user, ) from core.entities.provider_entities import CredentialConfiguration +from extensions.ext_database import db from fields.base import ResponseModel from graphon.model_runtime.entities.model_entities import ModelType, ParameterRule from graphon.model_runtime.errors.validate import CredentialsValidateFailedError @@ -297,6 +298,7 @@ class ModelProviderModelApi(Resource): model_type=args.model_type, configs=args.load_balancing.configs, config_from=args.config_from or "", + session=db.session(), ) if args.load_balancing.enabled: @@ -356,6 +358,7 @@ class ModelProviderModelCredentialApi(Resource): provider=provider, model=args.model, model_type=args.model_type, + session=db.session(), config_from=args.config_from or "", ) diff --git a/api/controllers/console/workspace/plugin.py b/api/controllers/console/workspace/plugin.py index 682aa5b6190..c7644af2b48 100644 --- a/api/controllers/console/workspace/plugin.py +++ b/api/controllers/console/workspace/plugin.py @@ -38,6 +38,7 @@ from core.tools.builtin_tool.providers._positions import BuiltinToolProviderSort from core.tools.entities.common_entities import I18nObject from core.tools.entities.tool_entities import ToolProviderType from core.tools.tool_manager import ToolManager +from extensions.ext_database import db from fields.base import ResponseModel from graphon.model_runtime.utils.encoders import jsonable_encoder from libs.helper import dump_response @@ -973,7 +974,7 @@ class PluginChangePermissionApi(Resource): args = ParserPermissionChange.model_validate(console_ns.payload) set_permission_result = PluginPermissionService.change_permission( - tenant_id, args.install_permission, args.debug_permission + tenant_id, args.install_permission, args.debug_permission, session=db.session() ) if not set_permission_result: return jsonable_encoder({"success": False, "message": "Failed to set permission"}) @@ -989,7 +990,7 @@ class PluginFetchPermissionApi(Resource): @account_initialization_required @with_current_tenant_id def get(self, tenant_id: str): - permission = PluginPermissionService.get_permission(tenant_id) + permission = PluginPermissionService.get_permission(tenant_id, session=db.session()) if not permission: return jsonable_encoder( { @@ -1094,6 +1095,7 @@ class PluginChangeAutoUpgradeApi(Resource): auto_upgrade.exclude_plugins, auto_upgrade.include_plugins, category=args.category, + session=db.session(), ) if not set_auto_upgrade_strategy_result: return jsonable_encoder({"success": False, "message": "Failed to set auto upgrade strategy"}) @@ -1111,7 +1113,7 @@ class PluginFetchAutoUpgradeApi(Resource): @with_current_tenant_id def get(self, tenant_id: str): args = ParserAutoUpgradeFetch.model_validate(request.args.to_dict(flat=True)) - auto_upgrade = PluginAutoUpgradeService.get_strategy(tenant_id, args.category) + auto_upgrade = PluginAutoUpgradeService.get_strategy(tenant_id, args.category, session=db.session()) auto_upgrade_dict = ( _auto_upgrade_settings_to_dict(auto_upgrade) if auto_upgrade @@ -1140,7 +1142,11 @@ class PluginAutoUpgradeExcludePluginApi(Resource): args = ParserExcludePlugin.model_validate(console_ns.payload) return jsonable_encoder( - {"success": PluginAutoUpgradeService.exclude_plugin(tenant_id, args.plugin_id, args.category)} + { + "success": PluginAutoUpgradeService.exclude_plugin( + tenant_id, args.plugin_id, args.category, session=db.session() + ) + } ) diff --git a/api/controllers/console/workspace/rbac.py b/api/controllers/console/workspace/rbac.py index a155bb0cf0b..39ad12080f3 100644 --- a/api/controllers/console/workspace/rbac.py +++ b/api/controllers/console/workspace/rbac.py @@ -14,6 +14,7 @@ from controllers.console import console_ns from controllers.console.wraps import RBACPermission, RBACResourceScope, rbac_permission_required from core.db.session_factory import session_factory from core.rbac import RBACResourceWhitelistScope +from extensions.ext_database import db from libs.login import current_account_with_tenant, login_required from models import Account from services.enterprise import rbac_service as svc @@ -564,6 +565,7 @@ class RBACMyPermissionsApi(Resource): account_id, app_id=request.args.get("app_id") or None, dataset_id=request.args.get("dataset_id") or None, + session=db.session(), ) ) @@ -902,7 +904,7 @@ class RBACMemberRolesApi(Resource): @console_ns.response(200, "Success", console_ns.models[svc.MemberRolesResponse.__name__]) def get(self, member_id): tenant_id, account_id = _current_ids() - return _dump(svc.RBACService.MemberRoles.get(tenant_id, account_id, str(member_id))) + return _dump(svc.RBACService.MemberRoles.get(tenant_id, account_id, str(member_id), session=db.session())) @login_required @console_ns.expect(console_ns.models[_ReplaceMemberRolesRequest.__name__]) @@ -916,6 +918,7 @@ class RBACMemberRolesApi(Resource): account_id, str(member_id), role_ids=list(request.role_ids), + session=db.session(), ) ) diff --git a/api/controllers/console/workspace/snippets.py b/api/controllers/console/workspace/snippets.py index c849336401c..18407b62d52 100644 --- a/api/controllers/console/workspace/snippets.py +++ b/api/controllers/console/workspace/snippets.py @@ -188,7 +188,7 @@ class CustomizedSnippetsApi(Resource): snippet_service = _snippet_service() snippets, total, has_more = snippet_service.get_snippets( tenant_id=current_tenant_id, - session=db.session, + session=db.session(), page=query.page, limit=query.limit, keyword=query.keyword, diff --git a/api/controllers/console/workspace/tool_providers.py b/api/controllers/console/workspace/tool_providers.py index 7a3f158b0c3..30eeec2fcc5 100644 --- a/api/controllers/console/workspace/tool_providers.py +++ b/api/controllers/console/workspace/tool_providers.py @@ -459,6 +459,7 @@ class ToolBuiltinProviderGetCredentialsApi(Resource): BuiltinToolManageService.get_builtin_tool_provider_credentials( tenant_id=tenant_id, provider_name=provider, + session=db.session(), user=user, include_credential_ids=query.include_credential_ids or None, ) @@ -1064,6 +1065,7 @@ class ToolBuiltinProviderGetCredentialInfoApi(Resource): BuiltinToolManageService.get_builtin_tool_provider_credential_info( tenant_id=tenant_id, provider=provider, + session=db.session(), user=user, include_credential_ids=query.include_credential_ids or None, ) diff --git a/api/controllers/console/workspace/workspace.py b/api/controllers/console/workspace/workspace.py index 0630281de75..23ce116b349 100644 --- a/api/controllers/console/workspace/workspace.py +++ b/api/controllers/console/workspace/workspace.py @@ -223,7 +223,7 @@ class TenantListApi(Resource): def get(self, current_tenant_id: str, current_user: Account): tenant_rows: list[tuple[Tenant, TenantAccountJoin]] = [ (tenant, membership) - for tenant, membership in TenantService.get_workspaces_for_account(db.session, current_user.id) + for tenant, membership in TenantService.get_workspaces_for_account(current_user.id, session=db.session()) if tenant.status == TenantStatus.NORMAL ] tenants = [tenant for tenant, _ in tenant_rows] @@ -306,16 +306,19 @@ class TenantApi(Resource): raise ValueError("No current tenant") if tenant.status == TenantStatus.ARCHIVE: - tenants = TenantService.get_join_tenants(current_user, session=db.session) + tenants = TenantService.get_join_tenants(current_user, session=db.session()) # if there is any tenant, switch to the first one if len(tenants) > 0: - TenantService.switch_tenant(current_user, tenants[0].id, session=db.session) + TenantService.switch_tenant(current_user, tenants[0].id, session=db.session()) tenant = tenants[0] # else, raise Unauthorized else: raise Unauthorized("workspace is archived") - return dump_response(TenantInfoResponse, WorkspaceService.get_tenant_info(tenant)), HTTPStatus.OK + return ( + dump_response(TenantInfoResponse, WorkspaceService.get_tenant_info(tenant, session=db.session())), + HTTPStatus.OK, + ) @console_ns.route("/workspaces/switch") @@ -332,7 +335,7 @@ class SwitchWorkspaceApi(Resource): # Check whether the tenant_id belongs to the current account. try: - TenantService.switch_tenant(current_user, args.tenant_id, session=db.session) + TenantService.switch_tenant(current_user, args.tenant_id, session=db.session()) except Exception: raise AccountNotLinkTenantError("Account not link tenant") @@ -341,7 +344,7 @@ class SwitchWorkspaceApi(Resource): raise ValueError("Tenant not found") return SwitchWorkspaceResponse( - result="success", new_tenant=WorkspaceService.get_tenant_info(new_tenant) + result="success", new_tenant=WorkspaceService.get_tenant_info(new_tenant, session=db.session()) ).model_dump(mode="json") @@ -372,7 +375,7 @@ class CustomConfigWorkspaceApi(Resource): db.session.commit() return WorkspaceTenantResultResponse( - result="success", tenant=WorkspaceService.get_tenant_info(tenant) + result="success", tenant=WorkspaceService.get_tenant_info(tenant, session=db.session()) ).model_dump(mode="json") @@ -438,7 +441,7 @@ class WorkspaceInfoApi(Resource): db.session.commit() return WorkspaceTenantResultResponse( - result="success", tenant=WorkspaceService.get_tenant_info(tenant) + result="success", tenant=WorkspaceService.get_tenant_info(tenant, session=db.session()) ).model_dump(mode="json") diff --git a/api/controllers/files/agent_drive_archive.py b/api/controllers/files/agent_drive_archive.py index afa6ac79483..8ecec2e9a4c 100644 --- a/api/controllers/files/agent_drive_archive.py +++ b/api/controllers/files/agent_drive_archive.py @@ -8,6 +8,7 @@ from werkzeug.exceptions import Forbidden, NotFound from controllers.common.file_response import enforce_download_for_html from controllers.common.schema import register_schema_models from controllers.files import files_ns +from extensions.ext_database import db from models.agent import AgentDriveFileKind from services.agent_drive_service import AgentDriveError, AgentDriveService @@ -54,6 +55,7 @@ class AgentDriveArchiveMemberApi(Resource): archive_file_kind=args.archive_file_kind, archive_file_id=args.archive_file_id, member_path=args.member_path, + session=db.session(), ) except AgentDriveError as exc: raise NotFound(exc.message) from exc diff --git a/api/controllers/inner_api/app/dsl.py b/api/controllers/inner_api/app/dsl.py index 915a11dcddc..9fd111f86dc 100644 --- a/api/controllers/inner_api/app/dsl.py +++ b/api/controllers/inner_api/app/dsl.py @@ -98,6 +98,7 @@ class EnterpriseAppDSLExport(Resource): data = AppDslService.export_dsl( app_model=app_model, + session=db.session(), include_secret=include_secret, ) diff --git a/api/controllers/inner_api/plugin/agent_drive.py b/api/controllers/inner_api/plugin/agent_drive.py index 0cdb9dab35f..e06720a8e99 100644 --- a/api/controllers/inner_api/plugin/agent_drive.py +++ b/api/controllers/inner_api/plugin/agent_drive.py @@ -17,6 +17,7 @@ from controllers.console.wraps import setup_required from controllers.inner_api import inner_api_ns from controllers.inner_api.plugin.wraps import get_user from controllers.inner_api.wraps import plugin_inner_api_only +from extensions.ext_database import db from services.agent_drive_service import ( AgentDriveError, AgentDriveService, @@ -53,6 +54,7 @@ class AgentDriveManifestApi(Resource): agent_id=agent_id, prefix=request.args.get("prefix", ""), include_download_url=include_download_url, + session=db.session(), ) except AgentDriveError as exc: return _error_response(exc) @@ -71,7 +73,7 @@ class AgentDriveSkillsApi(Resource): tenant_id = (request.args.get("tenant_id") or "").strip() if not tenant_id: raise AgentDriveError("missing_tenant_id", "tenant_id is required", status_code=400) - items = AgentDriveService().list_skills(tenant_id=tenant_id, agent_id=agent_id) + items = AgentDriveService().list_skills(tenant_id=tenant_id, agent_id=agent_id, session=db.session()) except AgentDriveError as exc: return _error_response(exc) return {"items": items} @@ -96,6 +98,7 @@ class AgentDriveCommitApi(Resource): user_id=user.id, agent_id=agent_id, items=body.items, + session=db.session(), ) except AgentDriveError as exc: return _error_response(exc) diff --git a/api/controllers/inner_api/workspace/workspace.py b/api/controllers/inner_api/workspace/workspace.py index 1f25eb576d3..b3a571112f6 100644 --- a/api/controllers/inner_api/workspace/workspace.py +++ b/api/controllers/inner_api/workspace/workspace.py @@ -47,8 +47,8 @@ class EnterpriseWorkspace(Resource): if account is None: return {"message": "owner account not found."}, 404 - tenant = TenantService.create_tenant(args.name, is_from_dashboard=True, session=db.session) - TenantService.create_tenant_member(tenant, account, db.session, role="owner") + tenant = TenantService.create_tenant(args.name, is_from_dashboard=True, session=db.session()) + TenantService.create_tenant_member(tenant, account, db.session(), role="owner") tenant_was_created.send(tenant) @@ -84,7 +84,7 @@ class EnterpriseWorkspaceNoOwnerEmail(Resource): def post(self): args = WorkspaceOwnerlessPayload.model_validate(inner_api_ns.payload or {}) - tenant = TenantService.create_tenant(args.name, is_from_dashboard=True, session=db.session) + tenant = TenantService.create_tenant(args.name, is_from_dashboard=True, session=db.session()) tenant_was_created.send(tenant) diff --git a/api/controllers/openapi/account.py b/api/controllers/openapi/account.py index 8ad0b02f4a0..b4786f2ae25 100644 --- a/api/controllers/openapi/account.py +++ b/api/controllers/openapi/account.py @@ -45,8 +45,10 @@ class AccountApi(Resource): enforce(LIMIT_ME_PER_ACCOUNT, key=f"account:{auth_data.account_id}") account_id_str = str(auth_data.account_id) if auth_data.account_id else None - account = AccountService.get_account_by_id(db.session, account_id_str) if account_id_str else None - memberships = TenantService.get_account_memberships(db.session, account_id_str) if account_id_str else [] + account = AccountService.get_account_by_id(account_id_str, session=db.session()) if account_id_str else None + memberships = ( + TenantService.get_account_memberships(account_id_str, session=db.session()) if account_id_str else [] + ) default_ws_id = _pick_default_workspace(memberships) return AccountResponse( @@ -63,7 +65,7 @@ class AccountSessionsSelfApi(Resource): @auth_router.guard(scope=Scope.FULL, allowed_token_types=frozenset({TokenType.OAUTH_ACCOUNT})) @returns(200, RevokeResponse, description="Session revoked") def delete(self, *, auth_data: AuthData): - revoke_oauth_token(db.session, redis_client, str(auth_data.token_id)) + revoke_oauth_token(redis_client, str(auth_data.token_id), session=db.session()) return RevokeResponse(status="revoked") @@ -81,7 +83,7 @@ class AccountSessionsApi(Resource): page = query.page limit = query.limit - all_rows = list_active_sessions(db.session, ctx, now) + all_rows = list_active_sessions(ctx, now, session=db.session()) total = len(all_rows) sliced = all_rows[(page - 1) * limit : page * limit] @@ -117,10 +119,10 @@ class AccountSessionByIdApi(Resource): # 404 (not 403) on cross-subject so the endpoint doesn't leak # token IDs that belong to other subjects. - if not token_belongs_to_subject(db.session, session_id, ctx): + if not token_belongs_to_subject(session_id, ctx, session=db.session()): raise NotFound("session not found") - revoke_oauth_token(db.session, redis_client, session_id) + revoke_oauth_token(redis_client, session_id, session=db.session()) return RevokeResponse(status="revoked") diff --git a/api/controllers/openapi/app_dsl.py b/api/controllers/openapi/app_dsl.py index cea7127bd07..d06845dada4 100644 --- a/api/controllers/openapi/app_dsl.py +++ b/api/controllers/openapi/app_dsl.py @@ -145,6 +145,7 @@ class AppDslExportApi(Resource): try: data = AppDslService.export_dsl( app_model=app, + session=db.session(), include_secret=query.include_secret, workflow_id=query.workflow_id, ) diff --git a/api/controllers/openapi/apps.py b/api/controllers/openapi/apps.py index d4cb175ba5e..882b55b7041 100644 --- a/api/controllers/openapi/apps.py +++ b/api/controllers/openapi/apps.py @@ -66,13 +66,13 @@ class AppReadResource(Resource): if is_uuid: # ``str(parsed_uuid)`` normalises to the canonical dashed form. - app = AppService.get_visible_app_by_id(db.session, str(parsed_uuid)) + app = AppService.get_visible_app_by_id(str(parsed_uuid), session=db.session()) if app is None: raise NotFound("app not found") else: if not workspace_id: raise UnprocessableEntity("workspace_id is required for name-based lookup") - matches = AppService.find_visible_apps_by_name(db.session, name=app_id, tenant_id=workspace_id) + matches = AppService.find_visible_apps_by_name(name=app_id, tenant_id=workspace_id, session=db.session()) if len(matches) == 0: raise NotFound("app not found") if len(matches) > 1: @@ -177,7 +177,7 @@ class AppListApi(Resource): tenant_name: str | None = None if parsed_uuid is not None: - app: App | None = AppService.get_visible_app_by_id(db.session, str(parsed_uuid)) + app: App | None = AppService.get_visible_app_by_id(str(parsed_uuid), session=db.session()) if app is None or str(app.tenant_id) != workspace_id: return empty if not _is_listable(app): @@ -188,7 +188,7 @@ class AppListApi(Resource): str(app.id), str(app.maintainer) if app.maintainer else None, str(auth_data.account_id) ): return empty - tenant_name = TenantService.get_tenant_name(db.session, workspace_id) + tenant_name = TenantService.get_tenant_name(workspace_id, session=db.session()) item = AppListRow( id=str(app.id), name=app.name, @@ -215,13 +215,13 @@ class AppListApi(Resource): if apply_rbac_filter: access_filter.apply_to_params(params) - pagination = AppService().get_paginate_apps(str(auth_data.account_id), workspace_id, params, db.session) + pagination = AppService().get_paginate_apps(str(auth_data.account_id), workspace_id, params, db.session()) if pagination is None: return empty tenant_name = None if pagination.items: - tenant_name = TenantService.get_tenant_name(db.session, workspace_id) + tenant_name = TenantService.get_tenant_name(workspace_id, session=db.session()) items = [ AppListRow( diff --git a/api/controllers/openapi/apps_permitted_external.py b/api/controllers/openapi/apps_permitted_external.py index 5c6fdce5141..353a1ec1cb3 100644 --- a/api/controllers/openapi/apps_permitted_external.py +++ b/api/controllers/openapi/apps_permitted_external.py @@ -55,10 +55,10 @@ class PermittedExternalAppsListApi(Resource): return env apps_by_id: dict[str, App] = { - str(a.id): a for a in AppService.find_visible_apps_by_ids(db.session, page_result.app_ids) + str(a.id): a for a in AppService.find_visible_apps_by_ids(page_result.app_ids, session=db.session()) } tenant_ids = list({str(a.tenant_id) for a in apps_by_id.values()}) - tenants_by_id = {str(t.id): t for t in TenantService.get_tenants_by_ids(db.session, tenant_ids)} + tenants_by_id = {str(t.id): t for t in TenantService.get_tenants_by_ids(tenant_ids, session=db.session())} items: list[AppListRow] = [] for app_id in page_result.app_ids: diff --git a/api/controllers/openapi/auth/prepare.py b/api/controllers/openapi/auth/prepare.py index 6704b27decc..96cf9a8858f 100644 --- a/api/controllers/openapi/auth/prepare.py +++ b/api/controllers/openapi/auth/prepare.py @@ -23,7 +23,7 @@ def load_app(data: AuthData) -> None: uuid.UUID(app_id) except ValueError: raise NotFound("app not found") - app = AppService.get_app_by_id(db.session, app_id) + app = AppService.get_app_by_id(app_id, session=db.session()) if not app or app.status != AppStatus.NORMAL: raise NotFound("app not found") data.app = app @@ -34,7 +34,7 @@ def load_tenant(data: AuthData) -> None: return if data.app is None: raise InternalServerError("pipeline_invariant_violated: app not loaded before load_tenant") - tenant = TenantService.get_tenant_by_id(db.session, str(data.app.tenant_id)) + tenant = TenantService.get_tenant_by_id(str(data.app.tenant_id), session=db.session()) if tenant is None or tenant.status == TenantStatus.ARCHIVE: raise Forbidden("workspace unavailable") data.tenant = tenant @@ -50,7 +50,7 @@ def load_tenant_from_request(data: AuthData) -> None: uuid.UUID(workspace_id) except ValueError: raise NotFound("workspace not found") - tenant = TenantService.get_tenant_by_id(db.session, workspace_id) + tenant = TenantService.get_tenant_by_id(workspace_id, session=db.session()) if tenant is None or tenant.status == TenantStatus.ARCHIVE: raise NotFound("workspace not found") data.tenant = tenant @@ -59,7 +59,7 @@ def load_tenant_from_request(data: AuthData) -> None: def load_account(data: AuthData) -> None: if data.caller is not None: return - account = AccountService.get_account_by_id(db.session, str(data.account_id)) + account = AccountService.get_account_by_id(str(data.account_id), session=db.session()) if account is None: raise Unauthorized("account not found") if data.tenant: @@ -75,7 +75,7 @@ def load_workspace_role(data: AuthData) -> None: return if data.caller is not None and getattr(data.caller, "status", None) != AccountStatus.ACTIVE: return - role = TenantService.get_account_role_in_tenant(db.session, str(data.account_id), str(data.tenant.id)) + role = TenantService.get_account_role_in_tenant(str(data.account_id), str(data.tenant.id), session=db.session()) if role is None: return data.tenant_role = role diff --git a/api/controllers/openapi/auth/verify.py b/api/controllers/openapi/auth/verify.py index b5f10f66b34..b6ef95e3ea3 100644 --- a/api/controllers/openapi/auth/verify.py +++ b/api/controllers/openapi/auth/verify.py @@ -82,7 +82,7 @@ def check_app_api_enabled(data: AuthData) -> None: def check_app_access(data: AuthData) -> None: if data.tenant is None: return - if not TenantService.account_belongs_to_tenant(db.session, data.account_id, data.tenant.id): + if not TenantService.account_belongs_to_tenant(data.account_id, data.tenant.id, session=db.session()): raise Forbidden("subject_no_app_access") @@ -127,5 +127,5 @@ def _resolve_user_id(data: AuthData) -> str | None: return str(data.account_id) if data.account_id is not None else None if data.external_identity is None: return None - account = AccountService.get_account_by_email(db.session, data.external_identity.email) + account = AccountService.get_account_by_email(data.external_identity.email, session=db.session()) return str(account.id) if account is not None else None diff --git a/api/controllers/openapi/oauth_device.py b/api/controllers/openapi/oauth_device.py index cee187daaf3..3ba5f2ee207 100644 --- a/api/controllers/openapi/oauth_device.py +++ b/api/controllers/openapi/oauth_device.py @@ -247,7 +247,6 @@ class DeviceApproveApi(Resource): raise BadRequest(description=str(e)) from None ttl_days = oauth_ttl_days(tenant_id=tenant) mint = mint_oauth_token( - db.session, redis_client, subject_email=account.email, subject_issuer=ACCOUNT_ISSUER_SENTINEL, @@ -256,6 +255,7 @@ class DeviceApproveApi(Resource): device_label=state.device_label, prefix=profile.prefix, ttl_days=ttl_days, + session=db.session(), ) poll_payload = _build_account_poll_payload(account, tenant, mint) @@ -342,7 +342,7 @@ def _audit_cross_ip_if_needed(state) -> None: def _build_account_poll_payload(account, tenant, mint) -> PollPayload: - rows = TenantService.get_workspaces_for_account(db.session, str(account.id)) + rows = TenantService.get_workspaces_for_account(str(account.id), session=db.session()) workspaces = [WorkspacePayload(id=str(t.id), name=t.name, role=getattr(m, "role", "")) for t, m in rows] # Prefer active session tenant → DB-flagged current join → first membership. default_ws_id = None diff --git a/api/controllers/openapi/oauth_device_sso.py b/api/controllers/openapi/oauth_device_sso.py index 79538f48059..fbf7bfa6295 100644 --- a/api/controllers/openapi/oauth_device_sso.py +++ b/api/controllers/openapi/oauth_device_sso.py @@ -194,7 +194,7 @@ def _sso_complete_impl(): if state.status is not DeviceFlowStatus.PENDING: return _device_error_redirect("sso_failed", user_code) - if AccountService.has_active_account_with_email(db.session, claims.email): + if AccountService.has_active_account_with_email(claims.email, session=db.session()): _emit_external_rejection_audit( state, _RejectedClaims(subject_email=claims.email, subject_issuer=claims.issuer), @@ -274,7 +274,7 @@ def approve_external(): if state.status is not DeviceFlowStatus.PENDING: raise Conflict("user_code_not_pending") - if AccountService.has_active_account_with_email(db.session, claims.subject_email): + if AccountService.has_active_account_with_email(claims.subject_email, session=db.session()): _emit_external_rejection_audit(state, claims, reason="email_belongs_to_dify_account") raise Forbidden("email_belongs_to_dify_account") @@ -293,7 +293,6 @@ def approve_external(): ttl_days = oauth_ttl_days(tenant_id=None) mint = mint_oauth_token( - db.session, redis_client, subject_email=claims.subject_email, subject_issuer=claims.subject_issuer, @@ -302,6 +301,7 @@ def approve_external(): device_label=state.device_label, prefix=profile.prefix, ttl_days=ttl_days, + session=db.session(), ) # SSO branch of the shared PollPayload contract: account/workspace diff --git a/api/controllers/openapi/workspaces.py b/api/controllers/openapi/workspaces.py index c45c02e54d3..7f8eb0f7012 100644 --- a/api/controllers/openapi/workspaces.py +++ b/api/controllers/openapi/workspaces.py @@ -64,14 +64,14 @@ def _member_response(account: Account) -> MemberResponse: def _load_tenant(workspace_id: str) -> Tenant: - tenant = TenantService.get_tenant_by_id(db.session, workspace_id) + tenant = TenantService.get_tenant_by_id(workspace_id, session=db.session()) if tenant is None or tenant.status != TenantStatus.NORMAL: raise NotFound("workspace not found") return tenant def _load_account(account_id: object) -> Account: - account = AccountService.get_account_by_id(db.session, str(account_id)) if account_id else None + account = AccountService.get_account_by_id(str(account_id), session=db.session()) if account_id else None if account is None: raise RuntimeError("authenticated account_id has no Account row") return account @@ -94,7 +94,7 @@ class WorkspacesApi(Resource): @auth_router.guard(scope=Scope.WORKSPACE_READ, allowed_token_types=frozenset({TokenType.OAUTH_ACCOUNT})) @returns(200, WorkspaceListResponse, description="Workspace list") def get(self, *, auth_data: AuthData): - rows = TenantService.get_workspaces_for_account(db.session, str(auth_data.account_id)) + rows = TenantService.get_workspaces_for_account(str(auth_data.account_id), session=db.session()) return WorkspaceListResponse(workspaces=list(starmap(_workspace_summary, rows))) @@ -104,7 +104,7 @@ class WorkspaceByIdApi(Resource): @auth_router.guard(scope=Scope.WORKSPACE_READ, allowed_token_types=frozenset({TokenType.OAUTH_ACCOUNT})) @returns(200, WorkspaceDetailResponse, description="Workspace detail") def get(self, workspace_id: str, *, auth_data: AuthData): - row = TenantService.find_workspace_for_account(db.session, str(auth_data.account_id), workspace_id) + row = TenantService.find_workspace_for_account(str(auth_data.account_id), workspace_id, session=db.session()) # 404 (not 403) on non-member so workspace IDs don't leak across tenants. if row is None: raise NotFound("workspace not found") @@ -128,11 +128,11 @@ class WorkspaceSwitchApi(Resource): account = _load_account(auth_data.account_id) try: - TenantService.switch_tenant(account, workspace_id, session=db.session) + TenantService.switch_tenant(account, workspace_id, session=db.session()) except AccountNotLinkTenantError: raise NotFound("workspace not found") - row = TenantService.find_workspace_for_account(db.session, str(auth_data.account_id), workspace_id) + row = TenantService.find_workspace_for_account(str(auth_data.account_id), workspace_id, session=db.session()) if row is None: raise NotFound("workspace not found") tenant, membership = row @@ -152,7 +152,7 @@ class WorkspaceMembersApi(Resource): @accepts(query=MemberListQuery) def get(self, workspace_id: str, *, auth_data: AuthData, query: MemberListQuery): tenant = _load_tenant(workspace_id) - members = TenantService.get_tenant_members(tenant, session=db.session) + members = TenantService.get_tenant_members(tenant, session=db.session()) total = len(members) start = (query.page - 1) * query.limit page_items = members[start : start + query.limit] @@ -184,7 +184,7 @@ class WorkspaceMembersApi(Resource): language=None, role=body.role, inviter=inviter, - session=db.session, + session=db.session(), ) except AccountAlreadyInTenantError as exc: raise BadRequest(str(exc)) @@ -194,7 +194,7 @@ class WorkspaceMembersApi(Resource): raise BadRequest(str(exc)) normalized_email = body.email.lower() - member = AccountService.get_account_by_email_with_case_fallback(db.session, normalized_email) + member = AccountService.get_account_by_email_with_case_fallback(normalized_email, session=db.session()) if member is None: # invite_new_member just created or fetched this account. raise RuntimeError("invited member missing from DB after invite") @@ -229,12 +229,12 @@ class WorkspaceMemberApi(Resource): def delete(self, workspace_id: str, member_id: str, *, auth_data: AuthData): operator = _load_account(auth_data.account_id) tenant = _load_tenant(workspace_id) - member = AccountService.get_account_by_id(db.session, member_id) + member = AccountService.get_account_by_id(member_id, session=db.session()) if member is None: raise NotFound("member not found") try: - TenantService.remove_member_from_tenant(tenant, member, operator, session=db.session) + TenantService.remove_member_from_tenant(tenant, member, operator, session=db.session()) except CannotOperateSelfError as exc: raise BadRequest(str(exc)) except NoPermissionError as exc: @@ -254,12 +254,12 @@ class WorkspaceMemberApi(Resource): def patch(self, workspace_id: str, member_id: str, *, auth_data: AuthData, body: MemberRoleUpdatePayload): operator = _load_account(auth_data.account_id) tenant = _load_tenant(workspace_id) - member = AccountService.get_account_by_id(db.session, member_id) + member = AccountService.get_account_by_id(member_id, session=db.session()) if member is None: raise NotFound("member not found") try: - TenantService.update_member_role(tenant, member, body.role, operator, session=db.session) + TenantService.update_member_role(tenant, member, body.role, operator, session=db.session()) except CannotOperateSelfError as exc: raise BadRequest(str(exc)) except NoPermissionError as exc: diff --git a/api/controllers/service_api/app/annotation.py b/api/controllers/service_api/app/annotation.py index 0fbf8125ed9..126c67b5d61 100644 --- a/api/controllers/service_api/app/annotation.py +++ b/api/controllers/service_api/app/annotation.py @@ -201,7 +201,7 @@ class AnnotationListApi(Resource): query = AnnotationListQuery.model_validate(request.args.to_dict(flat=True)) annotation_list, total = AppAnnotationService.get_annotation_list_by_app_id( - app_model.id, query.page, query.limit, query.keyword + app_model.id, query.page, query.limit, query.keyword, session=db.session() ) annotation_models = TypeAdapter(list[Annotation]).validate_python(annotation_list, from_attributes=True) response = AnnotationList( @@ -243,7 +243,9 @@ class AnnotationListApi(Resource): """Create a new annotation.""" payload = AnnotationCreatePayload.model_validate(service_api_ns.payload or {}) insert_args: InsertAnnotationArgs = {"question": payload.question, "answer": payload.answer} - annotation = AppAnnotationService.insert_app_annotation_directly(insert_args, app_model.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + insert_args, app_model.id, session=db.session() + ) response = Annotation.model_validate(annotation, from_attributes=True) return response.model_dump(mode="json"), HTTPStatus.CREATED @@ -285,7 +287,7 @@ class AnnotationUpdateDeleteApi(Resource): update_args: UpdateAnnotationArgs = {"question": payload.question, "answer": payload.answer} app_ref = AppRefService.create_app_ref(app_model) annotation_ref = AppRefService.create_annotation_ref(app_ref, str(annotation_id)) - annotation = AppAnnotationService.update_app_annotation_directly(update_args, annotation_ref, db.session) + annotation = AppAnnotationService.update_app_annotation_directly(update_args, annotation_ref, db.session()) response = Annotation.model_validate(annotation, from_attributes=True) return response.model_dump(mode="json") @@ -316,5 +318,5 @@ class AnnotationUpdateDeleteApi(Resource): """Delete an annotation.""" app_ref = AppRefService.create_app_ref(app_model) annotation_ref = AppRefService.create_annotation_ref(app_ref, str(annotation_id)) - AppAnnotationService.delete_app_annotation(annotation_ref, db.session) + AppAnnotationService.delete_app_annotation(annotation_ref, db.session()) return "", 204 diff --git a/api/controllers/service_api/app/app.py b/api/controllers/service_api/app/app.py index 3ac44b12c66..60f83d7d070 100644 --- a/api/controllers/service_api/app/app.py +++ b/api/controllers/service_api/app/app.py @@ -11,6 +11,7 @@ from controllers.service_api.app.error import AgentNotPublishedError, AppUnavail from controllers.service_api.wraps import validate_app_token from core.app.app_config.common.parameters_mapping import get_parameters_from_feature_dict from core.app.apps.agent_app.errors import AgentAppGeneratorError, AgentAppNotPublishedError +from extensions.ext_database import db from fields.base import ResponseModel from models.model import App, AppMode from services.app_service import AppService @@ -122,7 +123,7 @@ class AppMetaApi(Resource): Returns metadata about the application including configuration and settings. """ - return AppService().get_app_meta(app_model) + return AppService().get_app_meta(app_model, session=db.session()) @service_api_ns.route("/info") diff --git a/api/controllers/service_api/app/audio.py b/api/controllers/service_api/app/audio.py index 53b31c8e6c4..68ab5f31ea5 100644 --- a/api/controllers/service_api/app/audio.py +++ b/api/controllers/service_api/app/audio.py @@ -188,7 +188,7 @@ class TextApi(Resource): ) response = AudioService.transcript_tts( app_model=app_model, - session=db.session, + session=db.session(), text=text, voice=voice, end_user=end_user.external_user_id, diff --git a/api/controllers/service_api/app/conversation.py b/api/controllers/service_api/app/conversation.py index 9b5533ea07a..a395dcb93fc 100644 --- a/api/controllers/service_api/app/conversation.py +++ b/api/controllers/service_api/app/conversation.py @@ -249,7 +249,7 @@ class ConversationDetailApi(Resource): conversation_id = str(c_id) try: - ConversationService.delete(app_model, conversation_id, end_user) + ConversationService.delete(app_model, conversation_id, end_user, session=db.session()) except services.errors.conversation.ConversationNotExistsError: raise NotFound("Conversation Not Exists.") return "", 204 @@ -299,7 +299,7 @@ class ConversationRenameApi(Resource): try: conversation = ConversationService.rename( - app_model, conversation_id, end_user, payload.name, payload.auto_generate + app_model, conversation_id, end_user, payload.name, payload.auto_generate, session=db.session() ) return ( TypeAdapter(SimpleConversation) @@ -356,7 +356,13 @@ class ConversationVariablesApi(Resource): try: pagination = ConversationService.get_conversational_variable( - app_model, conversation_id, end_user, query_args.limit, last_id, query_args.variable_name + app_model, + conversation_id, + end_user, + query_args.limit, + last_id, + query_args.variable_name, + session=db.session(), ) return ConversationVariableInfiniteScrollPaginationResponse.model_validate( pagination, from_attributes=True @@ -417,7 +423,7 @@ class ConversationVariableDetailApi(Resource): try: variable = ConversationService.update_conversation_variable( - app_model, conversation_id, variable_id_str, end_user, payload.value + app_model, conversation_id, variable_id_str, end_user, payload.value, session=db.session() ) return ConversationVariableResponse.model_validate(variable, from_attributes=True).model_dump(mode="json") except services.errors.conversation.ConversationNotExistsError: diff --git a/api/controllers/service_api/app/message.py b/api/controllers/service_api/app/message.py index 18d1c5d3254..3acb2c74872 100644 --- a/api/controllers/service_api/app/message.py +++ b/api/controllers/service_api/app/message.py @@ -15,6 +15,7 @@ from controllers.service_api.app.error import NotChatAppError from controllers.service_api.schema import expect_with_user from controllers.service_api.wraps import FetchUserArg, WhereisUserArg, validate_app_token from core.app.entities.app_invoke_entities import InvokeFrom +from extensions.ext_database import db from fields.base import ResponseModel from fields.conversation_fields import ResultResponse from fields.message_fields import MessageInfiniteScrollPagination, MessageListItem @@ -109,7 +110,7 @@ class MessageListApi(Resource): try: pagination = MessageService.pagination_by_first_id( - app_model, end_user, conversation_id, first_id, query_args.limit + app_model, end_user, conversation_id, first_id, query_args.limit, session=db.session() ) adapter = TypeAdapter(MessageListItem) items = [adapter.validate_python(message, from_attributes=True) for message in pagination.data] @@ -167,6 +168,7 @@ class MessageFeedbackApi(Resource): user=end_user, rating=FeedbackRating(payload.rating) if payload.rating else None, content=payload.content, + session=db.session(), ) except MessageNotExistsError: raise NotFound("Message Not Exists.") @@ -208,7 +210,9 @@ class AppGetFeedbacksApi(Resource): Returns paginated list of all feedback submitted for messages in this app. """ query_args = FeedbackListQuery.model_validate(request.args.to_dict()) - feedbacks = MessageService.get_all_messages_feedbacks(app_model, page=query_args.page, limit=query_args.limit) + feedbacks = MessageService.get_all_messages_feedbacks( + app_model, page=query_args.page, limit=query_args.limit, session=db.session() + ) return {"data": feedbacks} @@ -258,7 +262,11 @@ class MessageSuggestedApi(Resource): try: questions = MessageService.get_suggested_questions_after_answer( - app_model=app_model, user=end_user, message_id=message_id_str, invoke_from=InvokeFrom.SERVICE_API + app_model=app_model, + user=end_user, + message_id=message_id_str, + invoke_from=InvokeFrom.SERVICE_API, + session=db.session(), ) except MessageNotExistsError: raise NotFound("Message Not Exists.") diff --git a/api/controllers/service_api/dataset/dataset.py b/api/controllers/service_api/dataset/dataset.py index 56836f56895..66085ca0642 100644 --- a/api/controllers/service_api/dataset/dataset.py +++ b/api/controllers/service_api/dataset/dataset.py @@ -414,7 +414,7 @@ class DatasetListApi(DatasetApiResource): datasets, total = DatasetService.get_datasets( query.page, query.limit, - db.session, + db.session(), tenant_id, current_user, query.keyword, @@ -565,11 +565,11 @@ class DatasetApi(DatasetApiResource): ) def get(self, _, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) data = _dump_service_dataset_detail(dataset) @@ -601,7 +601,7 @@ class DatasetApi(DatasetApiResource): retrieval_model_dict["search_method"] = "keyword_search" if data.get("permission") == "partial_members": - part_users_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session) + part_users_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session()) data.update({"partial_member_list": part_users_list}) return _dump_service_dataset_with_partial_members(data), 200 @@ -640,7 +640,7 @@ class DatasetApi(DatasetApiResource): @with_session def patch(self, session: Session, _, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") @@ -681,10 +681,10 @@ class DatasetApi(DatasetApiResource): dataset, str(payload.permission) if payload.permission else None, payload.partial_member_list, - db.session, + session=db.session(), ) - dataset = DatasetService.update_dataset(session, dataset_id_str, update_data, current_user) + dataset = DatasetService.update_dataset(dataset_id_str, update_data, current_user, session=session) if dataset is None: raise NotFound("Dataset not found.") @@ -695,13 +695,13 @@ class DatasetApi(DatasetApiResource): if payload.partial_member_list and payload.permission == DatasetPermissionEnum.PARTIAL_TEAM: DatasetPermissionService.update_partial_member_list( - tenant_id, dataset_id_str, payload.partial_member_list, db.session + tenant_id, dataset_id_str, payload.partial_member_list, db.session() ) # clear partial member list when permission is only_me or all_team_members elif payload.permission in {DatasetPermissionEnum.ONLY_ME, DatasetPermissionEnum.ALL_TEAM}: - DatasetPermissionService.clear_partial_member_list(dataset_id_str, db.session) + DatasetPermissionService.clear_partial_member_list(dataset_id_str, db.session()) - partial_member_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session) + partial_member_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session()) result_data.update({"partial_member_list": partial_member_list}) return _dump_service_dataset_with_partial_members(result_data), 200 @@ -754,8 +754,8 @@ class DatasetApi(DatasetApiResource): dataset_id_str = str(dataset_id) try: - if DatasetService.delete_dataset(dataset_id_str, current_user, db.session): - DatasetPermissionService.clear_partial_member_list(dataset_id_str, db.session) + if DatasetService.delete_dataset(dataset_id_str, current_user, db.session()): + DatasetPermissionService.clear_partial_member_list(dataset_id_str, db.session()) return "", 204 else: raise NotFound("Dataset not found.") @@ -820,14 +820,14 @@ class DocumentStatusApi(DatasetApiResource): InvalidActionError: If the action is invalid or cannot be performed. """ dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") # Check user's permission try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -839,7 +839,7 @@ class DocumentStatusApi(DatasetApiResource): document_ids = data.get("document_ids", []) try: - DocumentService.batch_update_document_status(dataset, document_ids, action, current_user, db.session) + DocumentService.batch_update_document_status(dataset, document_ids, action, current_user, db.session()) except services.errors.document.DocumentIndexingError as e: raise InvalidActionError(str(e)) except ValueError as e: @@ -876,7 +876,7 @@ class DatasetTagsApi(DatasetApiResource): assert isinstance(current_user, Account) cid = current_user.current_tenant_id assert cid is not None - tags = TagService.get_tags(db.session(), "knowledge", cid) + tags = TagService.get_tags("knowledge", cid, session=db.session()) return dump_response(KnowledgeTagListResponse, tags), 200 @service_api_ns.doc( @@ -909,7 +909,7 @@ class DatasetTagsApi(DatasetApiResource): raise Forbidden() payload = TagCreatePayload.model_validate(service_api_ns.payload or {}) - tag = TagService.save_tags(SaveTagPayload(name=payload.name, type=TagType.KNOWLEDGE), db.session) + tag = TagService.save_tags(SaveTagPayload(name=payload.name, type=TagType.KNOWLEDGE), db.session()) response = dump_response( KnowledgeTagResponse, @@ -948,10 +948,10 @@ class DatasetTagsApi(DatasetApiResource): payload = TagUpdatePayload.model_validate(service_api_ns.payload or {}) tag_id = payload.tag_id tag = TagService.update_tags( - UpdateTagServicePayload(name=payload.name), tag_id, db.session, tag_type=TagType.KNOWLEDGE + UpdateTagServicePayload(name=payload.name), tag_id, db.session(), tag_type=TagType.KNOWLEDGE ) - binding_count = TagService.get_tag_binding_count(tag_id, db.session, tag_type=TagType.KNOWLEDGE) + binding_count = TagService.get_tag_binding_count(tag_id, db.session(), tag_type=TagType.KNOWLEDGE) response = dump_response( KnowledgeTagResponse, @@ -981,7 +981,7 @@ class DatasetTagsApi(DatasetApiResource): def delete(self, _): """Delete a knowledge type tag.""" payload = TagDeletePayload.model_validate(service_api_ns.payload or {}) - TagService.delete_tag(payload.tag_id, db.session, tag_type=TagType.KNOWLEDGE) + TagService.delete_tag(payload.tag_id, db.session(), tag_type=TagType.KNOWLEDGE) return "", 204 @@ -1015,7 +1015,7 @@ class DatasetTagBindingApi(DatasetApiResource): payload = TagBindingPayload.model_validate(service_api_ns.payload or {}) TagService.save_tag_binding( TagBindingCreatePayload(tag_ids=payload.tag_ids, target_id=payload.target_id, type=TagType.KNOWLEDGE), - db.session, + db.session(), ) return "", 204 @@ -1050,7 +1050,7 @@ class DatasetTagUnbindingApi(DatasetApiResource): payload = TagUnbindingPayload.model_validate(service_api_ns.payload or {}) TagService.delete_tag_binding( TagBindingDeletePayload(tag_ids=payload.tag_ids, target_id=payload.target_id, type=TagType.KNOWLEDGE), - db.session, + db.session(), ) return "", 204 @@ -1086,7 +1086,7 @@ class DatasetTagsBindingStatusApi(DatasetApiResource): assert isinstance(current_user, Account) assert current_user.current_tenant_id is not None tags = TagService.get_tags_by_target_id( - "knowledge", current_user.current_tenant_id, str(dataset_id), db.session + "knowledge", current_user.current_tenant_id, str(dataset_id), db.session() ) tags_list = [{"id": tag.id, "name": tag.name} for tag in tags] return dump_response(DatasetBoundTagListResponse, {"data": tags_list, "total": len(tags)}), 200 diff --git a/api/controllers/service_api/dataset/document.py b/api/controllers/service_api/dataset/document.py index 4c083d3d50f..5e5919a7048 100644 --- a/api/controllers/service_api/dataset/document.py +++ b/api/controllers/service_api/dataset/document.py @@ -401,7 +401,7 @@ def _create_document_by_text(tenant_id: str, dataset_id: UUID) -> tuple[Mapping[ account=current_user, dataset_process_rule=dataset.latest_process_rule if "process_rule" not in args else None, created_from="api", - session=db.session, + session=db.session(), ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -461,7 +461,7 @@ def _update_document_by_text(tenant_id: str, dataset_id: UUID, document_id: UUID account=current_user, dataset_process_rule=dataset.latest_process_rule if "process_rule" not in args else None, created_from="api", - session=db.session, + session=db.session(), ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -759,7 +759,7 @@ class DocumentAddByFileApi(DatasetApiResource): account=dataset.created_by_account, dataset_process_rule=dataset_process_rule, created_from="api", - session=db.session, + session=db.session(), ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -836,7 +836,7 @@ def _update_document_by_file(tenant_id: str, dataset_id: UUID, document_id: UUID account=dataset.created_by_account, dataset_process_rule=dataset.latest_process_rule if "process_rule" not in args else None, created_from="api", - session=db.session, + session=db.session(), ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -955,6 +955,7 @@ class DocumentListApi(DatasetApiResource): documents=documents, dataset=dataset, tenant_id=tenant_id, + session=db.session(), ) response = { @@ -1007,7 +1008,7 @@ class DocumentBatchDownloadZipApi(DatasetApiResource): document_ids=[str(document_id) for document_id in payload.document_ids], tenant_id=str(tenant_id), current_user=current_user, - session=db.session, + session=db.session(), ) with ExitStack() as stack: @@ -1064,7 +1065,7 @@ class DocumentIndexingStatusApi(DatasetApiResource): if not dataset: raise NotFound("Dataset not found.") # get documents - documents = DocumentService.get_batch_documents(dataset_id_str, batch, db.session) + documents = DocumentService.get_batch_documents(dataset_id_str, batch, db.session()) if not documents: raise NotFound("Documents not found.") documents_status = [] @@ -1140,7 +1141,7 @@ class DocumentDownloadApi(DatasetApiResource): @cloud_edition_billing_rate_limit_check("knowledge", "dataset") def get(self, tenant_id, dataset_id: UUID, document_id: UUID): dataset = self.get_dataset(str(dataset_id), str(tenant_id)) - document = DocumentService.get_document(dataset.id, str(document_id), session=db.session) + document = DocumentService.get_document(dataset.id, str(document_id), session=db.session()) if not document: raise NotFound("Document not found.") @@ -1148,7 +1149,7 @@ class DocumentDownloadApi(DatasetApiResource): if document.tenant_id != str(tenant_id): raise Forbidden("No permission.") - return {"url": DocumentService.get_document_download_url(document, db.session)} + return {"url": DocumentService.get_document_download_url(document, db.session())} @service_api_ns.route("/datasets//documents/") @@ -1196,7 +1197,7 @@ class DocumentApi(DatasetApiResource): dataset = self.get_dataset(dataset_id_str, tenant_id) - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") @@ -1216,12 +1217,13 @@ class DocumentApi(DatasetApiResource): document_id=document_id_str, dataset_id=dataset_id_str, tenant_id=tenant_id, + session=db.session(), ) if metadata == "only": response = {"id": document.id, "doc_type": document.doc_type, "doc_metadata": document.doc_metadata_details} elif metadata == "without": - dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, db.session) + dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, db.session()) document_process_rules = document.dataset_process_rule.to_dict() if document.dataset_process_rule else {} data_source_info = document.data_source_detail_dict response = { @@ -1256,7 +1258,7 @@ class DocumentApi(DatasetApiResource): "need_summary": document.need_summary if document.need_summary is not None else False, } else: - dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, db.session) + dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, db.session()) document_process_rules = document.dataset_process_rule.to_dict() if document.dataset_process_rule else {} data_source_info = document.data_source_detail_dict response = { @@ -1351,7 +1353,7 @@ class DocumentApi(DatasetApiResource): if not dataset: raise ValueError("Dataset does not exist.") - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) # 404 if document not found if document is None: @@ -1363,7 +1365,7 @@ class DocumentApi(DatasetApiResource): try: # delete document - DocumentService.delete_document(document, db.session) + DocumentService.delete_document(document, db.session()) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Cannot delete document during indexing.") diff --git a/api/controllers/service_api/dataset/metadata.py b/api/controllers/service_api/dataset/metadata.py index aec3b06a91e..1d793583cc2 100644 --- a/api/controllers/service_api/dataset/metadata.py +++ b/api/controllers/service_api/dataset/metadata.py @@ -81,12 +81,12 @@ class DatasetMetadataCreateServiceApi(DatasetApiResource): metadata_args = MetadataArgs.model_validate(service_api_ns.payload or {}) dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) - metadata = MetadataService.create_metadata(db.session(), dataset_id_str, metadata_args) + metadata = MetadataService.create_metadata(dataset_id_str, metadata_args, session=db.session()) return dump_response(DatasetMetadataResponse, metadata), 201 @service_api_ns.doc( @@ -116,10 +116,10 @@ class DatasetMetadataCreateServiceApi(DatasetApiResource): def get(self, tenant_id, dataset_id: UUID): """Get all metadata for a dataset.""" dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - metadata = MetadataService.get_dataset_metadatas(db.session(), dataset) + metadata = MetadataService.get_dataset_metadatas(dataset, session=db.session()) return dump_response(DatasetMetadataListResponse, metadata), 200 @@ -154,12 +154,14 @@ class DatasetMetadataServiceApi(DatasetApiResource): dataset_id_str = str(dataset_id) metadata_id_str = str(metadata_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) - metadata = MetadataService.update_metadata_name(db.session(), dataset_id_str, metadata_id_str, payload.name) + metadata = MetadataService.update_metadata_name( + dataset_id_str, metadata_id_str, payload.name, session=db.session() + ) return dump_response(DatasetMetadataResponse, metadata), 200 @service_api_ns.doc( @@ -189,12 +191,12 @@ class DatasetMetadataServiceApi(DatasetApiResource): """Delete metadata.""" dataset_id_str = str(dataset_id) metadata_id_str = str(metadata_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) - MetadataService.delete_metadata(db.session(), dataset_id_str, metadata_id_str) + MetadataService.delete_metadata(dataset_id_str, metadata_id_str, session=db.session()) return "", 204 @@ -257,16 +259,16 @@ class DatasetMetadataBuiltInFieldActionServiceApi(DatasetApiResource): def post(self, tenant_id, dataset_id: UUID, action: Literal["enable", "disable"]): """Enable or disable built-in metadata field.""" dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) match action: case "enable": - MetadataService.enable_built_in_field(db.session(), dataset) + MetadataService.enable_built_in_field(dataset, session=db.session()) case "disable": - MetadataService.disable_built_in_field(db.session(), dataset) + MetadataService.disable_built_in_field(dataset, session=db.session()) return dump_response(DatasetMetadataActionResponse, {"result": "success"}), 200 @@ -303,13 +305,13 @@ class DocumentMetadataEditServiceApi(DatasetApiResource): def post(self, tenant_id, dataset_id: UUID): """Update metadata for multiple documents.""" dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) metadata_args = MetadataOperationData.model_validate(service_api_ns.payload or {}) - MetadataService.update_documents_metadata(db.session(), dataset, metadata_args) + MetadataService.update_documents_metadata(dataset, metadata_args, session=db.session()) return dump_response(DatasetMetadataActionResponse, {"result": "success"}), 200 diff --git a/api/controllers/service_api/dataset/rag_pipeline/rag_pipeline_workflow.py b/api/controllers/service_api/dataset/rag_pipeline/rag_pipeline_workflow.py index f0f953462c9..35f3a4c01a0 100644 --- a/api/controllers/service_api/dataset/rag_pipeline/rag_pipeline_workflow.py +++ b/api/controllers/service_api/dataset/rag_pipeline/rag_pipeline_workflow.py @@ -159,7 +159,7 @@ class DatasourcePluginsApi(DatasetApiResource): query = query_params_from_request(DatasourcePluginsQuery) - rag_pipeline_service: RagPipelineService = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) datasource_plugins: list[dict[Any, Any]] = rag_pipeline_service.get_datasource_plugins( tenant_id=tenant_id, dataset_id=dataset_id_str, is_published=query.is_published ) @@ -204,7 +204,7 @@ class DatasourceNodeRunApi(DatasetApiResource): payload = DatasourceNodeRunPayload.model_validate(service_api_ns.payload or {}) assert isinstance(current_user, Account) - rag_pipeline_service: RagPipelineService = RagPipelineService() + rag_pipeline_service: RagPipelineService = RagPipelineService(db.session()) pipeline: Pipeline = rag_pipeline_service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset_id_str) datasource_node_run_api_entity = DatasourceNodeRunApiEntity.model_validate( { @@ -272,7 +272,7 @@ class PipelineRunApi(DatasetApiResource): dataset_id_str = str(dataset_id) # Verify dataset ownership stmt = select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id_str) - dataset = db.session.scalar(stmt) + dataset = session.scalar(stmt) if not dataset: raise NotFound("Dataset not found.") @@ -281,8 +281,8 @@ class PipelineRunApi(DatasetApiResource): if not isinstance(current_user, Account): raise Forbidden() - rag_pipeline_service: RagPipelineService = RagPipelineService() - pipeline: Pipeline = rag_pipeline_service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset_id_str) + rag_pipeline_service = RagPipelineService(session) + pipeline = rag_pipeline_service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset_id_str) try: response: dict[Any, Any] | Generator[str, Any, None] = PipelineGenerateService.generate( session=session, diff --git a/api/controllers/service_api/dataset/segment.py b/api/controllers/service_api/dataset/segment.py index 41fbc709fdd..e911c454c9e 100644 --- a/api/controllers/service_api/dataset/segment.py +++ b/api/controllers/service_api/dataset/segment.py @@ -137,7 +137,7 @@ def _get_segment_for_document( raise NotFound("Document not found.") segment_ref = DatasetRefService.create_segment_ref(document_ref, segment_id) - segment = SegmentService.get_segment_by_ref(segment_ref) + segment = SegmentService.get_segment_by_ref(segment_ref, db.session()) if not segment: raise NotFound("Segment not found.") return segment_ref, segment @@ -191,7 +191,7 @@ class SegmentApi(DatasetApiResource): raise NotFound("Dataset not found.") document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") if document.indexing_status != "completed": @@ -227,13 +227,13 @@ class SegmentApi(DatasetApiResource): for args_item in segment_items: SegmentService.segment_create_args_validate(args_item, document) segments = cast( - list[DocumentSegment], SegmentService.multi_create_segment(segment_items, document, dataset, db.session) + list[DocumentSegment], SegmentService.multi_create_segment(segment_items, document, dataset, db.session()) ) segment_ids = [segment.id for segment in segments] summaries: dict[str, str | None] = {} if segment_ids: summary_records = SummaryIndexService.get_segments_summaries( - segment_ids=segment_ids, dataset_id=dataset_id_str + segment_ids=segment_ids, dataset_id=dataset_id_str, session=db.session() ) summaries = {chunk_id: record.summary_content for chunk_id, record in summary_records.items()} response = { @@ -285,7 +285,7 @@ class SegmentApi(DatasetApiResource): raise NotFound("Dataset not found.") document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") # check embedding model setting @@ -317,7 +317,7 @@ class SegmentApi(DatasetApiResource): summaries: dict[str, str | None] = {} if segment_ids: summary_records = SummaryIndexService.get_segments_summaries( - segment_ids=segment_ids, dataset_id=dataset_id_str + segment_ids=segment_ids, dataset_id=dataset_id_str, session=db.session() ) summaries = {chunk_id: record.summary_content for chunk_id, record in summary_records.items()} @@ -367,12 +367,12 @@ class DatasetSegmentApi(DatasetApiResource): DatasetService.check_dataset_model_setting(dataset) document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") segment_id_str = str(segment_id) _, segment = _get_segment_for_document(dataset, document, segment_id_str) - SegmentService.delete_segment(segment, document, dataset, db.session) + SegmentService.delete_segment(segment, document, dataset, db.session()) return "", 204 @service_api_ns.doc( @@ -410,7 +410,7 @@ class DatasetSegmentApi(DatasetApiResource): DatasetService.check_dataset_model_setting(dataset) document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: @@ -434,8 +434,10 @@ class DatasetSegmentApi(DatasetApiResource): payload = SegmentUpdatePayload.model_validate(service_api_ns.payload or {}) - updated_segment = SegmentService.update_segment(payload.segment, segment, document, dataset, db.session) - summary = SummaryIndexService.get_segment_summary(segment_id=updated_segment.id, dataset_id=dataset_id_str) + updated_segment = SegmentService.update_segment(payload.segment, segment, document, dataset, db.session()) + summary = SummaryIndexService.get_segment_summary( + segment_id=updated_segment.id, dataset_id=dataset_id_str, session=db.session() + ) response = { "data": segment_response_with_summary(updated_segment, summary.summary_content if summary else None), "doc_form": document.doc_form, @@ -481,13 +483,15 @@ class DatasetSegmentApi(DatasetApiResource): DatasetService.check_dataset_model_setting(dataset) document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") segment_id_str = str(segment_id) _, segment = _get_segment_for_document(dataset, document, segment_id_str) - summary = SummaryIndexService.get_segment_summary(segment_id=segment.id, dataset_id=dataset_id_str) + summary = SummaryIndexService.get_segment_summary( + segment_id=segment.id, dataset_id=dataset_id_str, session=db.session() + ) response = { "data": segment_response_with_summary(segment, summary.summary_content if summary else None), "doc_form": document.doc_form, @@ -542,7 +546,7 @@ class ChildChunkApi(DatasetApiResource): document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") @@ -570,7 +574,7 @@ class ChildChunkApi(DatasetApiResource): payload = ChildChunkCreatePayload.model_validate(service_api_ns.payload or {}) try: - child_chunk = SegmentService.create_child_chunk(payload.content, segment, document, dataset, db.session) + child_chunk = SegmentService.create_child_chunk(payload.content, segment, document, dataset, db.session()) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) @@ -613,7 +617,7 @@ class ChildChunkApi(DatasetApiResource): document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") @@ -680,7 +684,7 @@ class DatasetChildChunkApi(DatasetApiResource): document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") @@ -689,12 +693,12 @@ class DatasetChildChunkApi(DatasetApiResource): child_chunk_id_str = str(child_chunk_id) # check child chunk - child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref) + child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref, db.session()) if not child_chunk: raise NotFound("Child chunk not found.") try: - SegmentService.delete_child_chunk(child_chunk, dataset, db.session) + SegmentService.delete_child_chunk(child_chunk, dataset, db.session()) except ChildChunkDeleteIndexServiceError as e: raise ChildChunkDeleteIndexError(str(e)) @@ -741,7 +745,7 @@ class DatasetChildChunkApi(DatasetApiResource): document_id_str = str(document_id) # get document - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") @@ -750,7 +754,7 @@ class DatasetChildChunkApi(DatasetApiResource): child_chunk_id_str = str(child_chunk_id) # get child chunk - child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref) + child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref, db.session()) if not child_chunk: raise NotFound("Child chunk not found.") @@ -759,7 +763,7 @@ class DatasetChildChunkApi(DatasetApiResource): try: child_chunk = SegmentService.update_child_chunk( - payload.content, child_chunk, segment, document, dataset, db.session + payload.content, child_chunk, segment, document, dataset, db.session() ) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) diff --git a/api/controllers/web/app.py b/api/controllers/web/app.py index 17ff05f7137..6804d072ef0 100644 --- a/api/controllers/web/app.py +++ b/api/controllers/web/app.py @@ -12,6 +12,7 @@ from controllers.common.agent_app_parameters import get_published_agent_app_feat from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models from core.app.app_config.common.parameters_mapping import get_parameters_from_feature_dict from core.app.apps.agent_app.errors import AgentAppGeneratorError, AgentAppNotPublishedError +from extensions.ext_database import db from libs.passport import PassportService from libs.token import extract_webapp_passport from models.model import App, AppMode, EndUser @@ -122,7 +123,7 @@ class AppMeta(WebApiResource): @web_ns.response(200, "Success", web_ns.models[AppMetaResponse.__name__]) def get(self, app_model: App, end_user: EndUser): """Get app meta""" - return AppService().get_app_meta(app_model) + return AppService().get_app_meta(app_model, session=db.session()) @web_ns.route("/webapp/access-mode") @@ -148,7 +149,7 @@ class AppAccessMode(Resource): app_id = args.app_id if args.app_code: - app_id = AppService.get_app_id_by_code(args.app_code) + app_id = AppService.get_app_id_by_code(args.app_code, session=db.session()) if not app_id: raise ValueError("appId or appCode must be provided") @@ -179,7 +180,9 @@ class AppWebAuthPermission(Resource): if not app_id or not app_code: raise ValueError("appId must be provided") - require_permission_check = WebAppAuthService.is_app_require_permission_check(app_id=app_id) + require_permission_check = WebAppAuthService.is_app_require_permission_check( + app_id=app_id, session=db.session() + ) if not require_permission_check: return {"result": True} @@ -200,6 +203,6 @@ class AppWebAuthPermission(Resource): return {"result": True} res = True - if WebAppAuthService.is_app_require_permission_check(app_id=app_id): + if WebAppAuthService.is_app_require_permission_check(app_id=app_id, session=db.session()): res = EnterpriseService.WebAppAuth.is_user_allowed_to_access_webapp(str(user_id), app_id) return {"result": res} diff --git a/api/controllers/web/audio.py b/api/controllers/web/audio.py index 47e72ff95a5..b7856f7dd90 100644 --- a/api/controllers/web/audio.py +++ b/api/controllers/web/audio.py @@ -141,7 +141,7 @@ class TextApi(WebApiResource): ) response = AudioService.transcript_tts( app_model=app_model, - session=db.session, + session=db.session(), text=text, voice=voice, end_user=end_user.external_user_id, diff --git a/api/controllers/web/completion.py b/api/controllers/web/completion.py index 343afd68f9a..c1a7d1f8d10 100644 --- a/api/controllers/web/completion.py +++ b/api/controllers/web/completion.py @@ -30,6 +30,7 @@ from core.errors.error import ( ProviderTokenNotInitError, QuotaExceededError, ) +from extensions.ext_database import db from graphon.model_runtime.errors.invoke import InvokeError from libs import helper from libs.helper import uuid_value @@ -219,7 +220,10 @@ class ChatApi(WebApiResource): # Eagerly validate conversation to avoid hanging on invalid conversation_id if payload.conversation_id: ConversationService.get_conversation( - app_model=app_model, conversation_id=payload.conversation_id, user=end_user + app_model=app_model, + conversation_id=payload.conversation_id, + user=end_user, + session=db.session(), ) response = AppGenerateService.generate( diff --git a/api/controllers/web/conversation.py b/api/controllers/web/conversation.py index 73461b1a294..09a3a508824 100644 --- a/api/controllers/web/conversation.py +++ b/api/controllers/web/conversation.py @@ -112,7 +112,7 @@ class ConversationApi(WebApiResource): conversation_id = str(c_id) try: - ConversationService.delete(app_model, conversation_id, end_user) + ConversationService.delete(app_model, conversation_id, end_user, session=db.session()) except ConversationNotExistsError: raise NotFound("Conversation Not Exists.") return "", 204 @@ -157,7 +157,7 @@ class ConversationRenameApi(WebApiResource): try: conversation = ConversationService.rename( - app_model, conversation_id, end_user, payload.name, payload.auto_generate + app_model, conversation_id, end_user, payload.name, payload.auto_generate, session=db.session() ) return ( TypeAdapter(SimpleConversation) @@ -192,7 +192,7 @@ class ConversationPinApi(WebApiResource): conversation_id = str(c_id) try: - WebConversationService.pin(app_model, conversation_id, end_user) + WebConversationService.pin(app_model, conversation_id, end_user, db.session()) except ConversationNotExistsError: raise NotFound("Conversation Not Exists.") @@ -221,6 +221,6 @@ class ConversationUnPinApi(WebApiResource): raise NotChatAppError() conversation_id = str(c_id) - WebConversationService.unpin(app_model, conversation_id, end_user) + WebConversationService.unpin(app_model, conversation_id, end_user, db.session()) return ResultResponse(result="success").model_dump(mode="json") diff --git a/api/controllers/web/forgot_password.py b/api/controllers/web/forgot_password.py index ecc91113c32..a9374555ed4 100644 --- a/api/controllers/web/forgot_password.py +++ b/api/controllers/web/forgot_password.py @@ -69,7 +69,7 @@ class ForgotPasswordSendEmailApi(Resource): else: language = "en-US" - account = AccountService.get_account_by_email_with_case_fallback(db.session, request_email) + account = AccountService.get_account_by_email_with_case_fallback(request_email, session=db.session()) if account is None: raise AuthenticationFailedError() else: @@ -168,7 +168,7 @@ class ForgotPasswordResetApi(Resource): email = reset_data.get("email", "") - account = AccountService.get_account_by_email_with_case_fallback(db.session, email) + account = AccountService.get_account_by_email_with_case_fallback(email, session=db.session()) if account: account = db.session.merge(account) diff --git a/api/controllers/web/login.py b/api/controllers/web/login.py index 011bb43b880..0aa42f43687 100644 --- a/api/controllers/web/login.py +++ b/api/controllers/web/login.py @@ -30,6 +30,7 @@ from controllers.console.wraps import ( ) from controllers.web import web_ns from controllers.web.wraps import decode_jwt_token +from extensions.ext_database import db from libs.helper import EmailStr, extract_remote_ip from libs.passport import PassportService from libs.password import valid_password @@ -104,7 +105,7 @@ class LoginApi(Resource): normalized_email = payload.email.lower() try: - account = WebAppAuthService.authenticate(payload.email, payload.password) + account = WebAppAuthService.authenticate(payload.email, payload.password, db.session()) except services.errors.account.AccountLoginError: _log_web_login_failure(email=normalized_email, reason=LoginFailureReason.ACCOUNT_BANNED) raise AccountBannedError() @@ -144,9 +145,9 @@ class LoginStatusApi(Resource): token = extract_webapp_access_token(request) if not app_code: return LoginStatusResponse(logged_in=bool(token), app_logged_in=False).model_dump(mode="json") - app_id = AppService.get_app_id_by_code(app_code) + app_id = AppService.get_app_id_by_code(app_code, session=db.session()) is_public = not dify_config.ENTERPRISE_ENABLED or not WebAppAuthService.is_app_require_permission_check( - app_id=app_id + app_id=app_id, session=db.session() ) user_logged_in = False @@ -211,7 +212,7 @@ class EmailCodeLoginSendEmailApi(Resource): else: language = "en-US" - account = WebAppAuthService.get_user_through_email(payload.email) + account = WebAppAuthService.get_user_through_email(payload.email, db.session()) if account is None: raise AuthenticationFailedError() token = WebAppAuthService.send_email_code_login_email(account=account, language=language) @@ -264,7 +265,7 @@ class EmailCodeLoginApi(Resource): WebAppAuthService.revoke_email_code_login_token(payload.token) try: - account = WebAppAuthService.get_user_through_email(token_email) + account = WebAppAuthService.get_user_through_email(token_email, db.session()) except Unauthorized as exc: _log_web_login_failure(email=user_email, reason=LoginFailureReason.ACCOUNT_BANNED) raise AccountBannedError() from exc diff --git a/api/controllers/web/message.py b/api/controllers/web/message.py index 691eba05491..45fea9a328e 100644 --- a/api/controllers/web/message.py +++ b/api/controllers/web/message.py @@ -25,6 +25,7 @@ from controllers.web.error import ( from controllers.web.wraps import WebApiResource from core.app.entities.app_invoke_entities import InvokeFrom from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError +from extensions.ext_database import db from fields.conversation_fields import ResultResponse from fields.message_fields import SuggestedQuestionsResponse, WebMessageInfiniteScrollPagination, WebMessageListItem from graphon.model_runtime.errors.invoke import InvokeError @@ -86,7 +87,7 @@ class MessageListApi(WebApiResource): try: pagination = MessageService.pagination_by_first_id( - app_model, end_user, query.conversation_id, query.first_id, query.limit + app_model, end_user, query.conversation_id, query.first_id, query.limit, session=db.session() ) adapter = TypeAdapter(WebMessageListItem) items = [adapter.validate_python(message, from_attributes=True) for message in pagination.data] @@ -141,6 +142,7 @@ class MessageFeedbackApi(WebApiResource): user=end_user, rating=FeedbackRating(payload.rating) if payload.rating else None, content=payload.content, + session=db.session(), ) except MessageNotExistsError: raise NotFound("Message Not Exists.") @@ -231,7 +233,11 @@ class MessageSuggestedQuestionApi(WebApiResource): try: questions = MessageService.get_suggested_questions_after_answer( - app_model=app_model, user=end_user, message_id=message_id_str, invoke_from=InvokeFrom.WEB_APP + app_model=app_model, + user=end_user, + message_id=message_id_str, + invoke_from=InvokeFrom.WEB_APP, + session=db.session(), ) # questions is a list of strings, not a list of Message objects except MessageNotExistsError: diff --git a/api/controllers/web/passport.py b/api/controllers/web/passport.py index c11ce824731..4b0b25fb971 100644 --- a/api/controllers/web/passport.py +++ b/api/controllers/web/passport.py @@ -62,7 +62,7 @@ class PassportResource(Resource): raise Unauthorized("X-App-Code header is missing.") if system_features.webapp_auth.enabled: enterprise_user_decoded = decode_enterprise_webapp_user_id(access_token) - app_auth_type = WebAppAuthService.get_app_auth_type(app_code=app_code) + app_auth_type = WebAppAuthService.get_app_auth_type(app_code=app_code, session=db.session()) if app_auth_type != WebAppAuthType.PUBLIC: if not enterprise_user_decoded: raise WebAppAuthRequiredError() diff --git a/api/controllers/web/saved_message.py b/api/controllers/web/saved_message.py index 6e59a85e2b0..d61ffd545c3 100644 --- a/api/controllers/web/saved_message.py +++ b/api/controllers/web/saved_message.py @@ -44,7 +44,7 @@ class SavedMessageListApi(WebApiResource): query = SavedMessageListQuery.model_validate(raw_args) pagination = SavedMessageService.pagination_by_last_id( - db.session(), app_model, end_user, query.last_id, query.limit + app_model, end_user, query.last_id, query.limit, session=db.session() ) adapter = TypeAdapter(SavedMessageItem) items = [adapter.validate_python(message, from_attributes=True) for message in pagination.data] @@ -80,7 +80,7 @@ class SavedMessageListApi(WebApiResource): payload = SavedMessageCreatePayload.model_validate(web_ns.payload or {}) try: - SavedMessageService.save(db.session(), app_model, end_user, payload.message_id) + SavedMessageService.save(app_model, end_user, payload.message_id, session=db.session()) except MessageNotExistsError: raise NotFound("Message Not Exists.") @@ -108,6 +108,6 @@ class SavedMessageApi(WebApiResource): if app_model.mode != "completion": raise NotCompletionAppError() - SavedMessageService.delete(db.session(), app_model, end_user, message_id_str) + SavedMessageService.delete(app_model, end_user, message_id_str, session=db.session()) return "", 204 diff --git a/api/controllers/web/wraps.py b/api/controllers/web/wraps.py index ccc9c0f8f60..eff4b70ff0f 100644 --- a/api/controllers/web/wraps.py +++ b/api/controllers/web/wraps.py @@ -70,7 +70,7 @@ def decode_jwt_token(app_code: str | None = None, user_id: str | None = None) -> app_web_auth_enabled = False webapp_settings = None if system_features.webapp_auth.enabled: - app_id = AppService.get_app_id_by_code(app_code) + app_id = AppService.get_app_id_by_code(app_code, session=db.session()) webapp_settings = EnterpriseService.WebAppAuth.get_app_access_mode_by_id(app_id) if not webapp_settings: raise NotFound("Web app settings not found.") @@ -86,7 +86,7 @@ def decode_jwt_token(app_code: str | None = None, user_id: str | None = None) -> if system_features.webapp_auth.enabled: if not app_code: raise Unauthorized("Please re-login to access the web app.") - app_id = AppService.get_app_id_by_code(app_code) + app_id = AppService.get_app_id_by_code(app_code, session=db.session()) app_web_auth_enabled = ( EnterpriseService.WebAppAuth.get_app_access_mode_by_id(app_id=app_id).access_mode != WebAppAccessMode.PUBLIC @@ -129,8 +129,10 @@ def _validate_user_accessibility( if not webapp_settings: raise WebAppAuthRequiredError("Web app settings not found.") - if WebAppAuthService.is_app_require_permission_check(access_mode=webapp_settings.access_mode): - app_id = AppService.get_app_id_by_code(app_code) + if WebAppAuthService.is_app_require_permission_check( + access_mode=webapp_settings.access_mode, session=db.session() + ): + app_id = AppService.get_app_id_by_code(app_code, session=db.session()) if not EnterpriseService.WebAppAuth.is_user_allowed_to_access_webapp(user_id, app_id): raise WebAppAuthAccessDeniedError() diff --git a/api/core/app/app_config/easy_ui_based_app/dataset/manager.py b/api/core/app/app_config/easy_ui_based_app/dataset/manager.py index 140d4e6a2a6..0108e7d7c72 100644 --- a/api/core/app/app_config/easy_ui_based_app/dataset/manager.py +++ b/api/core/app/app_config/easy_ui_based_app/dataset/manager.py @@ -257,7 +257,7 @@ class DatasetConfigManager: @classmethod def is_dataset_exists(cls, tenant_id: str, dataset_id: str) -> bool: # verify if the dataset ID exists - dataset = DatasetService.get_dataset(dataset_id, db.session) + dataset = DatasetService.get_dataset(dataset_id, db.session()) if not dataset: return False diff --git a/api/core/app/apps/advanced_chat/app_generator.py b/api/core/app/apps/advanced_chat/app_generator.py index f52fd1046f8..75ada4fe888 100644 --- a/api/core/app/apps/advanced_chat/app_generator.py +++ b/api/core/app/apps/advanced_chat/app_generator.py @@ -156,7 +156,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): if conversation_id: try: conversation = ConversationService.get_conversation( - app_model=app_model, conversation_id=conversation_id, user=user + app_model=app_model, conversation_id=conversation_id, user=user, session=db.session() ) except ConversationNotExistsError: if invoke_from == InvokeFrom.SERVICE_API: diff --git a/api/core/app/apps/agent_app/app_generator.py b/api/core/app/apps/agent_app/app_generator.py index 9531f9092a4..f2f62496883 100644 --- a/api/core/app/apps/agent_app/app_generator.py +++ b/api/core/app/apps/agent_app/app_generator.py @@ -105,7 +105,7 @@ class AgentAppGenerator(MessageBasedAppGenerator): conversation_id = args.get("conversation_id") if conversation_id: conversation = ConversationService.get_conversation( - app_model=app_model, conversation_id=conversation_id, user=user + app_model=app_model, conversation_id=conversation_id, user=user, session=db.session() ) # Build the EasyUI-shaped config from the Agent Soul so the chat pipeline @@ -284,7 +284,7 @@ class AgentAppGenerator(MessageBasedAppGenerator): out of scope here — the message is persisted and can be re-fetched. """ conversation = ConversationService.get_conversation( - app_model=app_model, conversation_id=conversation_id, user=user + app_model=app_model, conversation_id=conversation_id, user=user, session=db.session() ) agent, agent_config_id, agent_config_version_kind, agent_soul = self._resolve_agent( app_model, diff --git a/api/core/app/apps/agent_chat/app_generator.py b/api/core/app/apps/agent_chat/app_generator.py index d640bcdc863..a3cc913abf3 100644 --- a/api/core/app/apps/agent_chat/app_generator.py +++ b/api/core/app/apps/agent_chat/app_generator.py @@ -108,7 +108,7 @@ class AgentChatAppGenerator(MessageBasedAppGenerator): conversation_id = args.get("conversation_id") if conversation_id: conversation = ConversationService.get_conversation( - app_model=app_model, conversation_id=conversation_id, user=user + app_model=app_model, conversation_id=conversation_id, user=user, session=db.session() ) # get app model config app_model_config = self._get_app_model_config(app_model=app_model, conversation=conversation) diff --git a/api/core/app/apps/chat/app_generator.py b/api/core/app/apps/chat/app_generator.py index 4873168b885..678525e0f77 100644 --- a/api/core/app/apps/chat/app_generator.py +++ b/api/core/app/apps/chat/app_generator.py @@ -105,7 +105,7 @@ class ChatAppGenerator(MessageBasedAppGenerator): conversation_id = args.get("conversation_id") if conversation_id: conversation = ConversationService.get_conversation( - app_model=app_model, conversation_id=conversation_id, user=user + app_model=app_model, conversation_id=conversation_id, user=user, session=db.session() ) # get app model config app_model_config = self._get_app_model_config(app_model=app_model, conversation=conversation) diff --git a/api/core/app/features/annotation_reply/annotation_reply.py b/api/core/app/features/annotation_reply/annotation_reply.py index 520ba7b85b3..9eff9747764 100644 --- a/api/core/app/features/annotation_reply/annotation_reply.py +++ b/api/core/app/features/annotation_reply/annotation_reply.py @@ -45,7 +45,7 @@ class AnnotationReplyFeature: embedding_model_name = collection_binding_detail.model_name dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding( - embedding_provider_name, embedding_model_name, db.session, CollectionBindingType.ANNOTATION + embedding_provider_name, embedding_model_name, db.session(), CollectionBindingType.ANNOTATION ) dataset = Dataset( @@ -66,7 +66,7 @@ class AnnotationReplyFeature: if documents and documents[0].metadata: annotation_id = documents[0].metadata["annotation_id"] score = documents[0].metadata["score"] - annotation = AppAnnotationService.get_annotation_by_id(annotation_id) + annotation = AppAnnotationService.get_annotation_by_id(annotation_id, session=db.session()) if annotation: if invoke_from in {InvokeFrom.SERVICE_API, InvokeFrom.WEB_APP}: from_source = ConversationFromSource.API @@ -84,6 +84,7 @@ class AnnotationReplyFeature: message.id, from_source, score, + session=db.session(), ) return annotation diff --git a/api/core/app/llm/quota.py b/api/core/app/llm/quota.py index 5bf3334a7b2..d26d5d8a998 100644 --- a/api/core/app/llm/quota.py +++ b/api/core/app/llm/quota.py @@ -125,6 +125,7 @@ def _deduct_used_llm_quota(*, tenant_id: str, provider: str, provider_configurat CreditPoolService.deduct_credits_capped( tenant_id=tenant_id, credits_required=used_quota, + session=db.session(), ) case ProviderQuotaType.PAID: from services.credit_pool_service import CreditPoolService @@ -133,6 +134,7 @@ def _deduct_used_llm_quota(*, tenant_id: str, provider: str, provider_configurat tenant_id=tenant_id, credits_required=used_quota, pool_type="paid", + session=db.session(), ) case ProviderQuotaType.FREE: _deduct_free_llm_quota( diff --git a/api/core/app/task_pipeline/message_cycle_manager.py b/api/core/app/task_pipeline/message_cycle_manager.py index 6b6437adac3..5ada7d0ba2d 100644 --- a/api/core/app/task_pipeline/message_cycle_manager.py +++ b/api/core/app/task_pipeline/message_cycle_manager.py @@ -154,7 +154,7 @@ class MessageCycleManager: :param event: event :return: """ - annotation = AppAnnotationService.get_annotation_by_id(event.message_annotation_id) + annotation = AppAnnotationService.get_annotation_by_id(event.message_annotation_id, session=db.session()) if annotation: account = annotation.account self._task_state.metadata.annotation_reply = AnnotationReply( diff --git a/api/core/callback_handler/index_tool_callback_handler.py b/api/core/callback_handler/index_tool_callback_handler.py index 26dc1a12a2c..d2024454a68 100644 --- a/api/core/callback_handler/index_tool_callback_handler.py +++ b/api/core/callback_handler/index_tool_callback_handler.py @@ -2,7 +2,7 @@ import logging from collections.abc import Sequence from sqlalchemy import select, update -from sqlalchemy.orm import Session, scoped_session, sessionmaker +from sqlalchemy.orm import Session, sessionmaker from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom from core.app.entities.app_invoke_entities import InvokeFrom @@ -30,7 +30,7 @@ class DatasetIndexToolCallbackHandler: self._user_id = user_id self._invoke_from = invoke_from - def on_query(self, query: str, dataset_id: str, session: scoped_session): + def on_query(self, query: str, dataset_id: str, session: Session): """ Handle query. """ @@ -52,7 +52,7 @@ class DatasetIndexToolCallbackHandler: with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as independent_session: independent_session.add(dataset_query) - def on_tool_end(self, documents: list[Document], session: scoped_session): + def on_tool_end(self, documents: list[Document], session: Session): """Handle tool end.""" # Use an independent session so hit-count updates do not # interfere with the caller's request-scoped session. diff --git a/api/core/llm_generator/llm_generator.py b/api/core/llm_generator/llm_generator.py index f97f9c38330..29a93fff815 100644 --- a/api/core/llm_generator/llm_generator.py +++ b/api/core/llm_generator/llm_generator.py @@ -6,6 +6,7 @@ from typing import Any, Literal, NotRequired, Protocol, TypedDict, cast import json_repair from sqlalchemy import select +from sqlalchemy.orm import Session from core.app.app_config.entities import ModelConfig from core.llm_generator.entities import RuleCodeGeneratePayload, RuleGeneratePayload, RuleStructuredOutputPayload @@ -117,7 +118,9 @@ def _parse_string_list(text: str) -> list[str]: class WorkflowServiceInterface(Protocol): - def get_draft_workflow(self, app_model: App, workflow_id: str | None = None) -> Workflow | None: + def get_draft_workflow( + self, app_model: App, workflow_id: str | None = None, *, session: Session + ) -> Workflow | None: pass def get_node_last_run(self, app_model: App, workflow: Workflow, node_id: str) -> WorkflowNodeExecutionModel | None: @@ -758,7 +761,7 @@ class LLMGenerator: app: App | None = session.scalar(select(App).where(App.id == flow_id, App.tenant_id == tenant_id).limit(1)) if not app: raise ValueError("App not found.") - workflow = workflow_service.get_draft_workflow(app_model=app) + workflow = workflow_service.get_draft_workflow(app_model=app, session=session) if not workflow: raise ValueError("Workflow not found for the given app model.") last_run = workflow_service.get_node_last_run(app_model=app, workflow=workflow, node_id=node_id) diff --git a/api/core/mcp/server/streamable_http.py b/api/core/mcp/server/streamable_http.py index 964f3211db0..7fd03788c7e 100644 --- a/api/core/mcp/server/streamable_http.py +++ b/api/core/mcp/server/streamable_http.py @@ -207,11 +207,11 @@ def handle_call_tool( raise ValueError("End user not found") response = AppGenerateService.generate( - session, - app, - end_user, - args, - InvokeFrom.SERVICE_API, + session=session, + app_model=app, + user=end_user, + args=args, + invoke_from=InvokeFrom.SERVICE_API, streaming=app.mode == AppMode.AGENT_CHAT, ) diff --git a/api/core/provider_manager.py b/api/core/provider_manager.py index e2c710923b5..ebfe77e8f30 100644 --- a/api/core/provider_manager.py +++ b/api/core/provider_manager.py @@ -1544,10 +1544,12 @@ class ProviderManager: trail_pool = CreditPoolService.get_pool( tenant_id=tenant_id, pool_type=ProviderQuotaType.TRIAL, + session=db.session(), ) paid_pool = CreditPoolService.get_pool( tenant_id=tenant_id, pool_type=ProviderQuotaType.PAID, + session=db.session(), ) else: trail_pool = None diff --git a/api/core/rag/datasource/retrieval_service.py b/api/core/rag/datasource/retrieval_service.py index 50381f5e75c..3b20f8bc530 100644 --- a/api/core/rag/datasource/retrieval_service.py +++ b/api/core/rag/datasource/retrieval_service.py @@ -199,7 +199,7 @@ class RetrievalService: metadata_filtering_conditions: dict[str, Any] | None = None, ): stmt = select(Dataset).where(Dataset.id == dataset_id) - dataset = db.session.scalar(stmt) + dataset = session.scalar(stmt) if not dataset: return [] metadata_condition = ( @@ -208,12 +208,12 @@ class RetrievalService: else None ) all_documents = ExternalDatasetService.fetch_external_knowledge_retrieval( - session, - dataset.tenant_id, - dataset_id, - query, - external_retrieval_model or {}, + tenant_id=dataset.tenant_id, + dataset_id=dataset_id, + query=query, + external_retrieval_parameters=external_retrieval_model or {}, metadata_condition=metadata_condition, + session=session, ) return all_documents diff --git a/api/core/rag/index_processor/processor/paragraph_index_processor.py b/api/core/rag/index_processor/processor/paragraph_index_processor.py index dd173207b09..b31c1bb634b 100644 --- a/api/core/rag/index_processor/processor/paragraph_index_processor.py +++ b/api/core/rag/index_processor/processor/paragraph_index_processor.py @@ -5,7 +5,7 @@ import re import uuid from typing import Any, TypedDict, cast, override -from sqlalchemy.orm import scoped_session +from sqlalchemy.orm import Session logger = logging.getLogger(__name__) @@ -162,10 +162,10 @@ class ParagraphIndexProcessor(BaseIndexProcessor): ).all() segment_ids = [segment.id for segment in segments] if segment_ids: - SummaryIndexService.delete_summaries_for_segments(dataset, segment_ids) + SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=segment_ids) else: # Delete all summaries for the dataset - SummaryIndexService.delete_summaries_for_segments(dataset, None) + SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=None) if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: vector = Vector(dataset) @@ -226,7 +226,7 @@ class ParagraphIndexProcessor(BaseIndexProcessor): all_multimodal_documents.append(file_document) doc.attachments = attachments else: - account = AccountService.load_user(document.created_by, db.session) + account = AccountService.load_user(document.created_by, db.session()) if not account: raise ValueError("Invalid account") doc.attachments = self._get_content_files(doc, current_user=account) @@ -414,12 +414,12 @@ class ParagraphIndexProcessor(BaseIndexProcessor): # First, try to get images from SegmentAttachmentBinding (preferred method) if segment_id: image_files = ParagraphIndexProcessor._extract_images_from_segment_attachments( - tenant_id, segment_id, db.session + tenant_id, segment_id, db.session() ) # If no images from attachments, fall back to extracting from text if not image_files: - image_files = ParagraphIndexProcessor._extract_images_from_text(tenant_id, text, db.session) + image_files = ParagraphIndexProcessor._extract_images_from_text(tenant_id, text, db.session()) # Build prompt messages prompt_messages = [] @@ -473,7 +473,7 @@ class ParagraphIndexProcessor(BaseIndexProcessor): return summary_content, usage @staticmethod - def _extract_images_from_text(tenant_id: str, text: str, session: scoped_session) -> list[File]: + def _extract_images_from_text(tenant_id: str, text: str, session: Session) -> list[File]: """ Extract images from markdown text and convert them to File objects. @@ -553,9 +553,7 @@ class ParagraphIndexProcessor(BaseIndexProcessor): return file_objects @staticmethod - def _extract_images_from_segment_attachments( - tenant_id: str, segment_id: str, session: scoped_session - ) -> list[File]: + def _extract_images_from_segment_attachments(tenant_id: str, segment_id: str, session: Session) -> list[File]: """ Extract images from SegmentAttachmentBinding table (preferred method). This matches how DatasetRetrieval gets segment attachments. diff --git a/api/core/rag/index_processor/processor/parent_child_index_processor.py b/api/core/rag/index_processor/processor/parent_child_index_processor.py index 78d8b7dcd53..aecb4154d6f 100644 --- a/api/core/rag/index_processor/processor/parent_child_index_processor.py +++ b/api/core/rag/index_processor/processor/parent_child_index_processor.py @@ -169,10 +169,10 @@ class ParentChildIndexProcessor(BaseIndexProcessor): ).all() segment_ids = [segment.id for segment in segments] if segment_ids: - SummaryIndexService.delete_summaries_for_segments(dataset, segment_ids) + SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=segment_ids) else: # Delete all summaries for the dataset - SummaryIndexService.delete_summaries_for_segments(dataset, None) + SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=None) if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: delete_child_chunks = kwargs.get("delete_child_chunks") or False @@ -291,7 +291,7 @@ class ParentChildIndexProcessor(BaseIndexProcessor): attachments.append(file_document) doc.attachments = attachments else: - account = AccountService.load_user(document.created_by, db.session) + account = AccountService.load_user(document.created_by, db.session()) if not account: raise ValueError("Invalid account") doc.attachments = self._get_content_files(doc, current_user=account) diff --git a/api/core/rag/index_processor/processor/qa_index_processor.py b/api/core/rag/index_processor/processor/qa_index_processor.py index 253acebc2c6..7b7443a621f 100644 --- a/api/core/rag/index_processor/processor/qa_index_processor.py +++ b/api/core/rag/index_processor/processor/qa_index_processor.py @@ -173,10 +173,10 @@ class QAIndexProcessor(BaseIndexProcessor): ).all() segment_ids = [segment.id for segment in segments] if segment_ids: - SummaryIndexService.delete_summaries_for_segments(dataset, segment_ids) + SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=segment_ids) else: # Delete all summaries for the dataset - SummaryIndexService.delete_summaries_for_segments(dataset, None) + SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=None) vector = Vector(dataset) if node_ids: diff --git a/api/core/rag/summary_index/summary_index.py b/api/core/rag/summary_index/summary_index.py index bff5f85decb..d9ce3879890 100644 --- a/api/core/rag/summary_index/summary_index.py +++ b/api/core/rag/summary_index/summary_index.py @@ -74,11 +74,16 @@ class SummaryIndex: def process_segment(segment_id: str) -> None: """Process a single segment in a thread with a fresh DB session.""" with session_factory.create_session() as session: + dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1)) + if dataset is None: + return segment = session.scalar(select(DocumentSegment).where(DocumentSegment.id == segment_id).limit(1)) if segment is None: return try: - SummaryIndexService.generate_and_vectorize_summary(segment, dataset, summary_index_setting) + SummaryIndexService.generate_and_vectorize_summary( + segment, dataset, summary_index_setting, session=session + ) except Exception: logger.exception( "Failed to generate summary for segment %s", diff --git a/api/core/tools/utils/dataset_retriever/dataset_multi_retriever_tool.py b/api/core/tools/utils/dataset_retriever/dataset_multi_retriever_tool.py index a3afe659563..c26523b9be5 100644 --- a/api/core/tools/utils/dataset_retriever/dataset_multi_retriever_tool.py +++ b/api/core/tools/utils/dataset_retriever/dataset_multi_retriever_tool.py @@ -80,7 +80,7 @@ class DatasetMultiRetrieverTool(DatasetRetrieverBaseTool): all_documents = rerank_runner.run(query, all_documents, self.score_threshold, self.top_k) for hit_callback in self.hit_callbacks: - hit_callback.on_tool_end(all_documents, db.session) + hit_callback.on_tool_end(all_documents, db.session()) document_score_list = {} for item in all_documents: @@ -167,7 +167,7 @@ class DatasetMultiRetrieverTool(DatasetRetrieverBaseTool): return [] for hit_callback in hit_callbacks: - hit_callback.on_query(query, dataset.id, db.session) + hit_callback.on_query(query, dataset.id, db.session()) # get retrieval model , if the model is not setting , using default retrieval_model = dataset.retrieval_model or default_retrieval_model diff --git a/api/core/tools/utils/dataset_retriever/dataset_retriever_tool.py b/api/core/tools/utils/dataset_retriever/dataset_retriever_tool.py index 247bd0705fc..d7e390ca877 100644 --- a/api/core/tools/utils/dataset_retriever/dataset_retriever_tool.py +++ b/api/core/tools/utils/dataset_retriever/dataset_retriever_tool.py @@ -65,7 +65,7 @@ class DatasetRetrieverTool(DatasetRetrieverBaseTool): if not dataset: return "" for hit_callback in self.hit_callbacks: - hit_callback.on_query(query, dataset.id, db.session) + hit_callback.on_query(query, dataset.id, db.session()) dataset_retrieval = DatasetRetrieval() metadata_filter_document_ids, metadata_condition = dataset_retrieval.get_metadata_filter_condition( session, @@ -162,7 +162,7 @@ class DatasetRetrieverTool(DatasetRetrieverBaseTool): else: documents = [] for hit_callback in self.hit_callbacks: - hit_callback.on_tool_end(documents, db.session) + hit_callback.on_tool_end(documents, db.session()) document_score_list = {} if dataset.indexing_technique != IndexTechniqueType.ECONOMY: for item in documents: diff --git a/api/core/workflow/nodes/agent_v2/dify_tools_builder.py b/api/core/workflow/nodes/agent_v2/dify_tools_builder.py index fc2719a6204..0e6fb3d1830 100644 --- a/api/core/workflow/nodes/agent_v2/dify_tools_builder.py +++ b/api/core/workflow/nodes/agent_v2/dify_tools_builder.py @@ -15,7 +15,6 @@ from dify_agent.layers.dify_plugin import ( DifyPluginToolsLayerConfig, ) from sqlalchemy import select -from sqlalchemy.orm import Session from core.agent.entities import AgentToolEntity from core.app.entities.app_invoke_entities import InvokeFrom @@ -132,7 +131,7 @@ def _list_provider_tool_names( def _resolve_mcp_provider_id(*, tenant_id: str, provider_id: str) -> str: """Normalize MCP provider ids to the runtime-facing server identifier.""" - service = MCPToolManageService(session=cast(Session, db.session)) + service = MCPToolManageService(session=db.session()) try: return service.get_provider_entity(provider_id, tenant_id, by_server_id=True).provider_id except ValueError: diff --git a/api/events/event_handlers/update_provider_when_message_created.py b/api/events/event_handlers/update_provider_when_message_created.py index 8dec5876a9b..15b40afdbf2 100644 --- a/api/events/event_handlers/update_provider_when_message_created.py +++ b/api/events/event_handlers/update_provider_when_message_created.py @@ -204,6 +204,7 @@ def _deduct_credit_pool_quota_capped(*, tenant_id: str, credits_required: int, p tenant_id=tenant_id, credits_required=credits_required, pool_type=pool_type, + session=db.session(), ) if deducted_credits < credits_required: logger.warning( diff --git a/api/extensions/ext_login.py b/api/extensions/ext_login.py index f6496c70a78..6515b22eb36 100644 --- a/api/extensions/ext_login.py +++ b/api/extensions/ext_login.py @@ -84,7 +84,7 @@ def load_user_from_request(request_from_flask_login: Request) -> LoginUser | Non if not user_id: raise Unauthorized("Invalid Authorization token.") - logged_in_account = AccountService.load_logged_in_account(account_id=user_id, session=db.session) + logged_in_account = AccountService.load_logged_in_account(account_id=user_id, session=db.session()) return logged_in_account elif request.blueprint == "openapi": # Account-branch device-flow approval routes (approve / deny / @@ -103,7 +103,7 @@ def load_user_from_request(request_from_flask_login: Request) -> LoginUser | Non source = decoded.get("token_source") if source or not user_id: return None - return AccountService.load_logged_in_account(account_id=user_id, session=db.session) + return AccountService.load_logged_in_account(account_id=user_id, session=db.session()) elif request.blueprint == "web": app_code = request.headers.get(HEADER_NAME_APP_CODE) webapp_token = extract_webapp_passport(app_code, request) if app_code else None diff --git a/api/services/account_service.py b/api/services/account_service.py index 1b9fd724a71..b5439467a23 100644 --- a/api/services/account_service.py +++ b/api/services/account_service.py @@ -16,7 +16,7 @@ from typing import Any, NotRequired, TypedDict, cast from pydantic import BaseModel, TypeAdapter, ValidationError from sqlalchemy import Row, delete, func, select, update -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from werkzeug.exceptions import Unauthorized from configs import dify_config @@ -188,12 +188,12 @@ class AccountService: raise ValueError(f"Builtin RBAC role not found for {role.value} in tenant {tenant_id}") @staticmethod - def get_workspace_permission_keys(tenant_id: str, account_id: str) -> set[str]: - permissions = RBACService.MyPermissions.get(tenant_id, account_id) + def get_workspace_permission_keys(tenant_id: str, account_id: str, *, session: Session) -> set[str]: + permissions = RBACService.MyPermissions.get(tenant_id, account_id, session=session) return set(getattr(getattr(permissions, "workspace", None), "permission_keys", []) or []) @staticmethod - def get_rbac_workspace_owner_account_id(tenant_id: str, actor_account_id: str) -> str: + def get_rbac_workspace_owner_account_id(tenant_id: str, actor_account_id: str, *, session: Session) -> str: """Return the account id bound to the workspace owner RBAC role.""" owner_role_id = AccountService._resolve_legacy_role_id( tenant_id=tenant_id, @@ -211,11 +211,14 @@ class AccountService: return owner_members[0].account_id @staticmethod - def is_rbac_workspace_owner(tenant_id: str, actor_account_id: str, member_account_id: str) -> bool: + def is_rbac_workspace_owner( + tenant_id: str, actor_account_id: str, member_account_id: str, *, session: Session + ) -> bool: roles = RBACService.MemberRoles.get( tenant_id=tenant_id, account_id=actor_account_id, member_account_id=member_account_id, + session=session, ).roles return any( role.is_builtin and role.category == "global_system_default" and role.role_tag == "owner" for role in roles @@ -246,7 +249,7 @@ class AccountService: ) @staticmethod - def _refresh_account_last_active(account: Account, session: scoped_session | Session) -> None: + def _refresh_account_last_active(account: Account, session: Session) -> None: now = naive_utc_now() refresh_before = now - ACCOUNT_LAST_ACTIVE_REFRESH_INTERVAL @@ -276,7 +279,7 @@ class AccountService: redis_client.delete(AccountService._get_account_refresh_token_key(account_id)) @staticmethod - def get_account_by_email(session: Session | scoped_session, email: str) -> Account | None: + def get_account_by_email(email: str, *, session: Session) -> Account | None: """Plain ``Account`` getter keyed by email. Case-sensitive — use :meth:`has_active_account_with_email` for the case-insensitive existence check that backs the SSO collision rule. @@ -284,7 +287,7 @@ class AccountService: return session.execute(select(Account).where(Account.email == email)).scalar_one_or_none() @staticmethod - def has_active_account_with_email(session: Session | scoped_session, email: str) -> bool: + def has_active_account_with_email(email: str, *, session: Session) -> bool: if not email: return False normalized = email.strip().lower() @@ -299,7 +302,7 @@ class AccountService: return row is not None @staticmethod - def get_account_by_id(session: Session | scoped_session, account_id: str) -> Account | None: + def get_account_by_id(account_id: str, *, session: Session) -> Account | None: """Plain ``Account`` getter — no banned check, no tenant rotation, no ``last_active_at`` write. Use this from read-only identity endpoints (``/openapi/v1/account``) where ``load_user``'s @@ -311,7 +314,7 @@ class AccountService: return session.get(Account, account_id) @staticmethod - def load_user(user_id: str, session: scoped_session | Session) -> None | Account: + def load_user(user_id: str, session: Session) -> None | Account: account = session.get(Account, user_id) if not account: return None @@ -363,9 +366,7 @@ class AccountService: return token @staticmethod - def authenticate( - email: str, password: str, invite_token: str | None = None, *, session: scoped_session | Session - ) -> Account: + def authenticate(email: str, password: str, invite_token: str | None = None, *, session: Session) -> Account: """authenticate account with email and password""" account = session.scalar(select(Account).where(Account.email == email).limit(1)) @@ -396,9 +397,7 @@ class AccountService: return account @staticmethod - def update_account_password( - account: Account, password: str, new_password: str, *, session: scoped_session | Session - ): + def update_account_password(account: Account, password: str, new_password: str, *, session: Session): """update account password""" if account.password and not compare_password(password, account.password, account.password_salt): raise CurrentPasswordIncorrectError("Current password is incorrect.") @@ -429,7 +428,7 @@ class AccountService: is_setup: bool | None = False, timezone: str | None = None, *, - session: scoped_session | Session, + session: Session, ) -> Account: """Create an account, preferring explicit user timezone over language-derived defaults.""" if not FeatureService.get_system_features().is_allow_register and not is_setup: @@ -487,7 +486,7 @@ class AccountService: password: str | None = None, timezone: str | None = None, *, - session: scoped_session | Session, + session: Session, ) -> Account: """Create an account and owner workspace.""" account = AccountService.create_account( @@ -544,12 +543,12 @@ class AccountService: return True @staticmethod - def delete_account(account: Account): + def delete_account(account: Account, *, session: Session): """Delete account. This method only adds a task to the queue for deletion.""" # Queue account deletion sync tasks for all workspaces BEFORE account deletion (enterprise only) from services.enterprise.account_deletion_sync import sync_account_deletion - sync_success = sync_account_deletion(account_id=account.id, source="account_deleted") + sync_success = sync_account_deletion(account_id=account.id, source="account_deleted", session=session) if not sync_success: logger.warning( "Enterprise account deletion sync failed for account %s; proceeding with local deletion.", @@ -560,7 +559,7 @@ class AccountService: delete_account_task.delay(account.id) @staticmethod - def link_account_integrate(provider: str, open_id: str, account: Account, *, session: scoped_session | Session): + def link_account_integrate(provider: str, open_id: str, account: Account, *, session: Session): """Link account integrate""" try: # Query whether there is an existing binding record for the same provider @@ -589,13 +588,13 @@ class AccountService: raise LinkAccountIntegrateError("Failed to link account.") from e @staticmethod - def close_account(account: Account, *, session: scoped_session | Session): + def close_account(account: Account, *, session: Session): """Close account""" account.status = AccountStatus.CLOSED session.commit() @staticmethod - def update_account(account: Account, *, session: scoped_session | Session, **kwargs): + def update_account(account: Account, *, session: Session, **kwargs): """Update account fields""" account = session.merge(account) for field, value in kwargs.items(): @@ -608,7 +607,7 @@ class AccountService: return account @staticmethod - def update_account_email(account: Account, email: str, session: scoped_session | Session) -> Account: + def update_account_email(account: Account, email: str, session: Session) -> Account: """Update account email""" account.email = email account_integrate = session.scalar( @@ -621,7 +620,7 @@ class AccountService: return account @staticmethod - def update_login_info(account: Account, session: scoped_session | Session, *, ip_address: str): + def update_login_info(account: Account, session: Session, *, ip_address: str): """Update last login time and ip""" account.last_login_at = naive_utc_now() account.last_login_ip = ip_address @@ -629,7 +628,7 @@ class AccountService: session.commit() @staticmethod - def login(account: Account, *, session: scoped_session | Session, ip_address: str | None = None) -> TokenPair: + def login(account: Account, *, session: Session, ip_address: str | None = None) -> TokenPair: if ip_address: AccountService.update_login_info(account=account, session=session, ip_address=ip_address) @@ -652,7 +651,7 @@ class AccountService: AccountService._delete_refresh_token(refresh_token.decode("utf-8"), account.id) @staticmethod - def refresh_token(refresh_token: str, *, session: scoped_session | Session) -> TokenPair: + def refresh_token(refresh_token: str, *, session: Session) -> TokenPair: # Verify the refresh token account_id = redis_client.get(AccountService._get_refresh_token_key(refresh_token)) if not account_id: @@ -673,7 +672,7 @@ class AccountService: return TokenPair(access_token=new_access_token, refresh_token=new_refresh_token, csrf_token=csrf_token) @staticmethod - def load_logged_in_account(*, account_id: str, session: scoped_session | Session): + def load_logged_in_account(*, account_id: str, session: Session): return AccountService.load_user(account_id, session) @classmethod @@ -1004,7 +1003,7 @@ class AccountService: return token @staticmethod - def get_account_by_email_with_case_fallback(session: Session | scoped_session, email: str) -> Account | None: + def get_account_by_email_with_case_fallback(email: str, *, session: Session) -> Account | None: """ Retrieve an account by email and fall back to the lowercase email if the original lookup fails. @@ -1026,7 +1025,7 @@ class AccountService: TokenManager.revoke_token(token, "email_code_login") @classmethod - def get_user_through_email(cls, email: str, *, session: scoped_session | Session): + def get_user_through_email(cls, email: str, *, session: Session): if dify_config.BILLING_ENABLED and BillingService.is_email_in_freeze(email): raise AccountRegisterError( description=( @@ -1234,7 +1233,7 @@ class AccountService: return False @staticmethod - def check_email_unique(email: str, *, session: scoped_session | Session) -> bool: + def check_email_unique(email: str, *, session: Session) -> bool: return session.scalar(select(Account).where(Account.email == email).limit(1)) is None @@ -1245,7 +1244,7 @@ class TenantService: is_setup: bool | None = False, is_from_dashboard: bool | None = False, *, - session: scoped_session | Session, + session: Session, ) -> Tenant: """Create tenant""" if ( @@ -1279,13 +1278,13 @@ class TenantService: from services.credit_pool_service import CreditPoolService - CreditPoolService.create_default_pool(tenant.id) + CreditPoolService.create_default_pool(tenant.id, session=session) return tenant @staticmethod def create_owner_tenant_if_not_exist( - account: Account, name: str | None = None, is_setup: bool | None = False, *, session: scoped_session | Session + account: Account, name: str | None = None, is_setup: bool | None = False, *, session: Session ): """Check if user have a workspace or not""" available_ta = session.scalar( @@ -1318,6 +1317,7 @@ class TenantService: account_id=account.id, member_account_id=account.id, role_ids=[owner_role_id], + session=session, ) account.current_tenant = tenant session.commit() @@ -1325,7 +1325,7 @@ class TenantService: @staticmethod def create_tenant_member( - tenant: Tenant, account: Account, session: scoped_session | Session, role: str = "normal" + tenant: Tenant, account: Account, session: Session, role: str = "normal" ) -> TenantAccountJoin: """Create tenant member""" if role == TenantAccountRole.OWNER: @@ -1350,7 +1350,7 @@ class TenantService: return ta @staticmethod - def get_join_tenants(account: Account, *, session: scoped_session | Session) -> list[Tenant]: + def get_join_tenants(account: Account, *, session: Session) -> list[Tenant]: """Get account join tenants""" return list( session.scalars( @@ -1361,10 +1361,7 @@ class TenantService: ) @staticmethod - def get_account_memberships( - session: Session | scoped_session, - account_id: str, - ) -> list[Row[tuple[TenantAccountJoin, Tenant]]]: + def get_account_memberships(account_id: str, *, session: Session) -> list[Row[tuple[TenantAccountJoin, Tenant]]]: """Return ``(TenantAccountJoin, Tenant)`` rows for every workspace the account belongs to. Unlike :meth:`get_join_tenants` this keeps the join row so callers can read ``role``/``current`` alongside the @@ -1385,10 +1382,7 @@ class TenantService: ) @staticmethod - def get_workspaces_for_account( - session: Session | scoped_session, - account_id: str, - ) -> list[Row[tuple[Tenant, TenantAccountJoin]]]: + def get_workspaces_for_account(account_id: str, *, session: Session) -> list[Row[tuple[Tenant, TenantAccountJoin]]]: """``(Tenant, TenantAccountJoin)`` rows for every workspace the account belongs to, ordered by ``Tenant.created_at`` ASC — the canonical ordering for ``/openapi/v1/workspaces``. @@ -1407,11 +1401,7 @@ class TenantService: ) @staticmethod - def account_belongs_to_tenant( - session: Session | scoped_session, - account_id: uuid.UUID | str | None, - tenant_id: str, - ) -> bool: + def account_belongs_to_tenant(account_id: uuid.UUID | str | None, tenant_id: str, *, session: Session) -> bool: """Existence check for ``TenantAccountJoin(account_id, tenant_id)``. Backs the CE-deployment membership fallback in ``controllers.openapi.auth.strategies.MembershipStrategy``. @@ -1431,9 +1421,7 @@ class TenantService: @staticmethod def get_account_role_in_tenant( - session: Session | scoped_session, - account_id: uuid.UUID | str | None, - tenant_id: str, + account_id: uuid.UUID | str | None, tenant_id: str, *, session: Session ) -> TenantAccountRole | None: """Return the caller's role in ``tenant_id``, or ``None`` if not a member. @@ -1459,7 +1447,7 @@ class TenantService: return TenantAccountRole(role) if role is not None else None @staticmethod - def get_tenant_by_id(session: Session | scoped_session, tenant_id: str) -> Tenant | None: + def get_tenant_by_id(tenant_id: str, *, session: Session) -> Tenant | None: """Plain ``session.get(Tenant, tenant_id)`` — no status filter. Callers map ``status == ARCHIVE`` to their own error code (the openapi auth pipeline raises 403 ``workspace unavailable``). @@ -1467,10 +1455,7 @@ class TenantService: return session.get(Tenant, tenant_id) @staticmethod - def get_tenants_by_ids( - session: Session | scoped_session, - tenant_ids: list[str], - ) -> list[Tenant]: + def get_tenants_by_ids(tenant_ids: list[str], *, session: Session) -> list[Tenant]: """Bulk ``Tenant`` fetch by primary-key list. Order is unspecified — callers index by ``tenant.id`` (e.g. for cross-tenant denorm in ``/openapi/v1/permitted-external-apps``). @@ -1483,7 +1468,7 @@ class TenantService: return list(session.execute(select(Tenant).where(Tenant.id.in_(tenant_ids))).scalars().all()) @staticmethod - def get_tenant_name(session: Session | scoped_session, tenant_id: str) -> str | None: + def get_tenant_name(tenant_id: str, *, session: Session) -> str | None: """Single-column tenant name read. Used by openapi list endpoints to denormalize ``workspace_name`` onto each row without dragging the full ``Tenant`` ORM entity through. @@ -1492,9 +1477,7 @@ class TenantService: @staticmethod def find_workspace_for_account( - session: Session | scoped_session, - account_id: str, - workspace_id: str, + account_id: str, workspace_id: str, *, session: Session ) -> Row[tuple[Tenant, TenantAccountJoin]] | None: """Single ``(Tenant, TenantAccountJoin)`` row scoped to the account's membership in ``workspace_id``. ``None`` on non-member @@ -1511,7 +1494,7 @@ class TenantService: ).first() @staticmethod - def get_current_tenant_by_account(account: Account, *, session: scoped_session | Session): + def get_current_tenant_by_account(account: Account, *, session: Session): """Get tenant by account and add the role""" tenant = account.current_tenant if not tenant: @@ -1529,7 +1512,7 @@ class TenantService: return tenant @staticmethod - def switch_tenant(account: Account, tenant_id: str | None = None, *, session: scoped_session | Session): + def switch_tenant(account: Account, tenant_id: str | None = None, *, session: Session): """Switch the current workspace for the account""" # Ensure tenant_id is provided @@ -1562,7 +1545,7 @@ class TenantService: session.commit() @staticmethod - def get_tenant_members(tenant: Tenant, *, session: scoped_session | Session) -> list[Account]: + def get_tenant_members(tenant: Tenant, *, session: Session) -> list[Account]: """Get tenant members""" stmt = ( select(Account, TenantAccountJoin.role) @@ -1581,7 +1564,7 @@ class TenantService: return updated_accounts @staticmethod - def get_dataset_operator_members(tenant: Tenant, *, session: scoped_session | Session) -> list[Account]: + def get_dataset_operator_members(tenant: Tenant, *, session: Session) -> list[Account]: """Get dataset admin members""" stmt = ( select(Account, TenantAccountJoin.role) @@ -1601,7 +1584,7 @@ class TenantService: return updated_accounts @staticmethod - def has_roles(tenant: Tenant, roles: list[TenantAccountRole], *, session: scoped_session | Session) -> bool: + def has_roles(tenant: Tenant, roles: list[TenantAccountRole], *, session: Session) -> bool: """Check if user has any of the given roles for a tenant""" if not all(isinstance(role, TenantAccountRole) for role in roles): raise ValueError("all roles must be TenantAccountRole") @@ -1619,9 +1602,7 @@ class TenantService: ) @staticmethod - def get_user_role( - account: Account, tenant: Tenant, *, session: scoped_session | Session - ) -> TenantAccountRole | None: + def get_user_role(account: Account, tenant: Tenant, *, session: Session) -> TenantAccountRole | None: """Get the role of the current account for a given tenant""" join = session.scalar( select(TenantAccountJoin) @@ -1631,13 +1612,13 @@ class TenantService: return TenantAccountRole(join.role) if join else None @staticmethod - def get_tenant_count(*, session: scoped_session | Session) -> int: + def get_tenant_count(*, session: Session) -> int: """Get tenant count""" return cast(int, session.scalar(select(func.count(Tenant.id)))) @staticmethod def check_member_permission( - tenant: Tenant, operator: Account, member: Account | None, action: str, *, session: scoped_session | Session + tenant: Tenant, operator: Account, member: Account | None, action: str, *, session: Session ): """Check member permission""" if action not in {"add", "remove", "update"}: @@ -1651,6 +1632,7 @@ class TenantService: workspace_permission_keys = AccountService.get_workspace_permission_keys( str(tenant.id), str(operator.id), + session=session, ) required_permission_key = ( "workspace.member.manage" if action in {"add", "remove"} else "workspace.role.manage" @@ -1661,7 +1643,9 @@ class TenantService: if ( action == "remove" and member - and AccountService.is_rbac_workspace_owner(str(tenant.id), str(operator.id), str(member.id)) + and AccountService.is_rbac_workspace_owner( + str(tenant.id), str(operator.id), str(member.id), session=session + ) ): raise NoPermissionError(f"No permission to {action} member.") return @@ -1691,9 +1675,7 @@ class TenantService: raise NoPermissionError(f"No permission to {action} member.") @staticmethod - def remove_member_from_tenant( - tenant: Tenant, account: Account, operator: Account, *, session: scoped_session | Session - ): + def remove_member_from_tenant(tenant: Tenant, account: Account, operator: Account, *, session: Session): """Remove member from tenant. Apps and datasets maintained by the removed member are reassigned to @@ -1722,7 +1704,9 @@ class TenantService: owner_id: str | None if dify_config.RBAC_ENABLED: - owner_id = AccountService.get_rbac_workspace_owner_account_id(str(tenant.id), str(operator.id)) + owner_id = AccountService.get_rbac_workspace_owner_account_id( + str(tenant.id), str(operator.id), session=session + ) else: owner_id = session.scalar( select(TenantAccountJoin.account_id) @@ -1796,9 +1780,7 @@ class TenantService: RBACService.MemberRoles.delete_rbac_bindings(tenant_id=tenant.id, account_id=account_id) @staticmethod - def update_member_role( - tenant: Tenant, member: Account, new_role: str, operator: Account, *, session: scoped_session | Session - ): + def update_member_role(tenant: Tenant, member: Account, new_role: str, operator: Account, *, session: Session): """Update member role""" TenantService.check_member_permission(tenant, operator, member, "update", session=session) new_tenant_role = TenantAccountRole(new_role) @@ -1841,6 +1823,7 @@ class TenantService: account_id=operator.id, member_account_id=str(current_owner_join.account_id), role_ids=[admin_role_id], + session=session, ) # Update the role of the target member @@ -1855,6 +1838,7 @@ class TenantService: account_id=operator.id, member_account_id=member.id, role_ids=[resolved_role_id], + session=session, ) else: target_member_join.role = new_tenant_role @@ -1867,11 +1851,11 @@ class TenantService: return tenant.custom_config_dict @staticmethod - def is_owner(account: Account, tenant: Tenant, *, session: scoped_session | Session) -> bool: + def is_owner(account: Account, tenant: Tenant, *, session: Session) -> bool: return TenantService.get_user_role(account, tenant, session=session) == TenantAccountRole.OWNER @staticmethod - def is_member(account: Account, tenant: Tenant, *, session: scoped_session | Session) -> bool: + def is_member(account: Account, tenant: Tenant, *, session: Session) -> bool: """Check if the account is a member of the tenant""" return TenantService.get_user_role(account, tenant, session=session) is not None @@ -1890,7 +1874,7 @@ class RegisterService: ip_address: str, language: str | None, *, - session: scoped_session | Session, + session: Session, ): """ Setup dify @@ -1943,7 +1927,7 @@ class RegisterService: create_workspace_required: bool | None = True, timezone: str | None = None, *, - session: scoped_session | Session, + session: Session, ) -> Account: """Register account""" session.begin_nested() @@ -2005,7 +1989,7 @@ class RegisterService: role: str = "normal", inviter: Account | None = None, *, - session: scoped_session | Session, + session: Session, ) -> str: if not inviter: raise ValueError("Inviter is required") @@ -2019,7 +2003,7 @@ class RegisterService: check_workspace_member_invite_permission(tenant.id) - account = AccountService.get_account_by_email_with_case_fallback(db.session, email) + account = AccountService.get_account_by_email_with_case_fallback(email, session=session) requires_setup = False if not account: @@ -2057,6 +2041,7 @@ class RegisterService: account_id=inviter.id, member_account_id=account.id, role_ids=[role], + session=session, ) if ta or dify_config.RBAC_ENABLED: raise AccountAlreadyInTenantError("Account already in tenant.") @@ -2068,6 +2053,7 @@ class RegisterService: account_id=inviter.id, member_account_id=account.id, role_ids=[role], + session=session, ) token = cls.generate_invite_token(tenant, account, role, requires_setup=requires_setup) @@ -2116,7 +2102,7 @@ class RegisterService: @classmethod def get_invitation_if_token_valid( - cls, workspace_id: str | None, email: str | None, token: str, *, session: scoped_session | Session + cls, workspace_id: str | None, email: str | None, token: str, *, session: Session ) -> InvitationDetailDict | None: invitation_data = cls.get_invitation_by_token(token, workspace_id, email) if not invitation_data: @@ -2169,7 +2155,7 @@ class RegisterService: @classmethod def get_invitation_with_case_fallback( - cls, workspace_id: str | None, email: str | None, token: str, *, session: scoped_session | Session + cls, workspace_id: str | None, email: str | None, token: str, *, session: Session ) -> InvitationDetailDict | None: invitation = cls.get_invitation_if_token_valid(workspace_id, email, token, session=session) if invitation or not email or email == email.lower(): diff --git a/api/services/agent/composer_service.py b/api/services/agent/composer_service.py index 28b86916b1d..4839012dd3a 100644 --- a/api/services/agent/composer_service.py +++ b/api/services/agent/composer_service.py @@ -4,6 +4,7 @@ from typing import Any from sqlalchemy import func, or_, select from sqlalchemy.exc import IntegrityError +from sqlalchemy.orm import Session from sqlalchemy.sql.elements import ColumnElement from extensions.ext_database import db @@ -104,23 +105,35 @@ def _agent_soul_config_json(agent_soul: AgentSoulConfig | dict[str, Any]) -> dic class AgentComposerService: @classmethod def load_workflow_composer( - cls, *, tenant_id: str, app_id: str, node_id: str, account_id: str | None = None, snapshot_id: str | None = None + cls, + *, + tenant_id: str, + app_id: str, + node_id: str, + account_id: str | None = None, + snapshot_id: str | None = None, + session: Session, ) -> dict[str, Any]: - workflow = cls._get_draft_workflow(tenant_id=tenant_id, app_id=app_id) - binding = cls._get_workflow_binding(tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id) + workflow = cls._get_draft_workflow(tenant_id=tenant_id, app_id=app_id, session=session) + binding = cls._get_workflow_binding( + tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id, session=session + ) if not binding: if snapshot_id: raise AgentVersionNotFoundError() return cls._empty_workflow_state(app_id=app_id, workflow_id=workflow.id, node_id=node_id) - agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=binding.agent_id) + agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=binding.agent_id, session=session) version = cls._workflow_composer_version( tenant_id=tenant_id, binding=binding, agent=agent, snapshot_id=snapshot_id, + session=session, + ) + return cls._serialize_workflow_state( + binding=binding, agent=agent, version=version, account_id=account_id, session=session ) - return cls._serialize_workflow_state(binding=binding, agent=agent, version=version, account_id=account_id) @classmethod def _workflow_composer_version( @@ -130,6 +143,7 @@ class AgentComposerService: binding: WorkflowAgentNodeBinding, agent: Agent | None, snapshot_id: str | None, + session: Session, ) -> AgentConfigSnapshot | None: if snapshot_id: if agent is None: @@ -147,7 +161,7 @@ class AgentComposerService: raise AgentVersionNotFoundError() else: raise AgentVersionNotFoundError() - return cls._require_version(tenant_id=tenant_id, agent_id=agent.id, version_id=snapshot_id) + return cls._require_version(tenant_id=tenant_id, agent_id=agent.id, version_id=snapshot_id, session=session) version_id = ( agent.active_config_snapshot_id @@ -158,11 +172,19 @@ class AgentComposerService: tenant_id=tenant_id, agent_id=agent.id if agent else None, version_id=version_id, + session=session, ) @classmethod def save_workflow_composer( - cls, *, tenant_id: str, app_id: str, node_id: str, account_id: str, payload: ComposerSavePayload + cls, + *, + tenant_id: str, + app_id: str, + node_id: str, + account_id: str, + payload: ComposerSavePayload, + session: Session, ) -> dict[str, Any]: if payload.variant != ComposerVariant.WORKFLOW: raise ValueError("Workflow composer endpoint only accepts workflow variant") @@ -171,8 +193,10 @@ class AgentComposerService: _validate_composer_payload_for_strategy(payload) if payload.save_strategy in _PUBLISH_SAVE_STRATEGIES: cls.validate_knowledge_datasets(tenant_id=tenant_id, agent_soul=payload.agent_soul) - workflow = cls._get_draft_workflow(tenant_id=tenant_id, app_id=app_id) - binding = cls._get_workflow_binding(tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id) + workflow = cls._get_draft_workflow(tenant_id=tenant_id, app_id=app_id, session=session) + binding = cls._get_workflow_binding( + tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id, session=session + ) match payload.save_strategy: case ComposerSaveStrategy.NODE_JOB_ONLY: @@ -184,14 +208,15 @@ class AgentComposerService: account_id=account_id, binding=binding, payload=payload, + session=session, ) case ComposerSaveStrategy.SAVE_TO_CURRENT_VERSION: binding = cls._save_to_current_version( - tenant_id=tenant_id, account_id=account_id, binding=binding, payload=payload + tenant_id=tenant_id, account_id=account_id, binding=binding, payload=payload, session=session ) case ComposerSaveStrategy.SAVE_AS_NEW_VERSION: binding = cls._save_as_new_version( - tenant_id=tenant_id, account_id=account_id, binding=binding, payload=payload + tenant_id=tenant_id, account_id=account_id, binding=binding, payload=payload, session=session ) case ComposerSaveStrategy.SAVE_AS_NEW_AGENT: binding = cls._save_as_new_agent( @@ -202,14 +227,15 @@ class AgentComposerService: account_id=account_id, binding=binding, payload=payload, + session=session, ) case ComposerSaveStrategy.SAVE_TO_ROSTER: binding = cls._save_to_roster( - tenant_id=tenant_id, account_id=account_id, binding=binding, payload=payload + tenant_id=tenant_id, account_id=account_id, binding=binding, payload=payload, session=session ) - db.session.commit() - agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=binding.agent_id) + session.commit() + agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=binding.agent_id, session=session) version_id = ( agent.active_config_snapshot_id if agent and binding.binding_type == WorkflowAgentBindingType.ROSTER_AGENT @@ -219,12 +245,16 @@ class AgentComposerService: tenant_id=tenant_id, agent_id=agent.id if agent else None, version_id=version_id, + session=session, + ) + state = cls._serialize_workflow_state( + binding=binding, agent=agent, version=version, account_id=account_id, session=session ) - state = cls._serialize_workflow_state(binding=binding, agent=agent, version=version, account_id=account_id) state["validation"] = cls.collect_validation_findings( tenant_id=tenant_id, payload=payload, agent_id=binding.agent_id, + session=session, ) return state @@ -239,33 +269,38 @@ class AgentComposerService: source_agent_id: str, source_snapshot_id: str | None = None, idempotency_key: str | None = None, + session: Session, ) -> dict[str, Any]: - workflow = cls._get_draft_workflow(tenant_id=tenant_id, app_id=app_id) + workflow = cls._get_draft_workflow(tenant_id=tenant_id, app_id=app_id, session=session) binding = cls._require_binding( - cls._get_workflow_binding(tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id) + cls._get_workflow_binding(tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id, session=session) ) if binding.binding_type == WorkflowAgentBindingType.INLINE_AGENT and idempotency_key: - agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=binding.agent_id) + agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=binding.agent_id, session=session) version = cls._get_version_if_present( tenant_id=tenant_id, agent_id=agent.id if agent else None, version_id=binding.current_snapshot_id, + session=session, + ) + return cls._serialize_workflow_state( + binding=binding, agent=agent, version=version, account_id=account_id, session=session ) - return cls._serialize_workflow_state(binding=binding, agent=agent, version=version, account_id=account_id) if binding.binding_type != WorkflowAgentBindingType.ROSTER_AGENT: raise InvalidComposerConfigError("Workflow agent node must be bound to a roster agent.") if binding.agent_id != source_agent_id: raise InvalidComposerConfigError("Source agent does not match the current workflow node binding.") - source_agent = cls._require_agent(tenant_id=tenant_id, agent_id=source_agent_id) + source_agent = cls._require_agent(tenant_id=tenant_id, agent_id=source_agent_id, session=session) if source_agent.scope != AgentScope.ROSTER or source_agent.status != AgentStatus.ACTIVE: raise InvalidComposerConfigError("Source agent must be an active roster agent.") source_version = cls._require_version( tenant_id=tenant_id, agent_id=source_agent.id, version_id=source_agent.active_config_snapshot_id, + session=session, ) if source_snapshot_id and source_snapshot_id != source_version.id: raise AgentVersionConflictError() @@ -284,6 +319,7 @@ class AgentComposerService: icon_type=source_agent.icon_type, icon=source_agent.icon, icon_background=source_agent.icon_background, + session=session, ) cls._copy_agent_drive_rows( tenant_id=tenant_id, @@ -292,45 +328,48 @@ class AgentComposerService: account_id=account_id, agent_soul=agent_soul, node_job=WorkflowNodeJobConfig.model_validate(binding.node_job_config_dict), + session=session, ) binding.binding_type = WorkflowAgentBindingType.INLINE_AGENT binding.agent_id = inline_agent.id binding.current_snapshot_id = inline_agent.active_config_snapshot_id binding.updated_by = account_id - db.session.flush() - db.session.commit() + session.flush() + session.commit() version = cls._require_version( tenant_id=tenant_id, agent_id=inline_agent.id, version_id=inline_agent.active_config_snapshot_id, + session=session, ) return cls._serialize_workflow_state( - binding=binding, agent=inline_agent, version=version, account_id=account_id + binding=binding, agent=inline_agent, version=version, account_id=account_id, session=session ) @classmethod - def load_agent_app_composer(cls, *, tenant_id: str, app_id: str) -> dict[str, Any]: - agent = cls._require_agent_app_agent(tenant_id=tenant_id, app_id=app_id) - return cls._load_agent_composer_for_agent(tenant_id=tenant_id, agent=agent) + def load_agent_app_composer(cls, *, tenant_id: str, app_id: str, session: Session) -> dict[str, Any]: + agent = cls._require_agent_app_agent(tenant_id=tenant_id, app_id=app_id, session=session) + return cls._load_agent_composer_for_agent(tenant_id=tenant_id, agent=agent, session=session) @classmethod - def load_agent_composer(cls, *, tenant_id: str, agent_id: str) -> dict[str, Any]: - agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id) - return cls._load_agent_composer_for_agent(tenant_id=tenant_id, agent=agent) + def load_agent_composer(cls, *, tenant_id: str, agent_id: str, session: Session) -> dict[str, Any]: + agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id, session=session) + return cls._load_agent_composer_for_agent(tenant_id=tenant_id, agent=agent, session=session) @classmethod - def _load_agent_composer_for_agent(cls, *, tenant_id: str, agent: Agent) -> dict[str, Any]: + def _load_agent_composer_for_agent(cls, *, tenant_id: str, agent: Agent, session: Session) -> dict[str, Any]: draft = cls._get_or_create_agent_draft( tenant_id=tenant_id, agent=agent, draft_type=AgentConfigDraftType.DRAFT, account_id=None, created_by=agent.updated_by or agent.created_by, + session=session, ) version = cls._get_version_if_present( - tenant_id=tenant_id, agent_id=agent.id, version_id=agent.active_config_snapshot_id + tenant_id=tenant_id, agent_id=agent.id, version_id=agent.active_config_snapshot_id, session=session ) return { "variant": ComposerVariant.AGENT_APP.value, @@ -347,7 +386,13 @@ class AgentComposerService: @classmethod def save_agent_app_composer( - cls, *, tenant_id: str, app_id: str, account_id: str, payload: ComposerSavePayload + cls, + *, + tenant_id: str, + app_id: str, + account_id: str, + payload: ComposerSavePayload, + session: Session, ) -> dict[str, Any]: if payload.variant != ComposerVariant.AGENT_APP: raise ValueError("Agent App composer endpoint only accepts agent_app variant") @@ -360,7 +405,7 @@ class AgentComposerService: _backfill_cli_tool_ids(payload.agent_soul) _validate_composer_payload_for_strategy(payload) - agent = cls._get_agent_app_agent(tenant_id=tenant_id, app_id=app_id) + agent = cls._get_agent_app_agent(tenant_id=tenant_id, app_id=app_id, session=session) if not agent: agent = Agent( tenant_id=tenant_id, @@ -375,22 +420,29 @@ class AgentComposerService: created_by=account_id, updated_by=account_id, ) - db.session.add(agent) + session.add(agent) try: - db.session.flush() + session.flush() except IntegrityError as exc: - db.session.rollback() + session.rollback() raise AgentNameConflictError() from exc return cls._save_agent_composer_for_agent( tenant_id=tenant_id, agent=agent, account_id=account_id, payload=payload, + session=session, ) @classmethod def save_agent_composer( - cls, *, tenant_id: str, agent_id: str, account_id: str, payload: ComposerSavePayload + cls, + *, + tenant_id: str, + agent_id: str, + account_id: str, + payload: ComposerSavePayload, + session: Session, ) -> dict[str, Any]: if payload.variant != ComposerVariant.AGENT_APP: raise ValueError("Agent composer endpoint only accepts agent_app variant") @@ -402,17 +454,24 @@ class AgentComposerService: raise ValueError("agent_soul is required") _backfill_cli_tool_ids(payload.agent_soul) _validate_composer_payload_for_strategy(payload) - agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id) + agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id, session=session) return cls._save_agent_composer_for_agent( tenant_id=tenant_id, agent=agent, account_id=account_id, payload=payload, + session=session, ) @classmethod def _save_agent_composer_for_agent( - cls, *, tenant_id: str, agent: Agent, account_id: str, payload: ComposerSavePayload + cls, + *, + tenant_id: str, + agent: Agent, + account_id: str, + payload: ComposerSavePayload, + session: Session, ) -> dict[str, Any]: if payload.agent_soul is None: raise ValueError("agent_soul is required") @@ -423,20 +482,23 @@ class AgentComposerService: account_id=None, agent_soul=payload.agent_soul, account_id_for_audit=account_id, + session=session, ) agent.updated_by = account_id agent.active_config_is_published = cls._agent_soul_matches_active_config( tenant_id=tenant_id, agent=agent, agent_soul=payload.agent_soul, + session=session, ) - db.session.commit() - state = cls.load_agent_composer(tenant_id=tenant_id, agent_id=agent.id) + session.commit() + state = cls.load_agent_composer(tenant_id=tenant_id, agent_id=agent.id, session=session) state["validation"] = cls.collect_validation_findings( tenant_id=tenant_id, payload=payload, agent_id=agent.id, + session=session, ) return state @@ -447,6 +509,7 @@ class AgentComposerService: tenant_id: str, agent: Agent, agent_soul: AgentSoulConfig, + session: Session, ) -> bool: if not agent.active_config_snapshot_id: return False @@ -455,6 +518,7 @@ class AgentComposerService: tenant_id=tenant_id, agent_id=agent.id, version_id=agent.active_config_snapshot_id, + session=session, ) if not active_version: return False @@ -490,9 +554,15 @@ class AgentComposerService: @classmethod def publish_agent_app_draft( - cls, *, tenant_id: str, agent_id: str, account_id: str, version_note: str | None = None + cls, + *, + tenant_id: str, + agent_id: str, + account_id: str, + version_note: str | None = None, + session: Session, ) -> dict[str, Any]: - agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id) + agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id, session=session) if agent.scope != AgentScope.ROSTER or agent.source != AgentSource.AGENT_APP: raise AgentNotFoundError() draft = cls._get_or_create_agent_draft( @@ -501,6 +571,7 @@ class AgentComposerService: draft_type=AgentConfigDraftType.DRAFT, account_id=None, created_by=account_id, + session=session, ) agent_soul = AgentSoulConfig.model_validate(draft.config_snapshot_dict) ComposerConfigValidator.validate_publish_payload( @@ -522,6 +593,7 @@ class AgentComposerService: operation=AgentConfigRevisionOperation.PUBLISH_DRAFT, version_note=version_note, previous_snapshot_id=agent.active_config_snapshot_id, + session=session, ) agent.active_config_snapshot_id = version.id agent.active_config_has_model = agent_soul_has_model(agent_soul) @@ -529,7 +601,7 @@ class AgentComposerService: agent.updated_by = account_id draft.base_snapshot_id = version.id draft.updated_by = account_id - db.session.commit() + session.commit() return { "result": "success", "active_config_snapshot_id": version.id, @@ -539,21 +611,29 @@ class AgentComposerService: @classmethod def checkout_agent_app_build_draft( - cls, *, tenant_id: str, agent_id: str, account_id: str, force: bool = False + cls, + *, + tenant_id: str, + agent_id: str, + account_id: str, + force: bool = False, + session: Session, ) -> dict[str, Any]: - agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id) + agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id, session=session) normal_draft = cls._get_or_create_agent_draft( tenant_id=tenant_id, agent=agent, draft_type=AgentConfigDraftType.DRAFT, account_id=None, created_by=account_id, + session=session, ) build_draft = cls._get_agent_draft( tenant_id=tenant_id, agent_id=agent.id, draft_type=AgentConfigDraftType.DEBUG_BUILD, account_id=account_id, + session=session, ) if build_draft is not None and not force: return cls._serialize_build_draft_state(build_draft) @@ -566,20 +646,23 @@ class AgentComposerService: draft_owner_key=account_id, created_by=account_id, ) - db.session.add(build_draft) + session.add(build_draft) build_draft.base_snapshot_id = normal_draft.base_snapshot_id build_draft.config_snapshot = AgentSoulConfig.model_validate(normal_draft.config_snapshot_dict) build_draft.updated_by = account_id - db.session.commit() + session.commit() return cls._serialize_build_draft_state(build_draft) @classmethod - def load_agent_app_build_draft(cls, *, tenant_id: str, agent_id: str, account_id: str) -> dict[str, Any]: + def load_agent_app_build_draft( + cls, *, tenant_id: str, agent_id: str, account_id: str, session: Session + ) -> dict[str, Any]: build_draft = cls._get_agent_draft( tenant_id=tenant_id, agent_id=agent_id, draft_type=AgentConfigDraftType.DEBUG_BUILD, account_id=account_id, + session=session, ) if build_draft is None: raise AgentVersionNotFoundError() @@ -587,13 +670,19 @@ class AgentComposerService: @classmethod def save_agent_app_build_draft( - cls, *, tenant_id: str, agent_id: str, account_id: str, payload: ComposerSavePayload + cls, + *, + tenant_id: str, + agent_id: str, + account_id: str, + payload: ComposerSavePayload, + session: Session, ) -> dict[str, Any]: if payload.agent_soul is None: raise ValueError("agent_soul is required") _backfill_cli_tool_ids(payload.agent_soul) ComposerConfigValidator.validate_draft_save_payload(payload) - agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id) + agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id, session=session) build_draft = cls._save_agent_draft( tenant_id=tenant_id, agent=agent, @@ -601,18 +690,22 @@ class AgentComposerService: account_id=account_id, agent_soul=payload.agent_soul, account_id_for_audit=account_id, + session=session, ) - db.session.commit() + session.commit() return cls._serialize_build_draft_state(build_draft) @classmethod - def apply_agent_app_build_draft(cls, *, tenant_id: str, agent_id: str, account_id: str) -> dict[str, Any]: - agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id) + def apply_agent_app_build_draft( + cls, *, tenant_id: str, agent_id: str, account_id: str, session: Session + ) -> dict[str, Any]: + agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id, session=session) build_draft = cls._get_agent_draft( tenant_id=tenant_id, agent_id=agent.id, draft_type=AgentConfigDraftType.DEBUG_BUILD, account_id=account_id, + session=session, ) if build_draft is None: raise AgentVersionNotFoundError() @@ -625,28 +718,33 @@ class AgentComposerService: agent_soul=applied_agent_soul, account_id_for_audit=account_id, base_snapshot_id=build_draft.base_snapshot_id, + session=session, ) agent.active_config_is_published = cls._agent_soul_matches_active_config( tenant_id=tenant_id, agent=agent, agent_soul=applied_agent_soul, + session=session, ) agent.updated_by = account_id - db.session.delete(build_draft) - db.session.commit() + session.delete(build_draft) + session.commit() return {"result": "success", "draft": cls._serialize_draft(normal_draft)} @classmethod - def discard_agent_app_build_draft(cls, *, tenant_id: str, agent_id: str, account_id: str) -> dict[str, Any]: + def discard_agent_app_build_draft( + cls, *, tenant_id: str, agent_id: str, account_id: str, session: Session + ) -> dict[str, Any]: build_draft = cls._get_agent_draft( tenant_id=tenant_id, agent_id=agent_id, draft_type=AgentConfigDraftType.DEBUG_BUILD, account_id=account_id, + session=session, ) if build_draft is not None: - db.session.delete(build_draft) - db.session.commit() + session.delete(build_draft) + session.commit() return {"result": "success"} @classmethod @@ -656,6 +754,7 @@ class AgentComposerService: tenant_id: str, payload: ComposerSavePayload, agent_id: str | None = None, + session: Session, ) -> dict[str, Any]: """ENG-617 soft findings, with DB-backed dataset and drive mention checks.""" existing_knowledge_set_ids = ( @@ -673,6 +772,7 @@ class AgentComposerService: tenant_id=tenant_id, agent_id=agent_id, prompt=payload.agent_soul.prompt.system_prompt, + session=session, ) ) return findings @@ -696,9 +796,9 @@ class AgentComposerService: ) @classmethod - def resolve_bound_agent_id(cls, *, tenant_id: str, app_id: str) -> str | None: + def resolve_bound_agent_id(cls, *, tenant_id: str, app_id: str, session: Session) -> str | None: """The Agent App's bound roster agent id, if any (validate-endpoint context).""" - return db.session.scalar( + return session.scalar( select(Agent.id) .where( Agent.tenant_id == tenant_id, @@ -711,13 +811,17 @@ class AgentComposerService: ) @classmethod - def resolve_workflow_node_agent_id(cls, *, tenant_id: str, app_id: str, node_id: str) -> str | None: + def resolve_workflow_node_agent_id( + cls, *, tenant_id: str, app_id: str, node_id: str, session: Session + ) -> str | None: """The draft workflow node binding's agent id, if any (validate-endpoint context).""" try: - workflow = cls._get_draft_workflow(tenant_id=tenant_id, app_id=app_id) + workflow = cls._get_draft_workflow(tenant_id=tenant_id, app_id=app_id, session=session) except ValueError: return None - binding = cls._get_workflow_binding(tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id) + binding = cls._get_workflow_binding( + tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id, session=session + ) return binding.agent_id if binding else None @classmethod @@ -727,6 +831,7 @@ class AgentComposerService: tenant_id: str, agent_id: str, prompt: str, + session: Session, ) -> list[dict[str, str | None]]: """Soft warnings for missing drive-backed prompt mentions.""" from services.agent.prompt_mentions import MentionKind, parse_prompt_mentions @@ -744,7 +849,7 @@ class AgentComposerService: return [] existing_keys = set( - db.session.scalars( + session.scalars( select(AgentDriveFile.key).where( AgentDriveFile.tenant_id == tenant_id, AgentDriveFile.agent_id == agent_id, @@ -768,22 +873,32 @@ class AgentComposerService: return findings @classmethod - def get_workflow_candidates(cls, *, tenant_id: str, app_id: str, node_id: str, user_id: str) -> dict[str, Any]: + def get_workflow_candidates( + cls, + *, + tenant_id: str, + app_id: str, + node_id: str, + user_id: str, + session: Session, + ) -> dict[str, Any]: """Slash-menu data source for the workflow Agent node composer (ENG-615).""" from services.agent.composer_candidates import previous_node_output_candidates, soul_candidates try: - workflow = cls._get_draft_workflow(tenant_id=tenant_id, app_id=app_id) + workflow = cls._get_draft_workflow(tenant_id=tenant_id, app_id=app_id, session=session) except ValueError: workflow = None node_job: WorkflowNodeJobConfig | None = None agent_soul: AgentSoulConfig | None = None if workflow is not None: - binding = cls._get_workflow_binding(tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id) + binding = cls._get_workflow_binding( + tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id, session=session + ) if binding is not None: node_job = cls._parse_node_job(binding) - agent_soul = cls._load_binding_soul(tenant_id=tenant_id, binding=binding) + agent_soul = cls._load_binding_soul(tenant_id=tenant_id, binding=binding, session=session) truncated = False previous_outputs: list[dict[str, Any]] = [] @@ -794,7 +909,7 @@ class AgentComposerService: graph=workflow.graph_dict, node_id=node_id, declared_outputs_loader=lambda nid: cls._binding_declared_outputs( - tenant_id=tenant_id, workflow_id=workflow.id, node_id=nid + tenant_id=tenant_id, workflow_id=workflow.id, node_id=nid, session=session ), draft_variables_loader=lambda nid: cls._draft_node_variables( session=draft_variable_session, app_id=app_id, node_id=nid, user_id=user_id @@ -829,11 +944,13 @@ class AgentComposerService: return response.model_dump(mode="json") @classmethod - def get_agent_app_candidates(cls, *, tenant_id: str, agent_id: str, user_id: str) -> dict[str, Any]: + def get_agent_app_candidates( + cls, *, tenant_id: str, agent_id: str, user_id: str, session: Session + ) -> dict[str, Any]: """Slash-menu data source for the Agent App (Console) composer (ENG-615).""" from services.agent.composer_candidates import soul_candidates - agent_soul = cls._load_agent_soul(tenant_id=tenant_id, agent_id=agent_id) + agent_soul = cls._load_agent_soul(tenant_id=tenant_id, agent_id=agent_id, session=session) soul_lists, truncated = soul_candidates( agent_soul=agent_soul, dataset_lookup=lambda ids: get_tenant_knowledge_dataset_rows(tenant_id=tenant_id, dataset_ids=ids), @@ -858,18 +975,21 @@ class AgentComposerService: return None @classmethod - def _load_binding_soul(cls, *, tenant_id: str, binding: WorkflowAgentNodeBinding) -> AgentSoulConfig | None: - agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=binding.agent_id) + def _load_binding_soul( + cls, *, tenant_id: str, binding: WorkflowAgentNodeBinding, session: Session + ) -> AgentSoulConfig | None: + agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=binding.agent_id, session=session) version = cls._get_version_if_present( tenant_id=tenant_id, agent_id=agent.id if agent else None, version_id=binding.current_snapshot_id, + session=session, ) return cls._parse_soul_snapshot(version) @classmethod - def _load_agent_soul(cls, *, tenant_id: str, agent_id: str) -> AgentSoulConfig | None: - agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=agent_id) + def _load_agent_soul(cls, *, tenant_id: str, agent_id: str, session: Session) -> AgentSoulConfig | None: + agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=agent_id, session=session) if agent is None: return None draft = cls._get_or_create_agent_draft( @@ -878,6 +998,7 @@ class AgentComposerService: draft_type=AgentConfigDraftType.DRAFT, account_id=None, created_by=agent.updated_by or agent.created_by, + session=session, ) return AgentSoulConfig.model_validate(draft.config_snapshot_dict) @@ -893,9 +1014,11 @@ class AgentComposerService: @classmethod def _binding_declared_outputs( - cls, *, tenant_id: str, workflow_id: str, node_id: str + cls, *, tenant_id: str, workflow_id: str, node_id: str, session: Session ) -> list[DeclaredOutputConfig] | None: - binding = cls._get_workflow_binding(tenant_id=tenant_id, workflow_id=workflow_id, node_id=node_id) + binding = cls._get_workflow_binding( + tenant_id=tenant_id, workflow_id=workflow_id, node_id=node_id, session=session + ) if binding is None: return None node_job = cls._parse_node_job(binding) @@ -970,8 +1093,8 @@ class AgentComposerService: return tools @classmethod - def calculate_impact(cls, *, tenant_id: str, current_snapshot_id: str) -> dict[str, Any]: - snapshot = db.session.scalar( + def calculate_impact(cls, *, tenant_id: str, current_snapshot_id: str, session: Session) -> dict[str, Any]: + snapshot = session.scalar( select(AgentConfigSnapshot) .where( AgentConfigSnapshot.tenant_id == tenant_id, @@ -987,7 +1110,7 @@ class AgentComposerService: & (WorkflowAgentNodeBinding.binding_type == WorkflowAgentBindingType.ROSTER_AGENT) ) bindings = list( - db.session.scalars( + session.scalars( select(WorkflowAgentNodeBinding).where( WorkflowAgentNodeBinding.tenant_id == tenant_id, or_(*predicates), @@ -1018,6 +1141,7 @@ class AgentComposerService: account_id: str, binding: WorkflowAgentNodeBinding | None, payload: ComposerSavePayload, + session: Session, ) -> WorkflowAgentNodeBinding: node_job = payload.node_job or WorkflowNodeJobConfig() if binding: @@ -1030,6 +1154,7 @@ class AgentComposerService: account_id=account_id, binding=binding, payload=payload, + session=session, ) binding.node_job_config = node_job if payload.agent_soul is not None and binding.binding_type == WorkflowAgentBindingType.INLINE_AGENT: @@ -1037,6 +1162,7 @@ class AgentComposerService: tenant_id=tenant_id, agent_id=binding.agent_id, version_id=binding.current_snapshot_id, + session=session, ) version = cls._update_current_version( current_snapshot=current_snapshot, @@ -1044,8 +1170,9 @@ class AgentComposerService: agent_soul=payload.agent_soul, operation=AgentConfigRevisionOperation.SAVE_CURRENT_VERSION, version_note=payload.version_note, + session=session, ) - agent = cls._require_agent(tenant_id=tenant_id, agent_id=binding.agent_id) + agent = cls._require_agent(tenant_id=tenant_id, agent_id=binding.agent_id, session=session) if agent.scope != AgentScope.WORKFLOW_ONLY: raise ValueError("Inline workflow agent binding must point to a workflow-only agent") agent.active_config_snapshot_id = version.id @@ -1064,6 +1191,7 @@ class AgentComposerService: node_id=node_id, account_id=account_id, agent_soul=agent_soul, + session=session, ) binding = WorkflowAgentNodeBinding( tenant_id=tenant_id, @@ -1078,8 +1206,8 @@ class AgentComposerService: created_by=account_id, updated_by=account_id, ) - db.session.add(binding) - db.session.flush() + session.add(binding) + session.flush() return binding @classmethod @@ -1101,6 +1229,7 @@ class AgentComposerService: account_id: str, binding: WorkflowAgentNodeBinding, payload: ComposerSavePayload, + session: Session, ) -> WorkflowAgentNodeBinding: if payload.binding and (payload.binding.agent_id or payload.binding.current_snapshot_id): raise ValueError("Start from Scratch must not provide an existing inline agent binding.") @@ -1113,13 +1242,14 @@ class AgentComposerService: node_id=node_id, account_id=account_id, agent_soul=agent_soul, + session=session, ) binding.binding_type = WorkflowAgentBindingType.INLINE_AGENT binding.agent_id = agent.id binding.current_snapshot_id = agent.active_config_snapshot_id binding.node_job_config = payload.node_job or binding.node_job_config binding.updated_by = account_id - db.session.flush() + session.flush() return binding @classmethod @@ -1130,6 +1260,7 @@ class AgentComposerService: account_id: str, binding: WorkflowAgentNodeBinding | None, payload: ComposerSavePayload, + session: Session, ) -> WorkflowAgentNodeBinding: binding = cls._require_binding(binding) if payload.agent_soul is None: @@ -1138,6 +1269,7 @@ class AgentComposerService: tenant_id=tenant_id, agent_id=binding.agent_id, version_id=binding.current_snapshot_id, + session=session, ) version = cls._update_current_version( current_snapshot=current_snapshot, @@ -1145,8 +1277,9 @@ class AgentComposerService: agent_soul=payload.agent_soul, operation=AgentConfigRevisionOperation.SAVE_CURRENT_VERSION, version_note=payload.version_note, + session=session, ) - agent = cls._require_agent(tenant_id=tenant_id, agent_id=binding.agent_id) + agent = cls._require_agent(tenant_id=tenant_id, agent_id=binding.agent_id, session=session) agent.active_config_snapshot_id = version.id agent.active_config_has_model = agent_soul_has_model(payload.agent_soul) agent.active_config_is_published = True @@ -1165,6 +1298,7 @@ class AgentComposerService: account_id: str, binding: WorkflowAgentNodeBinding | None, payload: ComposerSavePayload, + session: Session, ) -> WorkflowAgentNodeBinding: binding = cls._require_binding(binding) if not binding.agent_id or payload.agent_soul is None: @@ -1176,8 +1310,9 @@ class AgentComposerService: agent_soul=payload.agent_soul, operation=AgentConfigRevisionOperation.SAVE_NEW_VERSION, version_note=payload.version_note, + session=session, ) - agent = cls._require_agent(tenant_id=tenant_id, agent_id=binding.agent_id) + agent = cls._require_agent(tenant_id=tenant_id, agent_id=binding.agent_id, session=session) agent.active_config_snapshot_id = version.id agent.active_config_has_model = agent_soul_has_model(payload.agent_soul) agent.active_config_is_published = True @@ -1199,6 +1334,7 @@ class AgentComposerService: account_id: str, binding: WorkflowAgentNodeBinding | None, payload: ComposerSavePayload, + session: Session, ) -> WorkflowAgentNodeBinding: if payload.agent_soul is None: raise ValueError("agent_soul is required") @@ -1215,6 +1351,7 @@ class AgentComposerService: agent_soul=payload.agent_soul, operation=AgentConfigRevisionOperation.SAVE_NEW_AGENT, version_note=payload.version_note, + session=session, ) node_job = payload.node_job or WorkflowNodeJobConfig() if not binding: @@ -1226,13 +1363,13 @@ class AgentComposerService: node_id=node_id, created_by=account_id, ) - db.session.add(binding) + session.add(binding) binding.binding_type = WorkflowAgentBindingType.ROSTER_AGENT binding.agent_id = agent.id binding.current_snapshot_id = agent.active_config_snapshot_id binding.node_job_config = node_job binding.updated_by = account_id - db.session.flush() + session.flush() return binding @classmethod @@ -1243,13 +1380,15 @@ class AgentComposerService: account_id: str, binding: WorkflowAgentNodeBinding | None, payload: ComposerSavePayload, + session: Session, ) -> WorkflowAgentNodeBinding: binding = cls._require_binding(binding) - source_agent = cls._require_agent(tenant_id=tenant_id, agent_id=binding.agent_id) + source_agent = cls._require_agent(tenant_id=tenant_id, agent_id=binding.agent_id, session=session) source_version = cls._require_version( tenant_id=tenant_id, agent_id=source_agent.id, version_id=binding.current_snapshot_id, + session=session, ) agent_soul = payload.agent_soul or AgentSoulConfig.model_validate(source_version.config_snapshot_dict) agent_name = payload.new_agent_name or source_agent.name @@ -1267,6 +1406,7 @@ class AgentComposerService: agent_soul=agent_soul, operation=AgentConfigRevisionOperation.SAVE_TO_ROSTER, version_note=payload.version_note, + session=session, ) cls._copy_agent_drive_rows( tenant_id=tenant_id, @@ -1275,6 +1415,7 @@ class AgentComposerService: account_id=account_id, agent_soul=agent_soul, node_job=payload.node_job or WorkflowNodeJobConfig.model_validate(binding.node_job_config_dict), + session=session, ) binding.binding_type = WorkflowAgentBindingType.ROSTER_AGENT binding.agent_id = roster_agent.id @@ -1300,8 +1441,9 @@ class AgentComposerService: icon_type: Any | None = None, icon: str | None = None, icon_background: str | None = None, + session: Session, ) -> Agent: - backing_app = AgentRosterService(db.session).create_hidden_backing_app_for_workflow_agent( + backing_app = AgentRosterService(session).create_hidden_backing_app_for_workflow_agent( tenant_id=tenant_id, account_id=account_id, name=name or f"Workflow Agent {node_id}", @@ -1329,8 +1471,8 @@ class AgentComposerService: created_by=account_id, updated_by=account_id, ) - db.session.add(agent) - db.session.flush() + session.add(agent) + session.flush() version = cls._create_config_version( tenant_id=tenant_id, agent_id=agent.id, @@ -1338,6 +1480,7 @@ class AgentComposerService: agent_soul=agent_soul, operation=AgentConfigRevisionOperation.CREATE_VERSION, version_note=None, + session=session, ) agent.active_config_snapshot_id = version.id agent.active_config_has_model = agent_soul_has_model(agent_soul) @@ -1354,6 +1497,7 @@ class AgentComposerService: account_id: str, agent_soul: AgentSoulConfig, node_job: WorkflowNodeJobConfig | None = None, + session: Session, ) -> None: exact_keys, prefixes = cls._drive_copy_scopes_from_agent_configs(agent_soul=agent_soul, node_job=node_job) predicates: list[ColumnElement[bool]] = [] @@ -1364,7 +1508,7 @@ class AgentComposerService: return source_rows = list( - db.session.scalars( + session.scalars( select(AgentDriveFile).where( AgentDriveFile.tenant_id == tenant_id, AgentDriveFile.agent_id == source_agent_id, @@ -1376,7 +1520,7 @@ class AgentComposerService: return existing_target_keys = set( - db.session.scalars( + session.scalars( select(AgentDriveFile.key).where( AgentDriveFile.tenant_id == tenant_id, AgentDriveFile.agent_id == target_agent_id, @@ -1387,7 +1531,7 @@ class AgentComposerService: for row in source_rows: if row.key in existing_target_keys: continue - db.session.add( + session.add( AgentDriveFile( tenant_id=tenant_id, agent_id=target_agent_id, @@ -1451,8 +1595,9 @@ class AgentComposerService: icon_type: AgentIconType | None = None, icon: str | None = None, icon_background: str | None = None, + session: Session, ) -> Agent: - account = cls._require_account(account_id=account_id) + account = cls._require_account(account_id=account_id, session=session) try: app = AppService().create_app( tenant_id, @@ -1466,12 +1611,13 @@ class AgentComposerService: icon_background=icon_background, ), account, + session=session, ) except IntegrityError as exc: - db.session.rollback() + session.rollback() raise AgentNameConflictError() from exc - agent = AgentRosterService(db.session).get_app_backing_agent(tenant_id=tenant_id, app_id=app.id) + agent = AgentRosterService(session).get_app_backing_agent(tenant_id=tenant_id, app_id=app.id) if agent is None: raise AgentNotFoundError() @@ -1479,6 +1625,7 @@ class AgentComposerService: tenant_id=tenant_id, agent_id=agent.id, version_id=agent.active_config_snapshot_id, + session=session, ) version = cls._update_current_version( current_snapshot=current_snapshot, @@ -1486,6 +1633,7 @@ class AgentComposerService: agent_soul=agent_soul, operation=operation, version_note=version_note, + session=session, ) agent.active_config_snapshot_id = version.id agent.active_config_has_model = agent_soul_has_model(agent_soul) @@ -1504,9 +1652,10 @@ class AgentComposerService: operation: AgentConfigRevisionOperation, version_note: str | None, previous_snapshot_id: str | None = None, + session: Session, ) -> AgentConfigSnapshot: next_version = ( - db.session.scalar( + session.scalar( select(func.max(AgentConfigSnapshot.version)).where( AgentConfigSnapshot.tenant_id == tenant_id, AgentConfigSnapshot.agent_id == agent_id, @@ -1522,20 +1671,20 @@ class AgentComposerService: version_note=version_note, created_by=account_id, ) - db.session.add(version) - db.session.flush() + session.add(version) + session.flush() revision = AgentConfigRevision( tenant_id=tenant_id, agent_id=agent_id, previous_snapshot_id=previous_snapshot_id, current_snapshot_id=version.id, - revision=cls._next_revision(tenant_id=tenant_id, agent_id=agent_id), + revision=cls._next_revision(tenant_id=tenant_id, agent_id=agent_id, session=session), operation=operation, version_note=version_note, created_by=account_id, ) - db.session.add(revision) - db.session.flush() + session.add(revision) + session.flush() return version @classmethod @@ -1547,6 +1696,7 @@ class AgentComposerService: agent_soul: AgentSoulConfig, operation: AgentConfigRevisionOperation, version_note: str | None, + session: Session, ) -> AgentConfigSnapshot: return cls._create_config_version( tenant_id=current_snapshot.tenant_id, @@ -1556,12 +1706,13 @@ class AgentComposerService: operation=operation, version_note=version_note, previous_snapshot_id=current_snapshot.id, + session=session, ) @classmethod - def _next_revision(cls, *, tenant_id: str, agent_id: str) -> int: + def _next_revision(cls, *, tenant_id: str, agent_id: str, session: Session) -> int: return ( - db.session.scalar( + session.scalar( select(func.max(AgentConfigRevision.revision)).where( AgentConfigRevision.tenant_id == tenant_id, AgentConfigRevision.agent_id == agent_id, @@ -1571,8 +1722,8 @@ class AgentComposerService: ) + 1 @classmethod - def _get_agent_app_agent(cls, *, tenant_id: str, app_id: str) -> Agent | None: - return db.session.scalar( + def _get_agent_app_agent(cls, *, tenant_id: str, app_id: str, session: Session) -> Agent | None: + return session.scalar( select(Agent) .where( Agent.tenant_id == tenant_id, @@ -1586,8 +1737,8 @@ class AgentComposerService: ) @classmethod - def _require_agent_app_agent(cls, *, tenant_id: str, app_id: str) -> Agent: - agent = cls._get_agent_app_agent(tenant_id=tenant_id, app_id=app_id) + def _require_agent_app_agent(cls, *, tenant_id: str, app_id: str, session: Session) -> Agent: + agent = cls._get_agent_app_agent(tenant_id=tenant_id, app_id=app_id, session=session) if agent is None: raise AgentNotFoundError() return agent @@ -1600,6 +1751,7 @@ class AgentComposerService: agent_id: str, draft_type: AgentConfigDraftType, account_id: str | None, + session: Session, ) -> AgentConfigDraft | None: stmt = select(AgentConfigDraft).where( AgentConfigDraft.tenant_id == tenant_id, @@ -1610,7 +1762,7 @@ class AgentComposerService: stmt = stmt.where(AgentConfigDraft.account_id == account_id) else: stmt = stmt.where(AgentConfigDraft.account_id.is_(None)) - return db.session.scalar(stmt.order_by(AgentConfigDraft.updated_at.desc()).limit(1)) + return session.scalar(stmt.order_by(AgentConfigDraft.updated_at.desc()).limit(1)) @classmethod def _get_or_create_agent_draft( @@ -1621,12 +1773,14 @@ class AgentComposerService: draft_type: AgentConfigDraftType, account_id: str | None, created_by: str | None, + session: Session, ) -> AgentConfigDraft: draft = cls._get_agent_draft( tenant_id=tenant_id, agent_id=agent.id, draft_type=draft_type, account_id=account_id, + session=session, ) if draft is not None: return draft @@ -1634,6 +1788,7 @@ class AgentComposerService: tenant_id=tenant_id, agent_id=agent.id, version_id=agent.active_config_snapshot_id, + session=session, ) agent_soul = ( AgentSoulConfig.model_validate(base_snapshot.config_snapshot_dict) @@ -1651,8 +1806,8 @@ class AgentComposerService: created_by=created_by, updated_by=created_by, ) - db.session.add(draft) - db.session.flush() + session.add(draft) + session.flush() return draft @classmethod @@ -1666,6 +1821,7 @@ class AgentComposerService: agent_soul: AgentSoulConfig, account_id_for_audit: str, base_snapshot_id: str | None = None, + session: Session, ) -> AgentConfigDraft: draft = cls._get_or_create_agent_draft( tenant_id=tenant_id, @@ -1673,6 +1829,7 @@ class AgentComposerService: draft_type=draft_type, account_id=account_id, created_by=account_id_for_audit, + session=session, ) draft.config_snapshot = agent_soul if base_snapshot_id is not None: @@ -1682,7 +1839,7 @@ class AgentComposerService: draft.updated_by = account_id_for_audit if draft_type == AgentConfigDraftType.DRAFT and account_id is None: agent.active_config_is_published = False - db.session.flush() + session.flush() return draft @classmethod @@ -1710,8 +1867,8 @@ class AgentComposerService: } @classmethod - def _get_draft_workflow(cls, *, tenant_id: str, app_id: str) -> Workflow: - workflow = db.session.scalar( + def _get_draft_workflow(cls, *, tenant_id: str, app_id: str, session: Session) -> Workflow: + workflow = session.scalar( select(Workflow) .where( Workflow.tenant_id == tenant_id, @@ -1726,13 +1883,13 @@ class AgentComposerService: @classmethod def _get_workflow_binding( - cls, *, tenant_id: str, workflow_id: str, node_id: str + cls, *, tenant_id: str, workflow_id: str, node_id: str, session: Session ) -> WorkflowAgentNodeBinding | None: # Composer always operates against the draft workflow row, so this lookup # is scoped to ``workflow_version="draft"``. Published bindings are # materialized by WorkflowAgentPublishService.copy_agent_node_bindings_to_published # and are not edited through the Composer. - return db.session.scalar( + return session.scalar( select(WorkflowAgentNodeBinding) .where( WorkflowAgentNodeBinding.tenant_id == tenant_id, @@ -1750,32 +1907,34 @@ class AgentComposerService: return binding @classmethod - def _require_agent(cls, *, tenant_id: str, agent_id: str | None) -> Agent: + def _require_agent(cls, *, tenant_id: str, agent_id: str | None, session: Session) -> Agent: if not agent_id: raise AgentNotFoundError() - agent = db.session.scalar(select(Agent).where(Agent.tenant_id == tenant_id, Agent.id == agent_id).limit(1)) + agent = session.scalar(select(Agent).where(Agent.tenant_id == tenant_id, Agent.id == agent_id).limit(1)) if not agent: raise AgentNotFoundError() return agent @classmethod - def _require_account(cls, *, account_id: str) -> Account: - account = db.session.get(Account, account_id) + def _require_account(cls, *, account_id: str, session: Session) -> Account: + account = session.get(Account, account_id) if not account: raise ValueError("Account not found") return account @classmethod - def _get_agent_if_present(cls, *, tenant_id: str, agent_id: str | None) -> Agent | None: + def _get_agent_if_present(cls, *, tenant_id: str, agent_id: str | None, session: Session) -> Agent | None: if not agent_id: return None - return db.session.scalar(select(Agent).where(Agent.tenant_id == tenant_id, Agent.id == agent_id).limit(1)) + return session.scalar(select(Agent).where(Agent.tenant_id == tenant_id, Agent.id == agent_id).limit(1)) @classmethod - def _require_version(cls, *, tenant_id: str, agent_id: str | None, version_id: str | None) -> AgentConfigSnapshot: + def _require_version( + cls, *, tenant_id: str, agent_id: str | None, version_id: str | None, session: Session + ) -> AgentConfigSnapshot: if not agent_id or not version_id: raise AgentVersionNotFoundError() - version = db.session.scalar( + version = session.scalar( select(AgentConfigSnapshot) .where( AgentConfigSnapshot.tenant_id == tenant_id, @@ -1790,11 +1949,11 @@ class AgentComposerService: @classmethod def _get_version_if_present( - cls, *, tenant_id: str, agent_id: str | None, version_id: str | None + cls, *, tenant_id: str, agent_id: str | None, version_id: str | None, session: Session ) -> AgentConfigSnapshot | None: if not agent_id or not version_id: return None - return db.session.scalar( + return session.scalar( select(AgentConfigSnapshot) .where( AgentConfigSnapshot.tenant_id == tenant_id, @@ -1853,6 +2012,7 @@ class AgentComposerService: agent: Agent | None, version: AgentConfigSnapshot | None, account_id: str | None = None, + session: Session, ) -> dict[str, Any]: locked = bool(agent and agent.scope == AgentScope.ROSTER) save_options = [ComposerSaveStrategy.NODE_JOB_ONLY.value] @@ -1871,9 +2031,10 @@ class AgentComposerService: binding=binding, agent=agent, account_id=account_id, + session=session, ) debug_conversation_message_count = ( - AgentRosterService(db.session).count_agent_app_debug_conversation_messages( + AgentRosterService(session).count_agent_app_debug_conversation_messages( conversation_id=debug_conversation_id ) if debug_conversation_id @@ -1906,7 +2067,9 @@ class AgentComposerService: # this is the same list (so callers don't need to special-case). "effective_declared_outputs": cls._serialize_effective_outputs(cls._declared_outputs_from_binding(binding)), "save_options": save_options, - "impact_summary": cls.calculate_impact(tenant_id=binding.tenant_id, current_snapshot_id=version.id) + "impact_summary": cls.calculate_impact( + tenant_id=binding.tenant_id, current_snapshot_id=version.id, session=session + ) if version else None, "app_id": binding.app_id, @@ -1927,6 +2090,7 @@ class AgentComposerService: binding: WorkflowAgentNodeBinding, agent: Agent | None, account_id: str | None, + session: Session, ) -> str | None: if ( not account_id @@ -1938,7 +2102,7 @@ class AgentComposerService: from services.agent.roster_service import AgentRosterService - return AgentRosterService(db.session).get_or_create_agent_app_debug_conversation_id( + return AgentRosterService(session).get_or_create_agent_app_debug_conversation_id( tenant_id=tenant_id, agent_id=agent.id, account_id=account_id, diff --git a/api/services/agent/roster_service.py b/api/services/agent/roster_service.py index a34fca67105..de76b3c4eb2 100644 --- a/api/services/agent/roster_service.py +++ b/api/services/agent/roster_service.py @@ -826,6 +826,7 @@ class AgentRosterService: max_active_requests=source_app.max_active_requests, ), account, + session=self._session, ) target_app.enable_site = source_app.enable_site diff --git a/api/services/agent/skill_standardize_service.py b/api/services/agent/skill_standardize_service.py index cc2ba4b9bdc..2639f7a9a18 100644 --- a/api/services/agent/skill_standardize_service.py +++ b/api/services/agent/skill_standardize_service.py @@ -18,6 +18,8 @@ from __future__ import annotations import re from typing import Any +from sqlalchemy.orm import Session + from core.tools.tool_file_manager import ToolFileManager from services.agent.skill_package_service import SkillPackageService from services.agent_drive_service import AgentDriveService, DriveCommitItem, DriveFileRef, DriveSkillMetadata @@ -59,6 +61,7 @@ class SkillStandardizeService: tenant_id: str, user_id: str, agent_id: str, + session: Session, ) -> dict[str, Any]: """Create two ToolFiles, commit two drive-owned keys, and return skill metadata. @@ -113,6 +116,7 @@ class SkillStandardizeService: value_owned_by_drive=True, ), ], + session=session, ) self.last_committed_items = committed_items diff --git a/api/services/agent/skill_tool_inference_service.py b/api/services/agent/skill_tool_inference_service.py index a6d5e6b2de9..7ce53dd4666 100644 --- a/api/services/agent/skill_tool_inference_service.py +++ b/api/services/agent/skill_tool_inference_service.py @@ -19,6 +19,7 @@ from typing import Any import json_repair from pydantic import BaseModel, Field, ValidationError +from sqlalchemy.orm import Session from core.errors.error import ProviderTokenNotInitError from core.model_manager import ModelManager @@ -91,8 +92,8 @@ class SkillToolInferenceService: def __init__(self, *, drive_service: AgentDriveService | None = None) -> None: self._drive = drive_service or AgentDriveService() - def infer(self, *, tenant_id: str, agent_id: str, slug: str) -> dict[str, Any]: - skill_md = self._load_skill_md(tenant_id=tenant_id, agent_id=agent_id, slug=slug) + def infer(self, *, tenant_id: str, agent_id: str, slug: str, session: Session) -> dict[str, Any]: + skill_md = self._load_skill_md(tenant_id=tenant_id, agent_id=agent_id, slug=slug, session=session) user_prompt = f"SKILL.md of skill '{slug}':\n\n{skill_md}" @@ -115,9 +116,11 @@ class SkillToolInferenceService: tool.inferred_from = slug return result.model_dump(mode="json") - def _load_skill_md(self, *, tenant_id: str, agent_id: str, slug: str) -> str: + def _load_skill_md(self, *, tenant_id: str, agent_id: str, slug: str, session: Session) -> str: try: - preview = self._drive.preview(tenant_id=tenant_id, agent_id=agent_id, key=f"{slug}/SKILL.md") + preview = self._drive.preview( + tenant_id=tenant_id, agent_id=agent_id, key=f"{slug}/SKILL.md", session=session + ) except AgentDriveError as exc: if exc.code == "drive_key_not_found": raise SkillToolInferenceError( diff --git a/api/services/agent_app_feature_service.py b/api/services/agent_app_feature_service.py index 5fd794bb10f..d336cdf29de 100644 --- a/api/services/agent_app_feature_service.py +++ b/api/services/agent_app_feature_service.py @@ -13,7 +13,7 @@ from __future__ import annotations from typing import Any, cast -from sqlalchemy.orm import scoped_session +from sqlalchemy.orm import Session from core.app.app_config.common.sensitive_word_avoidance.manager import SensitiveWordAvoidanceConfigManager from core.app.app_config.features.opening_statement.manager import OpeningStatementConfigManager @@ -74,7 +74,7 @@ class AgentAppFeatureConfigService: app_model: App, account: Account, config: dict[str, Any], - session: scoped_session, + session: Session, ) -> AppModelConfig: """Persist the presentation features as a new app_model_config version. diff --git a/api/services/agent_app_sandbox_service.py b/api/services/agent_app_sandbox_service.py index b1652d628e9..3f5a0bf41b2 100644 --- a/api/services/agent_app_sandbox_service.py +++ b/api/services/agent_app_sandbox_service.py @@ -19,12 +19,12 @@ from dify_agent.client import Client from dify_agent.protocol import RuntimeLayerSpec, SandboxLocator, build_sandbox_locator_from_layer_specs from pydantic import BaseModel, TypeAdapter from sqlalchemy import select +from sqlalchemy.orm import Session from configs import dify_config from core.app.apps.agent_app.session_store import AgentAppRuntimeSessionStore from core.app.file_access import DatabaseFileAccessController from core.app.workflow.file_runtime import DifyWorkflowFileRuntime -from core.db.session_factory import session_factory from factories import file_factory from models.agent import AgentRuntimeSessionOwnerType, WorkflowAgentRuntimeSession, WorkflowAgentRuntimeSessionStatus @@ -134,6 +134,7 @@ class WorkflowAgentSandboxService: node_id: str, node_execution_id: str | None, path: str, + session: Session, ): locator = self._resolve_locator( tenant_id=tenant_id, @@ -141,6 +142,7 @@ class WorkflowAgentSandboxService: workflow_run_id=workflow_run_id, node_id=node_id, node_execution_id=node_execution_id, + session=session, ) return self._client_factory().list_sandbox_files_sync(locator, path) @@ -153,6 +155,7 @@ class WorkflowAgentSandboxService: node_id: str, node_execution_id: str | None, path: str, + session: Session, ): locator = self._resolve_locator( tenant_id=tenant_id, @@ -160,6 +163,7 @@ class WorkflowAgentSandboxService: workflow_run_id=workflow_run_id, node_id=node_id, node_execution_id=node_execution_id, + session=session, ) return self._client_factory().read_sandbox_file_sync(locator, path) @@ -172,6 +176,7 @@ class WorkflowAgentSandboxService: node_id: str, node_execution_id: str | None, path: str, + session: Session, ) -> AgentSandboxUploadDownload: locator = self._resolve_locator( tenant_id=tenant_id, @@ -179,6 +184,7 @@ class WorkflowAgentSandboxService: workflow_run_id=workflow_run_id, node_id=node_id, node_execution_id=node_execution_id, + session=session, ) uploaded = self._client_factory().upload_sandbox_file_sync(locator, path) return _upload_download_response( @@ -194,6 +200,7 @@ class WorkflowAgentSandboxService: workflow_run_id: str, node_id: str, node_execution_id: str | None, + session: Session, ) -> SandboxLocator: """Resolve one workflow Agent sandbox from product-facing identifiers. @@ -216,8 +223,7 @@ class WorkflowAgentSandboxService: stmt = stmt.where(WorkflowAgentRuntimeSession.node_execution_id == node_execution_id) stmt = stmt.order_by(WorkflowAgentRuntimeSession.updated_at.desc()).limit(1) - with session_factory.create_session() as session: - row = session.scalar(stmt) + row = session.scalar(stmt) if row is None: raise AgentSandboxInspectorError( diff --git a/api/services/agent_drive_service.py b/api/services/agent_drive_service.py index d79aa1abe9c..eb375f997d1 100644 --- a/api/services/agent_drive_service.py +++ b/api/services/agent_drive_service.py @@ -41,7 +41,6 @@ from sqlalchemy.orm import Session from configs import dify_config from core.app.file_access.controller import DatabaseFileAccessController -from core.db.session_factory import session_factory from extensions.ext_storage import storage from factories import file_factory from libs.uuid_utils import uuidv7 @@ -195,38 +194,38 @@ class AgentDriveService: *, tenant_id: str, agent_id: str, + session: Session, prefix: str = "", include_download_url: bool = False, ) -> list[dict[str, Any]]: - with session_factory.create_session() as session: - self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) - stmt = ( - select(AgentDriveFile) - .where(AgentDriveFile.tenant_id == tenant_id, AgentDriveFile.agent_id == agent_id) - .order_by(AgentDriveFile.key) - ) - if prefix: - stmt = stmt.where(AgentDriveFile.key.startswith(prefix)) - rows = list(session.scalars(stmt)) - items: list[dict[str, Any]] = [] - for row in rows: - item: dict[str, Any] = { - "key": row.key, - "size": row.size, - "hash": row.hash, - "mime_type": row.mime_type, - "file_kind": row.file_kind.value, - "file_id": row.file_id, - "is_skill": row.is_skill, - "skill_metadata": row.skill_metadata, - "created_at": int(row.created_at.timestamp()) if row.created_at else None, - } - if include_download_url: - item["download_url"] = self._resolve_download_url( - tenant_id=tenant_id, file_kind=row.file_kind, file_id=row.file_id - ) - items.append(item) - return items + self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) + stmt = ( + select(AgentDriveFile) + .where(AgentDriveFile.tenant_id == tenant_id, AgentDriveFile.agent_id == agent_id) + .order_by(AgentDriveFile.key) + ) + if prefix: + stmt = stmt.where(AgentDriveFile.key.startswith(prefix)) + rows = list(session.scalars(stmt)) + items: list[dict[str, Any]] = [] + for row in rows: + item: dict[str, Any] = { + "key": row.key, + "size": row.size, + "hash": row.hash, + "mime_type": row.mime_type, + "file_kind": row.file_kind.value, + "file_id": row.file_id, + "is_skill": row.is_skill, + "skill_metadata": row.skill_metadata, + "created_at": int(row.created_at.timestamp()) if row.created_at else None, + } + if include_download_url: + item["download_url"] = self._resolve_download_url( + tenant_id=tenant_id, file_kind=row.file_kind, file_id=row.file_id + ) + items.append(item) + return items def commit( self, @@ -235,25 +234,25 @@ class AgentDriveService: user_id: str, agent_id: str, items: list[DriveCommitItem], + session: Session, ) -> list[dict[str, Any]]: if not items: raise AgentDriveError("empty_commit", "commit requires at least one item", status_code=400) committed: list[dict[str, Any]] = [] pending_storage_deletes: list[str] = [] - with session_factory.create_session() as session: - self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) - for item in items: - committed.append( - self._commit_one( - session, - tenant_id=tenant_id, - user_id=user_id, - agent_id=agent_id, - item=item, - pending_storage_deletes=pending_storage_deletes, - ) + self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) + for item in items: + committed.append( + self._commit_one( + session, + tenant_id=tenant_id, + user_id=user_id, + agent_id=agent_id, + item=item, + pending_storage_deletes=pending_storage_deletes, ) - session.commit() + ) + session.commit() for storage_key in pending_storage_deletes: self._delete_storage(storage_key) return committed @@ -263,6 +262,7 @@ class AgentDriveService: *, tenant_id: str, agent_id: str, + session: Session, prefix: str | None = None, key: str | None = None, ) -> list[str]: @@ -276,59 +276,57 @@ class AgentDriveService: raise AgentDriveError("invalid_delete_scope", "delete requires exactly one of prefix or key") removed_keys: list[str] = [] pending_storage_deletes: list[str] = [] - with session_factory.create_session() as session: - self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) - stmt = select(AgentDriveFile).where( - AgentDriveFile.tenant_id == tenant_id, - AgentDriveFile.agent_id == agent_id, - ) - if key is not None: - stmt = stmt.where(AgentDriveFile.key == normalize_drive_key(key)) - else: - stmt = stmt.where(AgentDriveFile.key.startswith(normalize_drive_key(prefix or ""))) - rows = list(session.scalars(stmt)) - for row in rows: - if row.value_owned_by_drive: - self._cleanup_value( - session, - tenant_id=tenant_id, - file_kind=row.file_kind, - file_id=row.file_id, - exclude_row_id=row.id, - pending_storage_deletes=pending_storage_deletes, - ) - removed_keys.append(row.key) - session.delete(row) - session.commit() + self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) + stmt = select(AgentDriveFile).where( + AgentDriveFile.tenant_id == tenant_id, + AgentDriveFile.agent_id == agent_id, + ) + if key is not None: + stmt = stmt.where(AgentDriveFile.key == normalize_drive_key(key)) + else: + stmt = stmt.where(AgentDriveFile.key.startswith(normalize_drive_key(prefix or ""))) + rows = list(session.scalars(stmt)) + for row in rows: + if row.value_owned_by_drive: + self._cleanup_value( + session, + tenant_id=tenant_id, + file_kind=row.file_kind, + file_id=row.file_id, + exclude_row_id=row.id, + pending_storage_deletes=pending_storage_deletes, + ) + removed_keys.append(row.key) + session.delete(row) + session.commit() for storage_key in pending_storage_deletes: self._delete_storage(storage_key) return removed_keys - def list_skills(self, *, tenant_id: str, agent_id: str) -> list[AgentDriveSkillInfo]: + def list_skills(self, *, tenant_id: str, agent_id: str, session: Session) -> list[AgentDriveSkillInfo]: """Return the drive-backed skill catalog derived from canonical ``SKILL.md`` rows.""" - with session_factory.create_session() as session: - self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) - skill_rows = list( - session.scalars( - select(AgentDriveFile) - .where( - AgentDriveFile.tenant_id == tenant_id, - AgentDriveFile.agent_id == agent_id, - AgentDriveFile.is_skill.is_(True), - ) - .order_by(AgentDriveFile.key) - ) - ) - archive_keys = set( - session.scalars( - select(AgentDriveFile.key).where( - AgentDriveFile.tenant_id == tenant_id, - AgentDriveFile.agent_id == agent_id, - AgentDriveFile.key.in_([self._skill_archive_key(row.key) for row in skill_rows]), - ) + self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) + skill_rows = list( + session.scalars( + select(AgentDriveFile) + .where( + AgentDriveFile.tenant_id == tenant_id, + AgentDriveFile.agent_id == agent_id, + AgentDriveFile.is_skill.is_(True), + ) + .order_by(AgentDriveFile.key) + ) + ) + archive_keys = set( + session.scalars( + select(AgentDriveFile.key).where( + AgentDriveFile.tenant_id == tenant_id, + AgentDriveFile.agent_id == agent_id, + AgentDriveFile.key.in_([self._skill_archive_key(row.key) for row in skill_rows]), ) ) + ) skills: list[AgentDriveSkillInfo] = [] for row in skill_rows: @@ -349,14 +347,20 @@ class AgentDriveService: ) return skills - def inspect_skill(self, *, tenant_id: str, agent_id: str, skill_path: str) -> AgentDriveSkillInspectInfo: + def inspect_skill( + self, *, tenant_id: str, agent_id: str, skill_path: str, session: Session + ) -> AgentDriveSkillInspectInfo: """Return the UI-facing skill inspect view for slash-menu hover/detail.""" skill_path = normalize_drive_key(skill_path) skill_md_key = skill_path if skill_path.endswith(_SKILL_MD_SUFFIX) else f"{skill_path}{_SKILL_MD_SUFFIX}" skill_path = self._skill_path_from_key(skill_md_key) catalog = next( - (item for item in self.list_skills(tenant_id=tenant_id, agent_id=agent_id) if item["path"] == skill_path), + ( + item + for item in self.list_skills(tenant_id=tenant_id, agent_id=agent_id, session=session) + if item["path"] == skill_path + ), None, ) if catalog is None: @@ -366,10 +370,11 @@ class AgentDriveService: tenant_id=tenant_id, agent_id=agent_id, skill_md_key=skill_md_key, + session=session, ) - drive_items = self.manifest(tenant_id=tenant_id, agent_id=agent_id, prefix=f"{skill_path}/") + drive_items = self.manifest(tenant_id=tenant_id, agent_id=agent_id, prefix=f"{skill_path}/", session=session) drive_keys = {item["key"] for item in drive_items} - preview = self.preview(tenant_id=tenant_id, agent_id=agent_id, key=skill_md_key) + preview = self.preview(tenant_id=tenant_id, agent_id=agent_id, key=skill_md_key, session=session) files, warnings = self._skill_file_entries( skill_path=skill_path, skill_md_key=skill_md_key, @@ -582,23 +587,24 @@ class AgentDriveService: ) from exc @staticmethod - def _manifest_files_from_skill_metadata(*, tenant_id: str, agent_id: str, skill_md_key: str) -> list[str] | None: - with session_factory.create_session() as session: - row = session.scalar( - select(AgentDriveFile).where( - AgentDriveFile.tenant_id == tenant_id, - AgentDriveFile.agent_id == agent_id, - AgentDriveFile.key == skill_md_key, - AgentDriveFile.is_skill.is_(True), - ) + def _manifest_files_from_skill_metadata( + *, tenant_id: str, agent_id: str, skill_md_key: str, session: Session + ) -> list[str] | None: + row = session.scalar( + select(AgentDriveFile).where( + AgentDriveFile.tenant_id == tenant_id, + AgentDriveFile.agent_id == agent_id, + AgentDriveFile.key == skill_md_key, + AgentDriveFile.is_skill.is_(True), ) - if row is None: - return None - try: - metadata = AgentDriveService._parse_skill_metadata(row.key, row.skill_metadata) - except Exception: - logger.warning("drive skill inspect: malformed skill metadata for %s", skill_md_key, exc_info=True) - return None + ) + if row is None: + return None + try: + metadata = AgentDriveService._parse_skill_metadata(row.key, row.skill_metadata) + except Exception: + logger.warning("drive skill inspect: malformed skill metadata for %s", skill_md_key, exc_info=True) + return None return [str(item) for item in (metadata.manifest_files or []) if str(item).strip()] or None @classmethod @@ -932,15 +938,15 @@ class AgentDriveService: archive_file_kind: AgentDriveFileKind, archive_file_id: str, member_path: str, + session: Session, ) -> bytes: member_path = normalize_drive_key(member_path) - with session_factory.create_session() as session: - storage_key = self._storage_key_for_ref( - session, - tenant_id=tenant_id, - file_kind=archive_file_kind, - file_id=archive_file_id, - ) + storage_key = self._storage_key_for_ref( + session, + tenant_id=tenant_id, + file_kind=archive_file_kind, + file_id=archive_file_id, + ) archive_bytes = b"".join(storage.load_stream(storage_key)) try: with zipfile.ZipFile(io.BytesIO(archive_bytes)) as archive: @@ -978,26 +984,25 @@ class AgentDriveService: return {"key": key, "size": size, "truncated": truncated, "binary": True, "text": None} return {"key": key, "size": size, "truncated": truncated, "binary": False, "text": text} - def preview(self, *, tenant_id: str, agent_id: str, key: str) -> dict[str, Any]: + def preview(self, *, tenant_id: str, agent_id: str, key: str, session: Session) -> dict[str, Any]: """Truncated text preview of one drive value (binary-safe, never 500s on size).""" - with session_factory.create_session() as session: - self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) - try: - row = self._require_row(session, tenant_id=tenant_id, agent_id=agent_id, key=key) - storage_key = self._storage_key_for_row(session, tenant_id=tenant_id, row=row) - size = row.size - response_key = row.key - archive_ref: tuple[AgentDriveFile, str] | None = None - except AgentDriveError: - archive_ref = self._archive_member_for_key( - session, - tenant_id=tenant_id, - agent_id=agent_id, - key=key, - ) - storage_key = None - size = None - response_key = normalize_drive_key(key) + self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) + try: + row = self._require_row(session, tenant_id=tenant_id, agent_id=agent_id, key=key) + storage_key = self._storage_key_for_row(session, tenant_id=tenant_id, row=row) + size = row.size + response_key = row.key + archive_ref: tuple[AgentDriveFile, str] | None = None + except AgentDriveError: + archive_ref = self._archive_member_for_key( + session, + tenant_id=tenant_id, + agent_id=agent_id, + key=key, + ) + storage_key = None + size = None + response_key = normalize_drive_key(key) if archive_ref is not None: archive_row, member_path = archive_ref @@ -1006,6 +1011,7 @@ class AgentDriveService: archive_file_kind=archive_row.file_kind, archive_file_id=archive_row.file_id, member_path=member_path, + session=session, ) return self._preview_bytes(key=response_key, size=len(payload), payload=payload) @@ -1026,47 +1032,47 @@ class AgentDriveService: archive_file_kind: AgentDriveFileKind, archive_file_id: str, member_path: str, + session: Session, ) -> dict[str, Any]: - with session_factory.create_session() as session: - self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) + self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) payload = self._load_archive_member_bytes( tenant_id=tenant_id, archive_file_kind=archive_file_kind, archive_file_id=archive_file_id, member_path=member_path, + session=session, ) return self._preview_bytes(key=normalize_drive_key(key), size=len(payload), payload=payload) - def download_url(self, *, tenant_id: str, agent_id: str, key: str) -> str: + def download_url(self, *, tenant_id: str, agent_id: str, key: str, session: Session) -> str: """External signed URL for a browser download of one drive value.""" - with session_factory.create_session() as session: - self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) - try: - row = self._require_row(session, tenant_id=tenant_id, agent_id=agent_id, key=key) - except AgentDriveError: - archive_row, member_path = self._archive_member_for_key( - session, - tenant_id=tenant_id, - agent_id=agent_id, - key=key, - ) - return self.sign_archive_member_url( - tenant_id=tenant_id, - agent_id=agent_id, - key=key, - archive_file_kind=archive_row.file_kind, - archive_file_id=archive_row.file_id, - member_path=member_path, - for_external=True, - as_attachment=True, - ) - url = self._resolve_download_url( + self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) + try: + row = self._require_row(session, tenant_id=tenant_id, agent_id=agent_id, key=key) + except AgentDriveError: + archive_row, member_path = self._archive_member_for_key( + session, tenant_id=tenant_id, - file_kind=row.file_kind, - file_id=row.file_id, + agent_id=agent_id, + key=key, + ) + return self.sign_archive_member_url( + tenant_id=tenant_id, + agent_id=agent_id, + key=key, + archive_file_kind=archive_row.file_kind, + archive_file_id=archive_row.file_id, + member_path=member_path, for_external=True, as_attachment=True, ) + url = self._resolve_download_url( + tenant_id=tenant_id, + file_kind=row.file_kind, + file_id=row.file_id, + for_external=True, + as_attachment=True, + ) if url is None: raise AgentDriveError("drive_key_not_found", "drive value cannot be resolved", status_code=404) return url @@ -1080,10 +1086,10 @@ class AgentDriveService: archive_file_kind: AgentDriveFileKind, archive_file_id: str, member_path: str, + session: Session, for_external: bool = True, ) -> str: - with session_factory.create_session() as session: - self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) + self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) return self.sign_archive_member_url( tenant_id=tenant_id, agent_id=agent_id, @@ -1211,14 +1217,15 @@ class AgentDriveService: archive_file_kind: AgentDriveFileKind, archive_file_id: str, member_path: str, + session: Session, ) -> tuple[bytes, str, str]: - with session_factory.create_session() as session: - self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) + self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) payload = self._load_archive_member_bytes( tenant_id=tenant_id, archive_file_kind=archive_file_kind, archive_file_id=archive_file_id, member_path=member_path, + session=session, ) mime_type = mimetypes.guess_type(member_path)[0] or "application/octet-stream" filename = normalize_drive_key(key).rsplit("/", 1)[-1] diff --git a/api/services/agent_service.py b/api/services/agent_service.py index d8f4e11e758..a201eeb0485 100644 --- a/api/services/agent_service.py +++ b/api/services/agent_service.py @@ -3,13 +3,13 @@ from typing import Any import pytz from sqlalchemy import select +from sqlalchemy.orm import Session import contexts from core.app.app_config.easy_ui_based_app.agent.manager import AgentConfigManager from core.plugin.impl.agent import PluginAgentClient from core.plugin.impl.exc import PluginDaemonClientSideError from core.tools.tool_manager import ToolManager -from extensions.ext_database import db from libs.login import current_user from models import Account from models.model import App, Conversation, EndUser, Message @@ -17,14 +17,14 @@ from models.model import App, Conversation, EndUser, Message class AgentService: @classmethod - def get_agent_logs(cls, app_model: App, conversation_id: str, message_id: str): + def get_agent_logs(cls, app_model: App, conversation_id: str, message_id: str, session: Session): """ Service to get agent logs """ contexts.plugin_tool_providers.set({}) contexts.plugin_tool_providers_lock.set(threading.Lock()) - conversation: Conversation | None = db.session.scalar( + conversation: Conversation | None = session.scalar( select(Conversation) .where( Conversation.id == conversation_id, @@ -36,7 +36,7 @@ class AgentService: if not conversation: raise ValueError(f"Conversation not found: {conversation_id}") - message: Message | None = db.session.scalar( + message: Message | None = session.scalar( select(Message) .where( Message.id == message_id, @@ -52,9 +52,9 @@ class AgentService: if conversation.from_end_user_id: # only select name field - executor_name = db.session.scalar(select(EndUser.name).where(EndUser.id == conversation.from_end_user_id)) + executor_name = session.scalar(select(EndUser.name).where(EndUser.id == conversation.from_end_user_id)) else: - executor_name = db.session.scalar(select(Account.name).where(Account.id == conversation.from_account_id)) + executor_name = session.scalar(select(Account.name).where(Account.id == conversation.from_account_id)) executor = executor_name or "Unknown" assert isinstance(current_user, Account) diff --git a/api/services/agent_tool_inner_service.py b/api/services/agent_tool_inner_service.py index 4420f1b66b0..633ca893007 100644 --- a/api/services/agent_tool_inner_service.py +++ b/api/services/agent_tool_inner_service.py @@ -41,7 +41,7 @@ from services.errors.agent_tool_inner import AgentToolInnerServiceError class AgentToolInnerService: """Invoke one API-owned Agent tool declaration, including explicit plugin-via-core calls.""" - def invoke(self, session: Session, request: AgentToolInvokeRequest) -> AgentToolInvokeResponse: + def invoke(self, request: AgentToolInvokeRequest, *, session: Session) -> AgentToolInvokeResponse: app = session.get(App, request.caller.app_id) if app is None: raise AgentToolInnerServiceError( diff --git a/api/services/annotation_service.py b/api/services/annotation_service.py index 03e445a938b..ccca621aab5 100644 --- a/api/services/annotation_service.py +++ b/api/services/annotation_service.py @@ -4,12 +4,12 @@ from typing import TypedDict import pandas as pd from sqlalchemy import delete, or_, select, update -from sqlalchemy.orm import scoped_session +from sqlalchemy.orm import Session from werkzeug.datastructures import FileStorage from werkzeug.exceptions import NotFound from core.helper.csv_sanitizer import CSVSanitizer -from extensions.ext_database import db +from extensions.ext_database import db # noqa: F401 from extensions.ext_redis import redis_client from libs.datetime_utils import naive_utc_now from libs.login import current_account_with_tenant @@ -91,7 +91,7 @@ class UpdateAnnotationSettingArgs(TypedDict): class AppAnnotationService: @staticmethod - def _get_annotation_by_ref(annotation_ref: AnnotationRef, session: scoped_session) -> MessageAnnotation | None: + def _get_annotation_by_ref(annotation_ref: AnnotationRef, session: Session) -> MessageAnnotation | None: return session.scalar( select(MessageAnnotation) .where( @@ -102,10 +102,12 @@ class AppAnnotationService: ) @classmethod - def up_insert_app_annotation_from_message(cls, args: UpsertAnnotationArgs, app_id: str) -> MessageAnnotation: + def up_insert_app_annotation_from_message( + cls, args: UpsertAnnotationArgs, app_id: str, *, session: Session + ) -> MessageAnnotation: # get app info current_user, current_tenant_id = current_account_with_tenant() - app = db.session.scalar( + app = session.scalar( select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) ) @@ -119,9 +121,7 @@ class AppAnnotationService: raw_message_id = args.get("message_id") if raw_message_id: message_id = str(raw_message_id) - message = db.session.scalar( - select(Message).where(Message.id == message_id, Message.app_id == app.id).limit(1) - ) + message = session.scalar(select(Message).where(Message.id == message_id, Message.app_id == app.id).limit(1)) if not message: raise NotFound("Message Not Exists.") @@ -155,10 +155,10 @@ class AppAnnotationService: question=question, account_id=current_user.id, ) - db.session.add(annotation) - db.session.commit() + session.add(annotation) + session.commit() - annotation_setting = db.session.scalar( + annotation_setting = session.scalar( select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == app_id).limit(1) ) assert current_tenant_id is not None @@ -213,10 +213,10 @@ class AppAnnotationService: return {"job_id": job_id, "job_status": "waiting"} @classmethod - def get_annotation_list_by_app_id(cls, app_id: str, page: int, limit: int, keyword: str): + def get_annotation_list_by_app_id(cls, app_id: str, page: int, limit: int, keyword: str, *, session: Session): # get app info _, current_tenant_id = current_account_with_tenant() - app = db.session.scalar( + app = session.scalar( select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) ) @@ -247,7 +247,7 @@ class AppAnnotationService: return annotations.items, annotations.total or 0 @classmethod - def export_annotation_list_by_app_id(cls, app_id: str): + def export_annotation_list_by_app_id(cls, app_id: str, *, session: Session): """ Export all annotations for an app with CSV injection protection. @@ -256,13 +256,13 @@ class AppAnnotationService: """ # get app info _, current_tenant_id = current_account_with_tenant() - app = db.session.scalar( + app = session.scalar( select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) ) if not app: raise NotFound("App not found") - annotations = db.session.scalars( + annotations = session.scalars( select(MessageAnnotation) .where(MessageAnnotation.app_id == app_id) .order_by(MessageAnnotation.created_at.desc()) @@ -280,10 +280,12 @@ class AppAnnotationService: return annotations @classmethod - def insert_app_annotation_directly(cls, args: InsertAnnotationArgs, app_id: str) -> MessageAnnotation: + def insert_app_annotation_directly( + cls, args: InsertAnnotationArgs, app_id: str, *, session: Session + ) -> MessageAnnotation: # get app info current_user, current_tenant_id = current_account_with_tenant() - app = db.session.scalar( + app = session.scalar( select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) ) @@ -297,10 +299,10 @@ class AppAnnotationService: annotation = MessageAnnotation( app_id=app.id, content=args["answer"], question=question, account_id=current_user.id ) - db.session.add(annotation) - db.session.commit() + session.add(annotation) + session.commit() # if annotation reply is enabled , add annotation to index - annotation_setting = db.session.scalar( + annotation_setting = session.scalar( select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == app_id).limit(1) ) if annotation_setting: @@ -315,7 +317,7 @@ class AppAnnotationService: @classmethod def update_app_annotation_directly( - cls, args: UpdateAnnotationArgs, annotation_ref: AnnotationRef, session: scoped_session + cls, args: UpdateAnnotationArgs, annotation_ref: AnnotationRef, session: Session ): annotation = cls._get_annotation_by_ref(annotation_ref, session) @@ -351,7 +353,7 @@ class AppAnnotationService: return annotation @classmethod - def delete_app_annotation(cls, annotation_ref: AnnotationRef, session: scoped_session): + def delete_app_annotation(cls, annotation_ref: AnnotationRef, session: Session): annotation = cls._get_annotation_by_ref(annotation_ref, session) if not annotation: @@ -384,9 +386,9 @@ class AppAnnotationService: ) @classmethod - def delete_app_annotations_in_batch(cls, app_ref: AppRef, annotation_ids: list[str]): + def delete_app_annotations_in_batch(cls, app_ref: AppRef, annotation_ids: list[str], *, session: Session): # Fetch annotations and their settings in a single query - annotations_to_delete = db.session.execute( + annotations_to_delete = session.execute( select(MessageAnnotation, AppAnnotationSetting) .outerjoin(AppAnnotationSetting, MessageAnnotation.app_id == AppAnnotationSetting.app_id) .where(MessageAnnotation.id.in_(annotation_ids), MessageAnnotation.app_id == app_ref.app_id) @@ -399,7 +401,7 @@ class AppAnnotationService: annotation_ids_to_delete = [annotation.id for annotation, _ in annotations_to_delete] # Step 2: Bulk delete hit histories in a single query - db.session.execute( + session.execute( delete(AppAnnotationHitHistory).where( AppAnnotationHitHistory.app_id == app_ref.app_id, AppAnnotationHitHistory.annotation_id.in_(annotation_ids_to_delete), @@ -414,7 +416,7 @@ class AppAnnotationService: ) # Step 4: Bulk delete annotations in a single query - delete_result = db.session.execute( + delete_result = session.execute( delete(MessageAnnotation).where( MessageAnnotation.id.in_(annotation_ids_to_delete), MessageAnnotation.app_id == app_ref.app_id, @@ -422,11 +424,11 @@ class AppAnnotationService: ) deleted_count = getattr(delete_result, "rowcount", 0) - db.session.commit() + session.commit() return {"deleted_count": deleted_count} @classmethod - def batch_import_app_annotations(cls, app_id: str, file: FileStorage): + def batch_import_app_annotations(cls, app_id: str, file: FileStorage, *, session: Session): """ Batch import annotations from CSV file with enhanced security checks. @@ -441,7 +443,7 @@ class AppAnnotationService: # get app info current_user, current_tenant_id = current_account_with_tenant() - app = db.session.scalar( + app = session.scalar( select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) ) @@ -560,8 +562,8 @@ class AppAnnotationService: return {"job_id": job_id, "job_status": "waiting", "record_count": len(result)} @classmethod - def get_annotation_hit_histories(cls, annotation_ref: AnnotationRef, page, limit): - annotation = cls._get_annotation_by_ref(annotation_ref, db.session) + def get_annotation_hit_histories(cls, annotation_ref: AnnotationRef, page, limit, *, session: Session): + annotation = cls._get_annotation_by_ref(annotation_ref, session) if not annotation: raise NotFound("Annotation not found") @@ -578,8 +580,8 @@ class AppAnnotationService: return annotation_hit_histories.items, annotation_hit_histories.total or 0 @classmethod - def get_annotation_by_id(cls, annotation_id: str) -> MessageAnnotation | None: - annotation = db.session.get(MessageAnnotation, annotation_id) + def get_annotation_by_id(cls, annotation_id: str, *, session: Session) -> MessageAnnotation | None: + annotation = session.get(MessageAnnotation, annotation_id) if not annotation: return None @@ -597,9 +599,11 @@ class AppAnnotationService: message_id: str, from_source: str, score: float, - ): + *, + session: Session, + ) -> None: # add hit count to annotation - db.session.execute( + session.execute( update(MessageAnnotation) .where(MessageAnnotation.id == annotation_id) .values(hit_count=MessageAnnotation.hit_count + 1) @@ -616,21 +620,23 @@ class AppAnnotationService: annotation_question=annotation_question, annotation_content=annotation_content, ) - db.session.add(annotation_hit_history) - db.session.commit() + session.add(annotation_hit_history) + session.commit() @classmethod - def get_app_annotation_setting_by_app_id(cls, app_id: str) -> AnnotationSettingDict | AnnotationSettingDisabledDict: + def get_app_annotation_setting_by_app_id( + cls, app_id: str, *, session: Session + ) -> AnnotationSettingDict | AnnotationSettingDisabledDict: _, current_tenant_id = current_account_with_tenant() # get app info - app = db.session.scalar( + app = session.scalar( select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) ) if not app: raise NotFound("App not found") - annotation_setting = db.session.scalar( + annotation_setting = session.scalar( select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == app_id).limit(1) ) if annotation_setting: @@ -656,18 +662,18 @@ class AppAnnotationService: @classmethod def update_app_annotation_setting( - cls, app_id: str, annotation_setting_id: str, args: UpdateAnnotationSettingArgs + cls, app_id: str, annotation_setting_id: str, args: UpdateAnnotationSettingArgs, *, session: Session ) -> AnnotationSettingDict: current_user, current_tenant_id = current_account_with_tenant() # get app info - app = db.session.scalar( + app = session.scalar( select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) ) if not app: raise NotFound("App not found") - annotation_setting = db.session.scalar( + annotation_setting = session.scalar( select(AppAnnotationSetting) .where( AppAnnotationSetting.app_id == app_id, @@ -680,8 +686,8 @@ class AppAnnotationService: annotation_setting.score_threshold = args["score_threshold"] annotation_setting.updated_user_id = current_user.id annotation_setting.updated_at = naive_utc_now() - db.session.add(annotation_setting) - db.session.commit() + session.add(annotation_setting) + session.commit() collection_binding_detail = annotation_setting.collection_binding_detail @@ -704,9 +710,9 @@ class AppAnnotationService: } @classmethod - def clear_all_annotations(cls, app_id: str): + def clear_all_annotations(cls, app_id: str, *, session: Session): _, current_tenant_id = current_account_with_tenant() - app = db.session.scalar( + app = session.scalar( select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) ) @@ -714,19 +720,19 @@ class AppAnnotationService: raise NotFound("App not found") # if annotation reply is enabled, delete annotation index - app_annotation_setting = db.session.scalar( + app_annotation_setting = session.scalar( select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == app_id).limit(1) ) - annotations_iter = db.session.scalars( + annotations_iter = session.scalars( select(MessageAnnotation).where(MessageAnnotation.app_id == app_id) ).yield_per(100) for annotation in annotations_iter: - hit_histories_iter = db.session.scalars( + hit_histories_iter = session.scalars( select(AppAnnotationHitHistory).where(AppAnnotationHitHistory.annotation_id == annotation.id) ).yield_per(100) for annotation_hit_history in hit_histories_iter: - db.session.delete(annotation_hit_history) + session.delete(annotation_hit_history) # if annotation reply is enabled, delete annotation index if app_annotation_setting: @@ -734,7 +740,7 @@ class AppAnnotationService: annotation.id, app_id, current_tenant_id, app_annotation_setting.collection_binding_id ) - db.session.delete(annotation) + session.delete(annotation) - db.session.commit() + session.commit() return {"result": "success"} diff --git a/api/services/api_based_extension_service.py b/api/services/api_based_extension_service.py index 25f554b6bdc..e855780d6a1 100644 --- a/api/services/api_based_extension_service.py +++ b/api/services/api_based_extension_service.py @@ -8,7 +8,7 @@ from models.api_based_extension import APIBasedExtension, APIBasedExtensionPoint class APIBasedExtensionService: @staticmethod - def get_all_by_tenant_id(session: Session, tenant_id: str) -> list[APIBasedExtension]: + def get_all_by_tenant_id(tenant_id: str, *, session: Session) -> list[APIBasedExtension]: extension_list = list( session.scalars( select(APIBasedExtension) @@ -23,7 +23,7 @@ class APIBasedExtensionService: return extension_list @classmethod - def save(cls, session: Session, extension_data: APIBasedExtension) -> APIBasedExtension: + def save(cls, extension_data: APIBasedExtension, *, session: Session) -> APIBasedExtension: cls._validation(session, extension_data) extension_data.api_key = encrypt_token(extension_data.tenant_id, extension_data.api_key) @@ -33,12 +33,12 @@ class APIBasedExtensionService: return extension_data @staticmethod - def delete(session: Session, extension_data: APIBasedExtension): + def delete(extension_data: APIBasedExtension, *, session: Session): session.delete(extension_data) session.commit() @staticmethod - def get_with_tenant_id(session: Session, tenant_id: str, api_based_extension_id: str) -> APIBasedExtension: + def get_with_tenant_id(tenant_id: str, api_based_extension_id: str, *, session: Session) -> APIBasedExtension: extension = session.scalar( select(APIBasedExtension) .where(APIBasedExtension.tenant_id == tenant_id, APIBasedExtension.id == api_based_extension_id) diff --git a/api/services/app_dsl_service.py b/api/services/app_dsl_service.py index d042ad69f88..e8c12586856 100644 --- a/api/services/app_dsl_service.py +++ b/api/services/app_dsl_service.py @@ -474,7 +474,7 @@ class AppDslService: ] workflow_service = WorkflowService() - current_draft_workflow = workflow_service.get_draft_workflow(app_model=app) + current_draft_workflow = workflow_service.get_draft_workflow(app_model=app, session=self._session) if current_draft_workflow: unique_hash = current_draft_workflow.unique_hash else: @@ -500,6 +500,7 @@ class AppDslService: account=account, environment_variables=environment_variables, conversation_variables=conversation_variables, + session=self._session, ) case AppMode.CHAT | AppMode.AGENT_CHAT | AppMode.COMPLETION: # Initialize model config @@ -521,7 +522,14 @@ class AppDslService: return app @classmethod - def export_dsl(cls, app_model: App, include_secret: bool = False, workflow_id: str | None = None) -> str: + def export_dsl( + cls, + app_model: App, + *, + session: Session, + include_secret: bool = False, + workflow_id: str | None = None, + ) -> str: """ Export app :param app_model: App instance @@ -548,7 +556,11 @@ class AppDslService: if app_mode in {AppMode.ADVANCED_CHAT, AppMode.WORKFLOW}: cls._append_workflow_export_data( - export_data=export_data, app_model=app_model, include_secret=include_secret, workflow_id=workflow_id + export_data=export_data, + app_model=app_model, + include_secret=include_secret, + workflow_id=workflow_id, + session=session, ) else: cls._append_model_config_export_data(export_data, app_model) @@ -557,7 +569,13 @@ class AppDslService: @classmethod def _append_workflow_export_data( - cls, *, export_data: dict[str, Any], app_model: App, include_secret: bool, workflow_id: str | None = None + cls, + *, + export_data: dict[str, Any], + app_model: App, + include_secret: bool, + session: Session, + workflow_id: str | None = None, ): """ Append workflow export data @@ -565,7 +583,7 @@ class AppDslService: :param app_model: App instance """ workflow_service = WorkflowService() - workflow = workflow_service.get_draft_workflow(app_model, workflow_id) + workflow = workflow_service.get_draft_workflow(app_model, workflow_id, session=session) if not workflow: raise WorkflowNotFoundError("Missing draft workflow configuration, please check.") diff --git a/api/services/app_generate_service.py b/api/services/app_generate_service.py index 3e2c3c96403..940cab5f678 100644 --- a/api/services/app_generate_service.py +++ b/api/services/app_generate_service.py @@ -120,11 +120,12 @@ class AppGenerateService: @trace_span(AppGenerateHandler) def generate( cls, - session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], invoke_from: InvokeFrom, + *, + session: Session, streaming: bool = True, root_node_id: str | None = None, ): @@ -141,13 +142,13 @@ class AppGenerateService: app_model=app_model, streaming=streaming, action=lambda rate_limit, request_id: cls._dispatch_generate( - session=session, app_model=app_model, user=user, args=args, invoke_from=invoke_from, streaming=streaming, root_node_id=root_node_id, + session=session, rate_limit=rate_limit, request_id=request_id, ), @@ -189,13 +190,13 @@ class AppGenerateService: def _dispatch_generate( cls, *, - session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], invoke_from: InvokeFrom, streaming: bool, root_node_id: str | None, + session: Session, rate_limit: RateLimit, request_id: str, ): @@ -251,7 +252,7 @@ class AppGenerateService: ) case AppMode.ADVANCED_CHAT: workflow_id = args.get("workflow_id") - workflow = cls._get_workflow(app_model, invoke_from, workflow_id) + workflow = cls._get_workflow(app_model, invoke_from, workflow_id, session=session) if streaming: # Streaming mode: subscribe to SSE and enqueue the execution on first subscriber @@ -308,7 +309,7 @@ class AppGenerateService: ) case AppMode.WORKFLOW: workflow_id = args.get("workflow_id") - workflow = cls._get_workflow(app_model, invoke_from, workflow_id) + workflow = cls._get_workflow(app_model, invoke_from, workflow_id, session=session) if streaming: with rate_limit_context(rate_limit, request_id): payload = AppExecutionParams.new( @@ -384,12 +385,21 @@ class AppGenerateService: return min(limits) if limits else 0 @classmethod - def generate_single_iteration(cls, app_model: App, user: Account, node_id: str, args: Any, streaming: bool = True): + def generate_single_iteration( + cls, + app_model: App, + user: Account, + node_id: str, + args: Any, + *, + session: Session, + streaming: bool = True, + ): match app_model.mode: case AppMode.COMPLETION | AppMode.CHAT | AppMode.AGENT_CHAT: raise ValueError(f"Invalid app mode {app_model.mode}") case AppMode.ADVANCED_CHAT: - workflow = cls._get_workflow(app_model, InvokeFrom.DEBUGGER) + workflow = cls._get_workflow(app_model, InvokeFrom.DEBUGGER, session=session) return AdvancedChatAppGenerator.convert_to_event_stream( AdvancedChatAppGenerator().single_iteration_generate( app_model=app_model, @@ -401,7 +411,7 @@ class AppGenerateService: ) ) case AppMode.WORKFLOW: - workflow = cls._get_workflow(app_model, InvokeFrom.DEBUGGER) + workflow = cls._get_workflow(app_model, InvokeFrom.DEBUGGER, session=session) return AdvancedChatAppGenerator.convert_to_event_stream( WorkflowAppGenerator().single_iteration_generate( app_model=app_model, @@ -419,13 +429,20 @@ class AppGenerateService: @classmethod def generate_single_loop( - cls, app_model: App, user: Account, node_id: str, args: LoopNodeRunPayload, streaming: bool = True + cls, + app_model: App, + user: Account, + node_id: str, + args: LoopNodeRunPayload, + *, + session: Session, + streaming: bool = True, ): match app_model.mode: case AppMode.COMPLETION | AppMode.CHAT | AppMode.AGENT_CHAT: raise ValueError(f"Invalid app mode {app_model.mode}") case AppMode.ADVANCED_CHAT: - workflow = cls._get_workflow(app_model, InvokeFrom.DEBUGGER) + workflow = cls._get_workflow(app_model, InvokeFrom.DEBUGGER, session=session) return AdvancedChatAppGenerator.convert_to_event_stream( AdvancedChatAppGenerator().single_loop_generate( app_model=app_model, @@ -437,7 +454,7 @@ class AppGenerateService: ) ) case AppMode.WORKFLOW: - workflow = cls._get_workflow(app_model, InvokeFrom.DEBUGGER) + workflow = cls._get_workflow(app_model, InvokeFrom.DEBUGGER, session=session) return AdvancedChatAppGenerator.convert_to_event_stream( WorkflowAppGenerator().single_loop_generate( app_model=app_model, @@ -456,11 +473,12 @@ class AppGenerateService: @classmethod def generate_more_like_this( cls, - session: Session, app_model: App, user: Account | EndUser, message_id: str, invoke_from: InvokeFrom, + *, + session: Session, streaming: bool = True, ) -> Mapping | Generator: """ @@ -482,7 +500,14 @@ class AppGenerateService: ) @classmethod - def _get_workflow(cls, app_model: App, invoke_from: InvokeFrom, workflow_id: str | None = None) -> Workflow: + def _get_workflow( + cls, + app_model: App, + invoke_from: InvokeFrom, + workflow_id: str | None = None, + *, + session: Session, + ) -> Workflow: """ Get workflow :param app_model: app model @@ -498,20 +523,22 @@ class AppGenerateService: _ = uuid.UUID(workflow_id) except ValueError: raise WorkflowIdFormatError(f"Invalid workflow_id format: '{workflow_id}'. ") - workflow = workflow_service.get_published_workflow_by_id(app_model=app_model, workflow_id=workflow_id) + workflow = workflow_service.get_published_workflow_by_id( + app_model=app_model, workflow_id=workflow_id, session=session + ) if not workflow: raise WorkflowNotFoundError(f"Workflow not found with id: {workflow_id}") return workflow if invoke_from == InvokeFrom.DEBUGGER: # fetch draft workflow by app_model - workflow = workflow_service.get_draft_workflow(app_model=app_model) + workflow = workflow_service.get_draft_workflow(app_model=app_model, session=session) if not workflow: raise ValueError("Workflow not initialized") else: # fetch published workflow by app_model - workflow = workflow_service.get_published_workflow(app_model=app_model) + workflow = workflow_service.get_published_workflow(app_model=app_model, session=session) if not workflow: raise ValueError("Workflow not published") diff --git a/api/services/app_service.py b/api/services/app_service.py index 08cd30974e3..139513e87ee 100644 --- a/api/services/app_service.py +++ b/api/services/app_service.py @@ -8,7 +8,7 @@ import sqlalchemy as sa from pydantic import BaseModel, Field from sqlalchemy import ColumnElement, select from sqlalchemy.exc import IntegrityError -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from configs import dify_config from constants.model_template import default_app_templates @@ -18,7 +18,7 @@ from core.model_manager import ModelManager from core.tools.tool_manager import ToolManager from core.tools.utils.configuration import ToolParameterConfigurationManager from events.app_event import app_was_created, app_was_deleted, app_was_updated -from extensions.ext_database import db +from extensions.ext_database import db # noqa: F401 from graphon.model_runtime.entities.model_entities import ModelPropertyKey, ModelType from graphon.model_runtime.model_providers.base.large_language_model import LargeLanguageModel from libs.datetime_utils import naive_utc_now @@ -80,7 +80,7 @@ class CreateAppParams(BaseModel): class AppService: @staticmethod def _build_app_list_filters( - user_id: str, tenant_id: str, params: AppListBaseParams, session: scoped_session + user_id: str, tenant_id: str, params: AppListBaseParams, session: Session ) -> list[sa.ColumnElement[bool]]: filters = [App.tenant_id == tenant_id, App.is_universal == False] @@ -153,13 +153,7 @@ class AppService: }[sort_by] @staticmethod - def get_starred_app_ids( - session: Session | scoped_session, - *, - tenant_id: str, - account_id: str, - app_ids: Sequence[str], - ) -> set[str]: + def get_starred_app_ids(*, tenant_id: str, account_id: str, app_ids: Sequence[str], session: Session) -> set[str]: """Return app IDs starred by this account within the tenant.""" if not app_ids: return set() @@ -174,38 +168,24 @@ class AppService: return set(starred_app_ids) @staticmethod - def get_app_by_id( - session: Session | scoped_session, - app_id: str, - ) -> App | None: + def get_app_by_id(app_id: str, *, session: Session) -> App | None: return session.get(App, app_id) @staticmethod - def get_visible_app_by_id( - session: Session | scoped_session, - app_id: str, - ) -> App | None: + def get_visible_app_by_id(app_id: str, *, session: Session) -> App | None: app = session.get(App, app_id) if not app or app.status != "normal" or not is_openapi_visible(app): return None return app @staticmethod - def find_visible_apps_by_ids( - session: Session | scoped_session, - app_ids: Sequence[str], - ) -> list[App]: + def find_visible_apps_by_ids(app_ids: Sequence[str], *, session: Session) -> list[App]: if not app_ids: return [] return list(session.execute(apply_openapi_gate(select(App).where(App.id.in_(list(app_ids))))).scalars().all()) @staticmethod - def find_visible_apps_by_name( - session: Session | scoped_session, - *, - name: str, - tenant_id: str, - ) -> list[App]: + def find_visible_apps_by_name(*, name: str, tenant_id: str, session: Session) -> list[App]: return list( session.execute( apply_openapi_gate( @@ -219,7 +199,7 @@ class AppService: ) def get_paginate_apps( - self, user_id: str, tenant_id: str, params: AppListParams, session: scoped_session + self, user_id: str, tenant_id: str, params: AppListParams, session: Session ) -> PaginatedResult | None: """ Get app list with pagination, filters, and explicit sort order. @@ -238,14 +218,12 @@ class AppService: sa.select(App).where(*filters).order_by(order_by), page=params.page, per_page=params.limit, + session=session, ) app_ids = [str(app.id) for app in app_models.items] starred_app_ids = self.get_starred_app_ids( - db.session, - tenant_id=tenant_id, - account_id=user_id, - app_ids=app_ids, + tenant_id=tenant_id, account_id=user_id, app_ids=app_ids, session=session ) for app in app_models.items: app.is_starred = str(app.id) in starred_app_ids @@ -253,7 +231,7 @@ class AppService: return app_models def get_paginate_starred_apps( - self, user_id: str, tenant_id: str, params: StarredAppListParams, session: scoped_session + self, user_id: str, tenant_id: str, params: StarredAppListParams, session: Session ) -> PaginatedResult | None: """ Get apps starred by the current account with pagination, filters, and explicit sort order. @@ -277,6 +255,7 @@ class AppService: .order_by(order_by), page=params.page, per_page=params.limit, + session=session, ) for app in app_models.items: @@ -285,7 +264,7 @@ class AppService: return app_models @staticmethod - def star_app(session: Session, *, app: App, account_id: str) -> None: + def star_app(*, app: App, account_id: str, session: Session) -> None: """Create the account's app star if it does not already exist.""" existing_star = session.scalar( select(AppStar) @@ -302,7 +281,7 @@ class AppService: session.add(AppStar(tenant_id=app.tenant_id, app_id=app.id, account_id=account_id)) @staticmethod - def unstar_app(session: Session, *, app: App, account_id: str) -> None: + def unstar_app(*, app: App, account_id: str, session: Session) -> None: """Remove the account's app star if present.""" existing_star = session.scalar( select(AppStar) @@ -318,7 +297,7 @@ class AppService: session.delete(existing_star) - def create_app(self, tenant_id: str, params: CreateAppParams, account: Account) -> App: + def create_app(self, tenant_id: str, params: CreateAppParams, account: Account, *, session: Session) -> App: """ Create app :param tenant_id: tenant id @@ -397,15 +376,15 @@ class AppService: app.maintainer = account.id app.updated_by = account.id - db.session.add(app) - db.session.flush() + session.add(app) + session.flush() if default_model_config: app_model_config = AppModelConfig( **default_model_config, app_id=app.id, created_by=account.id, updated_by=account.id ) - db.session.add(app_model_config) - db.session.flush() + session.add(app_model_config) + session.flush() app.app_model_config_id = app_model_config.id elif app_mode == AppMode.AGENT: @@ -418,8 +397,8 @@ class AppService: # left unset so App.is_agent stays False (this is the new Agent App # type, not a legacy function-call/react agent). agent_app_model_config = AppModelConfig(app_id=app.id, created_by=account.id, updated_by=account.id) - db.session.add(agent_app_model_config) - db.session.flush() + session.add(agent_app_model_config) + session.flush() app.app_model_config_id = agent_app_model_config.id @@ -431,7 +410,7 @@ class AppService: from services.agent.roster_service import AgentRosterService icon_type = AgentIconType(params.icon_type) if params.icon_type else None - AgentRosterService(db.session).create_backing_agent_for_app( + AgentRosterService(session).create_backing_agent_for_app( tenant_id=tenant_id, account_id=account.id, app_id=app.id, @@ -443,7 +422,7 @@ class AppService: icon_background=params.icon_background, ) - db.session.commit() + session.commit() app_was_created.send(app, account=account) enterprise_rbac_service.try_sync_creator_access_policy_member_bindings( @@ -542,10 +521,10 @@ class AppService: role: NotRequired[str | None] @staticmethod - def _get_backing_agent_for_update(app: App) -> Agent | None: + def _get_backing_agent_for_update(app: App, *, session: Session) -> Agent | None: if app.mode != AppMode.AGENT: return None - return db.session.scalar( + return session.scalar( select(Agent).where( Agent.tenant_id == app.tenant_id, Agent.app_id == app.id, @@ -574,6 +553,7 @@ class AppService: icon_background: str | None = None, account_id: str | None = None, updated_at: datetime | None = None, + session: Session, ) -> None: """Keep the Roster identity aligned with its Agent App shell. @@ -584,7 +564,7 @@ class AppService: Role omission is intentional: ``role=None`` preserves the backing Agent's current role, while ``role=""`` explicitly clears it. """ - agent = self._get_backing_agent_for_update(app) + agent = self._get_backing_agent_for_update(app, session=session) if agent is None: return @@ -605,16 +585,16 @@ class AppService: agent.updated_at = updated_at @staticmethod - def _commit_app_identity_update(app: App) -> None: + def _commit_app_identity_update(app: App, *, session: Session) -> None: try: - db.session.commit() + session.commit() except IntegrityError as exc: - db.session.rollback() + session.rollback() if app.mode == AppMode.AGENT: raise AgentNameConflictError() from exc raise - def update_app(self, app: App, args: ArgsDict) -> App: + def update_app(self, app: App, args: ArgsDict, *, session: Session) -> App: """ Update app :param app: App instance @@ -649,14 +629,15 @@ class AppService: icon_background=app.icon_background, account_id=current_user.id, updated_at=app.updated_at, + session=session, ) - self._commit_app_identity_update(app) + self._commit_app_identity_update(app, session=session) app_was_updated.send(app) return app - def update_app_name(self, app: App, name: str) -> App: + def update_app_name(self, app: App, name: str, *, session: Session) -> App: """ Update app name :param app: App instance @@ -672,15 +653,22 @@ class AppService: name=app.name, account_id=current_user.id, updated_at=app.updated_at, + session=session, ) - self._commit_app_identity_update(app) + self._commit_app_identity_update(app, session=session) app_was_updated.send(app) return app def update_app_icon( - self, app: App, icon: str, icon_background: str, icon_type: IconType | str | None = None + self, + app: App, + icon: str, + icon_background: str, + icon_type: IconType | str | None = None, + *, + session: Session, ) -> App: """ Update app icon @@ -704,14 +692,15 @@ class AppService: icon_background=app.icon_background, account_id=current_user.id, updated_at=app.updated_at, + session=session, ) - db.session.commit() + session.commit() app_was_updated.send(app) return app - def update_app_site_status(self, app: App, enable_site: bool) -> App: + def update_app_site_status(self, app: App, enable_site: bool, *, session: Session) -> App: """ Update app site status :param app: App instance @@ -724,13 +713,13 @@ class AppService: app.enable_site = enable_site app.updated_by = current_user.id app.updated_at = naive_utc_now() - db.session.commit() + session.commit() app_was_updated.send(app) return app - def update_app_api_status(self, app: App, enable_api: bool) -> App: + def update_app_api_status(self, app: App, enable_api: bool, *, session: Session) -> App: """ Update app api status :param app: App instance @@ -744,20 +733,20 @@ class AppService: app.enable_api = enable_api app.updated_by = current_user.id app.updated_at = naive_utc_now() - db.session.commit() + session.commit() app_was_updated.send(app) return app - def delete_app(self, app: App): + def delete_app(self, app: App, *, session: Session) -> None: """ Delete app :param app: App instance """ app_was_deleted.send(app) - backing_agent = self._get_backing_agent_for_update(app) + backing_agent = self._get_backing_agent_for_update(app, session=session) if backing_agent is not None: now = naive_utc_now() account_id = getattr(current_user, "id", None) @@ -767,8 +756,8 @@ class AppService: backing_agent.updated_by = account_id backing_agent.updated_at = now - db.session.delete(app) - db.session.commit() + session.delete(app) + session.commit() # clean up web app settings if FeatureService.get_system_features().webapp_auth.enabled: @@ -780,7 +769,7 @@ class AppService: # Trigger asynchronous deletion of app and related data remove_app_and_related_data_task.delay(tenant_id=app.tenant_id, app_id=app.id) - def get_app_meta(self, app_model: App): + def get_app_meta(self, app_model: App, *, session: Session): """ Get app meta info :param app_model: app model @@ -833,7 +822,7 @@ class AppService: meta["tool_icons"][tool_name] = url_prefix + provider_id + "/icon" elif provider_type == "api": try: - provider: ApiToolProvider | None = db.session.get(ApiToolProvider, provider_id) + provider: ApiToolProvider | None = session.get(ApiToolProvider, provider_id) if provider is None: raise ValueError(f"provider not found for tool {tool_name}") meta["tool_icons"][tool_name] = json.loads(provider.icon) @@ -843,25 +832,25 @@ class AppService: return meta @staticmethod - def get_app_code_by_id(app_id: str) -> str: + def get_app_code_by_id(app_id: str, *, session: Session) -> str: """ Get app code by app id :param app_id: app id :return: app code """ - site = db.session.scalar(select(Site).where(Site.app_id == app_id).limit(1)) + site = session.scalar(select(Site).where(Site.app_id == app_id).limit(1)) if not site: raise ValueError(f"App with id {app_id} not found") return str(site.code) @staticmethod - def get_app_id_by_code(app_code: str) -> str: + def get_app_id_by_code(app_code: str, *, session: Session) -> str: """ Get app id by app code :param app_code: app code :return: app id """ - site = db.session.scalar(select(Site).where(Site.code == app_code).limit(1)) + site = session.scalar(select(Site).where(Site.code == app_code).limit(1)) if not site: raise ValueError(f"App with code {app_code} not found") return str(site.app_id) diff --git a/api/services/async_workflow_service.py b/api/services/async_workflow_service.py index ceda30e950f..601cad7557a 100644 --- a/api/services/async_workflow_service.py +++ b/api/services/async_workflow_service.py @@ -51,7 +51,7 @@ class AsyncWorkflowService: @classmethod def trigger_workflow_async( - cls, session: Session, user: Account | EndUser, trigger_data: TriggerData + cls, user: Account | EndUser, trigger_data: TriggerData, *, session: Session ) -> AsyncTriggerResponse: """ Universal entry point for async workflow execution - THIS METHOD WILL NOT BLOCK @@ -187,7 +187,7 @@ class AsyncWorkflowService: @classmethod def reinvoke_trigger( - cls, session: Session, user: Account | EndUser, workflow_trigger_log_id: str + cls, user: Account | EndUser, workflow_trigger_log_id: str, *, session: Session ) -> AsyncTriggerResponse: """ Re-invoke a previously failed or rate-limited trigger - THIS METHOD WILL NOT BLOCK @@ -231,7 +231,7 @@ class AsyncWorkflowService: session.commit() # Re-trigger workflow (this will create a new trigger log) - return cls.trigger_workflow_async(session, user, trigger_data) + return cls.trigger_workflow_async(user, trigger_data, session=session) @classmethod def get_trigger_log( @@ -309,7 +309,8 @@ class AsyncWorkflowService: workflow_service: WorkflowService, app_model: App, workflow_id: str | None = None, - session: Session | None = None, + *, + session: Session, ) -> Workflow: """ Get workflow for the app @@ -317,9 +318,7 @@ class AsyncWorkflowService: Args: app_model: App model instance workflow_id: Optional specific workflow ID - session: Reuse this SQLAlchemy session for the lookup when provided, - so the caller's explicit session bears the connection cost - instead of Flask's request-scoped ``db.session``. + session: SQLAlchemy session used for the workflow lookup. Returns: Workflow instance diff --git a/api/services/audio_service.py b/api/services/audio_service.py index 86c56e60a13..52c71edd576 100644 --- a/api/services/audio_service.py +++ b/api/services/audio_service.py @@ -6,7 +6,7 @@ from typing import cast from flask import Response, stream_with_context from sqlalchemy import select -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from werkzeug.datastructures import FileStorage from constants import AUDIO_EXTENSIONS @@ -32,7 +32,7 @@ logger = logging.getLogger(__name__) class AudioService: @staticmethod - def _get_message_by_ref(session: Session | scoped_session, message_ref: MessageRef) -> Message | None: + def _get_message_by_ref(session: Session, message_ref: MessageRef) -> Message | None: stmt = select(Message).where(Message.id == message_ref.message_id, Message.app_id == message_ref.app_id) if message_ref.end_user_id is not None: stmt = stmt.where(Message.from_end_user_id == message_ref.end_user_id) @@ -89,7 +89,7 @@ class AudioService: cls, app_model: App, *, - session: Session | scoped_session, + session: Session, text: str | None = None, voice: str | None = None, end_user: str | None = None, diff --git a/api/services/auth/api_key_auth_service.py b/api/services/auth/api_key_auth_service.py index 42f1d4d8d40..f9ad7cf27b0 100644 --- a/api/services/auth/api_key_auth_service.py +++ b/api/services/auth/api_key_auth_service.py @@ -11,7 +11,7 @@ from services.auth.api_key_auth_factory import ApiKeyAuthFactory class ApiKeyAuthService: @staticmethod - def get_provider_auth_list(session: Session, tenant_id: str): + def get_provider_auth_list(tenant_id: str, *, session: Session): data_source_api_key_bindings = session.scalars( select(DataSourceApiKeyAuthBinding).where( DataSourceApiKeyAuthBinding.tenant_id == tenant_id, DataSourceApiKeyAuthBinding.disabled.is_(False) @@ -20,7 +20,7 @@ class ApiKeyAuthService: return data_source_api_key_bindings @staticmethod - def create_provider_auth(session: Session, tenant_id: str, args: dict[str, Any]): + def create_provider_auth(tenant_id: str, args: dict[str, Any], *, session: Session): auth_result = ApiKeyAuthFactory(args["provider"], args["credentials"]).validate_credentials() if auth_result: # Encrypt the api key @@ -35,7 +35,7 @@ class ApiKeyAuthService: session.commit() @staticmethod - def get_auth_credentials(session: Session, tenant_id: str, category: str, provider: str): + def get_auth_credentials(tenant_id: str, category: str, provider: str, *, session: Session): data_source_api_key_bindings = session.scalar( select(DataSourceApiKeyAuthBinding).where( DataSourceApiKeyAuthBinding.tenant_id == tenant_id, @@ -52,7 +52,7 @@ class ApiKeyAuthService: return credentials @staticmethod - def delete_provider_auth(session: Session, tenant_id: str, binding_id: str): + def delete_provider_auth(tenant_id: str, binding_id: str, *, session: Session): data_source_api_key_binding = session.scalar( select(DataSourceApiKeyAuthBinding).where( DataSourceApiKeyAuthBinding.tenant_id == tenant_id, diff --git a/api/services/billing_service.py b/api/services/billing_service.py index ec391e51676..2ee7179f432 100644 --- a/api/services/billing_service.py +++ b/api/services/billing_service.py @@ -7,7 +7,7 @@ from typing import Any, Literal, NotRequired, TypedDict import httpx from pydantic import TypeAdapter from sqlalchemy import select -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from tenacity import retry, retry_if_exception_type, stop_before_delay, wait_fixed from werkzeug.exceptions import InternalServerError @@ -363,7 +363,7 @@ class BillingService: return response.json() @staticmethod - def is_tenant_owner_or_admin(session: Session | scoped_session, current_user: Account): + def is_tenant_owner_or_admin(current_user: Account, *, session: Session): tenant_id = current_user.current_tenant_id join: TenantAccountJoin | None = session.scalar( diff --git a/api/services/conversation_service.py b/api/services/conversation_service.py index 557ae8e89f3..7c3b8d451c5 100644 --- a/api/services/conversation_service.py +++ b/api/services/conversation_service.py @@ -8,16 +8,13 @@ from sqlalchemy.orm import Session from configs import dify_config from core.app.entities.app_invoke_entities import InvokeFrom -from core.db.session_factory import session_factory from core.llm_generator.llm_generator import LLMGenerator -from extensions.ext_database import db from factories import variable_factory from graphon.variables.types import SegmentType from libs.datetime_utils import naive_utc_now from libs.infinite_scroll_pagination import InfiniteScrollPagination from models import Account, ConversationVariable from models.model import App, Conversation, EndUser, Message -from services.conversation_variable_updater import ConversationVariableUpdater from services.errors.conversation import ( ConversationNotExistsError, ConversationVariableNotExistsError, @@ -122,24 +119,26 @@ class ConversationService: user: Account | EndUser | None, name: str | None, auto_generate: bool, + *, + session: Session, ): - conversation = cls.get_conversation(app_model, conversation_id, user) + conversation = cls.get_conversation(app_model, conversation_id, user, session=session) if auto_generate: - return cls.auto_generate_name(app_model, conversation) + return cls.auto_generate_name(app_model, conversation, session=session) else: if name is None: raise ValueError("name is required when auto_generate is false") conversation.name = name conversation.updated_at = naive_utc_now() - db.session.commit() + session.commit() return conversation @classmethod - def auto_generate_name(cls, app_model: App, conversation: Conversation): + def auto_generate_name(cls, app_model: App, conversation: Conversation, *, session: Session): # get conversation first message - message = db.session.scalar( + message = session.scalar( select(Message) .where(Message.app_id == app_model.id, Message.conversation_id == conversation.id) .order_by(Message.created_at.asc()) @@ -156,13 +155,15 @@ class ConversationService: ) conversation.name = name - db.session.commit() + session.commit() return conversation @classmethod - def get_conversation(cls, app_model: App, conversation_id: str, user: Account | EndUser | None): - conversation = db.session.scalar( + def get_conversation( + cls, app_model: App, conversation_id: str, user: Account | EndUser | None, *, session: Session + ): + conversation = session.scalar( select(Conversation) .where( Conversation.id == conversation_id, @@ -181,14 +182,14 @@ class ConversationService: return conversation @classmethod - def delete(cls, app_model: App, conversation_id: str, user: Account | EndUser | None): + def delete(cls, app_model: App, conversation_id: str, user: Account | EndUser | None, *, session: Session): """ Delete a conversation only if it belongs to the given user and app context. Raises: ConversationNotExistsError: When the conversation is not visible to the current user. """ - conversation = cls.get_conversation(app_model, conversation_id, user) + conversation = cls.get_conversation(app_model, conversation_id, user, session=session) try: logger.info( @@ -197,13 +198,13 @@ class ConversationService: conversation_id, ) - db.session.delete(conversation) - db.session.commit() + session.delete(conversation) + session.commit() delete_conversation_related_data.delay(conversation.id) except Exception as e: - db.session.rollback() + session.rollback() raise e @classmethod @@ -215,8 +216,10 @@ class ConversationService: limit: int, last_id: str | None, variable_name: str | None = None, + *, + session: Session, ) -> InfiniteScrollPagination: - conversation = cls.get_conversation(app_model, conversation_id, user) + conversation = cls.get_conversation(app_model, conversation_id, user, session=session) stmt = ( select(ConversationVariable) @@ -245,18 +248,17 @@ class ConversationService: ) ) - with session_factory.create_session() as session: - if last_id: - last_variable = session.scalar(stmt.where(ConversationVariable.id == last_id)) - if not last_variable: - raise ConversationVariableNotExistsError() + if last_id: + last_variable = session.scalar(stmt.where(ConversationVariable.id == last_id)) + if not last_variable: + raise ConversationVariableNotExistsError() - # Filter for variables created after the last_id - stmt = stmt.where(ConversationVariable.created_at > last_variable.created_at) + # Filter for variables created after the last_id + stmt = stmt.where(ConversationVariable.created_at > last_variable.created_at) - # Apply limit to query: fetch one extra row to determine has_more - query_stmt = stmt.limit(limit + 1) - rows = session.scalars(query_stmt).all() + # Apply limit to query: fetch one extra row to determine has_more + query_stmt = stmt.limit(limit + 1) + rows = session.scalars(query_stmt).all() has_more = False if len(rows) > limit: @@ -282,6 +284,8 @@ class ConversationService: variable_id: str, user: Account | EndUser | None, new_value: Any, + *, + session: Session, ): """ Update a conversation variable's value. @@ -302,7 +306,7 @@ class ConversationService: ConversationVariableTypeMismatchError: If the new value type doesn't match the variable's expected type """ # Verify conversation exists and user has access - conversation = cls.get_conversation(app_model, conversation_id, user) + conversation = cls.get_conversation(app_model, conversation_id, user, session=session) # Get the existing conversation variable stmt = ( @@ -312,48 +316,43 @@ class ConversationService: .where(ConversationVariable.id == variable_id) ) - with session_factory.create_session() as session: - existing_variable = session.scalar(stmt) - if not existing_variable: - raise ConversationVariableNotExistsError() + existing_variable = session.scalar(stmt) + if not existing_variable: + raise ConversationVariableNotExistsError() - # Convert existing variable to Variable object - current_variable = existing_variable.to_variable() + # Convert existing variable to Variable object + current_variable = existing_variable.to_variable() - # Validate that the new value type matches the expected variable type - expected_type = SegmentType(current_variable.value_type) + # Validate that the new value type matches the expected variable type + expected_type = SegmentType(current_variable.value_type) - # There is showing number in web ui but int in db - if expected_type == SegmentType.INTEGER: - expected_type = SegmentType.NUMBER + # There is showing number in web ui but int in db + if expected_type == SegmentType.INTEGER: + expected_type = SegmentType.NUMBER - if not expected_type.is_valid(new_value): - inferred_type = SegmentType.infer_segment_type(new_value) - raise ConversationVariableTypeMismatchError( - f"Type mismatch: variable '{current_variable.name}' expects {expected_type.value}, " - f"but got {inferred_type.value if inferred_type else 'unknown'} type" - ) + if not expected_type.is_valid(new_value): + inferred_type = SegmentType.infer_segment_type(new_value) + raise ConversationVariableTypeMismatchError( + f"Type mismatch: variable '{current_variable.name}' expects {expected_type.value}, " + f"but got {inferred_type.value if inferred_type else 'unknown'} type" + ) - # Create updated variable with new value only, preserving everything else - updated_variable_dict = { - "id": current_variable.id, - "name": current_variable.name, - "description": current_variable.description, - "value_type": current_variable.value_type, - "value": new_value, - "selector": current_variable.selector, - } + # Create updated variable with new value only, preserving everything else + updated_variable_dict = { + "id": current_variable.id, + "name": current_variable.name, + "description": current_variable.description, + "value_type": current_variable.value_type, + "value": new_value, + "selector": current_variable.selector, + } - updated_variable = variable_factory.build_conversation_variable_from_mapping(updated_variable_dict) + updated_variable = variable_factory.build_conversation_variable_from_mapping(updated_variable_dict) + existing_variable.data = updated_variable.model_dump_json() + session.commit() - # Use the conversation variable updater to persist the changes - updater = ConversationVariableUpdater(session_factory.get_session_maker()) - updater.update(conversation_id, updated_variable) - updater.flush() - - # Return the updated variable data - return { - "created_at": existing_variable.created_at, - "updated_at": naive_utc_now(), # Update timestamp - **updated_variable.model_dump(), - } + return { + "created_at": existing_variable.created_at, + "updated_at": naive_utc_now(), # Update timestamp + **updated_variable.model_dump(), + } diff --git a/api/services/credential_permission_service.py b/api/services/credential_permission_service.py index 2b1082d132b..d9ce5e7c502 100644 --- a/api/services/credential_permission_service.py +++ b/api/services/credential_permission_service.py @@ -1,7 +1,7 @@ from collections.abc import Sequence from sqlalchemy import or_, select -from sqlalchemy.orm import InstrumentedAttribute, Session, scoped_session +from sqlalchemy.orm import InstrumentedAttribute, Session from models.account import Account from models.credential_permission import CredentialPermission @@ -16,9 +16,7 @@ class CredentialPermissionService: """ @classmethod - def get_partial_member_list( - cls, session: Session | scoped_session, credential_id: str, credential_type: str - ) -> Sequence[str]: + def get_partial_member_list(cls, credential_id: str, credential_type: str, *, session: Session) -> Sequence[str]: """Return account_ids that have partial-member access to a credential.""" return session.scalars( select(CredentialPermission.account_id).where( diff --git a/api/services/credit_pool_service.py b/api/services/credit_pool_service.py index 94515309e79..afc49181185 100644 --- a/api/services/credit_pool_service.py +++ b/api/services/credit_pool_service.py @@ -12,9 +12,7 @@ from sqlalchemy import select from sqlalchemy.orm import Session from configs import dify_config -from core.db.session_factory import session_factory from core.errors.error import QuotaExceededError -from extensions.ext_database import db from extensions.ext_redis import redis_client from models import TenantCreditPool from models.enums import ProviderQuotaType @@ -66,7 +64,7 @@ class CreditPoolService: ) @classmethod - def create_default_pool(cls, tenant_id: str) -> TenantCreditPool: + def create_default_pool(cls, tenant_id: str, session: Session) -> TenantCreditPool: """create default credit pool for new tenant""" credit_pool = TenantCreditPool( tenant_id=tenant_id, @@ -74,22 +72,21 @@ class CreditPoolService: quota_used=0, pool_type=ProviderQuotaType.TRIAL, ) - db.session.add(credit_pool) - db.session.commit() + session.add(credit_pool) + session.commit() return credit_pool @classmethod - def get_pool(cls, tenant_id: str, pool_type: str = "trial") -> TenantCreditPool | None: + def get_pool(cls, tenant_id: str, pool_type: str = "trial", *, session: Session) -> TenantCreditPool | None: """get tenant credit pool""" - with session_factory.get_session_maker().begin() as session: - return session.scalar( - select(TenantCreditPool) - .where( - TenantCreditPool.tenant_id == tenant_id, - TenantCreditPool.pool_type == pool_type, - ) - .limit(1) + return session.scalar( + select(TenantCreditPool) + .where( + TenantCreditPool.tenant_id == tenant_id, + TenantCreditPool.pool_type == pool_type, ) + .limit(1) + ) @classmethod def check_credits_available( @@ -97,9 +94,11 @@ class CreditPoolService: tenant_id: str, credits_required: int, pool_type: str = "trial", + *, + session: Session, ) -> bool: """check if credits are available without deducting""" - pool = cls.get_pool(tenant_id, pool_type) + pool = cls.get_pool(tenant_id, pool_type, session=session) if not pool: return False return pool.remaining_credits >= credits_required @@ -110,25 +109,27 @@ class CreditPoolService: tenant_id: str, credits_required: int, pool_type: str = "trial", + *, + session: Session, ) -> int: """Deduct exactly the requested credits or raise without mutating the pool.""" if credits_required <= 0: return 0 def deduct() -> int: - with session_factory.get_session_maker().begin() as session: - pool = cls._get_locked_pool(session=session, tenant_id=tenant_id, pool_type=pool_type) - if not pool: - raise QuotaExceededError("Credit pool not found") + pool = cls._get_locked_pool(session=session, tenant_id=tenant_id, pool_type=pool_type) + if not pool: + raise QuotaExceededError("Credit pool not found") - remaining_credits = pool.remaining_credits - if remaining_credits <= 0: - raise QuotaExceededError("No credits remaining") - if remaining_credits < credits_required: - raise QuotaExceededError("Insufficient credits remaining") + remaining_credits = pool.remaining_credits + if remaining_credits <= 0: + raise QuotaExceededError("No credits remaining") + if remaining_credits < credits_required: + raise QuotaExceededError("Insufficient credits remaining") - pool.quota_used += credits_required - return credits_required + pool.quota_used += credits_required + session.commit() + return credits_required try: return cls._deduct_with_tenant_lock(tenant_id, deduct) @@ -144,24 +145,26 @@ class CreditPoolService: tenant_id: str, credits_required: int, pool_type: str = "trial", + *, + session: Session, ) -> int: """Deduct up to the available balance and return the actual deducted credits.""" if credits_required <= 0: return 0 def deduct() -> int: - with session_factory.get_session_maker().begin() as session: - pool = cls._get_locked_pool(session=session, tenant_id=tenant_id, pool_type=pool_type) - if not pool: - logger.warning("Credit pool not found, tenant_id=%s, pool_type=%s", tenant_id, pool_type) - return 0 + pool = cls._get_locked_pool(session=session, tenant_id=tenant_id, pool_type=pool_type) + if not pool: + logger.warning("Credit pool not found, tenant_id=%s, pool_type=%s", tenant_id, pool_type) + return 0 - deducted_credits = min(credits_required, pool.remaining_credits) - if deducted_credits <= 0: - return 0 + deducted_credits = min(credits_required, pool.remaining_credits) + if deducted_credits <= 0: + return 0 - pool.quota_used += deducted_credits - return deducted_credits + pool.quota_used += deducted_credits + session.commit() + return deducted_credits try: return cls._deduct_with_tenant_lock(tenant_id, deduct) diff --git a/api/services/data_migration/export_service.py b/api/services/data_migration/export_service.py index f5d214d230b..c0233006690 100644 --- a/api/services/data_migration/export_service.py +++ b/api/services/data_migration/export_service.py @@ -120,8 +120,8 @@ class MigrationExportService: self.package_service = package_service or MigrationPackageService() self.dependency_discovery_service = dependency_discovery_service or DependencyDiscoveryService() - def export(self, session: Session, selection: ExportSelection) -> ExportResult: - tenant = self._get_tenant(session, selection) + def export(self, selection: ExportSelection, *, session: Session) -> ExportResult: + tenant = self._get_tenant(selection, session=session) package = self.package_service.build_empty_package( source_tenant_id=tenant.id, source_tenant_name=tenant.name, @@ -131,10 +131,12 @@ class MigrationExportService: report_items: list[ResourceReportItem] = [] discovered_dependencies: list[DiscoveredDependency] = [] - apps = self._selected_apps(session, tenant.id, selection) + apps = self._selected_apps(tenant.id, selection, session=session) exported_app_ids = {app.id for app in apps} for app in apps: - dsl_content = AppDslService.export_dsl(app_model=app, include_secret=selection.include_secrets) + dsl_content = AppDslService.export_dsl( + app_model=app, session=session, include_secret=selection.include_secrets + ) package.workflows.append( { "id": app.id, @@ -157,7 +159,6 @@ class MigrationExportService: report_items=report_items, ) self._export_workflow_tools( - session, tenant, self._provider_ids( selection.additional_workflow_tools, discovered_dependencies, DependencyKind.WORKFLOW_TOOL @@ -166,9 +167,9 @@ class MigrationExportService: exported_workflow_tools=package.workflow_tools, dependencies=package.dependencies, report_items=report_items, + session=session, ) self._export_mcp_tools( - session, tenant_id=tenant.id, provider_ids=self._provider_ids( selection.additional_mcp_tools, @@ -179,6 +180,7 @@ class MigrationExportService: exported_mcp_tools=package.mcp_tools, dependencies=package.dependencies, report_items=report_items, + session=session, ) self._record_dependency_metadata( self._dependencies_by_kind(discovered_dependencies, DependencyKind.BUILTIN_OR_PLUGIN_TOOL), @@ -195,7 +197,7 @@ class MigrationExportService: ), ) - def _get_tenant(self, session: Session, selection: ExportSelection) -> Tenant: + def _get_tenant(self, selection: ExportSelection, *, session: Session) -> Tenant: if selection.source_tenant_id: tenant = session.get(Tenant, selection.source_tenant_id) if tenant is None: @@ -214,7 +216,7 @@ class MigrationExportService: ) return tenants[0] - def _selected_apps(self, session: Session, tenant_id: str, selection: ExportSelection) -> list[App]: + def _selected_apps(self, tenant_id: str, selection: ExportSelection, *, session: Session) -> list[App]: query = sa.select(App).where(App.tenant_id == tenant_id, App.mode.in_(SUPPORTED_APP_MODES)) if not selection.export_all_apps: if not selection.app_ids: @@ -267,7 +269,6 @@ class MigrationExportService: def _export_workflow_tools( self, - session: Session, tenant: Tenant, provider_ids: Iterable[str], *, @@ -275,11 +276,12 @@ class MigrationExportService: exported_workflow_tools: list[dict[str, Any]], dependencies: list[dict[str, Any]], report_items: list[ResourceReportItem], + session: Session, ) -> None: provider_ids = self._dedupe(provider_ids) if not provider_ids: return - owner = self._get_tenant_owner(session, tenant.id) + owner = self._get_tenant_owner(tenant.id, session=session) if owner is None: for provider_id in provider_ids: report_items.append( @@ -330,7 +332,7 @@ class MigrationExportService: ResourceReportItem(ResourceType.WORKFLOW_TOOL, provider_id, provider_id, "unresolved", str(exc)) ) - def _get_tenant_owner(self, session: Session, tenant_id: str) -> Account | None: + def _get_tenant_owner(self, tenant_id: str, *, session: Session) -> Account | None: return session.scalar( sa.select(Account) .join(TenantAccountJoin, Account.id == TenantAccountJoin.account_id) @@ -341,7 +343,6 @@ class MigrationExportService: def _export_mcp_tools( self, - session: Session, *, tenant_id: str, provider_ids: Iterable[str], @@ -349,6 +350,7 @@ class MigrationExportService: exported_mcp_tools: list[dict[str, Any]], dependencies: list[dict[str, Any]], report_items: list[ResourceReportItem], + session: Session, ) -> None: for provider_id in self._dedupe(provider_ids): if not include_secrets: @@ -359,7 +361,7 @@ class MigrationExportService: ) continue try: - provider = self._get_mcp_provider(session, tenant_id, provider_id) + provider = self._get_mcp_provider(tenant_id, provider_id, session=session) exported_mcp_tools.append(self._serialize_mcp_provider(provider)) report_items.append(ResourceReportItem(ResourceType.MCP_TOOL, provider_id, provider.name, "exported")) except Exception as exc: @@ -367,7 +369,7 @@ class MigrationExportService: ResourceReportItem(ResourceType.MCP_TOOL, provider_id, provider_id, "unresolved", str(exc)) ) - def _get_mcp_provider(self, session: Session, tenant_id: str, provider_id: str) -> MCPToolProvider: + def _get_mcp_provider(self, tenant_id: str, provider_id: str, *, session: Session) -> MCPToolProvider: predicates = [MCPToolProvider.server_identifier == provider_id] if self._is_uuid_string(provider_id): predicates.append(MCPToolProvider.id == provider_id) diff --git a/api/services/data_migration/import_service.py b/api/services/data_migration/import_service.py index 3eb251bbaef..b3354413ba1 100644 --- a/api/services/data_migration/import_service.py +++ b/api/services/data_migration/import_service.py @@ -82,11 +82,11 @@ class ImportTargetResolver: "Target tenant must be provided by --target-tenant, import config, or package metadata." ) - def resolve(self, session: Session, request: ImportRequest) -> ImportTarget: + def resolve(self, request: ImportRequest, *, session: Session) -> ImportTarget: target_tenant_name = self.select_target_tenant_name(request) package_target = request.package.metadata.target_tenant or {} if request.cli_target_tenant or request.config_target_tenant: - tenant = self._resolve_tenant_by_id_or_name(session, target_tenant_name) + tenant = self._resolve_tenant_by_id_or_name(target_tenant_name, session=session) elif package_target.get("id") and self._is_uuid(package_target["id"]): tenant = session.get(Tenant, package_target["id"]) if tenant is not None and package_target.get("name") and tenant.name != package_target.get("name"): @@ -94,7 +94,7 @@ class ImportTargetResolver: f"Target tenant id/name mismatch: {package_target['id']} / {package_target['name']}" ) else: - tenant = self._resolve_tenant_by_id_or_name(session, target_tenant_name) + tenant = self._resolve_tenant_by_id_or_name(target_tenant_name, session=session) if tenant is None: raise MigrationDataError(f"Target tenant not found: {target_tenant_name}") @@ -123,7 +123,7 @@ class ImportTargetResolver: operator_email=account.email, ) - def _resolve_tenant_by_id_or_name(self, session: Session, value: str) -> Tenant | None: + def _resolve_tenant_by_id_or_name(self, value: str, *, session: Session) -> Tenant | None: if self._is_uuid(value): tenant = session.get(Tenant, value) if tenant is not None: @@ -149,8 +149,8 @@ class MigrationImportService: def __init__(self, *, target_resolver: ImportTargetResolver | None = None) -> None: self.target_resolver = target_resolver or ImportTargetResolver() - def import_package(self, session: Session, request: ImportRequest) -> ImportResult: - target = self.target_resolver.resolve(session, request) + def import_package(self, request: ImportRequest, *, session: Session) -> ImportResult: + target = self.target_resolver.resolve(request, session=session) options = request.options_override or request.package.metadata.import_options report_items = [ ResourceReportItem( @@ -165,7 +165,6 @@ class MigrationImportService: id_mapping_details: list[ResourceIdMapping] = [] self._import_api_tools( - session, request.package, target, options, @@ -173,14 +172,16 @@ class MigrationImportService: id_mapping, id_mapping_details, self._source_api_provider_ids_by_name(request.package), + session=session, ) - self._import_mcp_tools(session, request.package, target, options, report_items, id_mapping, id_mapping_details) - self._preflight_dependency_only_mcp(session, request.package, target, report_items) + self._import_mcp_tools( + request.package, target, options, report_items, id_mapping, id_mapping_details, session=session + ) + self._preflight_dependency_only_mcp(request.package, target, report_items, session=session) workflow_tool_app_ids = self._workflow_tool_source_app_ids(request.package) imported_workflow_ids: set[str] = set() if workflow_tool_app_ids: self._import_workflows( - session, request.package, target, options, @@ -189,12 +190,12 @@ class MigrationImportService: id_mapping_details=id_mapping_details, imported_workflow_ids=imported_workflow_ids, only_app_ids=workflow_tool_app_ids, + session=session, ) self._import_workflow_tools( - session, request.package, target, options, id_mapping, id_mapping_details, report_items + request.package, target, options, id_mapping, id_mapping_details, report_items, session=session ) self._import_workflows( - session, request.package, target, options, @@ -203,6 +204,7 @@ class MigrationImportService: id_mapping_details=id_mapping_details, imported_workflow_ids=imported_workflow_ids, skip_app_ids=imported_workflow_ids, + session=session, ) return ImportResult( report_items=report_items, @@ -218,7 +220,6 @@ class MigrationImportService: def _import_workflows( self, - session: Session, package: MigrationPackage, target: ImportTarget, options: ImportOptions, @@ -228,6 +229,8 @@ class MigrationImportService: imported_workflow_ids: set[str] | None = None, only_app_ids: set[str] | None = None, skip_app_ids: set[str] | None = None, + *, + session: Session, ) -> None: account = session.get(Account, target.operator_id) tenant = session.get(Tenant, target.tenant_id) @@ -248,7 +251,7 @@ class MigrationImportService: id_mapping, ) existing_app = ( - self._find_existing_app(session, app_id, target.tenant_id) + self._find_existing_app(app_id, target.tenant_id, session=session) if options.id_strategy == IdStrategy.PRESERVE_ID else None ) @@ -270,13 +273,13 @@ class MigrationImportService: continue imported_app_id = self._import_workflow_app( - session=session, account=account, workflow_data=workflow_data, dsl_content=dsl_content, app_id=app_id, existing_app=existing_app, options=options, + session=session, ) if app_id: self._record_id_mappings( @@ -290,7 +293,7 @@ class MigrationImportService: if imported_workflow_ids is not None: imported_workflow_ids.add(app_id) if options.create_app_api_token_on_import: - self._create_or_reuse_app_api_token(session, imported_app_id, target.tenant_id) + self._create_or_reuse_app_api_token(imported_app_id, target.tenant_id, session=session) report_items.append( ResourceReportItem( ResourceType.WORKFLOW, @@ -311,15 +314,15 @@ class MigrationImportService: def _import_workflow_app( self, *, - session: Session, account: Account, workflow_data: dict[str, object], dsl_content: str, app_id: str | None, existing_app: App | None, options: ImportOptions, + session: Session, ) -> str: - import_service = AppDslService(session) + import_service = AppDslService(cast(Session, session)) if existing_app is not None: import_result = import_service.import_app( account=account, @@ -408,12 +411,12 @@ class MigrationImportService: def _should_preserve_source_app_id(self, options: ImportOptions) -> bool: return options.id_strategy == IdStrategy.PRESERVE_ID - def _find_existing_app(self, session: Session, app_id: str | None, tenant_id: str) -> App | None: + def _find_existing_app(self, app_id: str | None, tenant_id: str, *, session: Session) -> App | None: if not self._is_uuid_string(app_id): return None return session.scalar(sa.select(App).where(App.id == app_id, App.tenant_id == tenant_id)) - def _create_or_reuse_app_api_token(self, session: Session, app_id: str, tenant_id: str) -> None: + def _create_or_reuse_app_api_token(self, app_id: str, tenant_id: str, *, session: Session) -> None: existing = session.scalar( sa.select(ApiToken).where( ApiToken.type == ApiTokenType.APP, @@ -433,7 +436,6 @@ class MigrationImportService: def _import_api_tools( self, - session: Session, package: MigrationPackage, target: ImportTarget, options: ImportOptions, @@ -441,6 +443,8 @@ class MigrationImportService: id_mapping: dict[str, str], id_mapping_details: list[ResourceIdMapping], source_provider_ids_by_name: dict[str, set[str]], + *, + session: Session, ) -> None: for tool_data in package.tools: provider_name = self._required_string(tool_data, "provider_name", "api_tool") @@ -510,7 +514,7 @@ class MigrationImportService: icon=icon, ) status = "created" - target_provider = self._find_api_tool_provider(session, target.tenant_id, provider_name) + target_provider = self._find_api_tool_provider(target.tenant_id, provider_name, session=session) if target_provider is not None: self._record_id_mappings( id_mapping, @@ -522,7 +526,9 @@ class MigrationImportService: ) report_items.append(ResourceReportItem(ResourceType.API_TOOL, provider_name, provider_name, status)) - def _find_api_tool_provider(self, session: Session, tenant_id: str, provider_name: str) -> ApiToolProvider | None: + def _find_api_tool_provider( + self, tenant_id: str, provider_name: str, *, session: Session + ) -> ApiToolProvider | None: return session.scalar( sa.select(ApiToolProvider).where( ApiToolProvider.tenant_id == tenant_id, @@ -558,13 +564,14 @@ class MigrationImportService: def _import_workflow_tools( self, - session: Session, package: MigrationPackage, target: ImportTarget, options: ImportOptions, id_mapping: dict[str, str], id_mapping_details: list[ResourceIdMapping], report_items: list[ResourceReportItem], + *, + session: Session, ) -> None: if not package.workflow_tools: return @@ -574,7 +581,10 @@ class MigrationImportService: for workflow_tool_data in package.workflow_tools: app_id = self._optional_string(workflow_tool_data.get("app_id")) resolved_app_id = id_mapping.get(app_id or "", app_id) - if not resolved_app_id or self._find_existing_app(session, resolved_app_id, target.tenant_id) is None: + if ( + not resolved_app_id + or self._find_existing_app(resolved_app_id, target.tenant_id, session=session) is None + ): report_items.append( ResourceReportItem( ResourceType.WORKFLOW_TOOL, @@ -586,7 +596,7 @@ class MigrationImportService: ) continue try: - self._ensure_workflow_app_is_published(session, target, account, resolved_app_id) + self._ensure_workflow_app_is_published(target, account, resolved_app_id, session=session) except Exception as exc: report_items.append( ResourceReportItem( @@ -602,7 +612,7 @@ class MigrationImportService: tool_name = self._required_string(workflow_tool_data, "name", "workflow_tool") lookup_workflow_tool_id = workflow_tool_id if options.id_strategy == IdStrategy.PRESERVE_ID else None existing = self._find_existing_workflow_tool( - session, target.tenant_id, lookup_workflow_tool_id, tool_name, resolved_app_id + target.tenant_id, lookup_workflow_tool_id, tool_name, resolved_app_id, session=session ) if existing is not None and options.conflict_strategy == ConflictStrategy.FAIL: raise MigrationDataError(f"Workflow tool already exists and conflict_strategy=fail: {tool_name}") @@ -669,7 +679,7 @@ class MigrationImportService: ) status = "created" target_provider = self._find_existing_workflow_tool( - session, target.tenant_id, import_id or None, tool_name, resolved_app_id + target.tenant_id, import_id or None, tool_name, resolved_app_id, session=session ) if target_provider is None: raise MigrationDataError(f"Workflow tool was not created: {tool_name}") @@ -686,9 +696,9 @@ class MigrationImportService: report_items.append(ResourceReportItem(ResourceType.WORKFLOW_TOOL, identifier, tool_name, status)) def _ensure_workflow_app_is_published( - self, session: Session, target: ImportTarget, account: Account, app_id: str + self, target: ImportTarget, account: Account, app_id: str, *, session: Session ) -> None: - app = self._find_existing_app(session, app_id, target.tenant_id) + app = self._find_existing_app(app_id, target.tenant_id, session=session) if app is None: raise MigrationDataError(f"Referenced workflow app was not found in target tenant: {app_id}") if app.workflow_id: @@ -714,20 +724,23 @@ class MigrationImportService: def _import_mcp_tools( self, - session: Session, package: MigrationPackage, target: ImportTarget, options: ImportOptions, report_items: list[ResourceReportItem], id_mapping: dict[str, str], id_mapping_details: list[ResourceIdMapping], + *, + session: Session, ) -> None: for mcp_data in package.mcp_tools: name = self._required_string(mcp_data, "name", "mcp_tool") server_identifier = self._required_string(mcp_data, "server_identifier", "mcp_tool") provider_id = self._optional_string(mcp_data.get("id")) lookup_provider_id = provider_id if options.id_strategy == IdStrategy.PRESERVE_ID else None - existing = self._find_existing_mcp_tool(session, target.tenant_id, lookup_provider_id, server_identifier) + existing = self._find_existing_mcp_tool( + target.tenant_id, lookup_provider_id, server_identifier, session=session + ) if existing is not None and options.conflict_strategy == ConflictStrategy.FAIL: raise MigrationDataError(f"MCP tool already exists and conflict_strategy=fail: {name}") if existing is not None and options.conflict_strategy == ConflictStrategy.SKIP: @@ -743,7 +756,7 @@ class MigrationImportService: report_items.append(ResourceReportItem(ResourceType.MCP_TOOL, existing.id, name, "skipped")) continue - service = MCPToolManageService(session=session) + service = MCPToolManageService(session=cast(Session, session)) configuration = MCPConfiguration.model_validate(mcp_data.get("configuration") or {}) authentication = ( MCPAuthentication.model_validate(mcp_data["authentication"]) if mcp_data.get("authentication") else None @@ -784,7 +797,7 @@ class MigrationImportService: authentication=authentication, ) created_provider = self._find_existing_mcp_tool( - session, target.tenant_id, lookup_provider_id, server_identifier + target.tenant_id, lookup_provider_id, server_identifier, session=session ) if created_provider is None: raise MigrationDataError(f"MCP provider was not created: {name}") @@ -812,7 +825,12 @@ class MigrationImportService: provider.authed = True def _find_existing_mcp_tool( - self, session: Session, tenant_id: str, provider_id: str | None, server_identifier: str + self, + tenant_id: str, + provider_id: str | None, + server_identifier: str, + *, + session: Session, ) -> MCPToolProvider | None: predicates = [MCPToolProvider.server_identifier == server_identifier] if self._is_uuid_string(provider_id): @@ -831,7 +849,13 @@ class MigrationImportService: return True def _find_existing_workflow_tool( - self, session: Session, tenant_id: str, workflow_tool_id: str | None, tool_name: str, app_id: str + self, + tenant_id: str, + workflow_tool_id: str | None, + tool_name: str, + app_id: str, + *, + session: Session, ) -> WorkflowToolProvider | None: predicates = [WorkflowToolProvider.name == tool_name, WorkflowToolProvider.app_id == app_id] if self._is_uuid_string(workflow_tool_id): @@ -843,14 +867,21 @@ class MigrationImportService: ) def _preflight_dependency_only_mcp( - self, session: Session, package: MigrationPackage, target: ImportTarget, report_items: list[ResourceReportItem] + self, + package: MigrationPackage, + target: ImportTarget, + report_items: list[ResourceReportItem], + *, + session: Session, ) -> None: for dependency in package.dependencies: if dependency.get("kind") != DependencyKind.MCP_TOOL.value: continue provider_id = str(dependency.get("provider_id", dependency.get("id", ""))) provider_name = self._optional_string(dependency.get("provider_name") or dependency.get("name")) - existing = self._find_dependency_only_mcp_provider(session, target.tenant_id, provider_id, provider_name) + existing = self._find_dependency_only_mcp_provider( + target.tenant_id, provider_id, provider_name, session=session + ) report_name = f"mcp_tool {provider_name or getattr(existing, 'name', None) or provider_id}" if existing is not None: report_items.append( @@ -879,7 +910,12 @@ class MigrationImportService: ) def _find_dependency_only_mcp_provider( - self, session: Session, tenant_id: str, provider_id: str, provider_name: str | None + self, + tenant_id: str, + provider_id: str, + provider_name: str | None, + *, + session: Session, ) -> MCPToolProvider | None: predicates = [MCPToolProvider.server_identifier == provider_id] if self._is_uuid_string(provider_id): diff --git a/api/services/dataset_service.py b/api/services/dataset_service.py index b36926a32c5..dda5440f772 100644 --- a/api/services/dataset_service.py +++ b/api/services/dataset_service.py @@ -13,7 +13,7 @@ import sqlalchemy as sa from pydantic import BaseModel, ConfigDict, Field, ValidationError, field_validator from redis.exceptions import LockNotOwnedError from sqlalchemy import ColumnElement, delete, exists, func, select, update -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, NotFound from configs import dify_config @@ -110,13 +110,6 @@ from tasks.sync_website_document_indexing_task import sync_website_document_inde logger = logging.getLogger(__name__) -def _session_for_helpers(session: scoped_session | Session) -> Session: - """Return a concrete SQLAlchemy session for helpers that do not accept scoped_session.""" - if isinstance(session, scoped_session): - return session() - return session - - class ProcessRulesDict(TypedDict): mode: ProcessRuleMode rules: dict[str, Any] @@ -244,11 +237,11 @@ class _EstimateArgs(BaseModel): class DatasetService: @staticmethod - def _can_manage_all_datasets(tenant_id: str, account_id: str) -> bool: + def _can_manage_all_datasets(tenant_id: str, account_id: str, *, session: Session) -> bool: if not dify_config.RBAC_ENABLED: return False - permissions = enterprise_rbac_service.RBACService.MyPermissions.get(tenant_id, account_id) + permissions = enterprise_rbac_service.RBACService.MyPermissions.get(tenant_id, account_id, session=session) workspace_permission_keys = getattr(getattr(permissions, "workspace", None), "permission_keys", []) or [] return "dataset.create_and_management" in workspace_permission_keys @@ -256,7 +249,7 @@ class DatasetService: def get_datasets( page, per_page, - session: scoped_session | Session, + session: Session, tenant_id=None, user=None, search=None, @@ -291,7 +284,9 @@ class DatasetService: return [], 0 else: if dify_config.RBAC_ENABLED: - can_manage_all_datasets = DatasetService._can_manage_all_datasets(str(tenant_id), str(user.id)) + can_manage_all_datasets = DatasetService._can_manage_all_datasets( + str(tenant_id), str(user.id), session=session + ) should_show_all_datasets = include_all and can_manage_all_datasets else: should_show_all_datasets = user.current_role == TenantAccountRole.OWNER and include_all @@ -361,7 +356,7 @@ class DatasetService: return datasets.items, datasets.total @staticmethod - def get_process_rules(dataset_id, session: scoped_session | Session) -> ProcessRulesDict: + def get_process_rules(dataset_id, session: Session) -> ProcessRulesDict: # get the latest process rule dataset_process_rule = session.execute( select(DatasetProcessRule) @@ -406,7 +401,6 @@ class DatasetService: @staticmethod def create_empty_dataset( - session: Session, tenant_id: str, name: str, description: str | None, @@ -420,6 +414,8 @@ class DatasetService: embedding_model_name: str | None = None, retrieval_model: RetrievalModel | None = None, summary_index_setting: dict[str, Any] | None = None, + *, + session: Session, ): # check if dataset name already exists if session.scalar(select(Dataset).where(Dataset.name == name, Dataset.tenant_id == tenant_id).limit(1)): @@ -473,7 +469,7 @@ class DatasetService: if provider == "external" and external_knowledge_api_id: external_knowledge_api = ExternalDatasetService.get_external_knowledge_api( - session, external_knowledge_api_id, tenant_id + external_knowledge_api_id, tenant_id, session=session ) if not external_knowledge_api: raise ValueError("External API template not found.") @@ -501,7 +497,7 @@ class DatasetService: def create_empty_rag_pipeline_dataset( tenant_id: str, rag_pipeline_dataset_create_entity: RagPipelineDatasetCreateEntity, - session: scoped_session | Session, + session: Session, ): if rag_pipeline_dataset_create_entity.name: # check if dataset name already exists @@ -549,7 +545,7 @@ class DatasetService: return dataset @staticmethod - def get_dataset(dataset_id, session: scoped_session | Session) -> Dataset | None: + def get_dataset(dataset_id, session: Session) -> Dataset | None: dataset: Dataset | None = session.get(Dataset, dataset_id) return dataset @@ -632,7 +628,7 @@ class DatasetService: raise ValueError(ex.description) @staticmethod - def update_dataset(session: Session, dataset_id, data, user): + def update_dataset(dataset_id, data, user, *, session: Session): """ Update dataset configuration and settings. @@ -672,7 +668,7 @@ class DatasetService: return DatasetService._update_internal_dataset(dataset, data, user, session) @staticmethod - def _has_dataset_same_name(tenant_id: str, dataset_id: str, name: str, session: scoped_session | Session): + def _has_dataset_same_name(tenant_id: str, dataset_id: str, name: str, session: Session): dataset = session.scalar( select(Dataset) .where( @@ -725,7 +721,7 @@ class DatasetService: if not external_knowledge_api_id: raise ValueError("External knowledge api id is required.") # Ensure the referenced external API template exists and belongs to the dataset tenant. - ExternalDatasetService.get_external_knowledge_api(session, external_knowledge_api_id, dataset.tenant_id) + ExternalDatasetService.get_external_knowledge_api(external_knowledge_api_id, dataset.tenant_id, session=session) # Update metadata fields dataset.updated_by = user.id if user else None dataset.updated_at = naive_utc_now() @@ -743,7 +739,7 @@ class DatasetService: @staticmethod def _update_external_knowledge_binding( - dataset_id, external_knowledge_id, external_knowledge_api_id, session: scoped_session | Session + dataset_id, external_knowledge_id, external_knowledge_api_id, session: Session ): """ Update external knowledge binding configuration. @@ -770,7 +766,7 @@ class DatasetService: session.add(external_knowledge_binding) @staticmethod - def _update_internal_dataset(dataset, data, user, session: scoped_session | Session): + def _update_internal_dataset(dataset, data, user, session: Session): """ Update internal dataset configuration. @@ -836,9 +832,7 @@ class DatasetService: return dataset @staticmethod - def _update_pipeline_knowledge_base_node_data( - dataset: Dataset, updata_user_id: str, session: scoped_session | Session - ): + def _update_pipeline_knowledge_base_node_data(dataset: Dataset, updata_user_id: str, session: Session): """ Update pipeline knowledge base node data. """ @@ -850,7 +844,7 @@ class DatasetService: return try: - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(session) published_workflow = rag_pipeline_service.get_published_workflow(pipeline) draft_workflow = rag_pipeline_service.get_draft_workflow(pipeline) @@ -921,7 +915,7 @@ class DatasetService: raise @staticmethod - def _handle_indexing_technique_change(dataset, data, filtered_data, session: scoped_session | Session): + def _handle_indexing_technique_change(dataset, data, filtered_data, session: Session): """ Handle changes in indexing technique and configure embedding models accordingly. @@ -955,7 +949,7 @@ class DatasetService: return None @staticmethod - def _configure_embedding_model_for_high_quality(data, filtered_data, session: scoped_session | Session): + def _configure_embedding_model_for_high_quality(data, filtered_data, session: Session): """ Configure embedding model settings for high quality indexing. @@ -992,9 +986,7 @@ class DatasetService: raise ValueError(ex.description) @staticmethod - def _handle_embedding_model_update_when_technique_unchanged( - dataset, data, filtered_data, session: scoped_session | Session - ): + def _handle_embedding_model_update_when_technique_unchanged(dataset, data, filtered_data, session: Session): """ Handle embedding model updates when indexing technique remains the same. @@ -1043,7 +1035,7 @@ class DatasetService: del filtered_data["embedding_model"] @staticmethod - def _update_embedding_model_settings(dataset, data, filtered_data, session: scoped_session | Session): + def _update_embedding_model_settings(dataset, data, filtered_data, session: Session): """ Update embedding model settings with new values. @@ -1078,7 +1070,7 @@ class DatasetService: return None @staticmethod - def _apply_new_embedding_settings(dataset, data, filtered_data, session: scoped_session | Session): + def _apply_new_embedding_settings(dataset, data, filtered_data, session: Session): """ Apply new embedding model settings to the dataset. @@ -1176,7 +1168,11 @@ class DatasetService: @staticmethod def update_rag_pipeline_dataset_settings( - session: Session, dataset: Dataset, knowledge_configuration: KnowledgeConfiguration, has_published: bool = False + dataset: Dataset, + knowledge_configuration: KnowledgeConfiguration, + has_published: bool = False, + *, + session: Session, ): if not current_user or not current_user.current_tenant_id: raise ValueError("Current user or current tenant not found") @@ -1335,7 +1331,7 @@ class DatasetService: deal_dataset_index_update_task.delay(dataset.id, action) @staticmethod - def delete_dataset(dataset_id, user, session: scoped_session | Session): + def delete_dataset(dataset_id, user, session: Session): dataset = DatasetService.get_dataset(dataset_id, session) if dataset is None: @@ -1350,12 +1346,12 @@ class DatasetService: return True @staticmethod - def dataset_use_check(dataset_id, session: scoped_session | Session) -> bool: + def dataset_use_check(dataset_id, session: Session) -> bool: stmt = select(exists().where(AppDatasetJoin.dataset_id == dataset_id)) return session.execute(stmt).scalar_one() @staticmethod - def check_dataset_permission(dataset, user, session: scoped_session | Session): + def check_dataset_permission(dataset, user, session: Session): """Validate dataset access for a user, using the injected session for partial-member lookups.""" if dataset.tenant_id != user.current_tenant_id: logger.debug("User %s does not have permission to access dataset %s", user.id, dataset.id) @@ -1378,7 +1374,7 @@ class DatasetService: @staticmethod def check_dataset_operator_permission( - user: Account | None = None, dataset: Dataset | None = None, *, session: scoped_session | Session + user: Account | None = None, dataset: Dataset | None = None, *, session: Session ): if not dataset: raise ValueError("Dataset not found") @@ -1409,7 +1405,7 @@ class DatasetService: return dataset_queries.items, dataset_queries.total @staticmethod - def get_related_apps(dataset_id: str, session: scoped_session | Session): + def get_related_apps(dataset_id: str, session: Session): return session.scalars( select(AppDatasetJoin) .where(AppDatasetJoin.dataset_id == dataset_id) @@ -1417,7 +1413,7 @@ class DatasetService: ).all() @staticmethod - def update_dataset_api_status(dataset_id: str, status: bool, session: scoped_session | Session): + def update_dataset_api_status(dataset_id: str, status: bool, session: Session): dataset = DatasetService.get_dataset(dataset_id, session) if dataset is None: raise NotFound("Dataset not found.") @@ -1429,7 +1425,7 @@ class DatasetService: session.commit() @staticmethod - def get_dataset_auto_disable_logs(dataset_id: str, session: scoped_session | Session) -> AutoDisableLogsDict: + def get_dataset_auto_disable_logs(dataset_id: str, session: Session) -> AutoDisableLogsDict: assert isinstance(current_user, Account) assert current_user.current_tenant_id is not None features = FeatureService.get_features(current_user.current_tenant_id, exclude_vector_space=True) @@ -1628,9 +1624,7 @@ class DocumentService: } @staticmethod - def get_document( - dataset_id: str, document_id: str | None = None, *, session: scoped_session | Session - ) -> Document | None: + def get_document(dataset_id: str, document_id: str | None = None, *, session: Session) -> Document | None: """Fetch a document by id within a dataset using the caller-provided session.""" if document_id: document = session.scalar( @@ -1641,9 +1635,7 @@ class DocumentService: return None @staticmethod - def get_documents_by_ids( - dataset_id: str, document_ids: Sequence[str], session: scoped_session | Session - ) -> Sequence[Document]: + def get_documents_by_ids(dataset_id: str, document_ids: Sequence[str], session: Session) -> Sequence[Document]: """Fetch documents for a dataset in a single batch query.""" if not document_ids: return [] @@ -1661,7 +1653,7 @@ class DocumentService: def update_documents_need_summary( dataset_id: str, document_ids: Sequence[str], - session: scoped_session | Session, + session: Session, need_summary: bool = True, ) -> int: """ @@ -1705,7 +1697,7 @@ class DocumentService: return updated_count @staticmethod - def get_document_download_url(document: Document, session: scoped_session | Session) -> str: + def get_document_download_url(document: Document, session: Session) -> str: """ Return a signed download URL for an upload-file document. """ @@ -1717,6 +1709,7 @@ class DocumentService: documents: Sequence[Document], dataset: Dataset, tenant_id: str, + session: Session, ) -> None: """ Enrich documents with summary_index_status based on dataset summary index settings. @@ -1728,6 +1721,7 @@ class DocumentService: documents: List of Document instances to enrich dataset: Dataset instance containing summary_index_setting tenant_id: Tenant ID for summary status lookup + session: SQLAlchemy session used to read summary status records """ # Check if dataset has summary index enabled has_summary_index = dataset.summary_index_setting and dataset.summary_index_setting.get("enable") is True @@ -1745,6 +1739,7 @@ class DocumentService: document_ids=document_ids_need_summary, dataset_id=dataset.id, tenant_id=tenant_id, + session=session, ) # Add summary_index_status to each document @@ -1763,7 +1758,7 @@ class DocumentService: document_ids: Sequence[str], tenant_id: str, current_user: Account, - session: scoped_session | Session, + session: Session, ) -> tuple[list[UploadFile], str]: """ Resolve upload files for batch ZIP downloads and generate a client-visible filename. @@ -1814,7 +1809,7 @@ class DocumentService: return str(upload_file_id) @staticmethod - def _get_upload_file_for_upload_file_document(document: Document, session: scoped_session | Session) -> UploadFile: + def _get_upload_file_for_upload_file_document(document: Document, session: Session) -> UploadFile: """ Load the `UploadFile` row for an upload-file document. """ @@ -1823,9 +1818,7 @@ class DocumentService: invalid_source_message="Document does not have an uploaded file to download.", missing_file_message="Uploaded file not found.", ) - upload_files_by_id = FileService.get_upload_files_by_ids( - _session_for_helpers(session), document.tenant_id, [upload_file_id] - ) + upload_files_by_id = FileService.get_upload_files_by_ids(document.tenant_id, [upload_file_id], session=session) upload_file = upload_files_by_id.get(upload_file_id) if not upload_file: raise NotFound("Uploaded file not found.") @@ -1837,7 +1830,7 @@ class DocumentService: dataset_id: str, document_ids: Sequence[str], tenant_id: str, - session: scoped_session | Session, + session: Session, ) -> dict[str, UploadFile]: """ Batch load upload files keyed by document id for ZIP downloads. @@ -1865,9 +1858,7 @@ class DocumentService: upload_file_ids.append(upload_file_id) upload_file_ids_by_document_id[document_id] = upload_file_id - upload_files_by_id = FileService.get_upload_files_by_ids( - _session_for_helpers(session), tenant_id, upload_file_ids - ) + upload_files_by_id = FileService.get_upload_files_by_ids(tenant_id, upload_file_ids, session=session) missing_upload_file_ids: set[str] = set(upload_file_ids) - set(upload_files_by_id.keys()) if missing_upload_file_ids: raise NotFound("Only uploaded-file documents can be downloaded as ZIP.") @@ -1878,13 +1869,13 @@ class DocumentService: } @staticmethod - def get_document_by_id(document_id: str, session: scoped_session | Session) -> Document | None: + def get_document_by_id(document_id: str, session: Session) -> Document | None: document = session.get(Document, document_id) return document @staticmethod - def get_document_by_ids(document_ids: list[str], session: scoped_session | Session) -> Sequence[Document]: + def get_document_by_ids(document_ids: list[str], session: Session) -> Sequence[Document]: documents = session.scalars( select(Document).where( Document.id.in_(document_ids), @@ -1896,7 +1887,7 @@ class DocumentService: return documents @staticmethod - def get_document_by_dataset_id(dataset_id: str, session: scoped_session | Session) -> Sequence[Document]: + def get_document_by_dataset_id(dataset_id: str, session: Session) -> Sequence[Document]: documents = session.scalars( select(Document).where( Document.dataset_id == dataset_id, @@ -1907,7 +1898,7 @@ class DocumentService: return documents @staticmethod - def get_working_documents_by_dataset_id(dataset_id: str, session: scoped_session | Session) -> Sequence[Document]: + def get_working_documents_by_dataset_id(dataset_id: str, session: Session) -> Sequence[Document]: documents = session.scalars( select(Document).where( Document.dataset_id == dataset_id, @@ -1920,7 +1911,7 @@ class DocumentService: return documents @staticmethod - def get_error_documents_by_dataset_id(dataset_id: str, session: scoped_session | Session) -> Sequence[Document]: + def get_error_documents_by_dataset_id(dataset_id: str, session: Session) -> Sequence[Document]: documents = session.scalars( select(Document).where( Document.dataset_id == dataset_id, @@ -1930,7 +1921,7 @@ class DocumentService: return documents @staticmethod - def get_batch_documents(dataset_id: str, batch: str, session: scoped_session | Session) -> Sequence[Document]: + def get_batch_documents(dataset_id: str, batch: str, session: Session) -> Sequence[Document]: assert isinstance(current_user, Account) documents = session.scalars( select(Document).where( @@ -1943,7 +1934,7 @@ class DocumentService: return documents @staticmethod - def get_document_file_detail(file_id: str, session: scoped_session | Session): + def get_document_file_detail(file_id: str, session: Session): file_detail = session.get(UploadFile, file_id) return file_detail @@ -1955,7 +1946,7 @@ class DocumentService: return False @staticmethod - def delete_document(document, session: scoped_session | Session): + def delete_document(document, session: Session): # trigger document_was_deleted signal file_id = None if document.data_source_type == DataSourceType.UPLOAD_FILE: @@ -1975,7 +1966,7 @@ class DocumentService: dataset_ref: DatasetRef, document_ids: list[str], doc_form: str | None, - session: scoped_session | Session, + session: Session, ): # Check if document_ids is not empty to avoid WHERE false condition if not document_ids or len(document_ids) == 0: @@ -2006,7 +1997,7 @@ class DocumentService: 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: scoped_session | Session) -> Document: + def rename_document(dataset_id: str, document_id: str, name: str, session: Session) -> Document: assert isinstance(current_user, Account) dataset = DatasetService.get_dataset(dataset_id, session) @@ -2041,7 +2032,7 @@ class DocumentService: return document @staticmethod - def pause_document(document, session: scoped_session | Session): + def pause_document(document, session: Session): if document.indexing_status not in { IndexingStatus.WAITING, IndexingStatus.PARSING, @@ -2063,7 +2054,7 @@ class DocumentService: redis_client.setnx(indexing_cache_key, "True") @staticmethod - def recover_document(document, session: scoped_session | Session): + def recover_document(document, session: Session): if not document.is_paused: raise DocumentIndexingError() # update document to be recover @@ -2080,7 +2071,7 @@ class DocumentService: recover_document_indexing_task.delay(document.dataset_id, document.id) @staticmethod - def retry_document(dataset_id: str, documents: list[Document], session: scoped_session | Session): + def retry_document(dataset_id: str, documents: list[Document], session: Session): for document in documents: # add retry flag retry_indexing_cache_key = f"document_{document.id}_is_retried" @@ -2100,7 +2091,7 @@ class DocumentService: retry_document_indexing_task.delay(dataset_id, document_ids, current_user.id) @staticmethod - def sync_website_document(dataset_id: str, document: Document, session: scoped_session | Session): + def sync_website_document(dataset_id: str, document: Document, session: Session): # add sync flag sync_indexing_cache_key = f"document_{document.id}_is_sync" cache_result = redis_client.get(sync_indexing_cache_key) @@ -2120,7 +2111,7 @@ class DocumentService: sync_website_document_indexing_task.delay(dataset_id, document.id) @staticmethod - def get_documents_position(dataset_id, session: scoped_session | Session): + def get_documents_position(dataset_id, session: Session): document = session.scalar( select(Document).where(Document.dataset_id == dataset_id).order_by(Document.position.desc()).limit(1) ) @@ -2137,7 +2128,7 @@ class DocumentService: dataset_process_rule: DatasetProcessRule | None = None, created_from: str = DocumentCreatedFrom.WEB, *, - session: scoped_session | Session, + session: Session, ) -> tuple[list[Document], str]: # check doc_form DatasetService.check_doc_form(dataset, knowledge_config.doc_form) @@ -2793,7 +2784,7 @@ class DocumentService: return document @staticmethod - def get_tenant_documents_count(session: scoped_session | Session): + def get_tenant_documents_count(*, session: Session): assert isinstance(current_user, Account) documents_count = ( @@ -2817,7 +2808,7 @@ class DocumentService: dataset_process_rule: DatasetProcessRule | None = None, created_from: str = DocumentCreatedFrom.WEB, *, - session: scoped_session | Session, + session: Session, ): assert isinstance(current_user, Account) @@ -2944,7 +2935,7 @@ class DocumentService: @staticmethod def save_document_without_dataset_id( - tenant_id: str, knowledge_config: KnowledgeConfig, account: Account, session: scoped_session | Session + tenant_id: str, knowledge_config: KnowledgeConfig, account: Account, session: Session ): assert isinstance(current_user, Account) assert current_user.current_tenant_id is not None @@ -3133,7 +3124,7 @@ class DocumentService: document_ids: list[str], action: Literal["enable", "disable", "archive", "un_archive"], user, - session: scoped_session | Session, + session: Session, ): """ Batch update document status. @@ -3216,7 +3207,7 @@ class DocumentService: document = update_info["document"] indexing_cache_key = f"document_{document.id}_indexing" redis_client.setex(indexing_cache_key, 600, 1) - except Exception as e: + except Exception: # Log the error but do not rollback the transaction logger.exception("Error setting cache for document %s", update_info["document"].id) # Raise any propagation error after all updates @@ -3340,9 +3331,7 @@ class SegmentService: raise ValueError(f"Exceeded maximum attachment limit of {single_chunk_attachment_limit}") @classmethod - def create_segment( - cls, args: dict[str, Any], document: Document, dataset: Dataset, session: scoped_session | Session - ): + def create_segment(cls, args: dict[str, Any], document: Document, dataset: Dataset, session: Session): assert isinstance(current_user, Account) assert current_user.current_tenant_id is not None @@ -3408,7 +3397,13 @@ class SegmentService: try: keywords = args.get("keywords") keywords_list = [keywords] if keywords is not None else None - VectorService.create_segments_vector(keywords_list, [segment_document], dataset, document.doc_form) + VectorService.create_segments_vector( + keywords_list, + [segment_document], + dataset, + document.doc_form, + session, + ) except Exception as e: logger.exception("create segment index failed") segment_document.enabled = False @@ -3422,9 +3417,7 @@ class SegmentService: pass @classmethod - def multi_create_segment( - cls, segments: list, document: Document, dataset: Dataset, session: scoped_session | Session - ): + def multi_create_segment(cls, segments: list, document: Document, dataset: Dataset, session: Session): assert isinstance(current_user, Account) assert current_user.current_tenant_id is not None @@ -3498,7 +3491,11 @@ class SegmentService: try: # save vector index VectorService.create_segments_vector( - keywords_list, pre_segment_data_list, dataset, document.doc_form + keywords_list, + pre_segment_data_list, + dataset, + document.doc_form, + session, ) except Exception as e: logger.exception("create segment index failed") @@ -3519,7 +3516,7 @@ class SegmentService: segment: DocumentSegment, document: Document, dataset: Dataset, - session: scoped_session | Session, + session: Session, ): assert isinstance(current_user, Account) assert current_user.current_tenant_id is not None @@ -3597,7 +3594,13 @@ class SegmentService: processing_rule = session.get(DatasetProcessRule, document.dataset_process_rule_id) if processing_rule: VectorService.generate_child_chunks( - segment, document, dataset, embedding_model_instance, processing_rule, True + segment, + document, + dataset, + embedding_model_instance, + processing_rule, + session, + True, ) elif document.doc_form in (IndexStructureType.PARAGRAPH_INDEX, IndexStructureType.QA_INDEX): if args.enabled or keyword_changed: @@ -3628,7 +3631,12 @@ class SegmentService: from services.summary_index_service import SummaryIndexService try: - SummaryIndexService.update_summary_for_segment(segment, dataset, args.summary) + SummaryIndexService.update_summary_for_segment( + segment, + dataset, + args.summary, + session=session, + ) except Exception: logger.exception("Failed to update summary for segment %s", segment.id) # Don't fail the entire update if summary update fails @@ -3697,7 +3705,13 @@ class SegmentService: processing_rule = session.get(DatasetProcessRule, document.dataset_process_rule_id) if processing_rule: VectorService.generate_child_chunks( - segment, document, dataset, embedding_model_instance, processing_rule, True + segment, + document, + dataset, + embedding_model_instance, + processing_rule, + session, + True, ) elif document.doc_form in (IndexStructureType.PARAGRAPH_INDEX, IndexStructureType.QA_INDEX): # update segment vector index @@ -3728,7 +3742,10 @@ class SegmentService: try: SummaryIndexService.generate_and_vectorize_summary( - segment, dataset, dataset.summary_index_setting + segment, + dataset, + dataset.summary_index_setting, + session=session, ) logger.info("Auto-regenerated summary for segment %s after content change", segment.id) except Exception: @@ -3743,7 +3760,12 @@ class SegmentService: from services.summary_index_service import SummaryIndexService try: - SummaryIndexService.update_summary_for_segment(segment, dataset, args.summary) + SummaryIndexService.update_summary_for_segment( + segment, + dataset, + args.summary, + session=session, + ) logger.info("Updated summary for segment %s with user-provided content", segment.id) except Exception: logger.exception("Failed to update summary for segment %s", segment.id) @@ -3760,7 +3782,10 @@ class SegmentService: try: SummaryIndexService.generate_and_vectorize_summary( - segment, dataset, dataset.summary_index_setting + segment, + dataset, + dataset.summary_index_setting, + session=session, ) logger.info( "Regenerated summary for segment %s after content change (summary unchanged)", @@ -3770,7 +3795,7 @@ class SegmentService: logger.exception("Failed to regenerate summary for segment %s", segment.id) # Don't fail the entire update if summary regeneration fails # update multimodel vector index - VectorService.update_multimodel_vector(segment, args.attachment_ids or [], dataset) + VectorService.update_multimodel_vector(segment, args.attachment_ids or [], dataset, session) except Exception as e: logger.exception("update segment index failed") segment.enabled = False @@ -3784,9 +3809,7 @@ class SegmentService: return new_segment @classmethod - def delete_segment( - cls, segment: DocumentSegment, document: Document, dataset: Dataset, session: scoped_session | Session - ): + def delete_segment(cls, segment: DocumentSegment, document: Document, dataset: Dataset, session: Session): indexing_cache_key = f"segment_{segment.id}_delete_indexing" cache_result = redis_client.get(indexing_cache_key) if cache_result is not None: @@ -3821,9 +3844,7 @@ class SegmentService: session.commit() @classmethod - def delete_segments( - cls, segment_ids: list, document: Document, dataset: Dataset, session: scoped_session | Session - ): + def delete_segments(cls, segment_ids: list, document: Document, dataset: Dataset, session: Session): assert current_user is not None # Check if segment_ids is not empty to avoid WHERE false condition if not segment_ids or len(segment_ids) == 0: @@ -3882,7 +3903,7 @@ class SegmentService: action: Literal["enable", "disable"], dataset: Dataset, document: Document, - session: scoped_session | Session, + session: Session, ): assert current_user is not None @@ -3948,7 +3969,7 @@ class SegmentService: segment: DocumentSegment, document: Document, dataset: Dataset, - session: scoped_session | Session, + session: Session, ) -> ChildChunk: assert isinstance(current_user, Account) @@ -3997,7 +4018,7 @@ class SegmentService: segment: DocumentSegment, document: Document, dataset: Dataset, - session: scoped_session | Session, + session: Session, ) -> list[ChildChunk]: assert isinstance(current_user, Account) child_chunks = session.scalars( @@ -4072,7 +4093,7 @@ class SegmentService: segment: DocumentSegment, document: Document, dataset: Dataset, - session: scoped_session | Session, + session: Session, ) -> ChildChunk: assert current_user is not None @@ -4092,7 +4113,7 @@ class SegmentService: return child_chunk @classmethod - def delete_child_chunk(cls, child_chunk: ChildChunk, dataset: Dataset, session: scoped_session | Session): + def delete_child_chunk(cls, child_chunk: ChildChunk, dataset: Dataset, session: Session): session.delete(child_chunk) try: VectorService.delete_child_chunk_vector(child_chunk, dataset) @@ -4124,9 +4145,7 @@ class SegmentService: return paginate_query(query, page=page, per_page=limit, max_per_page=100) @classmethod - def get_child_chunk_by_id( - cls, child_chunk_id: str, tenant_id: str, session: scoped_session | Session - ) -> ChildChunk | None: + def get_child_chunk_by_id(cls, child_chunk_id: str, tenant_id: str, session: Session) -> ChildChunk | None: """Get a child chunk by its ID.""" result = session.scalar( select(ChildChunk).where(ChildChunk.id == child_chunk_id, ChildChunk.tenant_id == tenant_id).limit(1) @@ -4134,9 +4153,11 @@ class SegmentService: return result if isinstance(result, ChildChunk) else None @classmethod - def get_child_chunk_by_segment_ref(cls, child_chunk_id: str, segment_ref: SegmentRef) -> ChildChunk | None: + def get_child_chunk_by_segment_ref( + cls, child_chunk_id: str, segment_ref: SegmentRef, session: Session + ) -> ChildChunk | None: """Get a child chunk through the full tenant/dataset/document/segment chain.""" - result = db.session.scalar( + result = session.scalar( select(ChildChunk) .where( ChildChunk.id == child_chunk_id, @@ -4178,9 +4199,7 @@ class SegmentService: return paginated_segments.items, paginated_segments.total @classmethod - def get_segment_by_id( - cls, segment_id: str, tenant_id: str, session: scoped_session | Session - ) -> DocumentSegment | None: + def get_segment_by_id(cls, segment_id: str, tenant_id: str, session: Session) -> DocumentSegment | None: """Get a segment by its ID.""" result = session.scalar( select(DocumentSegment) @@ -4190,9 +4209,9 @@ class SegmentService: return result if isinstance(result, DocumentSegment) else None @classmethod - def get_segment_by_ref(cls, segment_ref: SegmentRef) -> DocumentSegment | None: + def get_segment_by_ref(cls, segment_ref: SegmentRef, session: Session) -> DocumentSegment | None: """Get a segment through the full tenant/dataset/document ownership chain.""" - result = db.session.scalar( + result = session.scalar( select(DocumentSegment) .where( DocumentSegment.id == segment_ref.segment_id, @@ -4209,7 +4228,7 @@ class SegmentService: cls, document_id: str, dataset_id: str, - session: scoped_session | Session, + session: Session, status: str | None = None, enabled: bool | None = None, ) -> Sequence[DocumentSegment]: @@ -4242,7 +4261,7 @@ class SegmentService: class DatasetCollectionBindingService: @classmethod def get_dataset_collection_binding( - cls, provider_name: str, model_name: str, session: scoped_session | Session, collection_type: str = "dataset" + cls, provider_name: str, model_name: str, session: Session, collection_type: str = "dataset" ) -> DatasetCollectionBinding: dataset_collection_binding = session.scalar( select(DatasetCollectionBinding) @@ -4268,7 +4287,7 @@ class DatasetCollectionBindingService: @classmethod def get_dataset_collection_binding_by_id_and_type( - cls, collection_binding_id: str, session: scoped_session | Session, collection_type: str = "dataset" + cls, collection_binding_id: str, session: Session, collection_type: str = "dataset" ) -> DatasetCollectionBinding: dataset_collection_binding = session.scalar( select(DatasetCollectionBinding) @@ -4286,7 +4305,7 @@ class DatasetCollectionBindingService: class DatasetPermissionService: @classmethod - def get_dataset_partial_member_list(cls, dataset_id, session: scoped_session | Session): + def get_dataset_partial_member_list(cls, dataset_id, session: Session): user_list_query = session.scalars( select( DatasetPermission.account_id, @@ -4296,7 +4315,7 @@ class DatasetPermissionService: return user_list_query @classmethod - def update_partial_member_list(cls, tenant_id, dataset_id, user_list, session: scoped_session | Session): + def update_partial_member_list(cls, tenant_id, dataset_id, user_list, session: Session): try: session.execute(delete(DatasetPermission).where(DatasetPermission.dataset_id == dataset_id)) permissions = [] @@ -4315,7 +4334,7 @@ class DatasetPermissionService: raise e @classmethod - def check_permission(cls, session: Session, user, dataset, requested_permission, requested_partial_member_list): + def check_permission(cls, user, dataset, requested_permission, requested_partial_member_list, *, session: Session): if not user.is_dataset_editor: raise NoPermissionError("User does not have permission to edit this dataset.") @@ -4332,7 +4351,7 @@ class DatasetPermissionService: raise ValueError("Dataset operators cannot change the dataset permissions.") @classmethod - def clear_partial_member_list(cls, dataset_id, session: scoped_session | Session): + def clear_partial_member_list(cls, dataset_id, session: Session): try: session.execute(delete(DatasetPermission).where(DatasetPermission.dataset_id == dataset_id)) session.commit() diff --git a/api/services/datasource_provider_service.py b/api/services/datasource_provider_service.py index 12807a41f04..5de194dd262 100644 --- a/api/services/datasource_provider_service.py +++ b/api/services/datasource_provider_service.py @@ -446,12 +446,14 @@ class DatasourceProviderService: is not None ) - def is_tenant_oauth_params_enabled(self, tenant_id: str, datasource_provider_id: DatasourceProviderID) -> bool: + def is_tenant_oauth_params_enabled( + self, tenant_id: str, datasource_provider_id: DatasourceProviderID, *, session: Session + ) -> bool: """ check if tenant oauth params is enabled """ return ( - db.session.scalar( + session.scalar( select(func.count(DatasourceOauthTenantParamConfig.id)).where( DatasourceOauthTenantParamConfig.tenant_id == tenant_id, DatasourceOauthTenantParamConfig.provider == datasource_provider_id.provider_name, @@ -463,12 +465,17 @@ class DatasourceProviderService: ) > 0 def get_tenant_oauth_client( - self, tenant_id: str, datasource_provider_id: DatasourceProviderID, mask: bool = False + self, + tenant_id: str, + datasource_provider_id: DatasourceProviderID, + mask: bool = False, + *, + session: Session, ) -> Mapping[str, Any] | None: """ get tenant oauth client """ - tenant_oauth_client_params = db.session.scalar( + tenant_oauth_client_params = session.scalar( select(DatasourceOauthTenantParamConfig) .where( DatasourceOauthTenantParamConfig.tenant_id == tenant_id, @@ -547,7 +554,7 @@ class DatasourceProviderService: @staticmethod def generate_next_datasource_provider_name( - session: Session, tenant_id: str, provider_id: DatasourceProviderID, credential_type: CredentialType + tenant_id: str, provider_id: DatasourceProviderID, credential_type: CredentialType, *, session: Session ) -> str: db_providers = session.scalars( select(DatasourceProvider).where( @@ -800,6 +807,8 @@ class DatasourceProviderService: provider: str, plugin_id: str, user: "Account | None" = None, + *, + session: Session, ) -> list[dict]: """ list datasource credentials with obfuscated sensitive fields, @@ -829,11 +838,11 @@ class DatasourceProviderService: credential_type=CredPermType.DATASOURCE_PROVIDER, user=user, ) - datasource_providers: list[DatasourceProvider] = list(db.session.scalars(query).all()) + datasource_providers: list[DatasourceProvider] = list(session.scalars(query).all()) if not datasource_providers: return [] copy_credentials_list = [] - default_provider = db.session.execute( + default_provider = session.execute( select(DatasourceProvider.id) .where( DatasourceProvider.tenant_id == tenant_id, @@ -870,7 +879,7 @@ class DatasourceProviderService: return copy_credentials_list - def get_all_datasource_credentials(self, tenant_id: str) -> list[dict]: + def get_all_datasource_credentials(self, tenant_id: str, *, session: Session) -> list[dict]: """ get datasource credentials. @@ -883,7 +892,10 @@ class DatasourceProviderService: for datasource in datasources: datasource_provider_id = DatasourceProviderID(f"{datasource.plugin_id}/{datasource.provider}") credentials = self.list_datasource_credentials( - tenant_id=tenant_id, provider=datasource.provider, plugin_id=datasource.plugin_id + tenant_id=tenant_id, + provider=datasource.provider, + plugin_id=datasource.plugin_id, + session=session, ) redirect_uri = ( f"{dify_config.CONSOLE_API_URL}/console/api/oauth/plugin/{datasource_provider_id}/datasource/callback" @@ -912,10 +924,10 @@ class DatasourceProviderService: for credential_schema in datasource.declaration.oauth_schema.credentials_schema ], "oauth_custom_client_params": self.get_tenant_oauth_client( - tenant_id, datasource_provider_id, mask=True + tenant_id, datasource_provider_id, mask=True, session=session ), "is_oauth_custom_client_enabled": self.is_tenant_oauth_params_enabled( - tenant_id, datasource_provider_id + tenant_id, datasource_provider_id, session=session ), "is_system_oauth_params_exists": self.is_system_oauth_params_exist(datasource_provider_id), "redirect_uri": redirect_uri, @@ -926,7 +938,7 @@ class DatasourceProviderService: ) return datasource_credentials - def get_hard_code_datasource_credentials(self, tenant_id: str) -> list[dict]: + def get_hard_code_datasource_credentials(self, tenant_id: str, *, session: Session) -> list[dict]: """ get hard code datasource credentials. @@ -945,7 +957,10 @@ class DatasourceProviderService: ]: datasource_provider_id = DatasourceProviderID(f"{datasource.plugin_id}/{datasource.provider}") credentials = self.list_datasource_credentials( - tenant_id=tenant_id, provider=datasource.provider, plugin_id=datasource.plugin_id + tenant_id=tenant_id, + provider=datasource.provider, + plugin_id=datasource.plugin_id, + session=session, ) redirect_uri = "{}/console/api/oauth/plugin/{}/datasource/callback".format( dify_config.CONSOLE_API_URL, datasource_provider_id @@ -974,10 +989,10 @@ class DatasourceProviderService: for credential_schema in datasource.declaration.oauth_schema.credentials_schema ], "oauth_custom_client_params": self.get_tenant_oauth_client( - tenant_id, datasource_provider_id, mask=True + tenant_id, datasource_provider_id, mask=True, session=session ), "is_oauth_custom_client_enabled": self.is_tenant_oauth_params_enabled( - tenant_id, datasource_provider_id + tenant_id, datasource_provider_id, session=session ), "is_system_oauth_params_exists": self.is_system_oauth_params_exist(datasource_provider_id), "redirect_uri": redirect_uri, @@ -988,7 +1003,9 @@ class DatasourceProviderService: ) return datasource_credentials - def get_real_datasource_credentials(self, tenant_id: str, provider: str, plugin_id: str) -> list[dict]: + def get_real_datasource_credentials( + self, tenant_id: str, provider: str, plugin_id: str, *, session: Session + ) -> list[dict]: """ get datasource credentials. @@ -998,7 +1015,7 @@ class DatasourceProviderService: """ # Get all provider configurations of the current workspace datasource_providers: list[DatasourceProvider] = list( - db.session.scalars( + session.scalars( select(DatasourceProvider).where( DatasourceProvider.tenant_id == tenant_id, DatasourceProvider.provider == provider, @@ -1110,7 +1127,9 @@ class DatasourceProviderService: datasource_provider.encrypted_credentials = encrypted_credentials - def remove_datasource_credentials(self, tenant_id: str, auth_id: str, provider: str, plugin_id: str) -> None: + def remove_datasource_credentials( + self, tenant_id: str, auth_id: str, provider: str, plugin_id: str, *, session: Session + ) -> None: """ remove datasource credentials. @@ -1119,7 +1138,7 @@ class DatasourceProviderService: :param plugin_id: plugin id :return: """ - datasource_provider = db.session.scalar( + datasource_provider = session.scalar( select(DatasourceProvider) .where( DatasourceProvider.tenant_id == tenant_id, @@ -1130,5 +1149,5 @@ class DatasourceProviderService: .limit(1) ) if datasource_provider: - db.session.delete(datasource_provider) - db.session.commit() + session.delete(datasource_provider) + session.commit() diff --git a/api/services/enterprise/account_deletion_sync.py b/api/services/enterprise/account_deletion_sync.py index b5107fb0f66..89c4b80e670 100644 --- a/api/services/enterprise/account_deletion_sync.py +++ b/api/services/enterprise/account_deletion_sync.py @@ -5,9 +5,9 @@ from datetime import UTC, datetime from redis import RedisError from sqlalchemy import select +from sqlalchemy.orm import Session from configs import dify_config -from extensions.ext_database import db from extensions.ext_redis import redis_client from models.account import TenantAccountJoin @@ -87,7 +87,7 @@ def sync_workspace_member_removal(workspace_id: str, member_id: str, *, source: return _queue_task(workspace_id=workspace_id, member_id=member_id, source=source) -def sync_account_deletion(account_id: str, *, source: str) -> bool: +def sync_account_deletion(account_id: str, *, source: str, session: Session) -> bool: """ Sync full account deletion across all workspaces (enterprise only). @@ -97,6 +97,7 @@ def sync_account_deletion(account_id: str, *, source: str) -> bool: Args: account_id: The account ID being deleted source: Source of the sync request (e.g., "account_deleted") + session: SQLAlchemy session used to fetch workspace memberships Returns: bool: True if all tasks were queued (or skipped in community), False if any queueing failed @@ -105,9 +106,7 @@ def sync_account_deletion(account_id: str, *, source: str) -> bool: return True # Fetch all workspaces the account belongs to - workspace_joins = db.session.scalars( - select(TenantAccountJoin).where(TenantAccountJoin.account_id == account_id) - ).all() + workspace_joins = session.scalars(select(TenantAccountJoin).where(TenantAccountJoin.account_id == account_id)).all() # Queue sync task for each workspace success = True diff --git a/api/services/enterprise/rbac_service.py b/api/services/enterprise/rbac_service.py index b2e77156d3d..47ac8d5aeae 100644 --- a/api/services/enterprise/rbac_service.py +++ b/api/services/enterprise/rbac_service.py @@ -9,9 +9,9 @@ from flask import has_request_context, request from pydantic import AliasChoices, BaseModel, ConfigDict, Field, field_validator from sqlalchemy import select from sqlalchemy.exc import SQLAlchemyError +from sqlalchemy.orm import Session from configs import dify_config -from core.db.session_factory import session_factory from core.rbac import RBACResourceWhitelistScope from models import TenantAccountJoin, TenantAccountRole from services.enterprise.base import EnterpriseRequest @@ -565,25 +565,24 @@ def _legacy_member_roles_response( ) -def _legacy_my_permissions(tenant_id: str, account_id: str | None) -> MyPermissionsResponse: +def _legacy_my_permissions(tenant_id: str, account_id: str | None, *, session: Session) -> MyPermissionsResponse: if not account_id: return MyPermissionsResponse() try: - with session_factory.create_session() as session: - role = session.scalar( - select(TenantAccountJoin.role).where( - TenantAccountJoin.tenant_id == tenant_id, - TenantAccountJoin.account_id == account_id, - ) + role = session.scalar( + select(TenantAccountJoin.role).where( + TenantAccountJoin.tenant_id == tenant_id, + TenantAccountJoin.account_id == account_id, ) - if not role: - return MyPermissionsResponse() + ) + if not role: + return MyPermissionsResponse() - try: - tenant_role = TenantAccountRole(role) - except ValueError: - return MyPermissionsResponse() + try: + tenant_role = TenantAccountRole(role) + except ValueError: + return MyPermissionsResponse() except SQLAlchemyError: return MyPermissionsResponse() @@ -600,8 +599,10 @@ def _legacy_resource_permission_keys_batch( account_id: str | None, resource_ids: list[str], resource_type: RBACResourceType, + *, + session: Session, ) -> dict[str, list[str]]: - snapshot = _legacy_my_permissions(tenant_id, account_id) + snapshot = _legacy_my_permissions(tenant_id, account_id, session=session) if resource_type == RBACResourceType.APP: permission_keys = snapshot.app.default_permission_keys else: @@ -1597,7 +1598,9 @@ class RBACService: class MemberRoles: @staticmethod - def get(tenant_id: str, account_id: str | None, member_account_id: str) -> MemberRolesResponse: + def get( + tenant_id: str, account_id: str | None, member_account_id: str, *, session: Session + ) -> MemberRolesResponse: if dify_config.RBAC_ENABLED: data = _inner_call( "GET", @@ -1609,14 +1612,13 @@ class RBACService: rst = MemberRolesResponse.model_validate(data or {}) return rst else: - with session_factory.create_session() as session: - role = session.scalar( - select(TenantAccountJoin.role).where( - TenantAccountJoin.tenant_id == tenant_id, - TenantAccountJoin.account_id == member_account_id, - ) + role = session.scalar( + select(TenantAccountJoin.role).where( + TenantAccountJoin.tenant_id == tenant_id, + TenantAccountJoin.account_id == member_account_id, ) - return _legacy_member_roles_response(tenant_id, member_account_id, role) + ) + return _legacy_member_roles_response(tenant_id, member_account_id, role) @staticmethod def batch_get( @@ -1646,34 +1648,35 @@ class RBACService: account_id: str | None, member_account_id: str, role_ids: list[str], + *, + session: Session, ) -> MemberRolesResponse: if not dify_config.RBAC_ENABLED: if len(role_ids) != 1: raise ValueError("Legacy workspace member role update requires exactly one role.") tenant_role = TenantAccountRole(role_ids[0]) - with session_factory.create_session() as session: - target_member_join = session.scalar( + target_member_join = session.scalar( + select(TenantAccountJoin).where( + TenantAccountJoin.tenant_id == tenant_id, + TenantAccountJoin.account_id == member_account_id, + ) + ) + if not target_member_join: + raise ValueError("Member not in tenant.") + + if tenant_role == TenantAccountRole.OWNER: + current_owner_join = session.scalar( select(TenantAccountJoin).where( TenantAccountJoin.tenant_id == tenant_id, - TenantAccountJoin.account_id == member_account_id, + TenantAccountJoin.role == TenantAccountRole.OWNER, ) ) - if not target_member_join: - raise ValueError("Member not in tenant.") + if current_owner_join and current_owner_join.account_id != member_account_id: + current_owner_join.role = TenantAccountRole.ADMIN - if tenant_role == TenantAccountRole.OWNER: - current_owner_join = session.scalar( - select(TenantAccountJoin).where( - TenantAccountJoin.tenant_id == tenant_id, - TenantAccountJoin.role == TenantAccountRole.OWNER, - ) - ) - if current_owner_join and current_owner_join.account_id != member_account_id: - current_owner_join.role = TenantAccountRole.ADMIN - - target_member_join.role = tenant_role - session.commit() + target_member_join.role = tenant_role + session.commit() return _legacy_member_roles_response(tenant_id, member_account_id, tenant_role) @@ -1739,11 +1742,15 @@ class RBACService: tenant_id: str, account_id: str | None, app_ids: list[str], + *, + session: Session, ) -> dict[str, list[str]]: if not app_ids: return {} if not dify_config.RBAC_ENABLED: - return _legacy_resource_permission_keys_batch(tenant_id, account_id, app_ids, RBACResourceType.APP) + return _legacy_resource_permission_keys_batch( + tenant_id, account_id, app_ids, RBACResourceType.APP, session=session + ) data = _inner_call( "POST", f"{_INNER_PREFIX}/apps/permission-keys/batch", @@ -1759,12 +1766,14 @@ class RBACService: tenant_id: str, account_id: str | None, dataset_ids: list[str], + *, + session: Session, ) -> dict[str, list[str]]: if not dataset_ids: return {} if not dify_config.RBAC_ENABLED: return _legacy_resource_permission_keys_batch( - tenant_id, account_id, dataset_ids, RBACResourceType.DATASET + tenant_id, account_id, dataset_ids, RBACResourceType.DATASET, session=session ) data = _inner_call( "POST", @@ -1783,9 +1792,10 @@ class RBACService: *, app_id: str | None = None, dataset_id: str | None = None, + session: Session, ) -> MyPermissionsResponse: if not dify_config.RBAC_ENABLED: - return _legacy_my_permissions(tenant_id, account_id) + return _legacy_my_permissions(tenant_id, account_id, session=session) data = _inner_call( "GET", diff --git a/api/services/external_knowledge_service.py b/api/services/external_knowledge_service.py index 42e7eca29d7..cdd6c48342e 100644 --- a/api/services/external_knowledge_service.py +++ b/api/services/external_knowledge_service.py @@ -10,6 +10,7 @@ from sqlalchemy.orm import Session from constants import HIDDEN_VALUE from core.helper import ssrf_proxy from core.rag.entities import MetadataFilteringCondition +from extensions.ext_database import db # noqa: F401 from graphon.nodes.http_request.exc import InvalidHttpMethodError from libs.datetime_utils import naive_utc_now from libs.pagination import paginate_query @@ -57,7 +58,7 @@ class ExternalDatasetService: @staticmethod def create_external_knowledge_api( - tenant_id: str, user_id: str, args: dict[str, Any], session: Session + tenant_id: str, user_id: str, args: dict[str, Any], *, session: Session ) -> ExternalKnowledgeApis: settings = args.get("settings") if settings is None: @@ -105,7 +106,7 @@ class ExternalDatasetService: @staticmethod def get_external_knowledge_api( - session: Session, external_knowledge_api_id: str, tenant_id: str + external_knowledge_api_id: str, tenant_id: str, *, session: Session ) -> ExternalKnowledgeApis: external_knowledge_api: ExternalKnowledgeApis | None = session.scalar( select(ExternalKnowledgeApis) @@ -118,7 +119,12 @@ class ExternalDatasetService: @staticmethod def update_external_knowledge_api( - session: Session, tenant_id: str, user_id: str, external_knowledge_api_id: str, args + tenant_id: str, + user_id: str, + external_knowledge_api_id: str, + args: dict[str, Any], + *, + session: Session, ) -> ExternalKnowledgeApis: external_knowledge_api: ExternalKnowledgeApis | None = session.scalar( select(ExternalKnowledgeApis) @@ -131,9 +137,9 @@ class ExternalDatasetService: if settings and settings.get("api_key") == HIDDEN_VALUE and external_knowledge_api.settings_dict: settings["api_key"] = external_knowledge_api.settings_dict.get("api_key") - external_knowledge_api.name = args.get("name") - external_knowledge_api.description = args.get("description", "") - external_knowledge_api.settings = json.dumps(args.get("settings"), ensure_ascii=False) + external_knowledge_api.name = str(args.get("name")) + external_knowledge_api.description = str(args.get("description", "")) + external_knowledge_api.settings = json.dumps(settings, ensure_ascii=False) external_knowledge_api.updated_by = user_id external_knowledge_api.updated_at = naive_utc_now() session.commit() @@ -141,7 +147,7 @@ class ExternalDatasetService: return external_knowledge_api @staticmethod - def delete_external_knowledge_api(session: Session, tenant_id: str, external_knowledge_api_id: str): + def delete_external_knowledge_api(tenant_id: str, external_knowledge_api_id: str, *, session: Session) -> None: external_knowledge_api = session.scalar( select(ExternalKnowledgeApis) .where(ExternalKnowledgeApis.id == external_knowledge_api_id, ExternalKnowledgeApis.tenant_id == tenant_id) @@ -155,7 +161,7 @@ class ExternalDatasetService: @staticmethod def external_knowledge_api_use_check( - session: Session, external_knowledge_api_id: str, tenant_id: str + external_knowledge_api_id: str, tenant_id: str, *, session: Session ) -> tuple[bool, int]: """ Return usage for an external knowledge API within a single tenant. @@ -176,7 +182,7 @@ class ExternalDatasetService: @staticmethod def get_external_knowledge_binding_with_dataset_id( - session: Session, tenant_id: str, dataset_id: str + tenant_id: str, dataset_id: str, *, session: Session ) -> ExternalKnowledgeBindings: external_knowledge_binding: ExternalKnowledgeBindings | None = session.scalar( select(ExternalKnowledgeBindings) @@ -189,8 +195,12 @@ class ExternalDatasetService: @staticmethod def document_create_args_validate( - session: Session, tenant_id: str, external_knowledge_api_id: str, process_parameter: dict[str, Any] - ): + tenant_id: str, + external_knowledge_api_id: str, + process_parameter: dict[str, Any], + *, + session: Session, + ) -> None: external_knowledge_api = session.scalar( select(ExternalKnowledgeApis) .where(ExternalKnowledgeApis.id == external_knowledge_api_id, ExternalKnowledgeApis.tenant_id == tenant_id) @@ -264,7 +274,7 @@ class ExternalDatasetService: return ExternalKnowledgeApiSetting.model_validate(settings) @staticmethod - def create_external_dataset(tenant_id: str, user_id: str, args: dict[str, Any], session: Session) -> Dataset: + def create_external_dataset(tenant_id: str, user_id: str, args: dict[str, Any], *, session: Session) -> Dataset: # check if dataset name already exists if session.scalar( select(Dataset).where(Dataset.name == args.get("name"), Dataset.tenant_id == tenant_id).limit(1) @@ -314,12 +324,13 @@ class ExternalDatasetService: @staticmethod def fetch_external_knowledge_retrieval( - session: Session, tenant_id: str, dataset_id: str, query: str, external_retrieval_parameters: dict[str, Any], metadata_condition: MetadataFilteringCondition | None = None, + *, + session: Session, ): """Fetch retrieval records from an external knowledge provider. diff --git a/api/services/file_service.py b/api/services/file_service.py index e41d74ad3eb..ec69af4e80c 100644 --- a/api/services/file_service.py +++ b/api/services/file_service.py @@ -268,7 +268,7 @@ class FileService: @staticmethod def get_upload_files_by_ids( - session: Session, tenant_id: str, upload_file_ids: Sequence[str] + tenant_id: str, upload_file_ids: Sequence[str], *, session: Session ) -> dict[str, UploadFile]: """ Fetch `UploadFile` rows for a tenant in a single batch query. diff --git a/api/services/hit_testing_service.py b/api/services/hit_testing_service.py index 1b51a2d279b..1bfa4025fa0 100644 --- a/api/services/hit_testing_service.py +++ b/api/services/hit_testing_service.py @@ -4,7 +4,7 @@ import time from typing import Any, TypedDict, cast from sqlalchemy import select -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from core.app.app_config.entities import ModelConfig from core.rag.datasource.retrieval_service import DefaultRetrievalModelDict, RetrievalService @@ -56,9 +56,7 @@ class HitTestingService: } @classmethod - def _dump_retrieval_records( - cls, session: Session | scoped_session, records: list[RetrievalSegments] - ) -> list[dict[str, Any]]: + def _dump_retrieval_records(cls, session: Session, records: list[RetrievalSegments]) -> list[dict[str, Any]]: document_ids = { document_id for record in records @@ -105,7 +103,6 @@ class HitTestingService: @classmethod def retrieve( cls, - session: Session, dataset: Dataset, query: str, account: Account, @@ -113,6 +110,8 @@ class HitTestingService: external_retrieval_model: dict[str, Any], attachment_ids: list | None = None, limit: int = 10, + *, + session: Session, ): start = time.perf_counter() @@ -144,7 +143,7 @@ class HitTestingService: if metadata_filter_document_ids: document_ids_filter = metadata_filter_document_ids.get(dataset.id, []) if metadata_condition and not document_ids_filter: - return cls.compact_retrieve_response(session, query, []) + return cls.compact_retrieve_response(query, [], session=session) all_documents = RetrievalService.retrieve( retrieval_method=RetrievalMethod( resolved_retrieval_model.get("search_method", RetrievalMethod.SEMANTIC_SEARCH) @@ -186,17 +185,18 @@ class HitTestingService: session.add(dataset_query) session.commit() - return cls.compact_retrieve_response(session, query, all_documents) + return cls.compact_retrieve_response(query, all_documents, session=session) @classmethod def external_retrieve( cls, - session: Session, dataset: Dataset, query: str, account: Account, external_retrieval_model: dict[str, Any] | None = None, metadata_filtering_conditions: dict[str, Any] | None = None, + *, + session: Session, ): if dataset.provider != "external": return { @@ -233,7 +233,7 @@ class HitTestingService: @classmethod def compact_retrieve_response( - cls, session: Session | scoped_session, query: str, documents: list[Document] + cls, query: str, documents: list[Document], *, session: Session ) -> RetrieveResponseDict: records = RetrievalService.format_retrieval_documents(documents) diff --git a/api/services/message_service.py b/api/services/message_service.py index e8d1b6232bc..4fbeb61e1f7 100644 --- a/api/services/message_service.py +++ b/api/services/message_service.py @@ -3,7 +3,7 @@ from collections.abc import Sequence from typing import cast from sqlalchemy import select -from sqlalchemy.orm import sessionmaker +from sqlalchemy.orm import Session, sessionmaker from core.app.apps.advanced_chat.app_config_manager import AdvancedChatAppConfigManager from core.app.entities.app_invoke_entities import InvokeFrom @@ -70,6 +70,8 @@ class MessageService: first_id: str | None, limit: int, order: str = "asc", + *, + session: Session, ) -> InfiniteScrollPagination: if not user: return InfiniteScrollPagination(data=[], limit=limit, has_more=False) @@ -78,20 +80,20 @@ class MessageService: return InfiniteScrollPagination(data=[], limit=limit, has_more=False) conversation = ConversationService.get_conversation( - app_model=app_model, user=user, conversation_id=conversation_id + app_model=app_model, user=user, conversation_id=conversation_id, session=session ) fetch_limit = limit + 1 if first_id: - first_message = db.session.scalar( + first_message = session.scalar( select(Message).where(Message.conversation_id == conversation.id, Message.id == first_id).limit(1) ) if not first_message: raise FirstMessageNotExistsError() - history_messages = db.session.scalars( + history_messages = session.scalars( select(Message) .where( Message.conversation_id == conversation.id, @@ -102,7 +104,7 @@ class MessageService: .limit(fetch_limit) ).all() else: - history_messages = db.session.scalars( + history_messages = session.scalars( select(Message) .where(Message.conversation_id == conversation.id) .order_by(Message.created_at.desc()) @@ -130,6 +132,8 @@ class MessageService: limit: int, conversation_id: str | None = None, include_ids: list | None = None, + *, + session: Session, ) -> InfiniteScrollPagination: if not user: return InfiniteScrollPagination(data=[], limit=limit, has_more=False) @@ -140,7 +144,7 @@ class MessageService: if conversation_id is not None: conversation = ConversationService.get_conversation( - app_model=app_model, user=user, conversation_id=conversation_id + app_model=app_model, user=user, conversation_id=conversation_id, session=session ) stmt = stmt.where(Message.conversation_id == conversation.id) @@ -152,18 +156,18 @@ class MessageService: stmt = stmt.where(Message.id.in_(include_ids)) if last_id: - last_message = db.session.scalar(stmt.where(Message.id == last_id).limit(1)) + last_message = session.scalar(stmt.where(Message.id == last_id).limit(1)) if not last_message: raise LastMessageNotExistsError() - history_messages = db.session.scalars( + history_messages = session.scalars( stmt.where(Message.created_at < last_message.created_at, Message.id != last_message.id) .order_by(Message.created_at.desc()) .limit(fetch_limit) ).all() else: - history_messages = db.session.scalars(stmt.order_by(Message.created_at.desc()).limit(fetch_limit)).all() + history_messages = session.scalars(stmt.order_by(Message.created_at.desc()).limit(fetch_limit)).all() has_more = False if len(history_messages) > limit: @@ -181,16 +185,17 @@ class MessageService: user: Account | EndUser | None, rating: FeedbackRating | None, content: str | None, + session: Session, ): if not user: raise ValueError("user cannot be None") - message = cls.get_message(app_model=app_model, user=user, message_id=message_id) + message = cls.get_message(app_model=app_model, user=user, message_id=message_id, session=session) feedback = message.user_feedback if isinstance(user, EndUser) else message.admin_feedback if not rating and feedback: - db.session.delete(feedback) + session.delete(feedback) elif rating and feedback: feedback.rating = rating feedback.content = content @@ -208,17 +213,17 @@ class MessageService: from_end_user_id=(user.id if isinstance(user, EndUser) else None), from_account_id=(user.id if isinstance(user, Account) else None), ) - db.session.add(feedback) + session.add(feedback) - db.session.commit() + session.commit() return feedback @classmethod - def get_all_messages_feedbacks(cls, app_model: App, page: int, limit: int): + def get_all_messages_feedbacks(cls, app_model: App, page: int, limit: int, *, session: Session): """Get all feedbacks of an app""" offset = (page - 1) * limit - feedbacks = db.session.scalars( + feedbacks = session.scalars( select(MessageFeedback) .where(MessageFeedback.app_id == app_model.id) .order_by(MessageFeedback.created_at.desc(), MessageFeedback.id.desc()) @@ -229,8 +234,8 @@ class MessageService: return [record.to_dict() for record in feedbacks] @classmethod - def get_message(cls, app_model: App, user: Account | EndUser | None, message_id: str): - message = db.session.scalar( + def get_message(cls, app_model: App, user: Account | EndUser | None, message_id: str, *, session: Session): + message = session.scalar( select(Message) .where( Message.id == message_id, @@ -249,15 +254,21 @@ class MessageService: @classmethod def get_suggested_questions_after_answer( - cls, app_model: App, user: Account | EndUser | None, message_id: str, invoke_from: InvokeFrom + cls, + app_model: App, + user: Account | EndUser | None, + message_id: str, + invoke_from: InvokeFrom, + *, + session: Session, ) -> list[str]: if not user: raise ValueError("user cannot be None") - message = cls.get_message(app_model=app_model, user=user, message_id=message_id) + message = cls.get_message(app_model=app_model, user=user, message_id=message_id, session=session) conversation = ConversationService.get_conversation( - app_model=app_model, conversation_id=message.conversation_id, user=user + app_model=app_model, conversation_id=message.conversation_id, user=user, session=session ) model_manager = ModelManager.for_tenant(tenant_id=app_model.tenant_id) @@ -266,9 +277,9 @@ class MessageService: if app_model.mode == AppMode.ADVANCED_CHAT: workflow_service = WorkflowService() if invoke_from == InvokeFrom.DEBUGGER: - workflow = workflow_service.get_draft_workflow(app_model=app_model) + workflow = workflow_service.get_draft_workflow(app_model=app_model, session=session) else: - workflow = workflow_service.get_published_workflow(app_model=app_model) + workflow = workflow_service.get_published_workflow(app_model=app_model, session=session) if workflow is None: return [] @@ -288,7 +299,7 @@ class MessageService: ) else: if not conversation.override_model_configs: - app_model_config = db.session.scalar( + app_model_config = session.scalar( select(AppModelConfig) .where(AppModelConfig.id == conversation.app_model_config_id, AppModelConfig.app_id == app_model.id) .limit(1) diff --git a/api/services/metadata_service.py b/api/services/metadata_service.py index 4e83858ea0e..481eb3b2e29 100644 --- a/api/services/metadata_service.py +++ b/api/services/metadata_service.py @@ -23,11 +23,12 @@ logger = logging.getLogger(__name__) class MetadataService: @staticmethod def create_metadata( - session: Session, dataset_id: str, metadata_args: MetadataArgs, current_user: Account | None = None, # TODO: the service_api is not migrated yet current_tenant_id: str | None = None, + *, + session: Session, ) -> DatasetMetadata: # check if metadata name is too long if len(metadata_args.name) > 255: @@ -60,12 +61,13 @@ class MetadataService: @staticmethod def update_metadata_name( - session: Session, dataset_id: str, metadata_id: str, name: str, current_user: Account | None = None, current_tenant_id: str | None = None, # TODO: the service_api is not migrated yet + *, + session: Session, ) -> DatasetMetadata | None: # check if metadata name is too long if len(name) > 255: @@ -126,7 +128,7 @@ class MetadataService: redis_client.delete(lock_key) @staticmethod - def delete_metadata(session: Session, dataset_id: str, metadata_id: str): + def delete_metadata(dataset_id: str, metadata_id: str, *, session: Session): lock_key = f"dataset_metadata_lock_{dataset_id}" try: MetadataService.knowledge_base_metadata_lock_check(dataset_id, None) @@ -172,7 +174,7 @@ class MetadataService: ] @staticmethod - def enable_built_in_field(session: Session, dataset: Dataset): + def enable_built_in_field(dataset: Dataset, *, session: Session): if dataset.built_in_field_enabled: return lock_key = f"dataset_metadata_lock_{dataset.id}" @@ -201,7 +203,7 @@ class MetadataService: redis_client.delete(lock_key) @staticmethod - def disable_built_in_field(session: Session, dataset: Dataset): + def disable_built_in_field(dataset: Dataset, *, session: Session): if not dataset.built_in_field_enabled: return lock_key = f"dataset_metadata_lock_{dataset.id}" @@ -233,11 +235,12 @@ class MetadataService: @staticmethod def update_documents_metadata( - session: Session, dataset: Dataset, metadata_args: MetadataOperationData, current_user: Account | None = None, # TODO: the service_api is not migrated yet current_tenant_id: str | None = None, + *, + session: Session, ): current_user, current_tenant_id = resolve_account_fallback( current_user, current_tenant_id, fallback_tenant_id=dataset.tenant_id @@ -316,7 +319,7 @@ class MetadataService: redis_client.set(lock_key, 1, ex=3600) @staticmethod - def get_dataset_metadatas(session: Session, dataset: Dataset): + def get_dataset_metadatas(dataset: Dataset, *, session: Session): return { "doc_metadata": [ { diff --git a/api/services/model_load_balancing_service.py b/api/services/model_load_balancing_service.py index 2a9094a35f2..6eab1ffbe3e 100644 --- a/api/services/model_load_balancing_service.py +++ b/api/services/model_load_balancing_service.py @@ -3,6 +3,7 @@ import logging from typing import Any, TypedDict, cast from sqlalchemy import or_, select +from sqlalchemy.orm import Session from constants import HIDDEN_VALUE from core.entities.provider_configuration import ProviderConfiguration @@ -14,7 +15,6 @@ from core.helper.model_provider_cache import ( from core.model_manager import LBModelManager from core.plugin.impl.model_runtime_factory import create_plugin_model_assembly, create_plugin_provider_manager from core.provider_manager import ProviderConfigurationCacheSource, ProviderManager -from extensions.ext_database import db from graphon.model_runtime.entities.model_entities import ModelType from graphon.model_runtime.entities.provider_entities import ( ModelCredentialSchema, @@ -93,7 +93,13 @@ class ModelLoadBalancingService: provider_configuration.disable_model_load_balancing(model=model, model_type=ModelType(model_type)) def get_load_balancing_configs( - self, tenant_id: str, provider: str, model: str, model_type: str, config_from: str = "" + self, + tenant_id: str, + provider: str, + model: str, + model_type: str, + session: Session, + config_from: str = "", ) -> tuple[bool, list[LoadBalancingConfigSummaryDict]]: """ Get load balancing configurations. @@ -131,7 +137,7 @@ class ModelLoadBalancingService: # Get load balancing configurations load_balancing_configs = list( - db.session.scalars( + session.scalars( select(LoadBalancingModelConfig) .where( LoadBalancingModelConfig.tenant_id == tenant_id, @@ -158,7 +164,7 @@ class ModelLoadBalancingService: if not inherit_config_exists: # Initialize the inherit configuration - inherit_config = self._init_inherit_config(tenant_id, provider, model, model_type_enum) + inherit_config = self._init_inherit_config(tenant_id, provider, model, model_type_enum, session=session) # prepend the inherit configuration load_balancing_configs.insert(0, inherit_config) @@ -233,7 +239,13 @@ class ModelLoadBalancingService: return is_load_balancing_enabled, datas def get_load_balancing_config( - self, tenant_id: str, provider: str, model: str, model_type: str, config_id: str + self, + tenant_id: str, + provider: str, + model: str, + model_type: str, + config_id: str, + session: Session, ) -> LoadBalancingConfigDetailDict | None: """ Get load balancing configuration. @@ -256,7 +268,7 @@ class ModelLoadBalancingService: model_type_enum = ModelType(model_type) # Get load balancing configurations - load_balancing_model_config = db.session.scalar( + load_balancing_model_config = session.scalar( select(LoadBalancingModelConfig) .where( LoadBalancingModelConfig.tenant_id == tenant_id, @@ -296,7 +308,12 @@ class ModelLoadBalancingService: return result def _init_inherit_config( - self, tenant_id: str, provider: str, model: str, model_type: ModelType + self, + tenant_id: str, + provider: str, + model: str, + model_type: ModelType, + session: Session, ) -> LoadBalancingModelConfig: """ Initialize the inherit configuration. @@ -314,8 +331,8 @@ class ModelLoadBalancingService: model_name=model, name="__inherit__", ) - db.session.add(inherit_config) - db.session.commit() + session.add(inherit_config) + session.commit() ProviderManager.invalidate_configurations_cache( tenant_id, sources=(ProviderConfigurationCacheSource.PROVIDER_LOAD_BALANCING_CONFIGS,), @@ -324,7 +341,14 @@ class ModelLoadBalancingService: return inherit_config def update_load_balancing_configs( - self, tenant_id: str, provider: str, model: str, model_type: str, configs: list[dict], config_from: str + self, + tenant_id: str, + provider: str, + model: str, + model_type: str, + configs: list[dict], + config_from: str, + session: Session, ): """ Update load balancing configurations. @@ -350,7 +374,7 @@ class ModelLoadBalancingService: if not isinstance(configs, list): raise ValueError("Invalid load balancing configs") - current_load_balancing_configs = db.session.scalars( + current_load_balancing_configs = session.scalars( select(LoadBalancingModelConfig).where( LoadBalancingModelConfig.tenant_id == tenant_id, LoadBalancingModelConfig.provider_name == provider_configuration.provider.provider, @@ -377,7 +401,7 @@ class ModelLoadBalancingService: if credential_id: if config_from == "predefined-model": - credential_record = db.session.scalar( + credential_record = session.scalar( select(ProviderCredential) .where( ProviderCredential.id == credential_id, @@ -387,7 +411,7 @@ class ModelLoadBalancingService: .limit(1) ) else: - credential_record = db.session.scalar( + credential_record = session.scalar( select(ProviderModelCredential) .where( ProviderModelCredential.id == credential_id, @@ -440,7 +464,7 @@ class ModelLoadBalancingService: load_balancing_config.name = name load_balancing_config.enabled = enabled load_balancing_config.updated_at = naive_utc_now() - db.session.commit() + session.commit() ProviderManager.invalidate_configurations_cache( tenant_id, sources=(ProviderConfigurationCacheSource.PROVIDER_LOAD_BALANCING_CONFIGS,), @@ -496,8 +520,8 @@ class ModelLoadBalancingService: encrypted_config=json.dumps(credentials), ) - db.session.add(load_balancing_model_config) - db.session.commit() + session.add(load_balancing_model_config) + session.commit() ProviderManager.invalidate_configurations_cache( tenant_id, sources=(ProviderConfigurationCacheSource.PROVIDER_LOAD_BALANCING_CONFIGS,), @@ -506,8 +530,8 @@ class ModelLoadBalancingService: # get deleted config ids deleted_config_ids = set(current_load_balancing_configs_dict.keys()) - updated_config_ids for config_id in deleted_config_ids: - db.session.delete(current_load_balancing_configs_dict[config_id]) - db.session.commit() + session.delete(current_load_balancing_configs_dict[config_id]) + session.commit() ProviderManager.invalidate_configurations_cache( tenant_id, sources=(ProviderConfigurationCacheSource.PROVIDER_LOAD_BALANCING_CONFIGS,), @@ -522,6 +546,7 @@ class ModelLoadBalancingService: model: str, model_type: str, credentials: dict[str, Any], + session: Session, config_id: str | None = None, ): """ @@ -548,7 +573,7 @@ class ModelLoadBalancingService: load_balancing_model_config = None if config_id: # Get load balancing config - load_balancing_model_config = db.session.scalar( + load_balancing_model_config = session.scalar( select(LoadBalancingModelConfig) .where( LoadBalancingModelConfig.tenant_id == tenant_id, diff --git a/api/services/oauth_device_flow.py b/api/services/oauth_device_flow.py index 9ec5711890b..9e59b8c326a 100644 --- a/api/services/oauth_device_flow.py +++ b/api/services/oauth_device_flow.py @@ -13,7 +13,7 @@ from enum import StrEnum from typing import Any, NotRequired, TypedDict from sqlalchemy import and_, func, select, update -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from libs.oauth_bearer import TOKEN_CACHE_KEY_FMT, AuthContext, SubjectType from models.oauth import OAuthAccessToken @@ -335,9 +335,6 @@ def sha256_hex(token: str) -> str: def mint_oauth_token( - # Accept either Session or Flask-SQLAlchemy's request-scoped wrapper — - # the wrapper proxies the same execute/commit surface. - session: Session | scoped_session, redis_client, *, subject_email: str, @@ -347,6 +344,7 @@ def mint_oauth_token( device_label: str, prefix: str, ttl_days: int, + session: Session, ) -> MintResult: """Live row rotates in place via partial unique index ``uq_oauth_active_per_device``; hard-expired rows are excluded by the @@ -390,7 +388,7 @@ def mint_oauth_token( def _upsert( - session: Session | scoped_session, + session: Session, *, subject_email: str, subject_issuer: str | None, @@ -501,11 +499,7 @@ def subject_match_clauses(ctx: AuthContext) -> tuple[Any, ...]: ) -def list_active_sessions( - session: Session | scoped_session, - ctx: AuthContext, - now: datetime, -) -> list[OAuthAccessToken]: +def list_active_sessions(ctx: AuthContext, now: datetime, *, session: Session) -> list[OAuthAccessToken]: return list( session.execute( select(OAuthAccessToken) @@ -524,11 +518,7 @@ def list_active_sessions( ) -def token_belongs_to_subject( - session: Session | scoped_session, - token_id: str, - ctx: AuthContext, -) -> bool: +def token_belongs_to_subject(token_id: str, ctx: AuthContext, *, session: Session) -> bool: row = session.execute( select(OAuthAccessToken.id).where( and_( @@ -540,11 +530,7 @@ def token_belongs_to_subject( return row is not None -def revoke_oauth_token( - session: Session | scoped_session, - redis_client: Any, - token_id: str, -) -> None: +def revoke_oauth_token(redis_client: Any, token_id: str, *, session: Session) -> None: row = ( session.query(OAuthAccessToken.token_hash) .filter( diff --git a/api/services/oauth_server.py b/api/services/oauth_server.py index 5f3277c9525..47aa1bc99bf 100644 --- a/api/services/oauth_server.py +++ b/api/services/oauth_server.py @@ -2,7 +2,7 @@ import enum import uuid from sqlalchemy import select -from sqlalchemy.orm import sessionmaker +from sqlalchemy.orm import Session, sessionmaker from werkzeug.exceptions import BadRequest from extensions.ext_database import db @@ -83,7 +83,7 @@ class OAuthServerService: return token @staticmethod - def validate_oauth_access_token(client_id: str, token: str) -> Account | None: + def validate_oauth_access_token(client_id: str, token: str, session: Session) -> Account | None: redis_key = OAUTH_ACCESS_TOKEN_REDIS_KEY.format(client_id=client_id, token=token) user_account_id = redis_client.get(redis_key) if not user_account_id: @@ -91,4 +91,4 @@ class OAuthServerService: user_id_str = user_account_id.decode("utf-8") - return AccountService.load_user(user_id_str, db.session) + return AccountService.load_user(user_id_str, session) diff --git a/api/services/ops_service.py b/api/services/ops_service.py index 3ad42faf249..b6f17168b3c 100644 --- a/api/services/ops_service.py +++ b/api/services/ops_service.py @@ -1,23 +1,23 @@ from typing import Any from sqlalchemy import select +from sqlalchemy.orm import Session from core.ops.entities.config_entity import BaseTracingConfig from core.ops.ops_trace_manager import OpsTraceManager, TracingProviderConfigEntry, provider_config_map -from extensions.ext_database import db from models.model import App, TraceAppConfig class OpsService: @classmethod - def get_tracing_app_config(cls, app_id: str, tracing_provider: str): + def get_tracing_app_config(cls, app_id: str, tracing_provider: str, session: Session): """ Get tracing app config :param app_id: app id :param tracing_provider: tracing provider :return: """ - trace_config_data: TraceAppConfig | None = db.session.scalar( + trace_config_data: TraceAppConfig | None = session.scalar( select(TraceAppConfig) .where(TraceAppConfig.app_id == app_id, TraceAppConfig.tracing_provider == tracing_provider) .limit(1) @@ -27,7 +27,7 @@ class OpsService: return None # decrypt_token and obfuscated_token - app = db.session.get(App, app_id) + app = session.get(App, app_id) if not app: return None tenant_id = app.tenant_id @@ -137,7 +137,9 @@ class OpsService: return trace_config_data.to_dict() @classmethod - def create_tracing_app_config(cls, app_id: str, tracing_provider: str, tracing_config: dict[str, Any]): + def create_tracing_app_config( + cls, app_id: str, tracing_provider: str, tracing_config: dict[str, Any], session: Session + ): """ Create tracing app config :param app_id: app id @@ -184,7 +186,7 @@ class OpsService: project_url = None # check if trace config already exists - trace_config_data: TraceAppConfig | None = db.session.scalar( + trace_config_data: TraceAppConfig | None = session.scalar( select(TraceAppConfig) .where(TraceAppConfig.app_id == app_id, TraceAppConfig.tracing_provider == tracing_provider) .limit(1) @@ -194,7 +196,7 @@ class OpsService: return None # get tenant id - app = db.session.get(App, app_id) + app = session.get(App, app_id) if not app: return None tenant_id = app.tenant_id @@ -206,13 +208,15 @@ class OpsService: tracing_provider=tracing_provider, tracing_config=tracing_config, ) - db.session.add(trace_config_data) - db.session.commit() + session.add(trace_config_data) + session.commit() return {"result": "success"} @classmethod - def update_tracing_app_config(cls, app_id: str, tracing_provider: str, tracing_config: dict[str, Any]): + def update_tracing_app_config( + cls, app_id: str, tracing_provider: str, tracing_config: dict[str, Any], session: Session + ): """ Update tracing app config :param app_id: app id @@ -226,7 +230,7 @@ class OpsService: raise ValueError(f"Invalid tracing provider: {tracing_provider}") # check if trace config already exists - current_trace_config = db.session.scalar( + current_trace_config = session.scalar( select(TraceAppConfig) .where(TraceAppConfig.app_id == app_id, TraceAppConfig.tracing_provider == tracing_provider) .limit(1) @@ -236,7 +240,7 @@ class OpsService: return None # get tenant id - app = db.session.get(App, app_id) + app = session.get(App, app_id) if not app: return None tenant_id = app.tenant_id @@ -251,19 +255,19 @@ class OpsService: raise ValueError("Invalid Credentials") current_trace_config.tracing_config = tracing_config - db.session.commit() + session.commit() return current_trace_config.to_dict() @classmethod - def delete_tracing_app_config(cls, app_id: str, tracing_provider: str): + def delete_tracing_app_config(cls, app_id: str, tracing_provider: str, session: Session): """ Delete tracing app config :param app_id: app id :param tracing_provider: tracing provider :return: """ - trace_config = db.session.scalar( + trace_config = session.scalar( select(TraceAppConfig) .where(TraceAppConfig.app_id == app_id, TraceAppConfig.tracing_provider == tracing_provider) .limit(1) @@ -272,7 +276,7 @@ class OpsService: if not trace_config: return None - db.session.delete(trace_config) - db.session.commit() + session.delete(trace_config) + session.commit() return True diff --git a/api/services/plugin/plugin_auto_upgrade_service.py b/api/services/plugin/plugin_auto_upgrade_service.py index f1e1918bdd2..79770063016 100644 --- a/api/services/plugin/plugin_auto_upgrade_service.py +++ b/api/services/plugin/plugin_auto_upgrade_service.py @@ -12,7 +12,6 @@ from hashlib import sha256 from sqlalchemy import select from sqlalchemy.orm import Session -from core.db.session_factory import session_factory from core.plugin.impl.plugin import PluginInstaller from models.account import ( TenantPluginAutoUpgradeCategory, @@ -141,6 +140,8 @@ class PluginAutoUpgradeService: @staticmethod def backfill_strategy_categories( tenant_id: str, + *, + session: Session, ) -> PluginAutoUpgradeBackfillResult: """Create missing category strategies and split include/exclude lists when needed. @@ -148,89 +149,85 @@ class PluginAutoUpgradeService: New category rows copy it first, then plugin lists are narrowed by real plugin category when the source strategy contains include/exclude IDs. """ - with session_factory.create_session() as session, session.begin(): - strategies = list( - session.scalars( - select(TenantPluginAutoUpgradeStrategy).where( - TenantPluginAutoUpgradeStrategy.tenant_id == tenant_id - ) - ).all() + strategies = list( + session.scalars( + select(TenantPluginAutoUpgradeStrategy).where(TenantPluginAutoUpgradeStrategy.tenant_id == tenant_id) + ).all() + ) + if not strategies: + return PluginAutoUpgradeBackfillResult(created_count=0, normalized=False) + + # Schema migration marks the historical workspace-level row as tool. + source_strategy = next( + (strategy for strategy in strategies if strategy.category == PluginCategory.TOOL), + strategies[0], + ) + source_has_default_strategy = PluginAutoUpgradeService._has_default_strategy(source_strategy) + strategies_by_category = {strategy.category: strategy for strategy in strategies} + exclude_plugins = source_strategy.exclude_plugins + include_plugins = source_strategy.include_plugins + should_split_plugin_lists = bool(exclude_plugins or include_plugins) + # Query daemon only for tenants that actually customized plugin lists. + plugin_categories = ( + PluginAutoUpgradeService._get_installed_plugin_categories(tenant_id) if should_split_plugin_lists else {} + ) + if should_split_plugin_lists: + PluginAutoUpgradeService._log_unknown_plugin_ids( + tenant_id, + "exclude_plugins", + exclude_plugins, + plugin_categories, ) - if not strategies: - return PluginAutoUpgradeBackfillResult(created_count=0, normalized=False) - - # Schema migration marks the historical workspace-level row as tool. - source_strategy = next( - (strategy for strategy in strategies if strategy.category == PluginCategory.TOOL), - strategies[0], + PluginAutoUpgradeService._log_unknown_plugin_ids( + tenant_id, + "include_plugins", + include_plugins, + plugin_categories, ) - source_has_default_strategy = PluginAutoUpgradeService._has_default_strategy(source_strategy) - strategies_by_category = {strategy.category: strategy for strategy in strategies} - exclude_plugins = source_strategy.exclude_plugins - include_plugins = source_strategy.include_plugins - should_split_plugin_lists = bool(exclude_plugins or include_plugins) - # Query daemon only for tenants that actually customized plugin lists. - plugin_categories = ( - PluginAutoUpgradeService._get_installed_plugin_categories(tenant_id) - if should_split_plugin_lists - else {} + + created_count = 0 + for category in PLUGIN_CATEGORIES: + strategy = strategies_by_category.get(category) + if strategy is None: + # Start from the legacy workspace-level behavior before narrowing lists. + strategy = TenantPluginAutoUpgradeStrategy( + tenant_id=tenant_id, + category=category, + strategy_setting=PluginAutoUpgradeService._strategy_setting_for_category( + source_strategy, category, source_has_default_strategy + ), + upgrade_time_of_day=PluginAutoUpgradeService._upgrade_time_of_day_for_category( + tenant_id, source_strategy, source_has_default_strategy + ), + upgrade_mode=source_strategy.upgrade_mode, + exclude_plugins=source_strategy.exclude_plugins.copy(), + include_plugins=source_strategy.include_plugins.copy(), + ) + session.add(strategy) + created_count += 1 + elif source_has_default_strategy: + strategy.strategy_setting = PluginAutoUpgradeService.default_strategy_setting_for_category( + strategy.category + ) + strategy.upgrade_time_of_day = PluginAutoUpgradeService.default_upgrade_time_of_day(tenant_id) + + if not should_split_plugin_lists: + continue + + # Narrow include/exclude lists to the current category after all rows exist. + strategy.exclude_plugins = PluginAutoUpgradeService._filter_plugin_ids_for_category( + exclude_plugins, + strategy.category, + plugin_categories, + ) + strategy.include_plugins = PluginAutoUpgradeService._filter_plugin_ids_for_category( + include_plugins, + strategy.category, + plugin_categories, ) - if should_split_plugin_lists: - PluginAutoUpgradeService._log_unknown_plugin_ids( - tenant_id, - "exclude_plugins", - exclude_plugins, - plugin_categories, - ) - PluginAutoUpgradeService._log_unknown_plugin_ids( - tenant_id, - "include_plugins", - include_plugins, - plugin_categories, - ) - created_count = 0 - for category in PLUGIN_CATEGORIES: - strategy = strategies_by_category.get(category) - if strategy is None: - # Start from the legacy workspace-level behavior before narrowing lists. - strategy = TenantPluginAutoUpgradeStrategy( - tenant_id=tenant_id, - category=category, - strategy_setting=PluginAutoUpgradeService._strategy_setting_for_category( - source_strategy, category, source_has_default_strategy - ), - upgrade_time_of_day=PluginAutoUpgradeService._upgrade_time_of_day_for_category( - tenant_id, source_strategy, source_has_default_strategy - ), - upgrade_mode=source_strategy.upgrade_mode, - exclude_plugins=source_strategy.exclude_plugins.copy(), - include_plugins=source_strategy.include_plugins.copy(), - ) - session.add(strategy) - created_count += 1 - elif source_has_default_strategy: - strategy.strategy_setting = PluginAutoUpgradeService.default_strategy_setting_for_category( - strategy.category - ) - strategy.upgrade_time_of_day = PluginAutoUpgradeService.default_upgrade_time_of_day(tenant_id) - - if not should_split_plugin_lists: - continue - - # Narrow include/exclude lists to the current category after all rows exist. - strategy.exclude_plugins = PluginAutoUpgradeService._filter_plugin_ids_for_category( - exclude_plugins, - strategy.category, - plugin_categories, - ) - strategy.include_plugins = PluginAutoUpgradeService._filter_plugin_ids_for_category( - include_plugins, - strategy.category, - plugin_categories, - ) - - return PluginAutoUpgradeBackfillResult(created_count=created_count, normalized=should_split_plugin_lists) + session.commit() + return PluginAutoUpgradeBackfillResult(created_count=created_count, normalized=should_split_plugin_lists) @staticmethod def _get_strategy( @@ -251,20 +248,18 @@ class PluginAutoUpgradeService: def get_strategy( tenant_id: str, category: PluginCategory, + *, + session: Session, ) -> TenantPluginAutoUpgradeStrategy | None: - with session_factory.create_session() as session: - return PluginAutoUpgradeService._get_strategy(session, tenant_id, category) + return PluginAutoUpgradeService._get_strategy(session, tenant_id, category) @staticmethod - def get_strategies(tenant_id: str) -> list[TenantPluginAutoUpgradeStrategy]: - with session_factory.create_session() as session: - return list( - session.scalars( - select(TenantPluginAutoUpgradeStrategy).where( - TenantPluginAutoUpgradeStrategy.tenant_id == tenant_id - ) - ).all() - ) + def get_strategies(tenant_id: str, *, session: Session) -> list[TenantPluginAutoUpgradeStrategy]: + return list( + session.scalars( + select(TenantPluginAutoUpgradeStrategy).where(TenantPluginAutoUpgradeStrategy.tenant_id == tenant_id) + ).all() + ) @staticmethod def _change_strategy( @@ -305,20 +300,22 @@ class PluginAutoUpgradeService: exclude_plugins: list[str], include_plugins: list[str], category: PluginCategory, + *, + session: Session, ) -> bool: - with session_factory.create_session() as session, session.begin(): - PluginAutoUpgradeService._change_strategy( - session, - tenant_id=tenant_id, - category=category, - strategy_setting=strategy_setting, - upgrade_time_of_day=upgrade_time_of_day, - upgrade_mode=upgrade_mode, - exclude_plugins=exclude_plugins, - include_plugins=include_plugins, - ) + PluginAutoUpgradeService._change_strategy( + session, + tenant_id=tenant_id, + category=category, + strategy_setting=strategy_setting, + upgrade_time_of_day=upgrade_time_of_day, + upgrade_mode=upgrade_mode, + exclude_plugins=exclude_plugins, + include_plugins=include_plugins, + ) - return True + session.commit() + return True @staticmethod def _exclude_plugin( @@ -363,13 +360,15 @@ class PluginAutoUpgradeService: tenant_id: str, plugin_id: str, category: PluginCategory, + *, + session: Session, ) -> bool: - with session_factory.create_session() as session, session.begin(): - PluginAutoUpgradeService._exclude_plugin( - session, - tenant_id, - category, - plugin_id, - ) + PluginAutoUpgradeService._exclude_plugin( + session, + tenant_id, + category, + plugin_id, + ) - return True + session.commit() + return True diff --git a/api/services/plugin/plugin_permission_service.py b/api/services/plugin/plugin_permission_service.py index 19f3de2e52c..339a6ccb89b 100644 --- a/api/services/plugin/plugin_permission_service.py +++ b/api/services/plugin/plugin_permission_service.py @@ -1,35 +1,36 @@ from sqlalchemy import select +from sqlalchemy.orm import Session -from core.db.session_factory import session_factory from models.account import TenantPluginDebugPermission, TenantPluginInstallPermission, TenantPluginPermission class PluginPermissionService: @staticmethod - def get_permission(tenant_id: str) -> TenantPluginPermission | None: - with session_factory.create_session() as session: - return session.scalar( - select(TenantPluginPermission).where(TenantPluginPermission.tenant_id == tenant_id).limit(1) - ) + def get_permission(tenant_id: str, *, session: Session) -> TenantPluginPermission | None: + return session.scalar( + select(TenantPluginPermission).where(TenantPluginPermission.tenant_id == tenant_id).limit(1) + ) @staticmethod def change_permission( tenant_id: str, install_permission: TenantPluginInstallPermission, debug_permission: TenantPluginDebugPermission, - ): - with session_factory.create_session() as session, session.begin(): - permission = session.scalar( - select(TenantPluginPermission).where(TenantPluginPermission.tenant_id == tenant_id).limit(1) + *, + session: Session, + ) -> bool: + permission = session.scalar( + select(TenantPluginPermission).where(TenantPluginPermission.tenant_id == tenant_id).limit(1) + ) + if not permission: + permission = TenantPluginPermission( + tenant_id=tenant_id, install_permission=install_permission, debug_permission=debug_permission ) - if not permission: - permission = TenantPluginPermission( - tenant_id=tenant_id, install_permission=install_permission, debug_permission=debug_permission - ) - session.add(permission) - else: - permission.install_permission = install_permission - permission.debug_permission = debug_permission + session.add(permission) + else: + permission.install_permission = install_permission + permission.debug_permission = debug_permission - return True + session.commit() + return True diff --git a/api/services/rag_pipeline/pipeline_generate_service.py b/api/services/rag_pipeline/pipeline_generate_service.py index e77ff9687ed..276bfaea158 100644 --- a/api/services/rag_pipeline/pipeline_generate_service.py +++ b/api/services/rag_pipeline/pipeline_generate_service.py @@ -17,12 +17,13 @@ class PipelineGenerateService: @classmethod def generate( cls, - session: Session, pipeline: Pipeline, user: Account | EndUser, args: Mapping[str, Any], invoke_from: InvokeFrom, streaming: bool = True, + *, + session: Session, ): """ Pipeline Content Generate @@ -34,10 +35,10 @@ class PipelineGenerateService: :return: """ try: - workflow = cls._get_workflow(pipeline, invoke_from) + workflow = cls._get_workflow(pipeline, invoke_from, session) if original_document_id := args.get("original_document_id"): # update document status to waiting - cls.update_document_status(original_document_id, session) + cls.update_document_status(original_document_id, session=session) return PipelineGenerator.convert_to_event_stream( PipelineGenerator().generate( pipeline=pipeline, @@ -64,9 +65,9 @@ class PipelineGenerateService: @classmethod def generate_single_iteration( - cls, pipeline: Pipeline, user: Account, node_id: str, args: Any, streaming: bool = True + cls, pipeline: Pipeline, user: Account, node_id: str, args: Any, session: Session, streaming: bool = True ): - workflow = cls._get_workflow(pipeline, InvokeFrom.DEBUGGER) + workflow = cls._get_workflow(pipeline, InvokeFrom.DEBUGGER, session) return PipelineGenerator.convert_to_event_stream( PipelineGenerator().single_iteration_generate( pipeline=pipeline, workflow=workflow, node_id=node_id, user=user, args=args, streaming=streaming @@ -74,8 +75,10 @@ class PipelineGenerateService: ) @classmethod - def generate_single_loop(cls, pipeline: Pipeline, user: Account, node_id: str, args: Any, streaming: bool = True): - workflow = cls._get_workflow(pipeline, InvokeFrom.DEBUGGER) + def generate_single_loop( + cls, pipeline: Pipeline, user: Account, node_id: str, args: Any, session: Session, streaming: bool = True + ): + workflow = cls._get_workflow(pipeline, InvokeFrom.DEBUGGER, session) return PipelineGenerator.convert_to_event_stream( PipelineGenerator().single_loop_generate( pipeline=pipeline, workflow=workflow, node_id=node_id, user=user, args=args, streaming=streaming @@ -83,14 +86,14 @@ class PipelineGenerateService: ) @classmethod - def _get_workflow(cls, pipeline: Pipeline, invoke_from: InvokeFrom) -> Workflow: + def _get_workflow(cls, pipeline: Pipeline, invoke_from: InvokeFrom, session: Session) -> Workflow: """ Get workflow :param pipeline: pipeline :param invoke_from: invoke from :return: """ - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(session) if invoke_from == InvokeFrom.DEBUGGER: # fetch draft workflow by app_model workflow = rag_pipeline_service.get_draft_workflow(pipeline=pipeline) @@ -107,7 +110,7 @@ class PipelineGenerateService: return workflow @classmethod - def update_document_status(cls, document_id: str, session: Session): + def update_document_status(cls, document_id: str, *, session: Session): """ Update document status to waiting :param document_id: document id diff --git a/api/services/rag_pipeline/pipeline_template/built_in/built_in_retrieval.py b/api/services/rag_pipeline/pipeline_template/built_in/built_in_retrieval.py index 6de0be33a4a..d56c239ace2 100644 --- a/api/services/rag_pipeline/pipeline_template/built_in/built_in_retrieval.py +++ b/api/services/rag_pipeline/pipeline_template/built_in/built_in_retrieval.py @@ -23,14 +23,15 @@ class BuiltInPipelineTemplateRetrieval(PipelineTemplateRetrievalBase): @override def get_pipeline_templates( - self, session: Session, language: str, current_tenant_id: str | None = None + self, language: str, current_tenant_id: str | None = None, *, session: Session ) -> dict[str, Any]: - del current_tenant_id + del current_tenant_id, session result = self.fetch_pipeline_templates_from_builtin(language) return result @override - def get_pipeline_template_detail(self, session: Session, template_id: str) -> dict[str, Any] | None: + def get_pipeline_template_detail(self, template_id: str, *, session: Session) -> dict[str, Any] | None: + del session result = self.fetch_pipeline_template_detail_from_builtin(template_id) return result diff --git a/api/services/rag_pipeline/pipeline_template/customized/customized_retrieval.py b/api/services/rag_pipeline/pipeline_template/customized/customized_retrieval.py index 4faaf342f66..3d6baefcc46 100644 --- a/api/services/rag_pipeline/pipeline_template/customized/customized_retrieval.py +++ b/api/services/rag_pipeline/pipeline_template/customized/customized_retrieval.py @@ -41,16 +41,16 @@ class CustomizedPipelineTemplateRetrieval(PipelineTemplateRetrievalBase): @override def get_pipeline_templates( - self, session: Session, language: str, current_tenant_id: str | None = None + self, language: str, current_tenant_id: str | None = None, *, session: Session ) -> dict[str, Any]: current_tenant_id = resolve_tenant_id_fallback(current_tenant_id) return self.fetch_pipeline_templates_from_customized( - session=session, tenant_id=current_tenant_id, language=language + tenant_id=current_tenant_id, language=language, session=session ) @override - def get_pipeline_template_detail(self, session: Session, template_id: str) -> dict[str, Any] | None: - return self.fetch_pipeline_template_detail_from_db(session, template_id) + def get_pipeline_template_detail(self, template_id: str, *, session: Session) -> dict[str, Any] | None: + return self.fetch_pipeline_template_detail_from_db(template_id, session=session) @override def get_type(self) -> str: @@ -58,7 +58,7 @@ class CustomizedPipelineTemplateRetrieval(PipelineTemplateRetrievalBase): @classmethod def fetch_pipeline_templates_from_customized( - cls, session: Session, tenant_id: str, language: str + cls, tenant_id: str, language: str, *, session: Session ) -> dict[str, Any]: """ Fetch pipeline templates from db. @@ -89,7 +89,7 @@ class CustomizedPipelineTemplateRetrieval(PipelineTemplateRetrievalBase): return {"pipeline_templates": recommended_pipelines_results} @classmethod - def fetch_pipeline_template_detail_from_db(cls, session: Session, template_id: str) -> dict[str, Any] | None: + def fetch_pipeline_template_detail_from_db(cls, template_id: str, *, session: Session) -> dict[str, Any] | None: """ Fetch pipeline template detail from db. :param template_id: Template ID diff --git a/api/services/rag_pipeline/pipeline_template/database/database_retrieval.py b/api/services/rag_pipeline/pipeline_template/database/database_retrieval.py index f6d2731e21a..d5c31ff74b2 100644 --- a/api/services/rag_pipeline/pipeline_template/database/database_retrieval.py +++ b/api/services/rag_pipeline/pipeline_template/database/database_retrieval.py @@ -41,21 +41,21 @@ class DatabasePipelineTemplateRetrieval(PipelineTemplateRetrievalBase): @override def get_pipeline_templates( - self, session: Session, language: str, current_tenant_id: str | None = None + self, language: str, current_tenant_id: str | None = None, *, session: Session ) -> dict[str, Any]: del current_tenant_id - return self.fetch_pipeline_templates_from_db(session, language) + return self.fetch_pipeline_templates_from_db(language, session=session) @override - def get_pipeline_template_detail(self, session: Session, template_id: str) -> dict[str, Any] | None: - return self.fetch_pipeline_template_detail_from_db(session, template_id) + def get_pipeline_template_detail(self, template_id: str, *, session: Session) -> dict[str, Any] | None: + return self.fetch_pipeline_template_detail_from_db(template_id, session=session) @override def get_type(self) -> str: return PipelineTemplateType.DATABASE @classmethod - def fetch_pipeline_templates_from_db(cls, session: Session, language: str) -> dict[str, Any]: + def fetch_pipeline_templates_from_db(cls, language: str, *, session: Session) -> dict[str, Any]: """ Fetch pipeline templates from db. :param language: language @@ -83,7 +83,7 @@ class DatabasePipelineTemplateRetrieval(PipelineTemplateRetrievalBase): return {"pipeline_templates": recommended_pipelines_results} @classmethod - def fetch_pipeline_template_detail_from_db(cls, session: Session, template_id: str) -> dict[str, Any] | None: + def fetch_pipeline_template_detail_from_db(cls, template_id: str, *, session: Session) -> dict[str, Any] | None: """ Fetch pipeline template detail from db. :param pipeline_id: Pipeline ID diff --git a/api/services/rag_pipeline/pipeline_template/pipeline_template_base.py b/api/services/rag_pipeline/pipeline_template/pipeline_template_base.py index ff53dc1f79e..c61ac6d60f2 100644 --- a/api/services/rag_pipeline/pipeline_template/pipeline_template_base.py +++ b/api/services/rag_pipeline/pipeline_template/pipeline_template_base.py @@ -7,9 +7,9 @@ class PipelineTemplateRetrievalBase(Protocol): """Interface for pipeline template retrieval.""" def get_pipeline_templates( - self, session: Session, language: str, current_tenant_id: str | None = None + self, language: str, current_tenant_id: str | None = None, *, session: Session ) -> dict[str, Any]: ... - def get_pipeline_template_detail(self, session: Session, template_id: str) -> dict[str, Any] | None: ... + def get_pipeline_template_detail(self, template_id: str, *, session: Session) -> dict[str, Any] | None: ... def get_type(self) -> str: ... diff --git a/api/services/rag_pipeline/pipeline_template/remote/remote_retrieval.py b/api/services/rag_pipeline/pipeline_template/remote/remote_retrieval.py index 7f9fe1b56ea..29acbd198b6 100644 --- a/api/services/rag_pipeline/pipeline_template/remote/remote_retrieval.py +++ b/api/services/rag_pipeline/pipeline_template/remote/remote_retrieval.py @@ -18,23 +18,25 @@ class RemotePipelineTemplateRetrieval(PipelineTemplateRetrievalBase): """ @override - def get_pipeline_template_detail(self, session: Session, template_id: str) -> dict[str, Any] | None: + def get_pipeline_template_detail(self, template_id: str, *, session: Session) -> dict[str, Any] | None: try: return self.fetch_pipeline_template_detail_from_dify_official(template_id) except Exception as e: logger.warning("fetch recommended app detail from dify official failed: %r, switch to database.", e) - return DatabasePipelineTemplateRetrieval.fetch_pipeline_template_detail_from_db(session, template_id) + return DatabasePipelineTemplateRetrieval.fetch_pipeline_template_detail_from_db( + template_id, session=session + ) @override def get_pipeline_templates( - self, session: Session, language: str, current_tenant_id: str | None = None + self, language: str, current_tenant_id: str | None = None, *, session: Session ) -> dict[str, Any]: del current_tenant_id try: return self.fetch_pipeline_templates_from_dify_official(language) except Exception as e: logger.warning("fetch pipeline templates from dify official failed: %r, switch to database.", e) - return DatabasePipelineTemplateRetrieval.fetch_pipeline_templates_from_db(session, language) + return DatabasePipelineTemplateRetrieval.fetch_pipeline_templates_from_db(language, session=session) @override def get_type(self) -> str: diff --git a/api/services/rag_pipeline/rag_pipeline.py b/api/services/rag_pipeline/rag_pipeline.py index 9e17a05be16..8bd3918eb15 100644 --- a/api/services/rag_pipeline/rag_pipeline.py +++ b/api/services/rag_pipeline/rag_pipeline.py @@ -27,7 +27,6 @@ from core.datasource.entities.datasource_entities import ( from core.datasource.online_document.online_document_plugin import OnlineDocumentDatasourcePlugin from core.datasource.online_drive.online_drive_plugin import OnlineDriveDatasourcePlugin from core.datasource.website_crawl.website_crawl_plugin import WebsiteCrawlDatasourcePlugin -from core.db.session_factory import session_factory from core.helper import marketplace from core.rag.entities import DatasourceCompletedEvent, DatasourceErrorEvent, DatasourceProcessingEvent from core.repositories.factory import DifyCoreRepositoryFactory, OrderConfig @@ -96,11 +95,13 @@ def _build_seeded_variable_pool(variables: Sequence[Variable]) -> VariablePool: class RagPipelineService: - def __init__(self, session_maker: sessionmaker | None = None): + _session: Session + + def __init__(self, session: Session, session_maker: sessionmaker | None = None): """Initialize RagPipelineService with repository dependencies.""" + self._session = session if session_maker is None: - session_maker = session_factory.get_session_maker() - self._session_maker = session_maker + session_maker = sessionmaker(bind=db.engine, expire_on_commit=False) self._node_execution_service_repo = DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository( session_maker ) @@ -109,15 +110,16 @@ class RagPipelineService: @classmethod def get_pipeline_templates( cls, - session: Session, type: str = "built-in", language: str = "en-US", current_tenant_id: str | None = None, + *, + session: Session, ) -> dict[str, Any]: if type == "built-in": mode = dify_config.HOSTED_FETCH_PIPELINE_TEMPLATES_MODE retrieval_instance = PipelineTemplateRetrievalFactory.get_pipeline_template_factory(mode)() - result = retrieval_instance.get_pipeline_templates(session, language, current_tenant_id) + result = retrieval_instance.get_pipeline_templates(language, current_tenant_id, session=session) if not result.get("pipeline_templates") and language != "en-US": template_retrieval = PipelineTemplateRetrievalFactory.get_built_in_pipeline_template_retrieval() result = template_retrieval.fetch_pipeline_templates_from_builtin("en-US") @@ -125,12 +127,12 @@ class RagPipelineService: else: mode = "customized" retrieval_instance = PipelineTemplateRetrievalFactory.get_pipeline_template_factory(mode)() - result = retrieval_instance.get_pipeline_templates(session, language, current_tenant_id) + result = retrieval_instance.get_pipeline_templates(language, current_tenant_id, session=session) return result @classmethod def get_pipeline_template_detail( - cls, session: Session, template_id: str, type: str = "built-in" + cls, template_id: str, type: str = "built-in", *, session: Session ) -> dict[str, Any] | None: """ Get pipeline template detail. @@ -143,7 +145,7 @@ class RagPipelineService: mode = dify_config.HOSTED_FETCH_PIPELINE_TEMPLATES_MODE retrieval_instance = PipelineTemplateRetrievalFactory.get_pipeline_template_factory(mode)() built_in_result: dict[str, Any] | None = retrieval_instance.get_pipeline_template_detail( - session, template_id + template_id, session=session ) if built_in_result is None: logger.warning( @@ -156,7 +158,7 @@ class RagPipelineService: mode = "customized" retrieval_instance = PipelineTemplateRetrievalFactory.get_pipeline_template_factory(mode)() customized_result: dict[str, Any] | None = retrieval_instance.get_pipeline_template_detail( - session, template_id + template_id, session=session ) return customized_result @@ -167,7 +169,8 @@ class RagPipelineService: template_info: PipelineTemplateInfoEntity, current_user: Account | None = None, current_tenant_id: str | None = None, - session: Session | None = None, + *, + session: Session, ): """ Update pipeline template. @@ -175,16 +178,6 @@ class RagPipelineService: :param template_info: template info """ current_user, current_tenant_id = resolve_account_fallback(current_user, current_tenant_id) - if session is None: - with session_factory.get_session_maker().begin() as new_session: - return cls.update_customized_pipeline_template( - template_id, - template_info, - current_user, - current_tenant_id, - session=new_session, - ) - customized_template: PipelineCustomizedTemplate | None = session.scalar( select(PipelineCustomizedTemplate) .where( @@ -213,21 +206,17 @@ class RagPipelineService: customized_template.description = template_info.description customized_template.icon = template_info.icon_info.model_dump() customized_template.updated_by = current_user.id + session.commit() return customized_template @classmethod def delete_customized_pipeline_template( - cls, template_id: str, current_tenant_id: str | None = None, session: Session | None = None + cls, template_id: str, current_tenant_id: str | None = None, *, session: Session ): """ Delete customized pipeline template. """ current_tenant_id = resolve_tenant_id_fallback(current_tenant_id) - if session is None: - with session_factory.get_session_maker().begin() as new_session: - cls.delete_customized_pipeline_template(template_id, current_tenant_id, session=new_session) - return - customized_template: PipelineCustomizedTemplate | None = session.scalar( select(PipelineCustomizedTemplate) .where( @@ -239,22 +228,22 @@ class RagPipelineService: if not customized_template: raise ValueError("Customized pipeline template not found.") session.delete(customized_template) + session.commit() def get_draft_workflow(self, pipeline: Pipeline) -> Workflow | None: """ Get draft workflow """ # fetch draft workflow by rag pipeline - with self._session_maker() as session: - workflow = session.scalar( - select(Workflow) - .where( - Workflow.tenant_id == pipeline.tenant_id, - Workflow.app_id == pipeline.id, - Workflow.version == "draft", - ) - .limit(1) + workflow = self._session.scalar( + select(Workflow) + .where( + Workflow.tenant_id == pipeline.tenant_id, + Workflow.app_id == pipeline.id, + Workflow.version == "draft", ) + .limit(1) + ) # return draft workflow return workflow @@ -268,31 +257,29 @@ class RagPipelineService: return None # fetch published workflow by workflow_id - with self._session_maker() as session: - workflow = session.scalar( - select(Workflow) - .where( - Workflow.tenant_id == pipeline.tenant_id, - Workflow.app_id == pipeline.id, - Workflow.id == pipeline.workflow_id, - ) - .limit(1) + workflow = self._session.scalar( + select(Workflow) + .where( + Workflow.tenant_id == pipeline.tenant_id, + Workflow.app_id == pipeline.id, + Workflow.id == pipeline.workflow_id, ) + .limit(1) + ) return workflow def get_published_workflow_by_id(self, pipeline: Pipeline, workflow_id: str) -> Workflow | None: """Fetch a published workflow snapshot by ID for restore operations.""" - with self._session_maker() as session: - workflow = session.scalar( - select(Workflow) - .where( - Workflow.tenant_id == pipeline.tenant_id, - Workflow.app_id == pipeline.id, - Workflow.id == workflow_id, - ) - .limit(1) + workflow = self._session.scalar( + select(Workflow) + .where( + Workflow.tenant_id == pipeline.tenant_id, + Workflow.app_id == pipeline.id, + Workflow.id == workflow_id, ) + .limit(1) + ) if workflow and workflow.version == Workflow.VERSION_DRAFT: raise IsDraftWorkflowError("source workflow must be published") return workflow @@ -350,51 +337,39 @@ class RagPipelineService: Sync draft workflow :raises WorkflowHashNotEqualError """ - with self._session_maker.begin() as session: - managed_pipeline = session.get(Pipeline, pipeline.id) - if not managed_pipeline: - raise ValueError("Pipeline not found") + # fetch draft workflow by app_model + workflow = self.get_draft_workflow(pipeline=pipeline) - # fetch draft workflow by app_model - workflow = session.scalar( - select(Workflow) - .where( - Workflow.tenant_id == managed_pipeline.tenant_id, - Workflow.app_id == managed_pipeline.id, - Workflow.version == "draft", - ) - .limit(1) + if workflow and workflow.unique_hash != unique_hash: + raise WorkflowHashNotEqualError() + + # create draft workflow if not found + if not workflow: + workflow = Workflow( + tenant_id=pipeline.tenant_id, + app_id=pipeline.id, + features="{}", + type=WorkflowType.RAG_PIPELINE.value, + version="draft", + graph=json.dumps(graph), + created_by=account.id, + environment_variables=environment_variables, + conversation_variables=conversation_variables, + rag_pipeline_variables=rag_pipeline_variables, ) - - if workflow and workflow.unique_hash != unique_hash: - raise WorkflowHashNotEqualError() - - # create draft workflow if not found - if not workflow: - workflow = Workflow( - tenant_id=managed_pipeline.tenant_id, - app_id=managed_pipeline.id, - features="{}", - type=WorkflowType.RAG_PIPELINE.value, - version="draft", - graph=json.dumps(graph), - created_by=account.id, - environment_variables=environment_variables, - conversation_variables=conversation_variables, - rag_pipeline_variables=rag_pipeline_variables, - ) - session.add(workflow) - session.flush() - managed_pipeline.workflow_id = workflow.id - pipeline.workflow_id = workflow.id - # update draft workflow if found - else: - workflow.graph = json.dumps(graph) - workflow.updated_by = account.id - workflow.updated_at = datetime.now(UTC).replace(tzinfo=None) - workflow.environment_variables = environment_variables - workflow.conversation_variables = conversation_variables - workflow.rag_pipeline_variables = rag_pipeline_variables + self._session.add(workflow) + self._session.flush() + pipeline.workflow_id = workflow.id + # update draft workflow if found + else: + workflow.graph = json.dumps(graph) + workflow.updated_by = account.id + workflow.updated_at = datetime.now(UTC).replace(tzinfo=None) + workflow.environment_variables = environment_variables + workflow.conversation_variables = conversation_variables + workflow.rag_pipeline_variables = rag_pipeline_variables + # commit db session changes + self._session.commit() # trigger workflow events TODO # app_draft_workflow_was_synced.send(pipeline, synced_draft_workflow=workflow) @@ -415,48 +390,26 @@ class RagPipelineService: the pipeline-specific flush/link step that wires a newly created draft back onto ``pipeline.workflow_id``. """ - with self._session_maker.begin() as session: - managed_pipeline = session.get(Pipeline, pipeline.id) - if not managed_pipeline: - raise ValueError("Pipeline not found") + source_workflow = self.get_published_workflow_by_id(pipeline=pipeline, workflow_id=workflow_id) + if not source_workflow: + raise WorkflowNotFoundError("Workflow not found.") - source_workflow = session.scalar( - select(Workflow) - .where( - Workflow.tenant_id == managed_pipeline.tenant_id, - Workflow.app_id == managed_pipeline.id, - Workflow.id == workflow_id, - ) - .limit(1) - ) - if source_workflow and source_workflow.version == Workflow.VERSION_DRAFT: - raise IsDraftWorkflowError("source workflow must be published") - if not source_workflow: - raise WorkflowNotFoundError("Workflow not found.") + draft_workflow = self.get_draft_workflow(pipeline=pipeline) + draft_workflow, is_new_draft = apply_published_workflow_snapshot_to_draft( + tenant_id=pipeline.tenant_id, + app_id=pipeline.id, + source_workflow=source_workflow, + draft_workflow=draft_workflow, + account=account, + updated_at_factory=lambda: datetime.now(UTC).replace(tzinfo=None), + ) - draft_workflow = session.scalar( - select(Workflow) - .where( - Workflow.tenant_id == managed_pipeline.tenant_id, - Workflow.app_id == managed_pipeline.id, - Workflow.version == Workflow.VERSION_DRAFT, - ) - .limit(1) - ) - draft_workflow, is_new_draft = apply_published_workflow_snapshot_to_draft( - tenant_id=managed_pipeline.tenant_id, - app_id=managed_pipeline.id, - source_workflow=source_workflow, - draft_workflow=draft_workflow, - account=account, - updated_at_factory=lambda: datetime.now(UTC).replace(tzinfo=None), - ) + if is_new_draft: + self._session.add(draft_workflow) + self._session.flush() + pipeline.workflow_id = draft_workflow.id - if is_new_draft: - session.add(draft_workflow) - session.flush() - managed_pipeline.workflow_id = draft_workflow.id - pipeline.workflow_id = draft_workflow.id + self._session.commit() return draft_workflow @@ -633,7 +586,7 @@ class RagPipelineService: workflow_node_execution.id ) - with self._session_maker.begin() as session: + with sessionmaker(bind=db.engine).begin() as session: draft_var_saver = DraftVariableSaver( session=session, app_id=pipeline.id, @@ -1050,22 +1003,23 @@ class RagPipelineService: dataset_id = get_system_segment(variable_pool, SystemVariableKey.DATASET_ID) pipeline_id = get_system_segment(variable_pool, SystemVariableKey.APP_ID) if document_id and dataset_id and pipeline_id: - with self._session_maker.begin() as session: - document = session.scalar( - select(Document) - .join(Dataset, Dataset.id == Document.dataset_id) - .where( - Document.id == document_id.value, - Document.tenant_id == tenant_id, - Document.dataset_id == dataset_id.value, - Dataset.tenant_id == tenant_id, - Dataset.pipeline_id == pipeline_id.value, - ) - .limit(1) + document = self._session.scalar( + select(Document) + .join(Dataset, Dataset.id == Document.dataset_id) + .where( + Document.id == document_id.value, + Document.tenant_id == tenant_id, + Document.dataset_id == dataset_id.value, + Dataset.tenant_id == tenant_id, + Dataset.pipeline_id == pipeline_id.value, ) - if document: - document.indexing_status = IndexingStatus.ERROR - document.error = error + .limit(1) + ) + if document: + document.indexing_status = IndexingStatus.ERROR + document.error = error + self._session.add(document) + self._session.commit() return workflow_node_execution @@ -1276,89 +1230,86 @@ class RagPipelineService: args: dict[str, Any], current_user: Account | None = None, current_tenant_id: str | None = None, + *, + session: Session, ): """ Publish customized pipeline template """ current_user, _ = resolve_account_fallback(current_user, current_tenant_id) - with session_factory.get_session_maker().begin() as session: - pipeline = session.get(Pipeline, pipeline_id) - if not pipeline: - raise ValueError("Pipeline not found") - if not pipeline.workflow_id: - raise ValueError("Pipeline workflow not found") - workflow = session.get(Workflow, pipeline.workflow_id) - if not workflow: - raise ValueError("Workflow not found") - dataset = pipeline.retrieve_dataset(session=session) - if not dataset: - raise ValueError("Dataset not found") + pipeline = session.get(Pipeline, pipeline_id) + if not pipeline: + raise ValueError("Pipeline not found") + if not pipeline.workflow_id: + raise ValueError("Pipeline workflow not found") + workflow = session.get(Workflow, pipeline.workflow_id) + if not workflow: + raise ValueError("Workflow not found") + dataset = pipeline.retrieve_dataset(session=session) + if not dataset: + raise ValueError("Dataset not found") - # check template name is exist - template_name = args.get("name") - if template_name: - template = session.scalar( - select(PipelineCustomizedTemplate) - .where( - PipelineCustomizedTemplate.name == template_name, - PipelineCustomizedTemplate.tenant_id == pipeline.tenant_id, - ) - .limit(1) - ) - if template: - raise ValueError("Template name is already exists") - - max_position = session.scalar( - select(func.max(PipelineCustomizedTemplate.position)).where( - PipelineCustomizedTemplate.tenant_id == pipeline.tenant_id + # check template name is exist + template_name = args.get("name") + if template_name: + template = session.scalar( + select(PipelineCustomizedTemplate) + .where( + PipelineCustomizedTemplate.name == template_name, + PipelineCustomizedTemplate.tenant_id == pipeline.tenant_id, ) + .limit(1) ) + if template: + raise ValueError("Template name is already exists") - from services.rag_pipeline.rag_pipeline_dsl_service import RagPipelineDslService - - rag_pipeline_dsl_service = RagPipelineDslService(session) - dsl = rag_pipeline_dsl_service.export_rag_pipeline_dsl(pipeline=pipeline, include_secret=True) - if args.get("icon_info") is None: - args["icon_info"] = {} - if args.get("description") is None: - raise ValueError("Description is required") - if args.get("name") is None: - raise ValueError("Name is required") - pipeline_customized_template = PipelineCustomizedTemplate( - name=args.get("name") or "", - description=args.get("description") or "", - icon=args.get("icon_info") or {}, - tenant_id=pipeline.tenant_id, - yaml_content=dsl, - install_count=0, - position=max_position + 1 if max_position else 1, - chunk_structure=dataset.chunk_structure, - language="en-US", - created_by=current_user.id, + max_position = session.scalar( + select(func.max(PipelineCustomizedTemplate.position)).where( + PipelineCustomizedTemplate.tenant_id == pipeline.tenant_id ) - session.add(pipeline_customized_template) + ) + + from services.rag_pipeline.rag_pipeline_dsl_service import RagPipelineDslService + + rag_pipeline_dsl_service = RagPipelineDslService(session) + dsl = rag_pipeline_dsl_service.export_rag_pipeline_dsl(pipeline=pipeline, include_secret=True) + if args.get("icon_info") is None: + args["icon_info"] = {} + if args.get("description") is None: + raise ValueError("Description is required") + if args.get("name") is None: + raise ValueError("Name is required") + pipeline_customized_template = PipelineCustomizedTemplate( + name=args.get("name") or "", + description=args.get("description") or "", + icon=args.get("icon_info") or {}, + tenant_id=pipeline.tenant_id, + yaml_content=dsl, + install_count=0, + position=max_position + 1 if max_position else 1, + chunk_structure=dataset.chunk_structure, + language="en-US", + created_by=current_user.id, + ) + session.add(pipeline_customized_template) + session.commit() def is_workflow_exist(self, pipeline: Pipeline) -> bool: - with self._session_maker() as session: - return ( - session.scalar( - select(func.count(Workflow.id)).where( - Workflow.tenant_id == pipeline.tenant_id, - Workflow.app_id == pipeline.id, - Workflow.version == Workflow.VERSION_DRAFT, - ) + return ( + self._session.scalar( + select(func.count(Workflow.id)).where( + Workflow.tenant_id == pipeline.tenant_id, + Workflow.app_id == pipeline.id, + Workflow.version == Workflow.VERSION_DRAFT, ) - or 0 - ) > 0 + ) + or 0 + ) > 0 def get_node_last_run( self, pipeline: Pipeline, workflow: Workflow, node_id: str ) -> WorkflowNodeExecutionModel | None: - node_execution_service_repo = DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository( - self._session_maker - ) - - node_exec = node_execution_service_repo.get_node_last_execution( + node_exec = self._node_execution_service_repo.get_node_last_execution( tenant_id=pipeline.tenant_id, app_id=pipeline.id, workflow_id=workflow.id, @@ -1431,7 +1382,7 @@ class RagPipelineService: # Convert node_execution to WorkflowNodeExecution after save workflow_node_execution_db_model = repository._to_db_model(workflow_node_execution) # type: ignore - with self._session_maker.begin() as session: + with sessionmaker(bind=db.engine).begin() as session: draft_var_saver = DraftVariableSaver( session=session, app_id=pipeline.id, @@ -1465,10 +1416,9 @@ class RagPipelineService: if type and type != "all": stmt = stmt.where(PipelineRecommendedPlugin.type == type) - with self._session_maker() as session: - pipeline_recommended_plugins = session.scalars( - stmt.order_by(PipelineRecommendedPlugin.position.asc()) - ).all() + pipeline_recommended_plugins = self._session.scalars( + stmt.order_by(PipelineRecommendedPlugin.position.asc()) + ).all() if not pipeline_recommended_plugins: return { @@ -1507,173 +1457,41 @@ class RagPipelineService: """ Retry error document """ - with self._session_maker() as session: - document_pipeline_execution_log = session.scalar( - select(DocumentPipelineExecutionLog) - .where(DocumentPipelineExecutionLog.document_id == document.id) - .limit(1) - ) - if not document_pipeline_execution_log: - raise ValueError("Document pipeline execution log not found") - pipeline = session.get(Pipeline, document_pipeline_execution_log.pipeline_id) - if not pipeline: - raise ValueError("Pipeline not found") - # convert to app config - workflow = session.scalar( - select(Workflow) - .where( - Workflow.tenant_id == pipeline.tenant_id, - Workflow.app_id == pipeline.id, - Workflow.id == pipeline.workflow_id, - ) - .limit(1) - ) - if not workflow: - raise ValueError("Workflow not found") - PipelineGenerator().generate( - pipeline=pipeline, - workflow=workflow, - user=user, - args={ - "inputs": document_pipeline_execution_log.input_data, - "start_node_id": document_pipeline_execution_log.datasource_node_id, - "datasource_type": document_pipeline_execution_log.datasource_type, - "datasource_info_list": [json.loads(document_pipeline_execution_log.datasource_info)], - "original_document_id": document.id, - }, - invoke_from=InvokeFrom.PUBLISHED_PIPELINE, - streaming=False, - call_depth=0, - workflow_thread_pool_id=None, - is_retry=True, - ) + document_pipeline_execution_log = self._session.scalar( + select(DocumentPipelineExecutionLog).where(DocumentPipelineExecutionLog.document_id == document.id).limit(1) + ) + if not document_pipeline_execution_log: + raise ValueError("Document pipeline execution log not found") + pipeline = self._session.get(Pipeline, document_pipeline_execution_log.pipeline_id) + if not pipeline: + raise ValueError("Pipeline not found") + # convert to app config + workflow = self.get_published_workflow(pipeline) + if not workflow: + raise ValueError("Workflow not found") + PipelineGenerator().generate( + pipeline=pipeline, + workflow=workflow, + user=user, + args={ + "inputs": document_pipeline_execution_log.input_data, + "start_node_id": document_pipeline_execution_log.datasource_node_id, + "datasource_type": document_pipeline_execution_log.datasource_type, + "datasource_info_list": [json.loads(document_pipeline_execution_log.datasource_info)], + "original_document_id": document.id, + }, + invoke_from=InvokeFrom.PUBLISHED_PIPELINE, + streaming=False, + call_depth=0, + workflow_thread_pool_id=None, + is_retry=True, + ) def get_datasource_plugins(self, tenant_id: str, dataset_id: str, is_published: bool) -> list[dict]: """ Get datasource plugins """ - with self._session_maker() as session: - dataset: Dataset | None = session.scalar( - select(Dataset) - .where( - Dataset.id == dataset_id, - Dataset.tenant_id == tenant_id, - ) - .limit(1) - ) - if not dataset: - raise ValueError("Dataset not found") - pipeline: Pipeline | None = session.scalar( - select(Pipeline) - .where( - Pipeline.id == dataset.pipeline_id, - Pipeline.tenant_id == tenant_id, - ) - .limit(1) - ) - if not pipeline: - raise ValueError("Pipeline not found") - - if is_published: - workflow = session.scalar( - select(Workflow) - .where( - Workflow.tenant_id == pipeline.tenant_id, - Workflow.app_id == pipeline.id, - Workflow.id == pipeline.workflow_id, - ) - .limit(1) - ) - else: - workflow = session.scalar( - select(Workflow) - .where( - Workflow.tenant_id == pipeline.tenant_id, - Workflow.app_id == pipeline.id, - Workflow.version == Workflow.VERSION_DRAFT, - ) - .limit(1) - ) - if not pipeline or not workflow: - raise ValueError("Pipeline or workflow not found") - - datasource_nodes = workflow.graph_dict.get("nodes", []) - datasource_plugins = [] - for datasource_node in datasource_nodes: - if datasource_node.get("data", {}).get("type") == "datasource": - datasource_node_data = datasource_node["data"] - if not datasource_node_data: - continue - - variables = workflow.rag_pipeline_variables - if variables: - variables_map = {item["variable"]: item for item in variables} - else: - variables_map = {} - - datasource_parameters = datasource_node_data.get("datasource_parameters", {}) - user_input_variables_keys = [] - user_input_variables = [] - - for _, value in datasource_parameters.items(): - if value.get("value") and isinstance(value.get("value"), str): - pattern = ( - r"\{\{#([a-zA-Z0-9_]{1,50}" - r"(?:\.[a-zA-Z0-9_][a-zA-Z0-9_]{0,29}){1,10})#\}\}" - ) - match = re.match(pattern, value["value"]) - if match: - full_path = match.group(1) - last_part = full_path.split(".")[-1] - user_input_variables_keys.append(last_part) - elif value.get("value") and isinstance(value.get("value"), list): - last_part = value.get("value")[-1] - user_input_variables_keys.append(last_part) - for key, value in variables_map.items(): - if key in user_input_variables_keys: - user_input_variables.append(value) - - # get credentials - datasource_provider_service: DatasourceProviderService = DatasourceProviderService() - credentials: list[dict[Any, Any]] = datasource_provider_service.list_datasource_credentials( - tenant_id=tenant_id, - provider=datasource_node_data.get("provider_name"), - plugin_id=datasource_node_data.get("plugin_id"), - ) - credential_info_list: list[Any] = [] - for credential in credentials: - credential_info_list.append( - { - "id": credential.get("id"), - "name": credential.get("name"), - "type": credential.get("type"), - "is_default": credential.get("is_default"), - } - ) - - datasource_plugins.append( - { - "node_id": datasource_node.get("id"), - "plugin_id": datasource_node_data.get("plugin_id"), - "provider_name": datasource_node_data.get("provider_name"), - "datasource_type": datasource_node_data.get("provider_type"), - "title": datasource_node_data.get("title"), - "user_input_variables": user_input_variables, - "credentials": credential_info_list, - } - ) - - return datasource_plugins - - def get_pipeline(self, tenant_id: str, dataset_id: str, session: Session | None = None) -> Pipeline: - """ - Get pipeline - """ - if session is None: - with self._session_maker() as new_session: - return self.get_pipeline(tenant_id, dataset_id, session=new_session) - - dataset: Dataset | None = session.scalar( + dataset: Dataset | None = self._session.scalar( select(Dataset) .where( Dataset.id == dataset_id, @@ -1683,7 +1501,106 @@ class RagPipelineService: ) if not dataset: raise ValueError("Dataset not found") - pipeline: Pipeline | None = session.scalar( + pipeline: Pipeline | None = self._session.scalar( + select(Pipeline) + .where( + Pipeline.id == dataset.pipeline_id, + Pipeline.tenant_id == tenant_id, + ) + .limit(1) + ) + if not pipeline: + raise ValueError("Pipeline not found") + + workflow: Workflow | None = None + if is_published: + workflow = self.get_published_workflow(pipeline=pipeline) + else: + workflow = self.get_draft_workflow(pipeline=pipeline) + if not pipeline or not workflow: + raise ValueError("Pipeline or workflow not found") + + datasource_nodes = workflow.graph_dict.get("nodes", []) + datasource_plugins = [] + for datasource_node in datasource_nodes: + if datasource_node.get("data", {}).get("type") == "datasource": + datasource_node_data = datasource_node["data"] + if not datasource_node_data: + continue + + variables = workflow.rag_pipeline_variables + if variables: + variables_map = {item["variable"]: item for item in variables} + else: + variables_map = {} + + datasource_parameters = datasource_node_data.get("datasource_parameters", {}) + user_input_variables_keys = [] + user_input_variables = [] + + for _, value in datasource_parameters.items(): + if value.get("value") and isinstance(value.get("value"), str): + pattern = r"\{\{#([a-zA-Z0-9_]{1,50}(?:\.[a-zA-Z0-9_][a-zA-Z0-9_]{0,29}){1,10})#\}\}" + match = re.match(pattern, value["value"]) + if match: + full_path = match.group(1) + last_part = full_path.split(".")[-1] + user_input_variables_keys.append(last_part) + elif value.get("value") and isinstance(value.get("value"), list): + last_part = value.get("value")[-1] + user_input_variables_keys.append(last_part) + for key, value in variables_map.items(): + if key in user_input_variables_keys: + user_input_variables.append(value) + + # get credentials + datasource_provider_service: DatasourceProviderService = DatasourceProviderService() + credentials: list[dict[Any, Any]] = datasource_provider_service.list_datasource_credentials( + tenant_id=tenant_id, + provider=datasource_node_data.get("provider_name"), + plugin_id=datasource_node_data.get("plugin_id"), + session=self._session, + ) + credential_info_list: list[Any] = [] + for credential in credentials: + credential_info_list.append( + { + "id": credential.get("id"), + "name": credential.get("name"), + "type": credential.get("type"), + "is_default": credential.get("is_default"), + } + ) + + datasource_plugins.append( + { + "node_id": datasource_node.get("id"), + "plugin_id": datasource_node_data.get("plugin_id"), + "provider_name": datasource_node_data.get("provider_name"), + "datasource_type": datasource_node_data.get("provider_type"), + "title": datasource_node_data.get("title"), + "user_input_variables": user_input_variables, + "credentials": credential_info_list, + } + ) + + return datasource_plugins + + def get_pipeline(self, tenant_id: str, dataset_id: str) -> Pipeline: + """ + Get pipeline + """ + dataset: Dataset | None = self._session.scalar( + select(Dataset) + .where( + Dataset.id == dataset_id, + Dataset.tenant_id == tenant_id, + ) + .limit(1) + ) + if not dataset: + raise ValueError("Dataset not found") + pipeline: Pipeline | None = self._session.scalar( select(Pipeline) .where( Pipeline.id == dataset.pipeline_id, diff --git a/api/services/rag_pipeline/rag_pipeline_dsl_service.py b/api/services/rag_pipeline/rag_pipeline_dsl_service.py index 5459c3e5f1f..d562a4b9adf 100644 --- a/api/services/rag_pipeline/rag_pipeline_dsl_service.py +++ b/api/services/rag_pipeline/rag_pipeline_dsl_service.py @@ -15,7 +15,7 @@ from Crypto.Util.Padding import pad, unpad from flask_login import current_user from pydantic import BaseModel from sqlalchemy import select -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from core.file import remote_fetcher from core.helper.name_generator import generate_incremental_name @@ -83,7 +83,7 @@ class RagPipelineDslService: when generated IDs are needed mid-operation; they never commit or rollback. """ - def __init__(self, session: Session | scoped_session): + def __init__(self, session: Session): self._session = session def import_rag_pipeline( diff --git a/api/services/rag_pipeline/rag_pipeline_transform_service.py b/api/services/rag_pipeline/rag_pipeline_transform_service.py index 1b922b3f7b9..6a7902c1908 100644 --- a/api/services/rag_pipeline/rag_pipeline_transform_service.py +++ b/api/services/rag_pipeline/rag_pipeline_transform_service.py @@ -96,7 +96,7 @@ class RagPipelineTransformService: # deal document data self._deal_document_data(dataset, session) - session.flush() + session.commit() return { "pipeline_id": pipeline.id, "dataset_id": dataset_id, @@ -194,6 +194,7 @@ class RagPipelineTransformService: def _create_pipeline( self, data: dict[str, Any], + *, session: Session, ) -> Pipeline: """Create a new app or update an existing one.""" @@ -291,7 +292,7 @@ class RagPipelineTransformService: logger.debug("Installing missing pipeline plugins %s", package_identifiers_to_install) PluginService.install_from_marketplace_pkg(tenant_id, package_identifiers_to_install) - def _transform_to_empty_pipeline(self, dataset: Dataset, session: Session): + def _transform_to_empty_pipeline(self, dataset: Dataset, *, session: Session): pipeline = Pipeline( tenant_id=dataset.tenant_id, name=dataset.name, @@ -306,7 +307,7 @@ class RagPipelineTransformService: dataset.updated_by = current_user.id dataset.updated_at = datetime.now(UTC).replace(tzinfo=None) session.add(dataset) - session.flush() + session.commit() return { "pipeline_id": pipeline.id, "dataset_id": dataset.id, diff --git a/api/services/recommend_app/buildin/buildin_retrieval.py b/api/services/recommend_app/buildin/buildin_retrieval.py index 03b72a4f57c..d29d754b67e 100644 --- a/api/services/recommend_app/buildin/buildin_retrieval.py +++ b/api/services/recommend_app/buildin/buildin_retrieval.py @@ -4,6 +4,7 @@ from pathlib import Path from typing import Any, override from flask import current_app +from sqlalchemy.orm import Session from services.recommend_app.database.database_retrieval import DatabaseRecommendAppRetrieval from services.recommend_app.recommend_app_base import RecommendAppRetrievalBase @@ -22,17 +23,19 @@ class BuildInRecommendAppRetrieval(RecommendAppRetrievalBase): return RecommendAppType.BUILDIN @override - def get_recommended_apps_and_categories(self, language: str): + def get_recommended_apps_and_categories(self, language: str, *, session: Session): + del session result = self.fetch_recommended_apps_from_builtin(language) return result @override - def get_learn_dify_apps(self, language: str): - result = DatabaseRecommendAppRetrieval.fetch_learn_dify_apps_from_db(language) + def get_learn_dify_apps(self, language: str, *, session: Session): + result = DatabaseRecommendAppRetrieval.fetch_learn_dify_apps_from_db(language, session=session) return result @override - def get_recommend_app_detail(self, app_id: str): + def get_recommend_app_detail(self, app_id: str, *, session: Session): + del session result = self.fetch_recommended_app_detail_from_builtin(app_id) return result diff --git a/api/services/recommend_app/database/database_retrieval.py b/api/services/recommend_app/database/database_retrieval.py index f6786175896..08d902fdeb5 100644 --- a/api/services/recommend_app/database/database_retrieval.py +++ b/api/services/recommend_app/database/database_retrieval.py @@ -1,9 +1,9 @@ from typing import Any, NotRequired, TypedDict, override from sqlalchemy import select +from sqlalchemy.orm import Session from constants.languages import languages -from extensions.ext_database import db from models.model import App, RecommendedApp from services.app_dsl_service import AppDslService from services.recommend_app.category_order import order_categories @@ -45,18 +45,18 @@ class DatabaseRecommendAppRetrieval(RecommendAppRetrievalBase): """ @override - def get_recommended_apps_and_categories(self, language: str) -> RecommendedAppsResultDict: - result = self.fetch_recommended_apps_from_db(language) + def get_recommended_apps_and_categories(self, language: str, *, session: Session) -> RecommendedAppsResultDict: + result = self.fetch_recommended_apps_from_db(language, session=session) return result @override - def get_learn_dify_apps(self, language: str) -> RecommendedAppsResultDict: - result = self.fetch_learn_dify_apps_from_db(language) + def get_learn_dify_apps(self, language: str, *, session: Session) -> RecommendedAppsResultDict: + result = self.fetch_learn_dify_apps_from_db(language, session=session) return result @override - def get_recommend_app_detail(self, app_id: str) -> RecommendedAppDetailDict | None: - result = self.fetch_recommended_app_detail_from_db(app_id) + def get_recommend_app_detail(self, app_id: str, *, session: Session) -> RecommendedAppDetailDict | None: + result = self.fetch_recommended_app_detail_from_db(app_id, session=session) return result @override @@ -64,42 +64,42 @@ class DatabaseRecommendAppRetrieval(RecommendAppRetrievalBase): return RecommendAppType.DATABASE @classmethod - def fetch_recommended_apps_from_db(cls, language: str) -> RecommendedAppsResultDict: + def fetch_recommended_apps_from_db(cls, language: str, *, session: Session) -> RecommendedAppsResultDict: """ Fetch recommended apps from db. :param language: language :return: """ - recommended_apps = cls._fetch_listed_recommended_apps(language) + recommended_apps = cls._fetch_listed_recommended_apps(language, session=session) if len(recommended_apps) == 0: - recommended_apps = cls._fetch_listed_recommended_apps(languages[0]) + recommended_apps = cls._fetch_listed_recommended_apps(languages[0], session=session) return cls._format_recommended_apps(recommended_apps, language) @classmethod - def fetch_learn_dify_apps_from_db(cls, language: str) -> RecommendedAppsResultDict: + def fetch_learn_dify_apps_from_db(cls, language: str, *, session: Session) -> RecommendedAppsResultDict: """ Fetch listed recommended apps explicitly marked for the Learn Dify section. :param language: language :return: """ - recommended_apps = cls._fetch_listed_recommended_apps(language, is_learn_dify=True) + recommended_apps = cls._fetch_listed_recommended_apps(language, session=session, is_learn_dify=True) if len(recommended_apps) == 0 and language != languages[0]: - recommended_apps = cls._fetch_listed_recommended_apps(languages[0], is_learn_dify=True) + recommended_apps = cls._fetch_listed_recommended_apps(languages[0], session=session, is_learn_dify=True) return cls._format_recommended_apps(recommended_apps, language) @classmethod def _fetch_listed_recommended_apps( - cls, language: str, *, is_learn_dify: bool | None = None + cls, language: str, *, session: Session, is_learn_dify: bool | None = None ) -> list[RecommendedApp]: filters = [RecommendedApp.is_listed.is_(True), RecommendedApp.language == language] if is_learn_dify is not None: filters.append(RecommendedApp.is_learn_dify.is_(is_learn_dify)) - return list(db.session.scalars(select(RecommendedApp).where(*filters)).all()) + return list(session.scalars(select(RecommendedApp).where(*filters)).all()) @classmethod def _format_recommended_apps( @@ -146,14 +146,14 @@ class DatabaseRecommendAppRetrieval(RecommendAppRetrievalBase): ) @classmethod - def fetch_recommended_app_detail_from_db(cls, app_id: str) -> RecommendedAppDetailDict | None: + def fetch_recommended_app_detail_from_db(cls, app_id: str, *, session: Session) -> RecommendedAppDetailDict | None: """ Fetch recommended app detail from db. :param app_id: App ID :return: """ # is in public recommended list - recommended_app = db.session.scalar( + recommended_app = session.scalar( select(RecommendedApp).where(RecommendedApp.is_listed == True, RecommendedApp.app_id == app_id).limit(1) ) @@ -161,7 +161,7 @@ class DatabaseRecommendAppRetrieval(RecommendAppRetrievalBase): return None # get app detail - app_model = db.session.get(App, app_id) + app_model = session.get(App, app_id) if not app_model or not app_model.is_public: return None @@ -171,5 +171,5 @@ class DatabaseRecommendAppRetrieval(RecommendAppRetrievalBase): icon=app_model.icon, icon_background=app_model.icon_background, mode=app_model.mode, - export_data=AppDslService.export_dsl(app_model=app_model), + export_data=AppDslService.export_dsl(app_model=app_model, session=session), ) diff --git a/api/services/recommend_app/recommend_app_base.py b/api/services/recommend_app/recommend_app_base.py index f819cc3a937..821ad476c42 100644 --- a/api/services/recommend_app/recommend_app_base.py +++ b/api/services/recommend_app/recommend_app_base.py @@ -1,13 +1,15 @@ from typing import Any, Protocol +from sqlalchemy.orm import Session + class RecommendAppRetrievalBase(Protocol): """Interface for recommend app retrieval.""" - def get_recommended_apps_and_categories(self, language: str) -> Any: ... + def get_recommended_apps_and_categories(self, language: str, *, session: Session) -> Any: ... - def get_learn_dify_apps(self, language: str) -> Any: ... + def get_learn_dify_apps(self, language: str, *, session: Session) -> Any: ... - def get_recommend_app_detail(self, app_id: str) -> Any: ... + def get_recommend_app_detail(self, app_id: str, *, session: Session) -> Any: ... def get_type(self) -> str: ... diff --git a/api/services/recommend_app/remote/remote_retrieval.py b/api/services/recommend_app/remote/remote_retrieval.py index 2e3222bb978..c676ec907e0 100644 --- a/api/services/recommend_app/remote/remote_retrieval.py +++ b/api/services/recommend_app/remote/remote_retrieval.py @@ -3,6 +3,7 @@ from typing import Any, override import httpx from flask import has_request_context, request +from sqlalchemy.orm import Session from configs import dify_config from services.recommend_app.buildin.buildin_retrieval import BuildInRecommendAppRetrieval @@ -33,7 +34,8 @@ class RemoteRecommendAppRetrieval(RecommendAppRetrievalBase): """ @override - def get_recommend_app_detail(self, app_id: str): + def get_recommend_app_detail(self, app_id: str, *, session: Session): + del session try: result = self.fetch_recommended_app_detail_from_dify_official(app_id) except Exception as e: @@ -42,7 +44,8 @@ class RemoteRecommendAppRetrieval(RecommendAppRetrievalBase): return result @override - def get_recommended_apps_and_categories(self, language: str): + def get_recommended_apps_and_categories(self, language: str, *, session: Session): + del session try: result = self.fetch_recommended_apps_from_dify_official(language) except Exception as e: @@ -51,12 +54,12 @@ class RemoteRecommendAppRetrieval(RecommendAppRetrievalBase): return result @override - def get_learn_dify_apps(self, language: str): + def get_learn_dify_apps(self, language: str, *, session: Session): try: result = self.fetch_learn_dify_apps_from_dify_official(language) except Exception as e: logger.warning("fetch learn dify apps from dify official failed: %s, switch to database.", e) - result = DatabaseRecommendAppRetrieval.fetch_learn_dify_apps_from_db(language) + result = DatabaseRecommendAppRetrieval.fetch_learn_dify_apps_from_db(language, session=session) return result @override diff --git a/api/services/recommended_app_service.py b/api/services/recommended_app_service.py index 2d247ba5b71..813aa74754c 100644 --- a/api/services/recommended_app_service.py +++ b/api/services/recommended_app_service.py @@ -1,7 +1,7 @@ from typing import Any from sqlalchemy import select -from sqlalchemy.orm import scoped_session +from sqlalchemy.orm import Session from configs import dify_config from models.model import AccountTrialAppRecord, TrialApp @@ -11,7 +11,7 @@ from services.recommend_app.recommend_app_factory import RecommendAppRetrievalFa class RecommendedAppService: @classmethod - def get_recommended_apps_and_categories(cls, session: scoped_session, language: str): + def get_recommended_apps_and_categories(cls, language: str, *, session: Session): """ Get recommended apps and categories. :param language: language @@ -19,7 +19,7 @@ class RecommendedAppService: """ mode = dify_config.HOSTED_FETCH_APP_TEMPLATES_MODE retrieval_instance = RecommendAppRetrievalFactory.get_recommend_app_factory(mode)() - result = retrieval_instance.get_recommended_apps_and_categories(language) + result = retrieval_instance.get_recommended_apps_and_categories(language, session=session) if not result.get("recommended_apps"): result = ( RecommendAppRetrievalFactory.get_buildin_recommend_app_retrieval().fetch_recommended_apps_from_builtin( @@ -35,7 +35,7 @@ class RecommendedAppService: return result @classmethod - def get_learn_dify_apps(cls, session: scoped_session, language: str) -> dict[str, Any]: + def get_learn_dify_apps(cls, language: str, *, session: Session) -> dict[str, Any]: """ Get recommended apps marked for the Learn Dify section. :param language: language @@ -43,7 +43,7 @@ class RecommendedAppService: """ mode = dify_config.HOSTED_FETCH_APP_TEMPLATES_MODE retrieval_instance = RecommendAppRetrievalFactory.get_recommend_app_factory(mode)() - result = retrieval_instance.get_learn_dify_apps(language) + result = retrieval_instance.get_learn_dify_apps(language, session=session) if FeatureService.get_system_features().enable_trial_app: for app in result["recommended_apps"]: @@ -52,7 +52,7 @@ class RecommendedAppService: return {"recommended_apps": result["recommended_apps"]} @classmethod - def get_recommend_app_detail(cls, session: scoped_session, app_id: str) -> dict[str, Any] | None: + def get_recommend_app_detail(cls, app_id: str, *, session: Session) -> dict[str, Any] | None: """ Get recommend app detail. :param app_id: app id @@ -60,7 +60,7 @@ class RecommendedAppService: """ mode = dify_config.HOSTED_FETCH_APP_TEMPLATES_MODE retrieval_instance = RecommendAppRetrievalFactory.get_recommend_app_factory(mode)() - result: dict[str, Any] | None = retrieval_instance.get_recommend_app_detail(app_id) + result: dict[str, Any] | None = retrieval_instance.get_recommend_app_detail(app_id, session=session) if result is None: return None if FeatureService.get_system_features().enable_trial_app: @@ -69,7 +69,7 @@ class RecommendedAppService: return result @classmethod - def add_trial_app_record(cls, session: scoped_session, app_id: str, account_id: str): + def add_trial_app_record(cls, app_id: str, account_id: str, *, session: Session): """ Add trial app record. :param app_id: app id @@ -88,6 +88,6 @@ class RecommendedAppService: session.commit() @staticmethod - def _can_trial_app(session: scoped_session, app_id: str) -> bool: + def _can_trial_app(session: Session, app_id: str) -> bool: trial_app_model = session.scalar(select(TrialApp).where(TrialApp.app_id == app_id).limit(1)) return trial_app_model is not None diff --git a/api/services/saved_message_service.py b/api/services/saved_message_service.py index 9a65429748e..6165d74333f 100644 --- a/api/services/saved_message_service.py +++ b/api/services/saved_message_service.py @@ -12,7 +12,7 @@ from services.message_service import MessageService class SavedMessageService: @classmethod def pagination_by_last_id( - cls, session: Session, app_model: App, user: Account | EndUser | None, last_id: str | None, limit: int + cls, app_model: App, user: Account | EndUser | None, last_id: str | None, limit: int, *, session: Session ) -> InfiniteScrollPagination: if not user: raise ValueError("User is required") @@ -28,11 +28,16 @@ class SavedMessageService: message_ids = [sm.message_id for sm in saved_messages] return MessageService.pagination_by_last_id( - app_model=app_model, user=user, last_id=last_id, limit=limit, include_ids=message_ids + app_model=app_model, + user=user, + last_id=last_id, + limit=limit, + include_ids=message_ids, + session=session, ) @classmethod - def save(cls, session: Session, app_model: App, user: Account | EndUser | None, message_id: str): + def save(cls, app_model: App, user: Account | EndUser | None, message_id: str, *, session: Session): if not user: return saved_message = session.scalar( @@ -49,7 +54,7 @@ class SavedMessageService: if saved_message: return - message = MessageService.get_message(app_model=app_model, user=user, message_id=message_id) + message = MessageService.get_message(app_model=app_model, user=user, message_id=message_id, session=session) saved_message = SavedMessage( app_id=app_model.id, @@ -62,7 +67,7 @@ class SavedMessageService: session.commit() @classmethod - def delete(cls, session: Session, app_model: App, user: Account | EndUser | None, message_id: str): + def delete(cls, app_model: App, user: Account | EndUser | None, message_id: str, *, session: Session): if not user: return saved_message = session.scalar( diff --git a/api/services/snippet_service.py b/api/services/snippet_service.py index a54c9f6a069..64c1ec12370 100644 --- a/api/services/snippet_service.py +++ b/api/services/snippet_service.py @@ -6,9 +6,8 @@ from datetime import UTC, datetime from typing import Any from sqlalchemy import delete, func, select -from sqlalchemy.orm import Session, scoped_session, sessionmaker +from sqlalchemy.orm import Session, sessionmaker -from core.db import session_factory from core.workflow.node_factory import LATEST_VERSION, NODE_TYPE_CLASSES_MAPPING from graphon.enums import BuiltinNodeTypes, NodeType from libs.infinite_scroll_pagination import InfiniteScrollPagination @@ -59,9 +58,8 @@ class SnippetService: session_maker = None if session is not None: session_maker = sessionmaker(bind=session.get_bind(), expire_on_commit=False) - elif session_maker is None: - session_maker = session_factory.get_session_maker() - assert session_maker is not None + if session_maker is None: + raise ValueError("SnippetService requires a session or session_maker.") self._session = session self._session_maker = session_maker self._node_execution_service_repo = DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository( @@ -192,7 +190,7 @@ class SnippetService: self, *, tenant_id: str, - session: scoped_session, + session: Session, page: int = 1, limit: int = 20, keyword: str | None = None, diff --git a/api/services/summary_index_service.py b/api/services/summary_index_service.py index 3e065653bdf..3adc18dd2d5 100644 --- a/api/services/summary_index_service.py +++ b/api/services/summary_index_service.py @@ -7,7 +7,7 @@ from datetime import UTC, datetime from typing import TypedDict, cast from sqlalchemy import select -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from core.db.session_factory import session_factory from core.model_manager import ModelManager @@ -94,6 +94,8 @@ class SummaryIndexService: dataset: Dataset, summary_content: str, status: SummaryStatus = SummaryStatus.GENERATING, + *, + session: Session, ) -> DocumentSegmentSummary: """ Create or update a DocumentSegmentSummary record. @@ -105,46 +107,48 @@ class SummaryIndexService: summary_content: Generated summary content status: Summary status (default: SummaryStatus.GENERATING) + Keyword Args: + session: SQLAlchemy session used for the summary record. + Returns: Created or updated DocumentSegmentSummary instance """ - with session_factory.create_session() as session: - # Check if summary record already exists - existing_summary = session.scalar( - select(DocumentSegmentSummary) - .where( - DocumentSegmentSummary.chunk_id == segment.id, - DocumentSegmentSummary.dataset_id == dataset.id, - ) - .limit(1) + # Check if summary record already exists + existing_summary = session.scalar( + select(DocumentSegmentSummary) + .where( + DocumentSegmentSummary.chunk_id == segment.id, + DocumentSegmentSummary.dataset_id == dataset.id, ) + .limit(1) + ) - if existing_summary: - # Update existing record - existing_summary.summary_content = summary_content - existing_summary.status = status - existing_summary.error = None # Clear any previous errors - # Re-enable if it was disabled - if not existing_summary.enabled: - existing_summary.enabled = True - existing_summary.disabled_at = None - existing_summary.disabled_by = None - session.add(existing_summary) - session.flush() - return existing_summary - else: - # Create new record (enabled by default) - summary_record = DocumentSegmentSummary( - dataset_id=dataset.id, - document_id=segment.document_id, - chunk_id=segment.id, - summary_content=summary_content, - status=status, - enabled=True, # Explicitly set enabled to True - ) - session.add(summary_record) - session.flush() - return summary_record + if existing_summary: + # Update existing record + existing_summary.summary_content = summary_content + existing_summary.status = status + existing_summary.error = None # Clear any previous errors + # Re-enable if it was disabled + if not existing_summary.enabled: + existing_summary.enabled = True + existing_summary.disabled_at = None + existing_summary.disabled_by = None + session.add(existing_summary) + session.flush() + return existing_summary + else: + # Create new record (enabled by default) + summary_record = DocumentSegmentSummary( + dataset_id=dataset.id, + document_id=segment.document_id, + chunk_id=segment.id, + summary_content=summary_content, + status=status, + enabled=True, # Explicitly set enabled to True + ) + session.add(summary_record) + session.flush() + return summary_record @staticmethod def vectorize_summary( @@ -641,6 +645,8 @@ class SummaryIndexService: segment: DocumentSegment, dataset: Dataset, summary_index_setting: SummaryIndexSettingDict, + *, + session: Session, ) -> DocumentSegmentSummary: """ Generate summary for a segment and vectorize it. @@ -651,106 +657,101 @@ class SummaryIndexService: dataset: Dataset containing the segment summary_index_setting: Summary index configuration + Keyword Args: + session: SQLAlchemy session used for summary record updates. + Returns: Created DocumentSegmentSummary instance Raises: ValueError: If summary generation fails """ - with session_factory.create_session() as session: - try: - # Get or refresh summary record in this session - summary_record_in_session = session.scalar( - select(DocumentSegmentSummary) - .where( - DocumentSegmentSummary.chunk_id == segment.id, - DocumentSegmentSummary.dataset_id == dataset.id, - ) - .limit(1) + try: + # Get or refresh summary record in this session + summary_record_in_session = session.scalar( + select(DocumentSegmentSummary) + .where( + DocumentSegmentSummary.chunk_id == segment.id, + DocumentSegmentSummary.dataset_id == dataset.id, ) + .limit(1) + ) - if not summary_record_in_session: - # If not found, create one - logger.warning("Summary record not found for segment %s, creating one", segment.id) - summary_record_in_session = DocumentSegmentSummary( - dataset_id=dataset.id, - document_id=segment.document_id, - chunk_id=segment.id, - summary_content="", - status=SummaryStatus.GENERATING, - enabled=True, - ) - session.add(summary_record_in_session) - session.flush() - - # Update status to "generating" - summary_record_in_session.status = SummaryStatus.GENERATING - summary_record_in_session.error = None - session.add(summary_record_in_session) - # Don't flush here - wait until after vectorization succeeds - - # Generate summary (returns summary_content and llm_usage) - summary_content, llm_usage = SummaryIndexService.generate_summary_for_segment( - segment, dataset, summary_index_setting + if not summary_record_in_session: + # If not found, create one + logger.warning("Summary record not found for segment %s, creating one", segment.id) + summary_record_in_session = DocumentSegmentSummary( + dataset_id=dataset.id, + document_id=segment.document_id, + chunk_id=segment.id, + summary_content="", + status=SummaryStatus.GENERATING, + enabled=True, ) - - # Update summary content - summary_record_in_session.summary_content = summary_content session.add(summary_record_in_session) - # Flush to ensure summary_content is saved before vectorize_summary queries it session.flush() - # Log LLM usage for summary generation - if llm_usage and llm_usage.total_tokens > 0: - logger.info( - "Summary generation for segment %s used %s tokens (prompt: %s, completion: %s)", - segment.id, - llm_usage.total_tokens, - llm_usage.prompt_tokens, - llm_usage.completion_tokens, - ) + # Update status to "generating" + summary_record_in_session.status = SummaryStatus.GENERATING + summary_record_in_session.error = None + session.add(summary_record_in_session) + # Don't flush here - wait until after vectorization succeeds - # Vectorize summary (will delete old vector if exists before creating new one) - # Pass the session-managed record to vectorize_summary - # vectorize_summary will update status to "completed" and tokens in its own session - # vectorize_summary will also ensure summary_content is preserved - try: - # Pass the session to vectorize_summary to avoid session isolation issues - SummaryIndexService.vectorize_summary(summary_record_in_session, segment, dataset, session=session) - # Refresh the object from database to get the updated status and tokens from vectorize_summary - session.refresh(summary_record_in_session) - # Commit the session - # (summary_record_in_session should have status="completed" and tokens from refresh) - session.commit() - logger.info("Successfully generated and vectorized summary for segment %s", segment.id) - return summary_record_in_session - except Exception as vectorize_error: - # If vectorization fails, update status to error in current session - logger.exception("Failed to vectorize summary for segment %s", segment.id) - summary_record_in_session.status = SummaryStatus.ERROR - summary_record_in_session.error = f"Vectorization failed: {str(vectorize_error)}" - session.add(summary_record_in_session) - session.commit() - raise + # Generate summary (returns summary_content and llm_usage) + summary_content, llm_usage = SummaryIndexService.generate_summary_for_segment( + segment, dataset, summary_index_setting + ) - except Exception as e: - logger.exception("Failed to generate summary for segment %s", segment.id) - # Update summary record with error status - summary_record_in_session = session.scalar( - select(DocumentSegmentSummary) - .where( - DocumentSegmentSummary.chunk_id == segment.id, - DocumentSegmentSummary.dataset_id == dataset.id, - ) - .limit(1) + # Update summary content + summary_record_in_session.summary_content = summary_content + session.add(summary_record_in_session) + # Flush to ensure summary_content is saved before vectorize_summary queries it + session.flush() + + # Log LLM usage for summary generation + if llm_usage and llm_usage.total_tokens > 0: + logger.info( + "Summary generation for segment %s used %s tokens (prompt: %s, completion: %s)", + segment.id, + llm_usage.total_tokens, + llm_usage.prompt_tokens, + llm_usage.completion_tokens, ) - if summary_record_in_session: - summary_record_in_session.status = SummaryStatus.ERROR - summary_record_in_session.error = str(e) - session.add(summary_record_in_session) - session.commit() + + try: + SummaryIndexService.vectorize_summary(summary_record_in_session, segment, dataset, session=session) + # vectorize_summary mutates status and token fields; refresh before returning the ORM object. + session.refresh(summary_record_in_session) + session.commit() + logger.info("Successfully generated and vectorized summary for segment %s", segment.id) + return summary_record_in_session + except Exception as vectorize_error: + # If vectorization fails, update status to error in current session + logger.exception("Failed to vectorize summary for segment %s", segment.id) + summary_record_in_session.status = SummaryStatus.ERROR + summary_record_in_session.error = f"Vectorization failed: {str(vectorize_error)}" + session.add(summary_record_in_session) + session.commit() raise + except Exception as e: + logger.exception("Failed to generate summary for segment %s", segment.id) + # Update summary record with error status + summary_record_in_session = session.scalar( + select(DocumentSegmentSummary) + .where( + DocumentSegmentSummary.chunk_id == segment.id, + DocumentSegmentSummary.dataset_id == dataset.id, + ) + .limit(1) + ) + if summary_record_in_session: + summary_record_in_session.status = SummaryStatus.ERROR + summary_record_in_session.error = str(e) + session.add(summary_record_in_session) + session.commit() + raise + @staticmethod def generate_summaries_for_document( dataset: Dataset, @@ -840,7 +841,7 @@ class SummaryIndexService: try: summary_record = SummaryIndexService.generate_and_vectorize_summary( - segment, dataset, summary_index_setting + segment, dataset, summary_index_setting, session=session ) summary_records.append(summary_record) except Exception as e: @@ -1048,6 +1049,8 @@ class SummaryIndexService: segment: DocumentSegment, dataset: Dataset, summary_content: str, + *, + session: Session, ) -> DocumentSegmentSummary | None: """ Update summary for a segment and re-vectorize it. @@ -1057,6 +1060,9 @@ class SummaryIndexService: dataset: Dataset containing the segment summary_content: New summary content + Keyword Args: + session: SQLAlchemy session used for summary record updates. + Returns: Updated DocumentSegmentSummary instance, or None if indexing technique is not high_quality """ @@ -1072,67 +1078,22 @@ class SummaryIndexService: if segment.document and segment.document.doc_form == "qa_model": return None - with session_factory.create_session() as session: - try: - # Check if summary_content is empty (whitespace-only strings are considered empty) - if not summary_content or not summary_content.strip(): - # If summary is empty, only delete existing summary vector and record - summary_record = session.scalar( - select(DocumentSegmentSummary) - .where( - DocumentSegmentSummary.chunk_id == segment.id, - DocumentSegmentSummary.dataset_id == dataset.id, - ) - .limit(1) - ) - - if summary_record: - # Delete old vector if exists - old_summary_node_id = summary_record.summary_index_node_id - if old_summary_node_id: - try: - vector = Vector(dataset) - vector.delete_by_ids([old_summary_node_id]) - except Exception as e: - logger.warning( - "Failed to delete old summary vector for segment %s: %s", - segment.id, - str(e), - ) - - # Delete summary record since summary is empty - session.delete(summary_record) - session.commit() - logger.info("Deleted summary for segment %s (empty content provided)", segment.id) - return None - else: - # No existing summary record, nothing to do - logger.info("No summary record found for segment %s, nothing to delete", segment.id) - return None - - # Find existing summary record - summary_record = session.scalar( - select(DocumentSegmentSummary) - .where( - DocumentSegmentSummary.chunk_id == segment.id, - DocumentSegmentSummary.dataset_id == dataset.id, - ) - .limit(1) + try: + summary_record = session.scalar( + select(DocumentSegmentSummary) + .where( + DocumentSegmentSummary.chunk_id == segment.id, + DocumentSegmentSummary.dataset_id == dataset.id, ) + .limit(1) + ) + # Check if summary_content is empty (whitespace-only strings are considered empty) + if not summary_content or not summary_content.strip(): + # If summary is empty, only delete existing summary vector and record if summary_record: - # Update existing summary + # Delete old vector if exists old_summary_node_id = summary_record.summary_index_node_id - - # Update summary content - summary_record.summary_content = summary_content - summary_record.status = SummaryStatus.GENERATING - summary_record.error = None # Clear any previous errors - session.add(summary_record) - # Flush to ensure summary_content is saved before vectorize_summary queries it - session.flush() - - # Delete old vector if exists (before vectorization) if old_summary_node_id: try: vector = Vector(dataset) @@ -1144,80 +1105,90 @@ class SummaryIndexService: str(e), ) - # Re-vectorize summary (this will update status to "completed" and tokens in its own session) - # vectorize_summary will also ensure summary_content is preserved - # Note: vectorize_summary may take time due to embedding API calls, but we need to complete it - # to ensure the summary is properly indexed - try: - # Pass the session to vectorize_summary to avoid session isolation issues - SummaryIndexService.vectorize_summary(summary_record, segment, dataset, session=session) - # Refresh the object from database to get the updated status and tokens from vectorize_summary - session.refresh(summary_record) - # Now commit the session (summary_record should have status="completed" and tokens from refresh) - session.commit() - logger.info("Successfully updated and re-vectorized summary for segment %s", segment.id) - return summary_record - except Exception as e: - # If vectorization fails, update status to error in current session - # Don't raise the exception - just log it and return the record with error status - # This allows the segment update to complete even if vectorization fails - summary_record.status = SummaryStatus.ERROR - summary_record.error = f"Vectorization failed: {str(e)}" - session.commit() - logger.exception("Failed to vectorize summary for segment %s", segment.id) - # Return the record with error status instead of raising - # The caller can check the status if needed - return summary_record - else: - # Create new summary record if doesn't exist - summary_record = SummaryIndexService.create_summary_record( - segment, dataset, summary_content, status=SummaryStatus.GENERATING - ) - # Re-vectorize summary (this will update status to "completed" and tokens in its own session) - # Note: summary_record was created in a different session, - # so we need to merge it into current session - try: - # Merge the record into current session first (since it was created in a different session) - summary_record = session.merge(summary_record) - # Pass the session to vectorize_summary - it will update the merged record - SummaryIndexService.vectorize_summary(summary_record, segment, dataset, session=session) - # Refresh to get updated status and tokens from database - session.refresh(summary_record) - # Commit the session to persist the changes - session.commit() - logger.info("Successfully created and vectorized summary for segment %s", segment.id) - return summary_record - except Exception as e: - # If vectorization fails, update status to error in current session - # Merge the record into current session first - error_record = session.merge(summary_record) - error_record.status = SummaryStatus.ERROR - error_record.error = f"Vectorization failed: {str(e)}" - session.commit() - logger.exception("Failed to vectorize summary for segment %s", segment.id) - # Return the record with error status instead of raising - return error_record - - except Exception as e: - logger.exception("Failed to update summary for segment %s", segment.id) - # Update summary record with error status if it exists - summary_record = session.scalar( - select(DocumentSegmentSummary) - .where( - DocumentSegmentSummary.chunk_id == segment.id, - DocumentSegmentSummary.dataset_id == dataset.id, - ) - .limit(1) - ) - if summary_record: - summary_record.status = SummaryStatus.ERROR - summary_record.error = str(e) - session.add(summary_record) + # Delete summary record since summary is empty + session.delete(summary_record) session.commit() - raise + logger.info("Deleted summary for segment %s (empty content provided)", segment.id) + return None + else: + # No existing summary record, nothing to do + logger.info("No summary record found for segment %s, nothing to delete", segment.id) + return None + + if summary_record: + # Update existing summary + old_summary_node_id = summary_record.summary_index_node_id + + # Update summary content + summary_record.summary_content = summary_content + summary_record.status = SummaryStatus.GENERATING + summary_record.error = None # Clear any previous errors + session.add(summary_record) + # Flush to ensure summary_content is saved before vectorize_summary queries it + session.flush() + + # Delete old vector if exists (before vectorization) + if old_summary_node_id: + try: + vector = Vector(dataset) + vector.delete_by_ids([old_summary_node_id]) + except Exception as e: + logger.warning( + "Failed to delete old summary vector for segment %s: %s", + segment.id, + str(e), + ) + else: + # Create new summary record if doesn't exist + summary_record = SummaryIndexService.create_summary_record( + segment, + dataset, + summary_content, + status=SummaryStatus.GENERATING, + session=session, + ) + + try: + # Vectorization must finish here so the manual summary is searchable immediately. + SummaryIndexService.vectorize_summary(summary_record, segment, dataset, session=session) + session.refresh(summary_record) + session.commit() + logger.info("Successfully updated and re-vectorized summary for segment %s", segment.id) + return summary_record + except Exception as e: + # If vectorization fails, update status to error in current session. + # Return the record with error status so callers can still finish segment updates. + summary_record.status = SummaryStatus.ERROR + summary_record.error = f"Vectorization failed: {str(e)}" + session.commit() + logger.exception("Failed to vectorize summary for segment %s", segment.id) + return summary_record + + except Exception as e: + logger.exception("Failed to update summary for segment %s", segment.id) + # Update summary record with error status if it exists + summary_record = session.scalar( + select(DocumentSegmentSummary) + .where( + DocumentSegmentSummary.chunk_id == segment.id, + DocumentSegmentSummary.dataset_id == dataset.id, + ) + .limit(1) + ) + if summary_record: + summary_record.status = SummaryStatus.ERROR + summary_record.error = str(e) + session.add(summary_record) + session.commit() + raise @staticmethod - def get_segment_summary(segment_id: str, dataset_id: str) -> DocumentSegmentSummary | None: + def get_segment_summary( + segment_id: str, + dataset_id: str, + *, + session: Session, + ) -> DocumentSegmentSummary | None: """ Get summary for a single segment. @@ -1225,22 +1196,29 @@ class SummaryIndexService: segment_id: Segment ID (chunk_id) dataset_id: Dataset ID + Keyword Args: + session: SQLAlchemy session used to read summary records. + Returns: DocumentSegmentSummary instance if found, None otherwise """ - with session_factory.create_session() as session: - return session.scalar( - select(DocumentSegmentSummary) - .where( - DocumentSegmentSummary.chunk_id == segment_id, - DocumentSegmentSummary.dataset_id == dataset_id, - DocumentSegmentSummary.enabled.is_(True), # Only return enabled summaries - ) - .limit(1) + return session.scalar( + select(DocumentSegmentSummary) + .where( + DocumentSegmentSummary.chunk_id == segment_id, + DocumentSegmentSummary.dataset_id == dataset_id, + DocumentSegmentSummary.enabled.is_(True), ) + .limit(1) + ) @staticmethod - def get_segments_summaries(segment_ids: list[str], dataset_id: str) -> dict[str, DocumentSegmentSummary]: + def get_segments_summaries( + segment_ids: list[str], + dataset_id: str, + *, + session: Session, + ) -> dict[str, DocumentSegmentSummary]: """ Get summaries for multiple segments. @@ -1248,26 +1226,31 @@ class SummaryIndexService: segment_ids: List of segment IDs (chunk_ids) dataset_id: Dataset ID + Keyword Args: + session: SQLAlchemy session used to read summary records. + Returns: Dictionary mapping segment_id to DocumentSegmentSummary (only enabled summaries) """ if not segment_ids: return {} - with session_factory.create_session() as session: - summary_records = session.scalars( - select(DocumentSegmentSummary).where( - DocumentSegmentSummary.chunk_id.in_(segment_ids), - DocumentSegmentSummary.dataset_id == dataset_id, - DocumentSegmentSummary.enabled.is_(True), # Only return enabled summaries - ) - ).all() - - return {summary.chunk_id: summary for summary in summary_records} + summaries = session.scalars( + select(DocumentSegmentSummary).where( + DocumentSegmentSummary.chunk_id.in_(segment_ids), + DocumentSegmentSummary.dataset_id == dataset_id, + DocumentSegmentSummary.enabled.is_(True), + ) + ).all() + return {summary.chunk_id: summary for summary in summaries} @staticmethod def get_document_summaries( - document_id: str, dataset_id: str, segment_ids: list[str] | None = None + document_id: str, + dataset_id: str, + segment_ids: list[str] | None = None, + *, + session: Session, ) -> list[DocumentSegmentSummary]: """ Get all summary records for a document. @@ -1277,23 +1260,31 @@ class SummaryIndexService: dataset_id: Dataset ID segment_ids: Optional list of segment IDs to filter by + Keyword Args: + session: SQLAlchemy session used to read summary records. + Returns: List of DocumentSegmentSummary instances (only enabled summaries) """ - with session_factory.create_session() as session: - stmt = select(DocumentSegmentSummary).where( - DocumentSegmentSummary.document_id == document_id, - DocumentSegmentSummary.dataset_id == dataset_id, - DocumentSegmentSummary.enabled.is_(True), # Only return enabled summaries - ) + stmt = select(DocumentSegmentSummary).where( + DocumentSegmentSummary.document_id == document_id, + DocumentSegmentSummary.dataset_id == dataset_id, + DocumentSegmentSummary.enabled.is_(True), + ) - if segment_ids: - stmt = stmt.where(DocumentSegmentSummary.chunk_id.in_(segment_ids)) + if segment_ids: + stmt = stmt.where(DocumentSegmentSummary.chunk_id.in_(segment_ids)) - return list(session.scalars(stmt).all()) + return list(session.scalars(stmt).all()) @staticmethod - def get_document_summary_index_status(document_id: str, dataset_id: str, tenant_id: str) -> str | None: + def get_document_summary_index_status( + document_id: str, + dataset_id: str, + tenant_id: str, + *, + session: Session, + ) -> str | None: """ Get summary_index_status for a single document. @@ -1302,26 +1293,28 @@ class SummaryIndexService: dataset_id: Dataset ID tenant_id: Tenant ID + Keyword Args: + session: SQLAlchemy session used to read summary status. + Returns: "SUMMARIZING" if there are pending summaries, None otherwise """ # Get all segments for this document (excluding qa_model and re_segment) - with session_factory.create_session() as session: - segment_ids = list( - session.scalars( - select(DocumentSegment.id).where( - DocumentSegment.document_id == document_id, - DocumentSegment.status != "re_segment", - DocumentSegment.tenant_id == tenant_id, - ) - ).all() - ) + segment_ids = list( + session.scalars( + select(DocumentSegment.id).where( + DocumentSegment.document_id == document_id, + DocumentSegment.status != "re_segment", + DocumentSegment.tenant_id == tenant_id, + ) + ).all() + ) if not segment_ids: return None # Get all summary records for these segments - summaries = SummaryIndexService.get_segments_summaries(segment_ids, dataset_id) + summaries = SummaryIndexService.get_segments_summaries(segment_ids, dataset_id, session=session) summary_status_map = {chunk_id: summary.status for chunk_id, summary in summaries.items()} # Check if there are any "not_started" or "generating" status summaries @@ -1335,7 +1328,11 @@ class SummaryIndexService: @staticmethod def get_documents_summary_index_status( - document_ids: list[str], dataset_id: str, tenant_id: str + document_ids: list[str], + dataset_id: str, + tenant_id: str, + *, + session: Session, ) -> dict[str, str | None]: """ Get summary_index_status for multiple documents. @@ -1345,6 +1342,9 @@ class SummaryIndexService: dataset_id: Dataset ID tenant_id: Tenant ID + Keyword Args: + session: SQLAlchemy session used to read summary status. + Returns: Dictionary mapping document_id to summary_index_status ("SUMMARIZING" or None) """ @@ -1352,14 +1352,13 @@ class SummaryIndexService: return {} # Get all segments for these documents (excluding qa_model and re_segment) - with session_factory.create_session() as session: - segments = session.execute( - select(DocumentSegment.id, DocumentSegment.document_id).where( - DocumentSegment.document_id.in_(document_ids), - DocumentSegment.status != "re_segment", - DocumentSegment.tenant_id == tenant_id, - ) - ).all() + segments = session.execute( + select(DocumentSegment.id, DocumentSegment.document_id).where( + DocumentSegment.document_id.in_(document_ids), + DocumentSegment.status != "re_segment", + DocumentSegment.tenant_id == tenant_id, + ) + ).all() # Group segments by document_id document_segments_map: dict[str, list[str]] = {} @@ -1371,7 +1370,7 @@ class SummaryIndexService: # Get all summary records for these segments all_segment_ids = [seg.id for seg in segments] - summaries = SummaryIndexService.get_segments_summaries(all_segment_ids, dataset_id) + summaries = SummaryIndexService.get_segments_summaries(all_segment_ids, dataset_id, session=session) summary_status_map = {chunk_id: summary.status for chunk_id, summary in summaries.items()} # Calculate summary_index_status for each document @@ -1407,7 +1406,7 @@ class SummaryIndexService: def get_document_summary_status_detail( document_id: str, dataset_id: str, - session: Session | scoped_session, + session: Session, ) -> DocumentSummaryStatusDetailDict: """ Get detailed summary status for a document. @@ -1448,6 +1447,7 @@ class SummaryIndexService: document_id=document_id, dataset_id=dataset_id, segment_ids=segment_ids, + session=session, ) # Create a mapping of chunk_id to summary diff --git a/api/services/tag_service.py b/api/services/tag_service.py index 2d89bafa920..f404ec0eb37 100644 --- a/api/services/tag_service.py +++ b/api/services/tag_service.py @@ -6,7 +6,7 @@ from flask_login import current_user from pydantic import BaseModel, Field from sqlalchemy import delete, func, select from sqlalchemy.engine import CursorResult -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from werkzeug.exceptions import NotFound from models.dataset import Dataset @@ -14,7 +14,6 @@ from models.enums import TagType from models.model import App, Tag, TagBinding from models.snippet import CustomizedSnippet -type _SessionLike = Session | scoped_session type _TagTypeLike = TagType | str @@ -41,7 +40,7 @@ class TagBindingDeletePayload(BaseModel): class TagService: @staticmethod - def get_tags(session: Session, tag_type: _TagTypeLike, current_tenant_id: str, keyword: str | None = None): + def get_tags(tag_type: _TagTypeLike, current_tenant_id: str, keyword: str | None = None, *, session: Session): stmt = ( select(Tag.id, Tag.type, Tag.name, func.count(TagBinding.id).label("binding_count")) .outerjoin(TagBinding, Tag.id == TagBinding.tag_id) @@ -61,7 +60,7 @@ class TagService: tag_type: _TagTypeLike, current_tenant_id: str, tag_ids: list[str], - session: _SessionLike, + session: Session, *, match_all: bool = False, ): @@ -107,7 +106,7 @@ class TagService: return tag_bindings @staticmethod - def get_tag_by_tag_name(tag_type: _TagTypeLike, current_tenant_id: str, tag_name: str, session: _SessionLike): + def get_tag_by_tag_name(tag_type: _TagTypeLike, current_tenant_id: str, tag_name: str, session: Session): if not tag_type or not tag_name: return [] tags = list( @@ -120,7 +119,7 @@ class TagService: return tags @staticmethod - def get_tags_by_target_id(tag_type: _TagTypeLike, current_tenant_id: str, target_id: str, session: _SessionLike): + def get_tags_by_target_id(tag_type: _TagTypeLike, current_tenant_id: str, target_id: str, session: Session): tags = session.scalars( select(Tag) .join(TagBinding, Tag.id == TagBinding.tag_id) @@ -135,7 +134,7 @@ class TagService: return tags or [] @staticmethod - def save_tags(payload: SaveTagPayload, session: _SessionLike) -> Tag: + def save_tags(payload: SaveTagPayload, session: Session) -> Tag: if TagService.get_tag_by_tag_name(payload.type, current_user.current_tenant_id, payload.name, session): raise ValueError("Tag name already exists") tag = Tag( @@ -151,7 +150,7 @@ class TagService: @staticmethod def update_tags( - payload: UpdateTagPayload, tag_id: str, session: _SessionLike, *, tag_type: TagType | None = None + payload: UpdateTagPayload, tag_id: str, session: Session, *, tag_type: TagType | None = None ) -> Tag: current_tenant_id = current_user.current_tenant_id stmt = select(Tag).where(Tag.id == tag_id, Tag.tenant_id == current_tenant_id) @@ -178,7 +177,7 @@ class TagService: return tag @staticmethod - def get_tag_binding_count(tag_id: str, session: _SessionLike, *, tag_type: TagType | None = None) -> int: + def get_tag_binding_count(tag_id: str, session: Session, *, tag_type: TagType | None = None) -> int: current_tenant_id = current_user.current_tenant_id stmt = ( select(func.count(TagBinding.id)) @@ -191,7 +190,7 @@ class TagService: return count @staticmethod - def delete_tag(tag_id: str, session: _SessionLike, *, tag_type: TagType | None = None): + def delete_tag(tag_id: str, session: Session, *, tag_type: TagType | None = None): current_tenant_id = current_user.current_tenant_id stmt = select(Tag).where(Tag.id == tag_id, Tag.tenant_id == current_tenant_id) if tag_type is not None: @@ -210,7 +209,7 @@ class TagService: session.commit() @staticmethod - def save_tag_binding(payload: TagBindingCreatePayload, session: _SessionLike): + def save_tag_binding(payload: TagBindingCreatePayload, session: Session): TagService.check_target_exists(payload.type, payload.target_id, session) valid_tag_ids = session.scalars( select(Tag.id).where( @@ -237,7 +236,7 @@ class TagService: session.commit() @staticmethod - def delete_tag_binding(payload: TagBindingDeletePayload, session: _SessionLike): + def delete_tag_binding(payload: TagBindingDeletePayload, session: Session): TagService.check_target_exists(payload.type, payload.target_id, session) result = cast( CursorResult, @@ -260,7 +259,7 @@ class TagService: session.commit() @staticmethod - def check_target_exists(type: _TagTypeLike, target_id: str, session: _SessionLike): + def check_target_exists(type: _TagTypeLike, target_id: str, session: Session): if type == "knowledge": dataset = session.scalar( select(Dataset) diff --git a/api/services/tools/builtin_tools_manage_service.py b/api/services/tools/builtin_tools_manage_service.py index e49ab8398f1..45480f71d1a 100644 --- a/api/services/tools/builtin_tools_manage_service.py +++ b/api/services/tools/builtin_tools_manage_service.py @@ -327,7 +327,7 @@ class BuiltinToolManageService: @staticmethod def generate_builtin_tool_provider_name( - session: Session, tenant_id: str, provider: str, credential_type: CredentialType + tenant_id: str, provider: str, credential_type: CredentialType, *, session: Session ) -> str: db_providers = session.scalars( select(BuiltinToolProvider) @@ -347,6 +347,7 @@ class BuiltinToolManageService: def get_builtin_tool_provider_credentials( tenant_id: str, provider_name: str, + session: Session, user: Account | None = None, include_credential_ids: list[str] | None = None, ) -> list[ToolProviderCredentialApiEntity]: @@ -367,7 +368,7 @@ class BuiltinToolManageService: from models.credential_permission import CredentialType as CredPermType from services.credential_permission_service import CredentialPermissionService - with db.session.no_autoflush: + with session.no_autoflush: base_filter = ( BuiltinToolProvider.tenant_id == tenant_id, BuiltinToolProvider.provider == provider_name, @@ -383,7 +384,7 @@ class BuiltinToolManageService: credential_type=CredPermType.BUILTIN_TOOL_PROVIDER, user=user, ) - visible_providers = list(db.session.scalars(visible_query).all()) + visible_providers = list(session.scalars(visible_query).all()) # Fetch any explicitly-included IDs that the visibility filter excluded. borrowed_ids: set[str] = set() @@ -397,7 +398,7 @@ class BuiltinToolManageService: .where(*base_filter, BuiltinToolProvider.id.in_(wanted_ids)) .order_by(*order) ) - borrowed_providers = list(db.session.scalars(borrowed_query).all()) + borrowed_providers = list(session.scalars(borrowed_query).all()) borrowed_ids = {p.id for p in borrowed_providers} providers = visible_providers + borrowed_providers @@ -427,7 +428,7 @@ class BuiltinToolManageService: if vis_str == "partial_members": credential_entity.partial_member_list = list( CredentialPermissionService.get_partial_member_list( - db.session, provider.id, CredPermType.BUILTIN_TOOL_PROVIDER + provider.id, CredPermType.BUILTIN_TOOL_PROVIDER, session=session ) ) if provider.id in borrowed_ids: @@ -439,6 +440,7 @@ class BuiltinToolManageService: def get_builtin_tool_provider_credential_info( tenant_id: str, provider: str, + session: Session, user: Account | None = None, include_credential_ids: list[str] | None = None, ) -> ToolProviderCredentialInfoApiEntity: @@ -450,6 +452,7 @@ class BuiltinToolManageService: credentials = BuiltinToolManageService.get_builtin_tool_provider_credentials( tenant_id, provider, + session=session, user=user, include_credential_ids=include_credential_ids, ) diff --git a/api/services/trigger/schedule_service.py b/api/services/trigger/schedule_service.py index a827222c1dc..495674248b1 100644 --- a/api/services/trigger/schedule_service.py +++ b/api/services/trigger/schedule_service.py @@ -26,10 +26,7 @@ logger = logging.getLogger(__name__) class ScheduleService: @staticmethod def create_schedule( - session: Session, - tenant_id: str, - app_id: str, - config: ScheduleConfig, + tenant_id: str, app_id: str, config: ScheduleConfig, *, session: Session ) -> WorkflowSchedulePlan: """ Create a new schedule with validated configuration. @@ -63,11 +60,7 @@ class ScheduleService: return schedule @staticmethod - def update_schedule( - session: Session, - schedule_id: str, - updates: SchedulePlanUpdate, - ) -> WorkflowSchedulePlan: + def update_schedule(schedule_id: str, updates: SchedulePlanUpdate, *, session: Session) -> WorkflowSchedulePlan: """ Update an existing schedule with validated configuration. @@ -110,10 +103,7 @@ class ScheduleService: return schedule @staticmethod - def delete_schedule( - session: Session, - schedule_id: str, - ) -> None: + def delete_schedule(schedule_id: str, *, session: Session) -> None: """ Delete a schedule plan. @@ -129,7 +119,7 @@ class ScheduleService: session.flush() @staticmethod - def get_tenant_owner(session: Session, tenant_id: str) -> Account: + def get_tenant_owner(tenant_id: str, *, session: Session) -> Account: """ Returns an account to execute scheduled workflows on behalf of the tenant. Prioritizes owner over admin to ensure proper authorization hierarchy. @@ -157,10 +147,7 @@ class ScheduleService: raise AccountNotFoundError(f"Account not found for tenant: {tenant_id}") @staticmethod - def update_next_run_at( - session: Session, - schedule_id: str, - ) -> datetime: + def update_next_run_at(schedule_id: str, *, session: Session) -> datetime: """ Advances the schedule to its next execution time after a successful trigger. Uses current time as base to prevent missing executions during delays. diff --git a/api/services/trigger/trigger_provider_service.py b/api/services/trigger/trigger_provider_service.py index b0a3de1cee8..8506c523a61 100644 --- a/api/services/trigger/trigger_provider_service.py +++ b/api/services/trigger/trigger_provider_service.py @@ -388,7 +388,7 @@ class TriggerProviderService: return subscription @classmethod - def delete_trigger_provider(cls, session: Session, tenant_id: str, subscription_id: str): + def delete_trigger_provider(cls, tenant_id: str, subscription_id: str, *, session: Session): """ Delete a trigger provider subscription within an existing session. diff --git a/api/services/trigger/trigger_subscription_operator_service.py b/api/services/trigger/trigger_subscription_operator_service.py index 5d7785549e6..491723c6ec2 100644 --- a/api/services/trigger/trigger_subscription_operator_service.py +++ b/api/services/trigger/trigger_subscription_operator_service.py @@ -40,12 +40,7 @@ class TriggerSubscriptionOperatorService: return list(subscribers) @classmethod - def delete_plugin_trigger_by_subscription( - cls, - session: Session, - tenant_id: str, - subscription_id: str, - ) -> None: + def delete_plugin_trigger_by_subscription(cls, tenant_id: str, subscription_id: str, *, session: Session) -> None: """Delete a plugin trigger by tenant_id and subscription_id within an existing session Args: diff --git a/api/services/trigger/webhook_service.py b/api/services/trigger/webhook_service.py index 23b3ac55b93..587048e2ccd 100644 --- a/api/services/trigger/webhook_service.py +++ b/api/services/trigger/webhook_service.py @@ -835,11 +835,7 @@ class WebhookService: # NOTE: don not use `with sessionmaker(bind=db.engine, expire_on_commit=False).begin()` # trigger_workflow_async need to handle multipe session commits internally with Session(db.engine, expire_on_commit=False) as session: - AsyncWorkflowService.trigger_workflow_async( - session, - end_user, - trigger_data, - ) + AsyncWorkflowService.trigger_workflow_async(end_user, trigger_data, session=session) quota_charge.commit() except Exception: quota_charge.refund() diff --git a/api/services/vector_service.py b/api/services/vector_service.py index 5b5088ec5a1..faf4fb085d6 100644 --- a/api/services/vector_service.py +++ b/api/services/vector_service.py @@ -1,6 +1,7 @@ import logging from sqlalchemy import delete, select +from sqlalchemy.orm import Session from core.model_manager import ModelInstance, ModelManager from core.rag.datasource.keyword.keyword_factory import Keyword @@ -11,7 +12,6 @@ from core.rag.index_processor.constant.index_type import IndexStructureType, Ind from core.rag.index_processor.index_processor_base import BaseIndexProcessor from core.rag.index_processor.index_processor_factory import IndexProcessorFactory from core.rag.models.document import AttachmentDocument, Document -from extensions.ext_database import db from graphon.model_runtime.entities.model_entities import ModelType from models import UploadFile from models.dataset import ChildChunk, Dataset, DatasetProcessRule, DocumentSegment, SegmentAttachmentBinding @@ -24,14 +24,20 @@ logger = logging.getLogger(__name__) class VectorService: @classmethod def create_segments_vector( - cls, keywords_list: list[list[str]] | None, segments: list[DocumentSegment], dataset: Dataset, doc_form: str + cls, + keywords_list: list[list[str]] | None, + segments: list[DocumentSegment], + dataset: Dataset, + doc_form: str, + session: Session, ): + """Create vector records for document segments using the caller's active DB session.""" documents: list[Document] = [] multimodal_documents: list[AttachmentDocument] = [] for segment in segments: if doc_form == IndexStructureType.PARENT_CHILD_INDEX: - dataset_document = db.session.get(DatasetDocument, segment.document_id) + dataset_document = session.get(DatasetDocument, segment.document_id) if not dataset_document: logger.warning( "Expected DatasetDocument record to exist, but none was found, document_id=%s, segment_id=%s", @@ -40,7 +46,7 @@ class VectorService: ) continue # get the process rule - processing_rule = db.session.get(DatasetProcessRule, dataset_document.dataset_process_rule_id) + processing_rule = session.get(DatasetProcessRule, dataset_document.dataset_process_rule_id) if not processing_rule: raise ValueError("No processing rule found.") # get embedding model instance @@ -63,7 +69,13 @@ class VectorService: else: raise ValueError("The knowledge base index technique is not high quality!") cls.generate_child_chunks( - segment, dataset_document, dataset, embedding_model_instance, processing_rule, False + segment, + dataset_document, + dataset, + embedding_model_instance, + processing_rule, + session, + False, ) else: rag_document = Document( @@ -136,8 +148,10 @@ class VectorService: dataset: Dataset, embedding_model_instance: ModelInstance, processing_rule: DatasetProcessRule, + session: Session, regenerate: bool = False, ): + """Generate child chunks and persist them with the caller's active DB session.""" index_processor = IndexProcessorFactory(dataset.doc_form).init_index_processor() assert segment.index_node_id if regenerate: @@ -184,8 +198,8 @@ class VectorService: type=SegmentType.AUTOMATIC, created_by=dataset_document.created_by, ) - db.session.add(child_segment) - db.session.commit() + session.add(child_segment) + session.commit() @classmethod def create_child_chunk_vector(cls, child_segment: ChildChunk, dataset: Dataset): @@ -255,7 +269,10 @@ class VectorService: vector.delete_by_ids([child_chunk.index_node_id]) @classmethod - def update_multimodel_vector(cls, segment: DocumentSegment, attachment_ids: list[str], dataset: Dataset): + def update_multimodel_vector( + cls, segment: DocumentSegment, attachment_ids: list[str], dataset: Dataset, session: Session + ): + """Update multimodal vectors and attachment bindings with the caller's active DB session.""" if dataset.indexing_technique != IndexTechniqueType.HIGH_QUALITY: return @@ -274,19 +291,17 @@ class VectorService: vector.delete_by_ids(old_attachment_ids) # Delete existing segment attachment bindings in one operation - db.session.execute( - delete(SegmentAttachmentBinding).where(SegmentAttachmentBinding.segment_id == segment.id) - ) + session.execute(delete(SegmentAttachmentBinding).where(SegmentAttachmentBinding.segment_id == segment.id)) if not attachment_ids: - db.session.commit() + session.commit() return # Bulk fetch upload files - only fetch needed fields - upload_file_list = db.session.scalars(select(UploadFile).where(UploadFile.id.in_(attachment_ids))).all() + upload_file_list = session.scalars(select(UploadFile).where(UploadFile.id.in_(attachment_ids))).all() if not upload_file_list: - db.session.commit() + session.commit() return # Create a mapping for quick lookup @@ -329,16 +344,16 @@ class VectorService: # Bulk insert all bindings at once if bindings: - db.session.add_all(bindings) + session.add_all(bindings) # Add documents to vector store if any if documents and dataset.is_multimodal: vector.create_multimodal(documents) # Single commit for all operations - db.session.commit() + session.commit() except Exception: logger.exception("Failed to update multimodal vector for segment %s", segment.id) - db.session.rollback() + session.rollback() raise diff --git a/api/services/web_conversation_service.py b/api/services/web_conversation_service.py index 2c8a3be8631..96d95d5f5ac 100644 --- a/api/services/web_conversation_service.py +++ b/api/services/web_conversation_service.py @@ -2,7 +2,6 @@ from sqlalchemy import select from sqlalchemy.orm import Session from core.app.entities.app_invoke_entities import InvokeFrom -from extensions.ext_database import db from libs.infinite_scroll_pagination import InfiniteScrollPagination from models import Account from models.enums import CreatorUserRole @@ -59,10 +58,10 @@ class WebConversationService: ) @classmethod - def pin(cls, app_model: App, conversation_id: str, user: Account | EndUser | None): + def pin(cls, app_model: App, conversation_id: str, user: Account | EndUser | None, session: Session): if not user: return - pinned_conversation = db.session.scalar( + pinned_conversation = session.scalar( select(PinnedConversation) .where( PinnedConversation.app_id == app_model.id, @@ -77,7 +76,7 @@ class WebConversationService: return conversation = ConversationService.get_conversation( - app_model=app_model, conversation_id=conversation_id, user=user + app_model=app_model, conversation_id=conversation_id, user=user, session=session ) pinned_conversation = PinnedConversation( @@ -87,14 +86,14 @@ class WebConversationService: created_by=user.id, ) - db.session.add(pinned_conversation) - db.session.commit() + session.add(pinned_conversation) + session.commit() @classmethod - def unpin(cls, app_model: App, conversation_id: str, user: Account | EndUser | None): + def unpin(cls, app_model: App, conversation_id: str, user: Account | EndUser | None, session: Session): if not user: return - pinned_conversation = db.session.scalar( + pinned_conversation = session.scalar( select(PinnedConversation) .where( PinnedConversation.app_id == app_model.id, @@ -108,5 +107,5 @@ class WebConversationService: if not pinned_conversation: return - db.session.delete(pinned_conversation) - db.session.commit() + session.delete(pinned_conversation) + session.commit() diff --git a/api/services/webapp_auth_service.py b/api/services/webapp_auth_service.py index 6ecc8eb8bc9..33267c53d5c 100644 --- a/api/services/webapp_auth_service.py +++ b/api/services/webapp_auth_service.py @@ -4,10 +4,10 @@ from datetime import UTC, datetime, timedelta from typing import Any from sqlalchemy import select +from sqlalchemy.orm import Session from werkzeug.exceptions import NotFound, Unauthorized from configs import dify_config -from extensions.ext_database import db from libs.helper import TokenManager from libs.passport import PassportService from libs.password import compare_password @@ -33,9 +33,9 @@ class WebAppAuthService: """Service for web app authentication.""" @staticmethod - def authenticate(email: str, password: str) -> Account: + def authenticate(email: str, password: str, session: Session) -> Account: """authenticate account with email and password""" - account = AccountService.get_account_by_email_with_case_fallback(db.session, email) + account = AccountService.get_account_by_email_with_case_fallback(email, session=session) if not account: raise AccountNotFoundError() @@ -54,8 +54,8 @@ class WebAppAuthService: return access_token @classmethod - def get_user_through_email(cls, email: str): - account = AccountService.get_account_by_email_with_case_fallback(db.session, email) + def get_user_through_email(cls, email: str, session: Session): + account = AccountService.get_account_by_email_with_case_fallback(email, session=session) if not account: return None @@ -93,11 +93,11 @@ class WebAppAuthService: TokenManager.revoke_token(token, "email_code_login") @classmethod - def create_end_user(cls, app_code, email) -> EndUser: - site = db.session.scalar(select(Site).where(Site.code == app_code).limit(1)) + def create_end_user(cls, app_code, email, session: Session) -> EndUser: + site = session.scalar(select(Site).where(Site.code == app_code).limit(1)) if not site: raise NotFound("Site not found.") - app_model = db.session.get(App, site.app_id) + app_model = session.get(App, site.app_id) if not app_model: raise NotFound("App not found.") end_user = EndUser( @@ -109,8 +109,8 @@ class WebAppAuthService: name="enterpriseuser", external_user_id="enterpriseuser", ) - db.session.add(end_user) - db.session.commit() + session.add(end_user) + session.commit() return end_user @@ -133,7 +133,7 @@ class WebAppAuthService: @classmethod def is_app_require_permission_check( - cls, app_code: str | None = None, app_id: str | None = None, access_mode: str | None = None + cls, app_code: str | None = None, app_id: str | None = None, access_mode: str | None = None, *, session: Session ) -> bool: """ Check if the app requires permission check based on its access mode. @@ -145,7 +145,7 @@ class WebAppAuthService: raise ValueError("Either app_code or app_id must be provided.") if app_code: - app_id = AppService.get_app_id_by_code(app_code) + app_id = AppService.get_app_id_by_code(app_code, session=session) if not app_id: raise ValueError("App ID could not be determined from the provided app_code.") @@ -155,7 +155,9 @@ class WebAppAuthService: return False @classmethod - def get_app_auth_type(cls, app_code: str | None = None, access_mode: str | None = None) -> WebAppAuthType: + def get_app_auth_type( + cls, app_code: str | None = None, access_mode: str | None = None, *, session: Session + ) -> WebAppAuthType: """ Get the authentication type for the app based on its access mode. """ @@ -171,8 +173,8 @@ class WebAppAuthService: return WebAppAuthType.EXTERNAL if app_code: - app_id = AppService.get_app_id_by_code(app_code) + app_id = AppService.get_app_id_by_code(app_code, session=session) webapp_settings = EnterpriseService.WebAppAuth.get_app_access_mode_by_id(app_id=app_id) - return cls.get_app_auth_type(access_mode=webapp_settings.access_mode) + return cls.get_app_auth_type(access_mode=webapp_settings.access_mode, session=session) raise ValueError("Could not determine app authentication type.") diff --git a/api/services/workflow/node_output_inspector_service.py b/api/services/workflow/node_output_inspector_service.py index 66dcfec591f..5d6a8f1c675 100644 --- a/api/services/workflow/node_output_inspector_service.py +++ b/api/services/workflow/node_output_inspector_service.py @@ -52,9 +52,9 @@ from typing import Any from pydantic import BaseModel, ConfigDict, Field from sqlalchemy import select +from sqlalchemy.orm import Session from core.app.file_access import DatabaseFileAccessController -from core.db.session_factory import session_factory from core.workflow.nodes.agent_v2.binding_resolver import ( WorkflowAgentBindingError, WorkflowAgentBindingResolver, @@ -410,8 +410,8 @@ class NodeOutputInspectorService: The service is dependency-light: it holds a single :class:`WorkflowAgentBindingResolver` so agent v2 nodes can map to their declared outputs without re-implementing binding lookup. All other I/O - uses the global session factory so workflow runs / executions stay on the - repo-default code path. + receives an explicit SQLAlchemy session from its caller so transaction + ownership stays at the controller/task boundary. Tenancy is enforced via ``app_model.tenant_id`` + ``app_model.id`` on every load — the same scope guard regardless of trigger source. @@ -422,9 +422,13 @@ class NodeOutputInspectorService: # ── public API ──────────────────────────────────────────────────────── - def snapshot_workflow_run(self, *, app_model: App, workflow_run_id: str) -> WorkflowRunSnapshotView: + def snapshot_workflow_run( + self, *, app_model: App, workflow_run_id: str, session: Session + ) -> WorkflowRunSnapshotView: """Build the per-node snapshot for one debug workflow run.""" - workflow_run, executions = self._load_run_and_executions(app_model=app_model, workflow_run_id=workflow_run_id) + workflow_run, executions = self._load_run_and_executions( + app_model=app_model, workflow_run_id=workflow_run_id, session=session + ) executions_by_node = self._index_executions_by_node(executions) graph_nodes = _graph_nodes(workflow_run) @@ -447,9 +451,11 @@ class NodeOutputInspectorService: node_outputs=node_views, ) - def node_detail(self, *, app_model: App, workflow_run_id: str, node_id: str) -> NodeOutputsView: + def node_detail(self, *, app_model: App, workflow_run_id: str, node_id: str, session: Session) -> NodeOutputsView: """Per-node Inspector entry — returns one ``NodeOutputsView``.""" - workflow_run, executions = self._load_run_and_executions(app_model=app_model, workflow_run_id=workflow_run_id) + workflow_run, executions = self._load_run_and_executions( + app_model=app_model, workflow_run_id=workflow_run_id, session=session + ) graph_nodes = _graph_nodes(workflow_run) raw_node = next((n for n in graph_nodes if str(n.get("id")) == node_id), None) if raw_node is None: @@ -474,9 +480,12 @@ class NodeOutputInspectorService: workflow_run_id: str, node_id: str, output_name: str, + session: Session, ) -> OutputPreviewView: """Full payload for one declared output (with signed file URL).""" - workflow_run, executions = self._load_run_and_executions(app_model=app_model, workflow_run_id=workflow_run_id) + workflow_run, executions = self._load_run_and_executions( + app_model=app_model, workflow_run_id=workflow_run_id, session=session + ) graph_nodes = _graph_nodes(workflow_run) raw_node = next((n for n in graph_nodes if str(n.get("id")) == node_id), None) if raw_node is None: @@ -536,7 +545,7 @@ class NodeOutputInspectorService: # ── DB loading ──────────────────────────────────────────────────────── def _load_run_and_executions( - self, *, app_model: App, workflow_run_id: str + self, *, app_model: App, workflow_run_id: str, session: Session ) -> tuple[WorkflowRun, Sequence[WorkflowNodeExecutionModel]]: """Fetch the ``WorkflowRun`` row + every execution that belongs to it. @@ -548,24 +557,23 @@ class NodeOutputInspectorService: deliberately not checked here — D-1 was lifted 2026-05-26 and the Inspector now serves both draft and published runs. """ - with session_factory.create_session() as session: - workflow_run = session.scalar( - select(WorkflowRun).where( - WorkflowRun.id == workflow_run_id, - WorkflowRun.app_id == app_model.id, - WorkflowRun.tenant_id == app_model.tenant_id, - ) + workflow_run = session.scalar( + select(WorkflowRun).where( + WorkflowRun.id == workflow_run_id, + WorkflowRun.app_id == app_model.id, + WorkflowRun.tenant_id == app_model.tenant_id, ) - if workflow_run is None: - raise NodeOutputInspectorError("workflow_run_not_found", "Workflow run not found.") + ) + if workflow_run is None: + raise NodeOutputInspectorError("workflow_run_not_found", "Workflow run not found.") - executions = session.scalars( - select(WorkflowNodeExecutionModel).where( - WorkflowNodeExecutionModel.workflow_run_id == workflow_run_id, - WorkflowNodeExecutionModel.tenant_id == app_model.tenant_id, - WorkflowNodeExecutionModel.app_id == app_model.id, - ) - ).all() + executions = session.scalars( + select(WorkflowNodeExecutionModel).where( + WorkflowNodeExecutionModel.workflow_run_id == workflow_run_id, + WorkflowNodeExecutionModel.tenant_id == app_model.tenant_id, + WorkflowNodeExecutionModel.app_id == app_model.id, + ) + ).all() return workflow_run, executions diff --git a/api/services/workflow/workflow_converter.py b/api/services/workflow/workflow_converter.py index e279f1daaa3..5f787bb51cd 100644 --- a/api/services/workflow/workflow_converter.py +++ b/api/services/workflow/workflow_converter.py @@ -2,6 +2,7 @@ import json from typing import Any, TypedDict from sqlalchemy import select +from sqlalchemy.orm import Session from core.app.app_config.entities import ( DatasetEntity, @@ -18,7 +19,6 @@ from core.helper import encrypter from core.prompt.simple_prompt_transform import SimplePromptTransform from core.prompt.utils.prompt_template_parser import PromptTemplateParser from events.app_event import app_was_created -from extensions.ext_database import db from graphon.file import FileUploadConfig from graphon.model_runtime.entities.llm_entities import LLMMode from graphon.model_runtime.utils.encoders import jsonable_encoder @@ -53,7 +53,14 @@ class WorkflowConverter: """ def convert_to_workflow( - self, app_model: App, account: Account, name: str, icon_type: str, icon: str, icon_background: str + self, + app_model: App, + account: Account, + name: str, + icon_type: str, + icon: str, + icon_background: str, + session: Session, ): """ Convert app to workflow @@ -77,7 +84,7 @@ class WorkflowConverter: raise ValueError("App model config is required") workflow = self.convert_app_model_config_to_workflow( - app_model=app_model, app_model_config=app_model.app_model_config, account_id=account.id + app_model=app_model, app_model_config=app_model.app_model_config, account_id=account.id, session=session ) # create new app @@ -97,17 +104,19 @@ class WorkflowConverter: new_app.created_by = account.id new_app.maintainer = account.id new_app.updated_by = account.id - db.session.add(new_app) - db.session.flush() + session.add(new_app) + session.flush() workflow.app_id = new_app.id - db.session.commit() + session.commit() app_was_created.send(new_app, account=account) return new_app - def convert_app_model_config_to_workflow(self, app_model: App, app_model_config: AppModelConfig, account_id: str): + def convert_app_model_config_to_workflow( + self, app_model: App, app_model_config: AppModelConfig, account_id: str, session: Session + ): """ Convert app model config to workflow mode :param app_model: App instance @@ -144,6 +153,7 @@ class WorkflowConverter: app_model=app_model, variables=app_config.variables, external_data_variables=app_config.external_data_variables, + session=session, ) for http_request_node in http_request_nodes: @@ -217,8 +227,8 @@ class WorkflowConverter: conversation_variables=[], ) - db.session.add(workflow) - db.session.commit() + session.add(workflow) + session.commit() return workflow @@ -262,7 +272,11 @@ class WorkflowConverter: } def _convert_to_http_request_node( - self, app_model: App, variables: list[VariableEntity], external_data_variables: list[ExternalDataVariableEntity] + self, + app_model: App, + variables: list[VariableEntity], + external_data_variables: list[ExternalDataVariableEntity], + session: Session, ) -> tuple[list[_NodeType], dict[str, str]]: """ Convert API Based Extension to HTTP Request Node @@ -290,7 +304,7 @@ class WorkflowConverter: # get api_based_extension api_based_extension = self._get_api_based_extension( - tenant_id=tenant_id, api_based_extension_id=api_based_extension_id + tenant_id=tenant_id, api_based_extension_id=api_based_extension_id, session=session ) # decrypt api_key @@ -650,14 +664,14 @@ class WorkflowConverter: else: return AppMode.ADVANCED_CHAT - def _get_api_based_extension(self, tenant_id: str, api_based_extension_id: str): + def _get_api_based_extension(self, tenant_id: str, api_based_extension_id: str, session: Session): """ Get API Based Extension :param tenant_id: tenant id :param api_based_extension_id: api based extension id :return: """ - api_based_extension = db.session.scalar( + api_based_extension = session.scalar( select(APIBasedExtension) .where(APIBasedExtension.tenant_id == tenant_id, APIBasedExtension.id == api_based_extension_id) .limit(1) diff --git a/api/services/workflow_collaboration_service.py b/api/services/workflow_collaboration_service.py index bec61ce666d..5c635d7d66a 100644 --- a/api/services/workflow_collaboration_service.py +++ b/api/services/workflow_collaboration_service.py @@ -9,8 +9,8 @@ from collections.abc import Mapping from typing import Any, override from sqlalchemy import select +from sqlalchemy.orm import Session -from core.db.session_factory import session_factory from models.account import Account from models.model import App from repositories.workflow_collaboration_repository import WorkflowCollaborationRepository, WorkflowSessionInfo @@ -94,20 +94,22 @@ class WorkflowCollaborationService: }, ) - def authorize_and_join_workflow_room(self, workflow_id: str, sid: str) -> tuple[str, bool] | None: + def authorize_and_join_workflow_room( + self, workflow_id: str, sid: str, *, session: Session + ) -> tuple[str, bool] | None: """ Join a collaboration room only after validating the socket session and tenant-scoped app access. The Socket.IO payload still calls the room key `workflow_id`, but the identifier is the workflow app's `App.id`. Returning `None` lets the controller reject the join before any Redis or room state is created. """ - session = self._socketio.get_session(sid) - user_id = session.get("user_id") - tenant_id = session.get("tenant_id") + socket_session = self._socketio.get_session(sid) + user_id = socket_session.get("user_id") + tenant_id = socket_session.get("tenant_id") if not user_id or not tenant_id: return None - if not self._can_access_workflow(workflow_id, str(tenant_id)): + if not self._can_access_workflow(workflow_id, str(tenant_id), session=session): logger.warning( "Workflow collaboration join rejected: workflow_id=%s tenant_id=%s user_id=%s sid=%s", workflow_id, @@ -121,8 +123,8 @@ class WorkflowCollaborationService: session_info: WorkflowSessionInfo = { "user_id": str(user_id), - "username": str(session.get("username", "Unknown")), - "avatar": session.get("avatar"), + "username": str(socket_session.get("username", "Unknown")), + "avatar": socket_session.get("avatar"), "sid": sid, "connected_at": int(time.time()), "server_id": self.server_id, @@ -140,10 +142,9 @@ class WorkflowCollaborationService: return str(user_id), is_leader - def _can_access_workflow(self, workflow_id: str, tenant_id: str) -> bool: + def _can_access_workflow(self, workflow_id: str, tenant_id: str, *, session: Session) -> bool: """Check room access without relying on Flask's app-context-bound scoped session.""" - with session_factory.create_session() as session: - app_id = session.scalar(select(App.id).where(App.id == workflow_id, App.tenant_id == tenant_id).limit(1)) + app_id = session.scalar(select(App.id).where(App.id == workflow_id, App.tenant_id == tenant_id).limit(1)) return app_id is not None def disconnect_session(self, sid: str) -> None: diff --git a/api/services/workflow_service.py b/api/services/workflow_service.py index 048b25c6bf9..95be0f7017a 100644 --- a/api/services/workflow_service.py +++ b/api/services/workflow_service.py @@ -7,7 +7,7 @@ from dataclasses import dataclass from typing import Any, cast from sqlalchemy import exists, select -from sqlalchemy.orm import Session, scoped_session, sessionmaker +from sqlalchemy.orm import Session, sessionmaker from configs import dify_config from core.app.apps.advanced_chat.app_config_manager import AdvancedChatAppConfigManager @@ -178,7 +178,7 @@ class WorkflowService: node_id=node_id, ) - def is_workflow_exist(self, app_model: App) -> bool: + def is_workflow_exist(self, app_model: App, *, session: Session) -> bool: stmt = select( exists().where( Workflow.tenant_id == app_model.tenant_id, @@ -186,23 +186,21 @@ class WorkflowService: Workflow.version == Workflow.VERSION_DRAFT, ) ) - return db.session.execute(stmt).scalar_one() + return session.execute(stmt).scalar_one() def get_draft_workflow( - self, app_model: App, workflow_id: str | None = None, session: Session | scoped_session | None = None + self, app_model: App, workflow_id: str | None = None, *, session: Session ) -> Workflow | None: """ Get draft workflow - When ``session`` is provided, reuse it so callers that already hold a - Session avoid checking out an extra request-scoped ``db.session`` - connection. Falls back to ``db.session`` for backward compatibility. + Reuses the caller's active session so workflow reads stay in the same + transaction as the surrounding request or task. """ if workflow_id: return self.get_published_workflow_by_id(app_model, workflow_id, session=session) # fetch draft workflow by app_model - bind = session if session is not None else db.session - workflow = bind.scalar( + workflow = session.scalar( select(Workflow) .where( Workflow.tenant_id == app_model.tenant_id, @@ -215,18 +213,14 @@ class WorkflowService: # return draft workflow return workflow - def get_published_workflow_by_id( - self, app_model: App, workflow_id: str, session: Session | scoped_session | None = None - ) -> Workflow | None: + def get_published_workflow_by_id(self, app_model: App, workflow_id: str, *, session: Session) -> Workflow | None: """ fetch published workflow by workflow_id - When ``session`` is provided, reuse it so callers that already hold a - Session avoid checking out an extra request-scoped ``db.session`` - connection. Falls back to ``db.session`` for backward compatibility. + Reuses the caller's active session so workflow reads stay in the same + transaction as the surrounding request or task. """ - bind = session if session is not None else db.session - workflow = bind.scalar( + workflow = session.scalar( select(Workflow) .where( Workflow.tenant_id == app_model.tenant_id, @@ -244,20 +238,18 @@ class WorkflowService: ) return workflow - def get_published_workflow(self, app_model: App, session: Session | None = None) -> Workflow | None: + def get_published_workflow(self, app_model: App, *, session: Session) -> Workflow | None: """ Get published workflow - When ``session`` is provided, reuse it so callers that already hold a - Session avoid checking out an extra request-scoped ``db.session`` - connection. Falls back to ``db.session`` for backward compatibility. + Reuses the caller's active session so workflow reads stay in the same + transaction as the surrounding request or task. """ if not app_model.workflow_id: return None - bind = session if session is not None else db.session - workflow = bind.scalar( + workflow = session.scalar( select(Workflow) .where( Workflow.tenant_id == app_model.tenant_id, @@ -269,7 +261,7 @@ class WorkflowService: return workflow - def get_accessible_app_ids(self, app_ids: Sequence[str], tenant_id: str) -> set[str]: + def get_accessible_app_ids(self, app_ids: Sequence[str], tenant_id: str, *, session: Session) -> set[str]: """ Return app IDs that belong to the given tenant. """ @@ -277,7 +269,7 @@ class WorkflowService: return set() stmt = select(App.id).where(App.id.in_(app_ids), App.tenant_id == tenant_id) - return {str(app_id) for app_id in db.session.scalars(stmt).all()} + return {str(app_id) for app_id in session.scalars(stmt).all()} def get_all_published_workflow( self, @@ -327,13 +319,14 @@ class WorkflowService: account: Account, environment_variables: Sequence[VariableBase], conversation_variables: Sequence[VariableBase], + session: Session, ) -> Workflow: """ Sync draft workflow :raises WorkflowHashNotEqualError """ # fetch draft workflow by app_model - workflow = self.get_draft_workflow(app_model=app_model) + workflow = self.get_draft_workflow(app_model=app_model, session=session) if workflow and workflow.unique_hash != unique_hash: raise WorkflowHashNotEqualError() @@ -357,7 +350,7 @@ class WorkflowService: environment_variables=environment_variables, conversation_variables=conversation_variables, ) - db.session.add(workflow) + session.add(workflow) # update draft workflow if found else: workflow.graph = json.dumps(graph) @@ -369,19 +362,19 @@ class WorkflowService: from services.agent.workflow_publish_service import WorkflowAgentPublishService - db.session.flush() + session.flush() WorkflowAgentPublishService.sync_agent_bindings_for_draft( - session=cast(Session, db.session), + session=session, draft_workflow=workflow, account_id=account.id, ) WorkflowAgentPublishService.validate_agent_nodes_for_draft_sync( - session=cast(Session, db.session), + session=session, draft_workflow=workflow, ) # commit db session changes - db.session.commit() + session.commit() # trigger app workflow events app_draft_workflow_was_synced.send(app_model, synced_draft_workflow=workflow) @@ -395,12 +388,13 @@ class WorkflowService: app_model: App, environment_variables: Sequence[VariableBase], account: Account, + session: Session, ): """ Update draft workflow environment variables """ # fetch draft workflow by app_model - workflow = self.get_draft_workflow(app_model=app_model) + workflow = self.get_draft_workflow(app_model=app_model, session=session) if not workflow: raise ValueError("No draft workflow found.") @@ -410,7 +404,7 @@ class WorkflowService: workflow.updated_at = naive_utc_now() # commit db session changes - db.session.commit() + session.commit() def update_draft_workflow_conversation_variables( self, @@ -418,12 +412,13 @@ class WorkflowService: app_model: App, conversation_variables: Sequence[VariableBase], account: Account, + session: Session, ): """ Update draft workflow conversation variables """ # fetch draft workflow by app_model - workflow = self.get_draft_workflow(app_model=app_model) + workflow = self.get_draft_workflow(app_model=app_model, session=session) if not workflow: raise ValueError("No draft workflow found.") @@ -433,7 +428,7 @@ class WorkflowService: workflow.updated_at = naive_utc_now() # commit db session changes - db.session.commit() + session.commit() def update_draft_workflow_features( self, @@ -441,12 +436,13 @@ class WorkflowService: app_model: App, features: dict, account: Account, + session: Session, ): """ Update draft workflow features """ # fetch draft workflow by app_model - workflow = self.get_draft_workflow(app_model=app_model) + workflow = self.get_draft_workflow(app_model=app_model, session=session) if not workflow: raise ValueError("No draft workflow found.") @@ -459,7 +455,7 @@ class WorkflowService: workflow.updated_at = naive_utc_now() # commit db session changes - db.session.commit() + session.commit() def restore_published_workflow_to_draft( self, @@ -467,20 +463,23 @@ class WorkflowService: app_model: App, workflow_id: str, account: Account, + session: Session, ) -> Workflow: """Restore a published workflow snapshot into the draft workflow. Secret environment variables are copied server-side from the selected published workflow so the normal draft sync flow stays stateless. """ - source_workflow = self.get_published_workflow_by_id(app_model=app_model, workflow_id=workflow_id) + source_workflow = self.get_published_workflow_by_id( + app_model=app_model, workflow_id=workflow_id, session=session + ) if not source_workflow: raise WorkflowNotFoundError("Workflow not found.") self.validate_features_structure(app_model=app_model, features=source_workflow.normalized_features_dict) self.validate_graph_structure(graph=source_workflow.graph_dict) - draft_workflow = self.get_draft_workflow(app_model=app_model) + draft_workflow = self.get_draft_workflow(app_model=app_model, session=session) draft_workflow, is_new_draft = apply_published_workflow_snapshot_to_draft( tenant_id=app_model.tenant_id, app_id=app_model.id, @@ -491,9 +490,9 @@ class WorkflowService: ) if is_new_draft: - db.session.add(draft_workflow) + session.add(draft_workflow) - db.session.commit() + session.commit() app_draft_workflow_was_synced.send(app_model, synced_draft_workflow=draft_workflow) return draft_workflow @@ -520,7 +519,7 @@ class WorkflowService: from services.feature_service import FeatureService if FeatureService.get_system_features().plugin_manager.enabled: - self._validate_workflow_credentials(draft_workflow) + self._validate_workflow_credentials(draft_workflow, session=session) # validate graph structure self.validate_graph_structure(graph=draft_workflow.graph_dict) @@ -577,7 +576,7 @@ class WorkflowService: # return new workflow return workflow - def _validate_workflow_credentials(self, workflow: Workflow) -> None: + def _validate_workflow_credentials(self, workflow: Workflow, *, session: Session) -> None: """ Validate all credentials in workflow nodes before publishing. @@ -609,7 +608,7 @@ class WorkflowService: ) else: # Check default workspace credential for this provider - self._check_default_tool_credential(workflow.tenant_id, provider) + self._check_default_tool_credential(workflow.tenant_id, provider, session=session) elif node_type == "agent": agent_params = node_data.get("agent_parameters", {}) @@ -622,7 +621,9 @@ class WorkflowService: # Validate load balancing credentials for agent model if load balancing is enabled agent_model_node_data = {"model": model_config} - self._validate_load_balancing_credentials(workflow, agent_model_node_data, node_id) + self._validate_load_balancing_credentials( + workflow, agent_model_node_data, node_id, session=session + ) # Validate agent tools tools = agent_params.get("tools", {}).get("value", []) @@ -636,7 +637,7 @@ class WorkflowService: check_credential_policy_compliance(credential_id, provider, PluginCredentialType.TOOL) else: - self._check_default_tool_credential(workflow.tenant_id, provider) + self._check_default_tool_credential(workflow.tenant_id, provider, session=session) elif node_type in ["llm", "knowledge_retrieval", "parameter_extractor", "question_classifier"]: model_config = node_data.get("model", {}) @@ -647,7 +648,7 @@ class WorkflowService: # Validate that the provider+model combination can fetch valid credentials self._validate_llm_model_config(workflow.tenant_id, provider, model_name) # Validate load balancing credentials if load balancing is enabled - self._validate_load_balancing_credentials(workflow, node_data, node_id) + self._validate_load_balancing_credentials(workflow, node_data, node_id, session=session) else: raise ValueError(f"Node {node_id} ({node_type}): Missing provider or model configuration") @@ -710,7 +711,7 @@ class WorkflowService: f"Failed to validate LLM model configuration (provider: {provider}, model: {model_name}): {str(e)}" ) - def _check_default_tool_credential(self, tenant_id: str, provider: str) -> None: + def _check_default_tool_credential(self, tenant_id: str, provider: str, *, session: Session) -> None: """ Check credential policy compliance for the default workspace credential of a tool provider. @@ -726,7 +727,7 @@ class WorkflowService: # Use the same fallback logic as runtime: get the first available credential # ordered by is_default DESC, created_at ASC (same as tool_manager.py) - default_provider = db.session.scalar( + default_provider = session.scalar( select(BuiltinToolProvider) .where( BuiltinToolProvider.tenant_id == tenant_id, @@ -753,7 +754,9 @@ class WorkflowService: except Exception as e: raise ValueError(f"Failed to validate default credential for tool provider {provider}: {str(e)}") - def _validate_load_balancing_credentials(self, workflow: Workflow, node_data: dict[str, Any], node_id: str) -> None: + def _validate_load_balancing_credentials( + self, workflow: Workflow, node_data: dict[str, Any], node_id: str, *, session: Session + ) -> None: """ Validate load balancing credentials for a workflow node. @@ -773,7 +776,9 @@ class WorkflowService: # Check if this model has load balancing enabled if self._is_load_balancing_enabled(workflow.tenant_id, provider, model_name): # Get all load balancing configurations for this model - load_balancing_configs = self._get_load_balancing_configs(workflow.tenant_id, provider, model_name) + load_balancing_configs = self._get_load_balancing_configs( + workflow.tenant_id, provider, model_name, session=session + ) # Validate each load balancing configuration try: for config in load_balancing_configs: @@ -817,7 +822,9 @@ class WorkflowService: # If we can't determine the status, assume load balancing is not enabled return False - def _get_load_balancing_configs(self, tenant_id: str, provider: str, model_name: str) -> list[dict[str, Any]]: + def _get_load_balancing_configs( + self, tenant_id: str, provider: str, model_name: str, *, session: Session + ) -> list[dict[str, Any]]: """ Get all load balancing configurations for a model. @@ -835,11 +842,17 @@ class WorkflowService: provider=provider, model=model_name, model_type="llm", # Load balancing is primarily used for LLM models + session=session, config_from="predefined-model", # Check both predefined and custom models ) _, custom_configs = model_load_balancing_service.get_load_balancing_configs( - tenant_id=tenant_id, provider=provider, model=model_name, model_type="llm", config_from="custom-model" + tenant_id=tenant_id, + provider=provider, + model=model_name, + model_type="llm", + session=session, + config_from="custom-model", ) all_configs = cast(list[dict[str, Any]], configs) + cast(list[dict[str, Any]], custom_configs) @@ -1047,6 +1060,7 @@ class WorkflowService: account: Account, node_id: str, inputs: Mapping[str, Any] | None = None, + session: Session, ) -> Mapping[str, Any]: """ Build a human input form preview for a draft workflow. @@ -1057,7 +1071,7 @@ class WorkflowService: node_id: Human input node ID. inputs: Values used to fill missing upstream variables referenced in form_content. """ - draft_workflow = self.get_draft_workflow(app_model=app_model) + draft_workflow = self.get_draft_workflow(app_model=app_model, session=session) if not draft_workflow: raise ValueError("Workflow not initialized") @@ -1104,6 +1118,7 @@ class WorkflowService: form_inputs: Mapping[str, Any], inputs: Mapping[str, Any] | None = None, action: str, + session: Session, ) -> Mapping[str, Any]: """ Submit a human input form preview for a draft workflow. @@ -1116,7 +1131,7 @@ class WorkflowService: inputs: Values used to fill missing upstream variables referenced in form_content. action: Selected action ID. """ - draft_workflow = self.get_draft_workflow(app_model=app_model) + draft_workflow = self.get_draft_workflow(app_model=app_model, session=session) if not draft_workflow: raise ValueError("Workflow not initialized") @@ -1189,8 +1204,9 @@ class WorkflowService: node_id: str, delivery_method_id: str, inputs: Mapping[str, Any] | None = None, + session: Session, ) -> None: - draft_workflow = self.get_draft_workflow(app_model=app_model) + draft_workflow = self.get_draft_workflow(app_model=app_model, session=session) if not draft_workflow: raise ValueError("Workflow not initialized") @@ -1529,7 +1545,7 @@ class WorkflowService: node_execution.status = WorkflowNodeExecutionStatus.FAILED node_execution.error = error - def convert_to_workflow(self, app_model: App, account: Account, args: dict[str, Any]) -> App: + def convert_to_workflow(self, app_model: App, account: Account, args: dict[str, Any], *, session: Session) -> App: """ Basic mode of chatbot app(expert mode) to workflow Completion App to Workflow App @@ -1553,6 +1569,7 @@ class WorkflowService: icon_type=args.get("icon_type", "emoji"), icon=args.get("icon", "🤖"), icon_background=args.get("icon_background", "#FFEAD5"), + session=session, ) return new_app diff --git a/api/services/workspace_service.py b/api/services/workspace_service.py index 180c077b88a..30853b2cc99 100644 --- a/api/services/workspace_service.py +++ b/api/services/workspace_service.py @@ -1,9 +1,9 @@ from flask_login import current_user from sqlalchemy import select +from sqlalchemy.orm import Session from configs import dify_config from enums.cloud_plan import CloudPlan -from extensions.ext_database import db from models.account import Tenant, TenantAccountJoin, TenantAccountRole from services.account_service import TenantService from services.feature_service import FeatureService @@ -11,7 +11,7 @@ from services.feature_service import FeatureService class WorkspaceService: @classmethod - def get_tenant_info(cls, tenant: Tenant): + def get_tenant_info(cls, tenant: Tenant, session: Session): if not tenant: return None tenant_info: dict[str, object] = { @@ -25,7 +25,7 @@ class WorkspaceService: } # Get role of user - tenant_account_join = db.session.scalar( + tenant_account_join = session.scalar( select(TenantAccountJoin) .where(TenantAccountJoin.tenant_id == tenant.id, TenantAccountJoin.account_id == current_user.id) .limit(1) @@ -37,7 +37,7 @@ class WorkspaceService: can_replace_logo = feature.can_replace_logo if can_replace_logo and TenantService.has_roles( - tenant, [TenantAccountRole.OWNER, TenantAccountRole.ADMIN], session=db.session + tenant, [TenantAccountRole.OWNER, TenantAccountRole.ADMIN], session=session ): base_url = dify_config.FILES_URL replace_webapp_logo = ( @@ -56,7 +56,7 @@ class WorkspaceService: from services.credit_pool_service import CreditPoolService - paid_pool = CreditPoolService.get_pool(tenant_id=tenant.id, pool_type="paid") + paid_pool = CreditPoolService.get_pool(tenant_id=tenant.id, pool_type="paid", session=session) # if the tenant is not on the sandbox plan and the paid pool is not full, use the paid pool if ( feature.billing.subscription.plan != CloudPlan.SANDBOX @@ -66,7 +66,7 @@ class WorkspaceService: tenant_info["trial_credits"] = paid_pool.quota_limit tenant_info["trial_credits_used"] = paid_pool.quota_used else: - trial_pool = CreditPoolService.get_pool(tenant_id=tenant.id, pool_type="trial") + trial_pool = CreditPoolService.get_pool(tenant_id=tenant.id, pool_type="trial", session=session) if trial_pool: tenant_info["trial_credits"] = trial_pool.quota_limit tenant_info["trial_credits_used"] = trial_pool.quota_used diff --git a/api/tasks/batch_create_segment_to_index_task.py b/api/tasks/batch_create_segment_to_index_task.py index 9f19b03544d..0f92af21dc8 100644 --- a/api/tasks/batch_create_segment_to_index_task.py +++ b/api/tasks/batch_create_segment_to_index_task.py @@ -177,7 +177,7 @@ def batch_create_segment_to_index_task( with session_factory.create_session() as session: dataset = session.get(Dataset, dataset_id) if dataset: - VectorService.create_segments_vector(None, document_segments, dataset, document_config["doc_form"]) + VectorService.create_segments_vector(None, document_segments, dataset, document_config["doc_form"], session) redis_client.setex(indexing_cache_key, 600, "completed") end_at = time.perf_counter() diff --git a/api/tasks/regenerate_summary_index_task.py b/api/tasks/regenerate_summary_index_task.py index 16b59fdbba8..5cb8d4281f0 100644 --- a/api/tasks/regenerate_summary_index_task.py +++ b/api/tasks/regenerate_summary_index_task.py @@ -259,9 +259,8 @@ def regenerate_summary_index_task( # Regenerate both summary content and vectors (for summary_model change) SummaryIndexService.generate_and_vectorize_summary( - segment, dataset, summary_index_setting + segment, dataset, summary_index_setting, session=session ) - session.commit() total_segments_processed += 1 except Exception as e: diff --git a/api/tasks/retry_document_indexing_task.py b/api/tasks/retry_document_indexing_task.py index fa02afda15f..dddb7715d22 100644 --- a/api/tasks/retry_document_indexing_task.py +++ b/api/tasks/retry_document_indexing_task.py @@ -101,8 +101,9 @@ def retry_document_indexing_task(dataset_id: str, document_ids: list[str], user_ session.commit() if dataset.runtime_mode == "rag_pipeline": - rag_pipeline_service = RagPipelineService() - rag_pipeline_service.retry_error_document(dataset, document, user) + with session_factory.create_session() as rag_session: + rag_pipeline_service = RagPipelineService(rag_session) + rag_pipeline_service.retry_error_document(dataset, document, user) else: indexing_runner = IndexingRunner() indexing_runner.run([document]) diff --git a/api/tasks/workflow_schedule_tasks.py b/api/tasks/workflow_schedule_tasks.py index 76386520000..38737f96e78 100644 --- a/api/tasks/workflow_schedule_tasks.py +++ b/api/tasks/workflow_schedule_tasks.py @@ -39,7 +39,7 @@ def run_schedule_trigger(schedule_id: str) -> None: if not schedule: raise ScheduleNotFoundError(f"Schedule {schedule_id} not found") - tenant_owner = ScheduleService.get_tenant_owner(session, schedule.tenant_id) + tenant_owner = ScheduleService.get_tenant_owner(schedule.tenant_id, session=session) if not tenant_owner: raise TenantOwnerNotFoundError(f"No owner or admin found for tenant {schedule.tenant_id}") diff --git a/api/tests/integration_tests/conftest.py b/api/tests/integration_tests/conftest.py index ea875e63fe8..25ee1974e92 100644 --- a/api/tests/integration_tests/conftest.py +++ b/api/tests/integration_tests/conftest.py @@ -84,7 +84,7 @@ def setup_account(request) -> Generator[Account, None, None]: password=secrets.token_hex(16), ip_address="localhost", language="en-US", - session=db.session, + session=db.session(), ) with _CACHED_APP.test_request_context(): diff --git a/api/tests/integration_tests/services/plugin/test_plugin_lifecycle.py b/api/tests/integration_tests/services/plugin/test_plugin_lifecycle.py index 23cbdd24b91..e6e681b426f 100644 --- a/api/tests/integration_tests/services/plugin/test_plugin_lifecycle.py +++ b/api/tests/integration_tests/services/plugin/test_plugin_lifecycle.py @@ -2,6 +2,7 @@ import pytest from sqlalchemy import delete, func, select from core.db.session_factory import session_factory +from extensions.ext_database import db from models import Tenant from models.account import ( TenantPluginAutoUpgradeCategory, @@ -39,17 +40,18 @@ def tenant(flask_req_ctx): class TestPluginPermissionLifecycle: def test_get_returns_none_for_new_tenant(self, tenant): - assert PluginPermissionService.get_permission(tenant) is None + assert PluginPermissionService.get_permission(tenant, session=db.session()) is None def test_change_creates_row(self, tenant): result = PluginPermissionService.change_permission( tenant, TenantPluginInstallPermission.ADMINS, TenantPluginDebugPermission.EVERYONE, + session=db.session, ) assert result is True - perm = PluginPermissionService.get_permission(tenant) + perm = PluginPermissionService.get_permission(tenant, session=db.session()) assert perm is not None assert perm.install_permission == TenantPluginInstallPermission.ADMINS assert perm.debug_permission == TenantPluginDebugPermission.EVERYONE @@ -59,13 +61,15 @@ class TestPluginPermissionLifecycle: tenant, TenantPluginInstallPermission.ADMINS, TenantPluginDebugPermission.NOBODY, + session=db.session, ) PluginPermissionService.change_permission( tenant, TenantPluginInstallPermission.EVERYONE, TenantPluginDebugPermission.ADMINS, + session=db.session, ) - perm = PluginPermissionService.get_permission(tenant) + perm = PluginPermissionService.get_permission(tenant, session=db.session()) assert perm is not None assert perm.install_permission == TenantPluginInstallPermission.EVERYONE assert perm.debug_permission == TenantPluginDebugPermission.ADMINS @@ -81,7 +85,7 @@ class TestPluginPermissionLifecycle: class TestPluginAutoUpgradeLifecycle: def test_get_returns_none_for_new_tenant(self, tenant): - assert PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY) is None + assert PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY, session=db.session()) is None def test_change_creates_row(self, tenant): result = PluginAutoUpgradeService.change_strategy( @@ -92,10 +96,11 @@ class TestPluginAutoUpgradeLifecycle: exclude_plugins=[], include_plugins=[], category=PLUGIN_CATEGORY, + session=db.session(), ) assert result is True - strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY) + strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY, session=db.session()) assert strategy is not None assert strategy.strategy_setting == TenantPluginAutoUpgradeStrategySetting.LATEST assert strategy.upgrade_time_of_day == 3 @@ -109,6 +114,7 @@ class TestPluginAutoUpgradeLifecycle: exclude_plugins=[], include_plugins=[], category=PLUGIN_CATEGORY, + session=db.session(), ) PluginAutoUpgradeService.change_strategy( tenant, @@ -118,9 +124,10 @@ class TestPluginAutoUpgradeLifecycle: exclude_plugins=[], include_plugins=["plugin-a"], category=PLUGIN_CATEGORY, + session=db.session(), ) - strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY) + strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY, session=db.session()) assert strategy is not None assert strategy.strategy_setting == TenantPluginAutoUpgradeStrategySetting.LATEST assert strategy.upgrade_time_of_day == 12 @@ -128,9 +135,9 @@ class TestPluginAutoUpgradeLifecycle: assert strategy.include_plugins == ["plugin-a"] def test_exclude_plugin_creates_strategy_when_none_exists(self, tenant): - PluginAutoUpgradeService.exclude_plugin(tenant, "my-plugin", PLUGIN_CATEGORY) + PluginAutoUpgradeService.exclude_plugin(tenant, "my-plugin", PLUGIN_CATEGORY, session=db.session()) - strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY) + strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY, session=db.session()) assert strategy is not None assert strategy.upgrade_mode == TenantPluginAutoUpgradeMode.EXCLUDE assert "my-plugin" in strategy.exclude_plugins @@ -144,10 +151,11 @@ class TestPluginAutoUpgradeLifecycle: exclude_plugins=["existing"], include_plugins=[], category=PLUGIN_CATEGORY, + session=db.session(), ) - PluginAutoUpgradeService.exclude_plugin(tenant, "new-plugin", PLUGIN_CATEGORY) + PluginAutoUpgradeService.exclude_plugin(tenant, "new-plugin", PLUGIN_CATEGORY, session=db.session()) - strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY) + strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY, session=db.session()) assert strategy is not None assert "existing" in strategy.exclude_plugins assert "new-plugin" in strategy.exclude_plugins @@ -161,10 +169,11 @@ class TestPluginAutoUpgradeLifecycle: exclude_plugins=["same-plugin"], include_plugins=[], category=PLUGIN_CATEGORY, + session=db.session(), ) - PluginAutoUpgradeService.exclude_plugin(tenant, "same-plugin", PLUGIN_CATEGORY) + PluginAutoUpgradeService.exclude_plugin(tenant, "same-plugin", PLUGIN_CATEGORY, session=db.session()) - strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY) + strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY, session=db.session()) assert strategy is not None assert strategy.exclude_plugins.count("same-plugin") == 1 @@ -177,10 +186,11 @@ class TestPluginAutoUpgradeLifecycle: exclude_plugins=[], include_plugins=["p1", "p2"], category=PLUGIN_CATEGORY, + session=db.session(), ) - PluginAutoUpgradeService.exclude_plugin(tenant, "p1", PLUGIN_CATEGORY) + PluginAutoUpgradeService.exclude_plugin(tenant, "p1", PLUGIN_CATEGORY, session=db.session()) - strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY) + strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY, session=db.session()) assert strategy is not None assert "p1" not in strategy.include_plugins assert "p2" in strategy.include_plugins @@ -194,10 +204,11 @@ class TestPluginAutoUpgradeLifecycle: exclude_plugins=[], include_plugins=[], category=PLUGIN_CATEGORY, + session=db.session(), ) - PluginAutoUpgradeService.exclude_plugin(tenant, "excluded-plugin", PLUGIN_CATEGORY) + PluginAutoUpgradeService.exclude_plugin(tenant, "excluded-plugin", PLUGIN_CATEGORY, session=db.session()) - strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY) + strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY, session=db.session()) assert strategy is not None assert strategy.upgrade_mode == TenantPluginAutoUpgradeMode.EXCLUDE assert "excluded-plugin" in strategy.exclude_plugins diff --git a/api/tests/integration_tests/services/test_node_output_inspector_service.py b/api/tests/integration_tests/services/test_node_output_inspector_service.py index 5a8c07e0434..c2253a20c15 100644 --- a/api/tests/integration_tests/services/test_node_output_inspector_service.py +++ b/api/tests/integration_tests/services/test_node_output_inspector_service.py @@ -219,6 +219,36 @@ def _stub_resolver(declared_outputs_payload: list[dict[str, Any]]): return _Resolver() +def _snapshot_workflow_run(service: NodeOutputInspectorService, *, app_model: Any, workflow_run_id: str): + with session_factory.create_session() as session: + return service.snapshot_workflow_run(app_model=app_model, workflow_run_id=workflow_run_id, session=session) + + +def _node_detail(service: NodeOutputInspectorService, *, app_model: Any, workflow_run_id: str, node_id: str): + with session_factory.create_session() as session: + return service.node_detail( + app_model=app_model, workflow_run_id=workflow_run_id, node_id=node_id, session=session + ) + + +def _output_preview( + service: NodeOutputInspectorService, + *, + app_model: Any, + workflow_run_id: str, + node_id: str, + output_name: str, +): + with session_factory.create_session() as session: + return service.output_preview( + app_model=app_model, + workflow_run_id=workflow_run_id, + node_id=node_id, + output_name=output_name, + session=session, + ) + + # ────────────────────────────────────────────────────────────────────────────── # Tests # ────────────────────────────────────────────────────────────────────────────── @@ -229,7 +259,8 @@ def test_snapshot_returns_agent_v2_declared_outputs_with_status_ready(seeded_run real ``WorkflowRun`` + ``WorkflowNodeExecutionModel`` rows.""" app_model, workflow_run, _ = seeded_run service = NodeOutputInspectorService(binding_resolver=_stub_resolver([{"name": "text", "type": "string"}])) - snapshot = service.snapshot_workflow_run( + snapshot = _snapshot_workflow_run( + service, app_model=app_model, workflow_run_id=workflow_run.id, ) @@ -256,7 +287,7 @@ def test_snapshot_404s_for_missing_run(fake_app_model): """Service raises ``workflow_run_not_found`` when the row doesn't exist.""" service = NodeOutputInspectorService(binding_resolver=_stub_resolver([])) with pytest.raises(NodeOutputInspectorError) as exc: - service.snapshot_workflow_run(app_model=fake_app_model, workflow_run_id=str(uuid.uuid4())) + _snapshot_workflow_run(service, app_model=fake_app_model, workflow_run_id=str(uuid.uuid4())) assert exc.value.code == "workflow_run_not_found" @@ -266,7 +297,7 @@ def test_snapshot_404s_for_cross_tenant_access(seeded_run): intruder = SimpleNamespace(id=str(uuid.uuid4()), tenant_id=str(uuid.uuid4())) service = NodeOutputInspectorService(binding_resolver=_stub_resolver([])) with pytest.raises(NodeOutputInspectorError) as exc: - service.snapshot_workflow_run(app_model=intruder, workflow_run_id=workflow_run.id) + _snapshot_workflow_run(service, app_model=intruder, workflow_run_id=workflow_run.id) assert exc.value.code == "workflow_run_not_found" @@ -286,7 +317,7 @@ def test_snapshot_404s_for_published_run_per_decision_d1(flask_req_ctx, fake_app try: service = NodeOutputInspectorService(binding_resolver=_stub_resolver([])) with pytest.raises(NodeOutputInspectorError) as exc: - service.snapshot_workflow_run(app_model=fake_app_model, workflow_run_id=run_id) + _snapshot_workflow_run(service, app_model=fake_app_model, workflow_run_id=run_id) assert exc.value.code == "published_run_inspector_not_implemented" finally: with session_factory.create_session() as session: @@ -328,7 +359,7 @@ def test_snapshot_surfaces_type_check_failure_from_metadata(flask_req_ctx, fake_ try: service = NodeOutputInspectorService(binding_resolver=_stub_resolver([{"name": "summary", "type": "string"}])) - snapshot = service.snapshot_workflow_run(app_model=fake_app_model, workflow_run_id=run_id) + snapshot = _snapshot_workflow_run(service, app_model=fake_app_model, workflow_run_id=run_id) output = snapshot.node_outputs[0].outputs[0] assert output.status == NodeOutputStatus.TYPE_CHECK_FAILED assert output.type_check is not None @@ -375,7 +406,7 @@ def test_snapshot_surfaces_output_check_failure_from_metadata(flask_req_ctx, fak "services.workflow.node_output_inspector_service.file_helpers.get_signed_file_url", return_value="https://signed.example/report", ): - snapshot = service.snapshot_workflow_run(app_model=fake_app_model, workflow_run_id=run_id) + snapshot = _snapshot_workflow_run(service, app_model=fake_app_model, workflow_run_id=run_id) output = snapshot.node_outputs[0].outputs[0] assert output.status == NodeOutputStatus.OUTPUT_CHECK_FAILED assert output.output_check is not None @@ -391,7 +422,8 @@ def test_snapshot_surfaces_output_check_failure_from_metadata(flask_req_ctx, fak def test_node_detail_serves_one_node(seeded_run): app_model, workflow_run, _ = seeded_run service = NodeOutputInspectorService(binding_resolver=_stub_resolver([{"name": "text", "type": "string"}])) - view = service.node_detail( + view = _node_detail( + service, app_model=app_model, workflow_run_id=workflow_run.id, node_id="agent-node-1", @@ -421,7 +453,8 @@ def test_output_preview_for_file_renders_signed_url(seeded_run, fake_app_model): "services.workflow.node_output_inspector_service.file_helpers.get_signed_file_url", return_value="https://signed.example/x.pdf", ): - preview = service.output_preview( + preview = _output_preview( + service, app_model=fake_app_model, workflow_run_id=workflow_run.id, node_id="agent-node-1", @@ -466,7 +499,7 @@ def test_keeps_latest_execution_per_node_by_index(flask_req_ctx, fake_app_model) try: service = NodeOutputInspectorService(binding_resolver=_stub_resolver([{"name": "text", "type": "string"}])) - snapshot = service.snapshot_workflow_run(app_model=fake_app_model, workflow_run_id=run_id) + snapshot = _snapshot_workflow_run(service, app_model=fake_app_model, workflow_run_id=run_id) assert snapshot.node_outputs[0].outputs[0].value_preview == "second attempt" finally: with session_factory.create_session() as session: diff --git a/api/tests/test_containers_integration_tests/controllers/console/app/test_app_apis.py b/api/tests/test_containers_integration_tests/controllers/console/app/test_app_apis.py index ae37d305670..df9d655fbbc 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/app/test_app_apis.py +++ b/api/tests/test_containers_integration_tests/controllers/console/app/test_app_apis.py @@ -500,7 +500,11 @@ class TestWorkflowDraftVariableEndpoints: api = workflow_draft_variable_module.WorkflowVariableCollectionApi() method = unwrap(api.get) - monkeypatch.setattr(workflow_draft_variable_module, "db", SimpleNamespace(engine=MagicMock())) + monkeypatch.setattr( + workflow_draft_variable_module, + "db", + SimpleNamespace(engine=MagicMock(), session=MagicMock()), + ) class DummySessionCtx: def __enter__(self): diff --git a/api/tests/test_containers_integration_tests/controllers/console/auth/test_data_source_bearer_auth.py b/api/tests/test_containers_integration_tests/controllers/console/auth/test_data_source_bearer_auth.py index e55b46d38bf..ef8c0add709 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/auth/test_data_source_bearer_auth.py +++ b/api/tests/test_containers_integration_tests/controllers/console/auth/test_data_source_bearer_auth.py @@ -85,7 +85,7 @@ def test_create_binding_successful( assert response.status_code == 200 assert response.get_json() == {"result": "success"} - create_auth.assert_called_once_with(ANY, tenant_id, payload) + create_auth.assert_called_once_with(tenant_id, payload, session=ANY) def test_create_binding_failure( diff --git a/api/tests/test_containers_integration_tests/controllers/console/auth/test_email_register.py b/api/tests/test_containers_integration_tests/controllers/console/auth/test_email_register.py index 109332e16c9..d893e9e6efb 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/auth/test_email_register.py +++ b/api/tests/test_containers_integration_tests/controllers/console/auth/test_email_register.py @@ -270,7 +270,7 @@ def test_get_account_by_email_with_case_fallback_falls_back_to_lowercase(): second_result.scalar_one_or_none.return_value = expected_account mock_session.execute.side_effect = [first_result, second_result] - result = AccountService.get_account_by_email_with_case_fallback(mock_session, "Case@Test.com") + result = AccountService.get_account_by_email_with_case_fallback("Case@Test.com", session=mock_session) assert result is expected_account assert mock_session.execute.call_count == 2 diff --git a/api/tests/test_containers_integration_tests/controllers/console/auth/test_forgot_password.py b/api/tests/test_containers_integration_tests/controllers/console/auth/test_forgot_password.py index 812aa299c1b..a7eba9d723c 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/auth/test_forgot_password.py +++ b/api/tests/test_containers_integration_tests/controllers/console/auth/test_forgot_password.py @@ -165,7 +165,7 @@ def test_get_account_by_email_with_case_fallback_falls_back_to_lowercase(): second_result.scalar_one_or_none.return_value = expected_account mock_session.execute.side_effect = [first_result, second_result] - result = AccountService.get_account_by_email_with_case_fallback(mock_session, "Mixed@Test.com") + result = AccountService.get_account_by_email_with_case_fallback("Mixed@Test.com", session=mock_session) assert result is expected_account assert mock_session.execute.call_count == 2 diff --git a/api/tests/test_containers_integration_tests/controllers/console/auth/test_oauth.py b/api/tests/test_containers_integration_tests/controllers/console/auth/test_oauth.py index 464e0134a2f..484ca71ca59 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/auth/test_oauth.py +++ b/api/tests/test_containers_integration_tests/controllers/console/auth/test_oauth.py @@ -494,7 +494,7 @@ class TestAccountGeneration: second_result.scalar_one_or_none.return_value = expected_account mock_session.execute.side_effect = [first_result, second_result] - result = AccountService.get_account_by_email_with_case_fallback(mock_session, "Case@Test.com") + result = AccountService.get_account_by_email_with_case_fallback("Case@Test.com", session=mock_session) assert result is expected_account assert mock_session.execute.call_count == 2 diff --git a/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py b/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py index c34810c97d0..6c73b0010ed 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py +++ b/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py @@ -5,7 +5,7 @@ from __future__ import annotations from collections.abc import Callable from inspect import unwrap from typing import cast -from unittest.mock import MagicMock, patch +from unittest.mock import ANY, MagicMock, patch from uuid import uuid4 import pytest @@ -90,55 +90,51 @@ class TestPipelineTemplateDetailApi: "graph": {"nodes": nodes, "edges": edges, "viewport": viewport}, } - service = MagicMock() - service.get_pipeline_template_detail.return_value = template - with ( app.test_request_context("/?type=built-in"), patch( - "controllers.console.datasets.rag_pipeline.rag_pipeline.RagPipelineService", - return_value=service, - ), + "controllers.console.datasets.rag_pipeline.rag_pipeline.RagPipelineService.get_pipeline_template_detail", + return_value=template, + ) as get_detail_mock, ): response, status = method(api, MagicMock(), "tpl-1") assert status == 200 assert response == {**template, "created_by": None} + get_detail_mock.assert_called_once_with("tpl-1", type="built-in", session=ANY) def test_get_returns_404_when_template_not_found(self, app: Flask) -> None: api = PipelineTemplateDetailApi() method = unwrap(api.get) - service = MagicMock() - service.get_pipeline_template_detail.return_value = None - with ( app.test_request_context("/?type=built-in"), patch( - "controllers.console.datasets.rag_pipeline.rag_pipeline.RagPipelineService", - return_value=service, - ), + "controllers.console.datasets.rag_pipeline.rag_pipeline.RagPipelineService.get_pipeline_template_detail", + return_value=None, + ) as get_detail_mock, ): with pytest.raises(NotFound): method(api, MagicMock(), "non-existent-id") + get_detail_mock.assert_called_once_with("non-existent-id", type="built-in", session=ANY) + def test_get_returns_404_for_customized_type_not_found(self, app: Flask) -> None: api = PipelineTemplateDetailApi() method = unwrap(api.get) - service = MagicMock() - service.get_pipeline_template_detail.return_value = None - with ( app.test_request_context("/?type=customized"), patch( - "controllers.console.datasets.rag_pipeline.rag_pipeline.RagPipelineService", - return_value=service, - ), + "controllers.console.datasets.rag_pipeline.rag_pipeline.RagPipelineService.get_pipeline_template_detail", + return_value=None, + ) as get_detail_mock, ): with pytest.raises(NotFound): method(api, MagicMock(), "non-existent-id") + get_detail_mock.assert_called_once_with("non-existent-id", type="customized", session=ANY) + class TestCustomizedPipelineTemplateApi: @pytest.fixture @@ -186,7 +182,7 @@ class TestCustomizedPipelineTemplateApi: ): response, status = method(api, tenant_id, "tpl-1") - delete_mock.assert_called_once_with("tpl-1", tenant_id) + delete_mock.assert_called_once_with("tpl-1", tenant_id, session=ANY) assert status == 204 assert response == "" diff --git a/api/tests/test_containers_integration_tests/controllers/console/test_api_based_extension.py b/api/tests/test_containers_integration_tests/controllers/console/test_api_based_extension.py index e60558040a5..4cca4c2170f 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/test_api_based_extension.py +++ b/api/tests/test_containers_integration_tests/controllers/console/test_api_based_extension.py @@ -97,13 +97,13 @@ def test_list_scopes_api_based_extensions_to_authenticated_tenant( assert account_create_response.status_code == 201 APIBasedExtensionService.save( - db_session_with_containers, APIBasedExtension( tenant_id=foreign_tenant_id, name="Foreign API", api_endpoint="https://foreign.example.com/hook", api_key="foreign-secret-12345", ), + session=db_session_with_containers, ) response = test_client_with_containers.get( diff --git a/api/tests/test_containers_integration_tests/controllers/openapi/test_account_sessions.py b/api/tests/test_containers_integration_tests/controllers/openapi/test_account_sessions.py index 4cdbec3e30e..4222f49a28a 100644 --- a/api/tests/test_containers_integration_tests/controllers/openapi/test_account_sessions.py +++ b/api/tests/test_containers_integration_tests/controllers/openapi/test_account_sessions.py @@ -30,7 +30,6 @@ def _mint_account_token( ) -> MintResult: """Mint a real, persisted ``dfoa_`` access token for ``account``.""" return mint_oauth_token( - db_session, redis_client, subject_email=account.email, subject_issuer=None, @@ -39,6 +38,7 @@ def _mint_account_token( device_label=device_label, prefix=PREFIX_OAUTH_ACCOUNT, ttl_days=14, + session=db_session, ) diff --git a/api/tests/test_containers_integration_tests/controllers/openapi/test_app_dsl.py b/api/tests/test_containers_integration_tests/controllers/openapi/test_app_dsl.py index 2b9feeede14..4d9bfb5ea17 100644 --- a/api/tests/test_containers_integration_tests/controllers/openapi/test_app_dsl.py +++ b/api/tests/test_containers_integration_tests/controllers/openapi/test_app_dsl.py @@ -96,7 +96,7 @@ def _app_and_account(db_session: Session, *, mode: str = "chat") -> tuple[App, A api_rph=100, api_rpm=10, ) - app_model = AppService().create_app(tenant.id, app_args, account) + app_model = AppService().create_app(tenant.id, app_args, account, session=db_session) return app_model, account diff --git a/api/tests/test_containers_integration_tests/controllers/openapi/test_app_run.py b/api/tests/test_containers_integration_tests/controllers/openapi/test_app_run.py index 8e9278ad244..df4f3873b1e 100644 --- a/api/tests/test_containers_integration_tests/controllers/openapi/test_app_run.py +++ b/api/tests/test_containers_integration_tests/controllers/openapi/test_app_run.py @@ -24,7 +24,7 @@ def _create_app(db_session: Session, account: Account, *, name: str = "Runner") icon="🤖", icon_background="#FF6B6B", ) - app_model = AppService().create_app(tenant.id, params, account) + app_model = AppService().create_app(tenant.id, params, account, session=db_session) db_session.commit() return app_model diff --git a/api/tests/test_containers_integration_tests/controllers/openapi/test_apps.py b/api/tests/test_containers_integration_tests/controllers/openapi/test_apps.py index 24580ae0a0e..ce1425d9e61 100644 --- a/api/tests/test_containers_integration_tests/controllers/openapi/test_apps.py +++ b/api/tests/test_containers_integration_tests/controllers/openapi/test_apps.py @@ -39,7 +39,7 @@ def _create_app( icon="🤖", icon_background="#FF6B6B", ) - app_model = AppService().create_app(tenant.id, params, account) + app_model = AppService().create_app(tenant.id, params, account, session=db_session) # The openapi surface gate keys off ``enable_api``; flip it explicitly so # the test states the visibility precondition rather than relying on the # template default. diff --git a/api/tests/test_containers_integration_tests/controllers/openapi/test_files.py b/api/tests/test_containers_integration_tests/controllers/openapi/test_files.py index 86cf70613c9..31a5485d2b3 100644 --- a/api/tests/test_containers_integration_tests/controllers/openapi/test_files.py +++ b/api/tests/test_containers_integration_tests/controllers/openapi/test_files.py @@ -25,7 +25,7 @@ def _create_app(db_session: Session, account: Account, *, name: str = "Uploader" icon="🤖", icon_background="#FF6B6B", ) - app_model = AppService().create_app(tenant.id, params, account) + app_model = AppService().create_app(tenant.id, params, account, session=db_session) db_session.commit() return app_model diff --git a/api/tests/test_containers_integration_tests/controllers/service_api/dataset/test_dataset.py b/api/tests/test_containers_integration_tests/controllers/service_api/dataset/test_dataset.py index 372157813cc..d670425be0c 100644 --- a/api/tests/test_containers_integration_tests/controllers/service_api/dataset/test_dataset.py +++ b/api/tests/test_containers_integration_tests/controllers/service_api/dataset/test_dataset.py @@ -734,7 +734,8 @@ class TestDatasetApiPatch: assert response["name"] == "Updated Dataset" assert response["partial_member_list"] == ["user-1"] mock_dataset_svc.update_dataset.assert_called_once() - session, _, update_data, _ = mock_dataset_svc.update_dataset.call_args.args + _, update_data, _ = mock_dataset_svc.update_dataset.call_args.args + session = mock_dataset_svc.update_dataset.call_args.kwargs["session"] assert isinstance(session, (Session, scoped_session)) assert update_data["name"] == "Updated Dataset" assert update_data["permission"] == "partial_members" @@ -1013,7 +1014,7 @@ class TestDatasetTagsApiGet: assert status == 200 assert response == [{"id": "tag-1", "name": "Test Tag", "type": "knowledge", "binding_count": "0"}] - mock_tag_svc.get_tags.assert_called_once_with(SessionMatcher(), "knowledge", "tenant-1") + mock_tag_svc.get_tags.assert_called_once_with("knowledge", "tenant-1", session=SessionMatcher()) @patch("controllers.service_api.dataset.dataset.current_user") def test_list_tags_from_db( diff --git a/api/tests/test_containers_integration_tests/controllers/web/test_web_forgot_password.py b/api/tests/test_containers_integration_tests/controllers/web/test_web_forgot_password.py index d568a1c0b04..cd754782df4 100644 --- a/api/tests/test_containers_integration_tests/controllers/web/test_web_forgot_password.py +++ b/api/tests/test_containers_integration_tests/controllers/web/test_web_forgot_password.py @@ -57,7 +57,7 @@ class TestForgotPasswordSendEmailApi: response = ForgotPasswordSendEmailApi().post() assert response == {"result": "success", "data": "token-123"} - mock_get_account.assert_called_once_with(ANY, "User@Example.com") + mock_get_account.assert_called_once_with("User@Example.com", session=ANY) mock_send_mail.assert_called_once_with(account=mock_account, email="user@example.com", language="zh-Hans") mock_extract_ip.assert_called_once() mock_rate_limit.assert_called_once_with("127.0.0.1") @@ -177,7 +177,7 @@ class TestForgotPasswordResetApi: response = ForgotPasswordResetApi().post() assert response == {"result": "success"} - mock_get_account.assert_called_once_with(ANY, "User@Example.com") + mock_get_account.assert_called_once_with("User@Example.com", session=ANY) mock_update_account.assert_called_once() mock_revoke_token.assert_called_once_with("token-123") diff --git a/api/tests/test_containers_integration_tests/controllers/web/test_wraps.py b/api/tests/test_containers_integration_tests/controllers/web/test_wraps.py index 3eab8ccbee5..aa85ac2ca7b 100644 --- a/api/tests/test_containers_integration_tests/controllers/web/test_wraps.py +++ b/api/tests/test_containers_integration_tests/controllers/web/test_wraps.py @@ -19,6 +19,8 @@ from controllers.web.wraps import ( decode_jwt_token, ) +pytestmark = pytest.mark.usefixtures("db_session_with_containers") + class TestValidateWebappToken: def test_enterprise_enabled_and_app_auth_requires_webapp_source(self) -> None: diff --git a/api/tests/test_containers_integration_tests/services/auth/test_api_key_auth_service.py b/api/tests/test_containers_integration_tests/services/auth/test_api_key_auth_service.py index e2f8c8fc703..e22aa102328 100644 --- a/api/tests/test_containers_integration_tests/services/auth/test_api_key_auth_service.py +++ b/api/tests/test_containers_integration_tests/services/auth/test_api_key_auth_service.py @@ -51,7 +51,7 @@ class TestApiKeyAuthService: self._create_binding(db_session_with_containers, tenant_id=tenant_id, category=category, provider=provider) db_session_with_containers.expire_all() - result = ApiKeyAuthService.get_provider_auth_list(db_session_with_containers, tenant_id) + result = ApiKeyAuthService.get_provider_auth_list(tenant_id, session=db_session_with_containers) assert len(result) >= 1 tenant_results = [r for r in result if r.tenant_id == tenant_id] @@ -61,7 +61,7 @@ class TestApiKeyAuthService: def test_get_provider_auth_list_empty( self, flask_app_with_containers: Flask, db_session_with_containers: Session, tenant_id ): - result = ApiKeyAuthService.get_provider_auth_list(db_session_with_containers, tenant_id) + result = ApiKeyAuthService.get_provider_auth_list(tenant_id, session=db_session_with_containers) tenant_results = [r for r in result if r.tenant_id == tenant_id] assert tenant_results == [] @@ -74,7 +74,7 @@ class TestApiKeyAuthService: ) db_session_with_containers.expire_all() - result = ApiKeyAuthService.get_provider_auth_list(db_session_with_containers, tenant_id) + result = ApiKeyAuthService.get_provider_auth_list(tenant_id, session=db_session_with_containers) tenant_results = [r for r in result if r.tenant_id == tenant_id] assert tenant_results == [] @@ -95,7 +95,7 @@ class TestApiKeyAuthService: mock_factory.return_value = mock_auth_instance mock_encrypter.encrypt_token.return_value = "encrypted_test_key_123" - ApiKeyAuthService.create_provider_auth(db_session_with_containers, tenant_id, mock_args) + ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=db_session_with_containers) mock_factory.assert_called_once() mock_auth_instance.validate_credentials.assert_called_once() @@ -118,7 +118,7 @@ class TestApiKeyAuthService: mock_auth_instance.validate_credentials.return_value = False mock_factory.return_value = mock_auth_instance - ApiKeyAuthService.create_provider_auth(db_session_with_containers, tenant_id, mock_args) + ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=db_session_with_containers) db_session_with_containers.expire_all() bindings = db_session_with_containers.query(DataSourceApiKeyAuthBinding).filter_by(tenant_id=tenant_id).all() @@ -142,7 +142,7 @@ class TestApiKeyAuthService: original_key = mock_args["credentials"]["config"]["api_key"] - ApiKeyAuthService.create_provider_auth(db_session_with_containers, tenant_id, mock_args) + ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=db_session_with_containers) assert mock_args["credentials"]["config"]["api_key"] == "encrypted_test_key_123" assert mock_args["credentials"]["config"]["api_key"] != original_key @@ -166,14 +166,18 @@ class TestApiKeyAuthService: ) db_session_with_containers.expire_all() - result = ApiKeyAuthService.get_auth_credentials(db_session_with_containers, tenant_id, category, provider) + result = ApiKeyAuthService.get_auth_credentials( + tenant_id, category, provider, session=db_session_with_containers + ) assert result == mock_credentials def test_get_auth_credentials_not_found( self, flask_app_with_containers: Flask, db_session_with_containers: Session, tenant_id, category, provider ): - result = ApiKeyAuthService.get_auth_credentials(db_session_with_containers, tenant_id, category, provider) + result = ApiKeyAuthService.get_auth_credentials( + tenant_id, category, provider, session=db_session_with_containers + ) assert result is None @@ -190,7 +194,9 @@ class TestApiKeyAuthService: ) db_session_with_containers.expire_all() - result = ApiKeyAuthService.get_auth_credentials(db_session_with_containers, tenant_id, category, provider) + result = ApiKeyAuthService.get_auth_credentials( + tenant_id, category, provider, session=db_session_with_containers + ) assert result == special_credentials assert result["config"]["api_key"] == "key_with_中文_and_special_chars_!@#$%" @@ -204,7 +210,7 @@ class TestApiKeyAuthService: binding_id = binding.id db_session_with_containers.expire_all() - ApiKeyAuthService.delete_provider_auth(db_session_with_containers, tenant_id, binding_id) + ApiKeyAuthService.delete_provider_auth(tenant_id, binding_id, session=db_session_with_containers) db_session_with_containers.expire_all() remaining = db_session_with_containers.query(DataSourceApiKeyAuthBinding).filter_by(id=binding_id).first() @@ -214,7 +220,7 @@ class TestApiKeyAuthService: self, flask_app_with_containers: Flask, db_session_with_containers: Session, tenant_id ): # Should not raise when binding not found - ApiKeyAuthService.delete_provider_auth(db_session_with_containers, tenant_id, str(uuid4())) + ApiKeyAuthService.delete_provider_auth(tenant_id, str(uuid4()), session=db_session_with_containers) def test_validate_api_key_auth_args_success(self, mock_args): ApiKeyAuthService.validate_api_key_auth_args(mock_args) @@ -291,13 +297,13 @@ class TestApiKeyAuthService: mock_session = MagicMock() mock_session.commit.side_effect = Exception("Database error") with pytest.raises(Exception, match="Database error"): - ApiKeyAuthService.create_provider_auth(mock_session, tenant_id, mock_args) + ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=mock_session) @patch("services.auth.api_key_auth_service.ApiKeyAuthFactory") def test_create_provider_auth_factory_exception(self, mock_factory: MagicMock, tenant_id, mock_args): mock_factory.side_effect = Exception("Factory error") with pytest.raises(Exception, match="Factory error"): - ApiKeyAuthService.create_provider_auth(MagicMock(), tenant_id, mock_args) + ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=MagicMock()) @patch("services.auth.api_key_auth_service.ApiKeyAuthFactory") @patch("services.auth.api_key_auth_service.encrypter") @@ -307,7 +313,7 @@ class TestApiKeyAuthService: mock_factory.return_value = mock_auth_instance mock_encrypter.encrypt_token.side_effect = Exception("Encryption error") with pytest.raises(Exception, match="Encryption error"): - ApiKeyAuthService.create_provider_auth(MagicMock(), tenant_id, mock_args) + ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=MagicMock()) def test_validate_api_key_auth_args_none_input(self): with pytest.raises(TypeError): diff --git a/api/tests/test_containers_integration_tests/services/auth/test_auth_integration.py b/api/tests/test_containers_integration_tests/services/auth/test_auth_integration.py index 9b86ab41f2b..cd3ed01cbfa 100644 --- a/api/tests/test_containers_integration_tests/services/auth/test_auth_integration.py +++ b/api/tests/test_containers_integration_tests/services/auth/test_auth_integration.py @@ -57,7 +57,7 @@ class TestAuthIntegration: mock_encrypt.return_value = "encrypted_fc_test_key_123" args = {"category": category, "provider": AuthType.FIRECRAWL, "credentials": firecrawl_credentials} - ApiKeyAuthService.create_provider_auth(db_session_with_containers, tenant_id_1, args) + ApiKeyAuthService.create_provider_auth(tenant_id_1, args, session=db_session_with_containers) mock_http.assert_called_once() call_args = mock_http.call_args @@ -101,15 +101,15 @@ class TestAuthIntegration: mock_encrypt.return_value = "encrypted_key" args1 = {"category": category, "provider": AuthType.FIRECRAWL, "credentials": firecrawl_credentials} - ApiKeyAuthService.create_provider_auth(db_session_with_containers, tenant_id_1, args1) + ApiKeyAuthService.create_provider_auth(tenant_id_1, args1, session=db_session_with_containers) args2 = {"category": category, "provider": AuthType.JINA, "credentials": jina_credentials} - ApiKeyAuthService.create_provider_auth(db_session_with_containers, tenant_id_2, args2) + ApiKeyAuthService.create_provider_auth(tenant_id_2, args2, session=db_session_with_containers) db_session_with_containers.expire_all() - result1 = ApiKeyAuthService.get_provider_auth_list(db_session_with_containers, tenant_id_1) - result2 = ApiKeyAuthService.get_provider_auth_list(db_session_with_containers, tenant_id_2) + result1 = ApiKeyAuthService.get_provider_auth_list(tenant_id_1, session=db_session_with_containers) + result2 = ApiKeyAuthService.get_provider_auth_list(tenant_id_2, session=db_session_with_containers) assert len(result1) == 1 assert result1[0].tenant_id == tenant_id_1 @@ -120,7 +120,7 @@ class TestAuthIntegration: self, flask_app_with_containers: Flask, db_session_with_containers: Session, tenant_id_2, category ): result = ApiKeyAuthService.get_auth_credentials( - db_session_with_containers, tenant_id_2, category, AuthType.FIRECRAWL + tenant_id_2, category, AuthType.FIRECRAWL, session=db_session_with_containers ) assert result is None @@ -163,7 +163,7 @@ class TestAuthIntegration: "provider": AuthType.FIRECRAWL, "credentials": {"auth_type": "bearer", "config": {"api_key": "fc_test_key_123"}}, } - ApiKeyAuthService.create_provider_auth(db.session(), tenant_id_1, thread_args) + ApiKeyAuthService.create_provider_auth(tenant_id_1, thread_args, session=db.session()) results.append("success") except Exception as e: exceptions.append(e) @@ -216,7 +216,7 @@ class TestAuthIntegration: args = {"category": category, "provider": AuthType.FIRECRAWL, "credentials": firecrawl_credentials} with pytest.raises(httpx.RequestError): - ApiKeyAuthService.create_provider_auth(db_session_with_containers, tenant_id_1, args) + ApiKeyAuthService.create_provider_auth(tenant_id_1, args, session=db_session_with_containers) db_session_with_containers.expire_all() bindings = db_session_with_containers.query(DataSourceApiKeyAuthBinding).filter_by(tenant_id=tenant_id_1).all() @@ -253,12 +253,12 @@ class TestAuthIntegration: mock_encrypt.return_value = "encrypted_key" args = {"category": category, "provider": AuthType.FIRECRAWL, "credentials": firecrawl_credentials} - ApiKeyAuthService.create_provider_auth(db_session_with_containers, tenant_id_1, args) + ApiKeyAuthService.create_provider_auth(tenant_id_1, args, session=db_session_with_containers) db_session_with_containers.expire_all() result = ApiKeyAuthService.get_auth_credentials( - db_session_with_containers, tenant_id_1, category, AuthType.FIRECRAWL + tenant_id_1, category, AuthType.FIRECRAWL, session=db_session_with_containers ) assert result is not None assert result["config"]["api_key"] == "encrypted_key" diff --git a/api/tests/test_containers_integration_tests/services/enterprise/test_account_deletion_sync.py b/api/tests/test_containers_integration_tests/services/enterprise/test_account_deletion_sync.py index 646a0592630..0a34733adeb 100644 --- a/api/tests/test_containers_integration_tests/services/enterprise/test_account_deletion_sync.py +++ b/api/tests/test_containers_integration_tests/services/enterprise/test_account_deletion_sync.py @@ -6,7 +6,7 @@ Redis queuing, error handling, and community vs enterprise behavior. from __future__ import annotations -from unittest.mock import patch +from unittest.mock import MagicMock, patch from uuid import uuid4 import pytest @@ -118,7 +118,7 @@ class TestSyncAccountDeletion: with patch("services.enterprise.account_deletion_sync.dify_config") as mock_config: mock_config.ENTERPRISE_ENABLED = False - result = sync_account_deletion(account_id=str(uuid4()), source="account_deleted") + result = sync_account_deletion(account_id=str(uuid4()), source="account_deleted", session=MagicMock()) assert result is True mock_queue_task.assert_not_called() @@ -137,7 +137,9 @@ class TestSyncAccountDeletion: with patch("services.enterprise.account_deletion_sync.dify_config") as mock_config: mock_config.ENTERPRISE_ENABLED = True - result = sync_account_deletion(account_id=account_id, source="account_deleted") + result = sync_account_deletion( + account_id=account_id, source="account_deleted", session=db_session_with_containers + ) assert result is True assert mock_queue_task.call_count == 3 @@ -151,7 +153,9 @@ class TestSyncAccountDeletion: with patch("services.enterprise.account_deletion_sync.dify_config") as mock_config: mock_config.ENTERPRISE_ENABLED = True - result = sync_account_deletion(account_id=str(uuid4()), source="account_deleted") + result = sync_account_deletion( + account_id=str(uuid4()), source="account_deleted", session=db_session_with_containers + ) assert result is True mock_queue_task.assert_not_called() @@ -176,7 +180,9 @@ class TestSyncAccountDeletion: with patch("services.enterprise.account_deletion_sync.dify_config") as mock_config: mock_config.ENTERPRISE_ENABLED = True - result = sync_account_deletion(account_id=account_id, source="account_deleted") + result = sync_account_deletion( + account_id=account_id, source="account_deleted", session=db_session_with_containers + ) assert result is False assert mock_queue_task.call_count == 3 @@ -196,7 +202,9 @@ class TestSyncAccountDeletion: with patch("services.enterprise.account_deletion_sync.dify_config") as mock_config: mock_config.ENTERPRISE_ENABLED = True - result = sync_account_deletion(account_id=account_id, source="account_deleted") + result = sync_account_deletion( + account_id=account_id, source="account_deleted", session=db_session_with_containers + ) assert result is False mock_queue_task.assert_called_once() diff --git a/api/tests/test_containers_integration_tests/services/plugin/test_plugin_permission_service.py b/api/tests/test_containers_integration_tests/services/plugin/test_plugin_permission_service.py index 0a8f49bc7a7..a52458ac972 100644 --- a/api/tests/test_containers_integration_tests/services/plugin/test_plugin_permission_service.py +++ b/api/tests/test_containers_integration_tests/services/plugin/test_plugin_permission_service.py @@ -2,7 +2,6 @@ from __future__ import annotations from uuid import uuid4 -import pytest from sqlalchemy import func, select from sqlalchemy.orm import Session @@ -38,7 +37,7 @@ class TestGetPermission: db_session_with_containers.add(permission) db_session_with_containers.commit() - result = PluginPermissionService.get_permission(tenant_id) + result = PluginPermissionService.get_permission(tenant_id, session=db_session_with_containers) assert result is not None assert result.id == permission.id @@ -46,9 +45,8 @@ class TestGetPermission: assert result.install_permission == TenantPluginInstallPermission.ADMINS assert result.debug_permission == TenantPluginDebugPermission.EVERYONE - @pytest.mark.usefixtures("flask_app_with_containers") - def test_returns_none_when_not_found(self) -> None: - result = PluginPermissionService.get_permission(_tenant_id()) + def test_returns_none_when_not_found(self, db_session_with_containers: Session) -> None: + result = PluginPermissionService.get_permission(_tenant_id(), session=db_session_with_containers) assert result is None @@ -63,6 +61,7 @@ class TestChangePermission: tenant_id, TenantPluginInstallPermission.EVERYONE, TenantPluginDebugPermission.EVERYONE, + session=db_session_with_containers, ) permission = _get_permission(db_session_with_containers, tenant_id) @@ -85,6 +84,7 @@ class TestChangePermission: tenant_id, TenantPluginInstallPermission.ADMINS, TenantPluginDebugPermission.ADMINS, + session=db_session_with_containers, ) permission = _get_permission(db_session_with_containers, tenant_id) diff --git a/api/tests/test_containers_integration_tests/services/rag_pipeline/test_rag_pipeline_service_db.py b/api/tests/test_containers_integration_tests/services/rag_pipeline/test_rag_pipeline_service_db.py index 2e7df67d266..75d127ce6b2 100644 --- a/api/tests/test_containers_integration_tests/services/rag_pipeline/test_rag_pipeline_service_db.py +++ b/api/tests/test_containers_integration_tests/services/rag_pipeline/test_rag_pipeline_service_db.py @@ -42,7 +42,9 @@ class TestRagPipelineServiceGetPipeline: yield db_session_with_containers.rollback() - def _make_service(self, flask_app_with_containers: Flask) -> RagPipelineService: + def _make_service( + self, flask_app_with_containers: Flask, db_session_with_containers: Session + ) -> RagPipelineService: with ( patch( "services.rag_pipeline.rag_pipeline.DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository", @@ -54,7 +56,7 @@ class TestRagPipelineServiceGetPipeline: ), ): session_factory = sessionmaker(bind=flask_app_with_containers.extensions["sqlalchemy"].engine) - return RagPipelineService(session_maker=session_factory) + return RagPipelineService(db_session_with_containers, session_maker=session_factory) def _create_pipeline(self, db_session: Session, tenant_id: str, created_by: str) -> Pipeline: pipeline = Pipeline( @@ -85,7 +87,7 @@ class TestRagPipelineServiceGetPipeline: self, db_session_with_containers: Session, flask_app_with_containers: Flask ) -> None: """get_pipeline raises ValueError when dataset does not exist.""" - service = self._make_service(flask_app_with_containers) + service = self._make_service(flask_app_with_containers, db_session_with_containers) with pytest.raises(ValueError, match="Dataset not found"): service.get_pipeline(tenant_id=str(uuid4()), dataset_id=str(uuid4())) @@ -99,10 +101,10 @@ class TestRagPipelineServiceGetPipeline: dataset = self._create_dataset(db_session_with_containers, tenant_id, created_by, pipeline_id=None) db_session_with_containers.flush() - service = self._make_service(flask_app_with_containers) + service = self._make_service(flask_app_with_containers, db_session_with_containers) with pytest.raises(ValueError, match="Pipeline not found"): - service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset.id, session=db_session_with_containers) + service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset.id) def test_get_pipeline_returns_pipeline_when_found( self, db_session_with_containers: Session, flask_app_with_containers: Flask @@ -115,9 +117,9 @@ class TestRagPipelineServiceGetPipeline: dataset = self._create_dataset(db_session_with_containers, tenant_id, created_by, pipeline_id=pipeline.id) db_session_with_containers.flush() - service = self._make_service(flask_app_with_containers) + service = self._make_service(flask_app_with_containers, db_session_with_containers) - result = service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset.id, session=db_session_with_containers) + result = service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset.id) assert result.id == pipeline.id @@ -185,7 +187,9 @@ class TestUpdateCustomizedPipelineTemplate: icon_info=IconInfo(icon="📄"), ) with pytest.raises(ValueError, match="Customized pipeline template not found"): - RagPipelineService.update_customized_pipeline_template(str(uuid4()), info, account, tenant_id) + RagPipelineService.update_customized_pipeline_template( + str(uuid4()), info, account, tenant_id, session=db_session_with_containers + ) def test_update_template_raises_on_duplicate_name( self, db_session_with_containers: Session, flask_app_with_containers: Flask @@ -264,4 +268,6 @@ class TestDeleteCustomizedPipelineTemplate: tenant_id = str(uuid4()) with pytest.raises(ValueError, match="Customized pipeline template not found"): - RagPipelineService.delete_customized_pipeline_template(str(uuid4()), tenant_id) + RagPipelineService.delete_customized_pipeline_template( + str(uuid4()), tenant_id, session=db_session_with_containers + ) diff --git a/api/tests/test_containers_integration_tests/services/recommend_app/test_database_retrieval.py b/api/tests/test_containers_integration_tests/services/recommend_app/test_database_retrieval.py index 0f7c790ba14..1c366d3ee32 100644 --- a/api/tests/test_containers_integration_tests/services/recommend_app/test_database_retrieval.py +++ b/api/tests/test_containers_integration_tests/services/recommend_app/test_database_retrieval.py @@ -1,6 +1,6 @@ from __future__ import annotations -from unittest.mock import patch +from unittest.mock import MagicMock, patch from uuid import uuid4 from flask import Flask @@ -82,8 +82,8 @@ class TestDatabaseRecommendAppRetrieval: "fetch_recommended_apps_from_db", return_value={"recommended_apps": [], "categories": []}, ) as mock_fetch: - result = DatabaseRecommendAppRetrieval().get_recommended_apps_and_categories("en-US") - mock_fetch.assert_called_once_with("en-US") + result = DatabaseRecommendAppRetrieval().get_recommended_apps_and_categories("en-US", session=MagicMock()) + mock_fetch.assert_called_once() assert result == {"recommended_apps": [], "categories": []} def test_get_recommend_app_detail_delegates(self): @@ -92,8 +92,8 @@ class TestDatabaseRecommendAppRetrieval: "fetch_recommended_app_detail_from_db", return_value={"id": "app-1"}, ) as mock_fetch: - result = DatabaseRecommendAppRetrieval().get_recommend_app_detail("app-1") - mock_fetch.assert_called_once_with("app-1") + result = DatabaseRecommendAppRetrieval().get_recommend_app_detail("app-1", session=MagicMock()) + mock_fetch.assert_called_once() assert result == {"id": "app-1"} @@ -112,7 +112,9 @@ class TestFetchRecommendedAppsFromDb: db_session_with_containers.expire_all() - result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db("en-US") + result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db( + "en-US", session=db_session_with_containers + ) app_ids = {r["app_id"] for r in result["recommended_apps"]} assert app1.id in app_ids @@ -135,7 +137,9 @@ class TestFetchRecommendedAppsFromDb: db_session_with_containers.expire_all() - result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db("en-US") + result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db( + "en-US", session=db_session_with_containers + ) recommended_app = next(item for item in result["recommended_apps"] if item["app_id"] == created_app.id) assert recommended_app["categories"] == ["writing", "assistant"] @@ -160,7 +164,9 @@ class TestFetchRecommendedAppsFromDb: db_session_with_containers.expire_all() - result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db("en-US") + result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db( + "en-US", session=db_session_with_containers + ) recommended_app = next(item for item in result["recommended_apps"] if item["app_id"] == created_app.id) assert "category" not in recommended_app @@ -177,7 +183,9 @@ class TestFetchRecommendedAppsFromDb: db_session_with_containers.expire_all() - result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db("fr-FR") + result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db( + "fr-FR", session=db_session_with_containers + ) app_ids = {r["app_id"] for r in result["recommended_apps"]} assert app1.id in app_ids @@ -190,7 +198,9 @@ class TestFetchRecommendedAppsFromDb: db_session_with_containers.expire_all() - result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db("en-US") + result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db( + "en-US", session=db_session_with_containers + ) app_ids = {r["app_id"] for r in result["recommended_apps"]} assert app1.id not in app_ids @@ -202,7 +212,9 @@ class TestFetchRecommendedAppsFromDb: db_session_with_containers.expire_all() - result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db("en-US") + result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db( + "en-US", session=db_session_with_containers + ) app_ids = {r["app_id"] for r in result["recommended_apps"]} assert app1.id not in app_ids @@ -235,7 +247,9 @@ class TestFetchRecommendedAppsFromDb: db_session_with_containers.expire_all() - result = DatabaseRecommendAppRetrieval.fetch_learn_dify_apps_from_db("en-US") + result = DatabaseRecommendAppRetrieval.fetch_learn_dify_apps_from_db( + "en-US", session=db_session_with_containers + ) app_ids = {r["app_id"] for r in result["recommended_apps"]} assert learn_dify_app.id in app_ids @@ -261,7 +275,9 @@ class TestFetchRecommendedAppsFromDb: db_session_with_containers.expire_all() - result = DatabaseRecommendAppRetrieval.fetch_learn_dify_apps_from_db("fr-FR") + result = DatabaseRecommendAppRetrieval.fetch_learn_dify_apps_from_db( + "fr-FR", session=db_session_with_containers + ) app_ids = {r["app_id"] for r in result["recommended_apps"]} assert learn_dify_app.id in app_ids @@ -269,7 +285,9 @@ class TestFetchRecommendedAppsFromDb: class TestFetchRecommendedAppDetailFromDb: def test_returns_none_when_not_listed(self, flask_app_with_containers: Flask, db_session_with_containers: Session): - result = DatabaseRecommendAppRetrieval.fetch_recommended_app_detail_from_db(str(uuid4())) + result = DatabaseRecommendAppRetrieval.fetch_recommended_app_detail_from_db( + str(uuid4()), session=db_session_with_containers + ) assert result is None @@ -282,7 +300,9 @@ class TestFetchRecommendedAppDetailFromDb: db_session_with_containers.expire_all() - result = DatabaseRecommendAppRetrieval.fetch_recommended_app_detail_from_db(app1.id) + result = DatabaseRecommendAppRetrieval.fetch_recommended_app_detail_from_db( + app1.id, session=db_session_with_containers + ) assert result is None @@ -298,7 +318,9 @@ class TestFetchRecommendedAppDetailFromDb: db_session_with_containers.expire_all() - result = DatabaseRecommendAppRetrieval.fetch_recommended_app_detail_from_db(app1.id) + result = DatabaseRecommendAppRetrieval.fetch_recommended_app_detail_from_db( + app1.id, session=db_session_with_containers + ) assert result is not None assert result["id"] == app1.id diff --git a/api/tests/test_containers_integration_tests/services/test_account_service.py b/api/tests/test_containers_integration_tests/services/test_account_service.py index 65a5b0a96bf..ac8ed39316b 100644 --- a/api/tests/test_containers_integration_tests/services/test_account_service.py +++ b/api/tests/test_containers_integration_tests/services/test_account_service.py @@ -1120,10 +1120,12 @@ class TestAccountService: mock_sync.return_value = True # Delete account - AccountService.delete_account(account) + AccountService.delete_account(account, session=db_session_with_containers) # Verify sync was called - mock_sync.assert_called_once_with(account_id=account.id, source="account_deleted") + mock_sync.assert_called_once_with( + account_id=account.id, source="account_deleted", session=db_session_with_containers + ) # Verify task was added to queue mock_delete_task.delay.assert_called_once_with(account.id) diff --git a/api/tests/test_containers_integration_tests/services/test_agent_service.py b/api/tests/test_containers_integration_tests/services/test_agent_service.py index 0ee0cb84e75..00b4a1563ff 100644 --- a/api/tests/test_containers_integration_tests/services/test_agent_service.py +++ b/api/tests/test_containers_integration_tests/services/test_agent_service.py @@ -132,7 +132,7 @@ class TestAgentService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Update the app model config to set agent_mode for agent-chat mode if app.mode == AppMode.AGENT_CHAT and app.app_model_config: @@ -295,7 +295,7 @@ class TestAgentService: agent_thoughts = self._create_test_agent_thoughts(db_session_with_containers, message) # Execute the method under test - result = AgentService.get_agent_logs(app, conversation.id, message.id) + result = AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) # Verify the result structure assert result is not None @@ -355,7 +355,7 @@ class TestAgentService: # Execute the method under test with non-existent conversation with pytest.raises(ValueError, match="Conversation not found"): - AgentService.get_agent_logs(app, fake.uuid4(), fake.uuid4()) + AgentService.get_agent_logs(app, fake.uuid4(), fake.uuid4(), db_session_with_containers) def test_get_agent_logs_message_not_found( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -371,7 +371,7 @@ class TestAgentService: # Execute the method under test with non-existent message with pytest.raises(ValueError, match="Message not found"): - AgentService.get_agent_logs(app, conversation.id, fake.uuid4()) + AgentService.get_agent_logs(app, conversation.id, fake.uuid4(), db_session_with_containers) def test_get_agent_logs_with_end_user( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -452,7 +452,7 @@ class TestAgentService: db_session_with_containers.commit() # Execute the method under test - result = AgentService.get_agent_logs(app, conversation.id, message.id) + result = AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) # Verify the result assert result is not None @@ -524,7 +524,7 @@ class TestAgentService: db_session_with_containers.commit() # Execute the method under test - result = AgentService.get_agent_logs(app, conversation.id, message.id) + result = AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) # Verify the result assert result is not None @@ -569,7 +569,7 @@ class TestAgentService: db_session_with_containers.commit() # Execute the method under test - result = AgentService.get_agent_logs(app, conversation.id, message.id) + result = AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) # Verify the result assert result is not None @@ -593,7 +593,7 @@ class TestAgentService: conversation, message = self._create_test_conversation_and_message(db_session_with_containers, app, account) # Execute the method under test - result = AgentService.get_agent_logs(app, conversation.id, message.id) + result = AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) # Verify the result assert result is not None @@ -655,7 +655,7 @@ class TestAgentService: # Execute the method under test with pytest.raises(ValueError, match="App model config not found"): - AgentService.get_agent_logs(app, conversation.id, message.id) + AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) def test_get_agent_logs_agent_config_not_found( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -674,7 +674,7 @@ class TestAgentService: # Execute the method under test with pytest.raises(ValueError, match="Agent config not found"): - AgentService.get_agent_logs(app, conversation.id, message.id) + AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) def test_list_agent_providers_success( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -804,7 +804,7 @@ class TestAgentService: db_session_with_containers.commit() # Execute the method under test - result = AgentService.get_agent_logs(app, conversation.id, message.id) + result = AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) # Verify the result assert result is not None @@ -899,7 +899,7 @@ class TestAgentService: db_session_with_containers.commit() # Execute the method under test - result = AgentService.get_agent_logs(app, conversation.id, message.id) + result = AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) # Verify the result assert result is not None @@ -927,7 +927,7 @@ class TestAgentService: mock_external_service_dependencies["current_user"].timezone = "Asia/Shanghai" # Execute the method under test - result = AgentService.get_agent_logs(app, conversation.id, message.id) + result = AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) # Verify the result assert result is not None @@ -968,7 +968,7 @@ class TestAgentService: db_session_with_containers.commit() # Execute the method under test - result = AgentService.get_agent_logs(app, conversation.id, message.id) + result = AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) # Verify the result assert result is not None @@ -1009,7 +1009,7 @@ class TestAgentService: db_session_with_containers.commit() # Execute the method under test - result = AgentService.get_agent_logs(app, conversation.id, message.id) + result = AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) # Verify the result - should handle malformed JSON gracefully assert result is not None diff --git a/api/tests/test_containers_integration_tests/services/test_annotation_service.py b/api/tests/test_containers_integration_tests/services/test_annotation_service.py index 94d72b19be8..2710df5e56c 100644 --- a/api/tests/test_containers_integration_tests/services/test_annotation_service.py +++ b/api/tests/test_containers_integration_tests/services/test_annotation_service.py @@ -101,7 +101,7 @@ class TestAnnotationService: # Create app app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Setup current_user mock self._mock_current_user(mock_external_service_dependencies, account.id, tenant.id) @@ -207,7 +207,9 @@ class TestAnnotationService: } # Insert annotation directly - annotation = AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + annotation_args, app.id, session=db_session_with_containers + ) # Verify annotation was created correctly assert annotation.app_id == app.id @@ -241,7 +243,9 @@ class TestAnnotationService: } with pytest.raises(ValueError): - AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + AppAnnotationService.insert_app_annotation_directly( + annotation_args, app.id, session=db_session_with_containers + ) def test_insert_app_annotation_directly_app_not_found( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -263,7 +267,9 @@ class TestAnnotationService: # Try to insert annotation with non-existent app with pytest.raises(NotFound, match="App not found"): - AppAnnotationService.insert_app_annotation_directly(annotation_args, non_existent_app_id) + AppAnnotationService.insert_app_annotation_directly( + annotation_args, non_existent_app_id, session=db_session_with_containers + ) def test_update_app_annotation_directly_success( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -279,7 +285,9 @@ class TestAnnotationService: "question": fake.sentence(), "answer": fake.text(max_nb_chars=200), } - annotation = AppAnnotationService.insert_app_annotation_directly(original_args, app.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + original_args, app.id, session=db_session_with_containers + ) # Update the annotation updated_args = { @@ -328,7 +336,9 @@ class TestAnnotationService: } # Insert annotation from message - annotation = AppAnnotationService.up_insert_app_annotation_from_message(annotation_args, app.id) + annotation = AppAnnotationService.up_insert_app_annotation_from_message( + annotation_args, app.id, session=db_session_with_containers + ) # Verify annotation was created correctly assert annotation.app_id == app.id @@ -361,7 +371,9 @@ class TestAnnotationService: "question": fake.sentence(), "answer": fake.text(max_nb_chars=200), } - initial_annotation = AppAnnotationService.up_insert_app_annotation_from_message(initial_args, app.id) + initial_annotation = AppAnnotationService.up_insert_app_annotation_from_message( + initial_args, app.id, session=db_session_with_containers + ) # Update the annotation updated_args = { @@ -369,7 +381,9 @@ class TestAnnotationService: "question": fake.sentence(), "answer": fake.text(max_nb_chars=200), } - updated_annotation = AppAnnotationService.up_insert_app_annotation_from_message(updated_args, app.id) + updated_annotation = AppAnnotationService.up_insert_app_annotation_from_message( + updated_args, app.id, session=db_session_with_containers + ) # Verify annotation was updated correctly (same ID) assert updated_annotation.id == initial_annotation.id @@ -402,7 +416,9 @@ class TestAnnotationService: # Try to insert annotation with non-existent app with pytest.raises(NotFound, match="App not found"): - AppAnnotationService.up_insert_app_annotation_from_message(annotation_args, non_existent_app_id) + AppAnnotationService.up_insert_app_annotation_from_message( + annotation_args, non_existent_app_id, session=db_session_with_containers + ) def test_get_annotation_list_by_app_id_success( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -420,12 +436,18 @@ class TestAnnotationService: "question": f"Question {i}: {fake.sentence()}", "answer": f"Answer {i}: {fake.text(max_nb_chars=200)}", } - annotation = AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + annotation_args, app.id, session=db_session_with_containers + ) annotations.append(annotation) # Get annotation list annotation_list, total = AppAnnotationService.get_annotation_list_by_app_id( - app.id, page=1, limit=10, keyword="" + app.id, + page=1, + limit=10, + keyword="", + session=db_session_with_containers, ) # Verify results @@ -452,18 +474,22 @@ class TestAnnotationService: "question": f"Question with {unique_keyword} keyword", "answer": f"Answer with {unique_keyword} keyword", } - AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id, session=db_session_with_containers) # Create another annotation without the keyword other_args = { "question": "Different question without special term", "answer": "Different answer without special content", } - AppAnnotationService.insert_app_annotation_directly(other_args, app.id) + AppAnnotationService.insert_app_annotation_directly(other_args, app.id, session=db_session_with_containers) # Search with keyword annotation_list, total = AppAnnotationService.get_annotation_list_by_app_id( - app.id, page=1, limit=10, keyword=unique_keyword + app.id, + page=1, + limit=10, + keyword=unique_keyword, + session=db_session_with_containers, ) # Verify only matching annotations are returned @@ -490,30 +516,42 @@ class TestAnnotationService: "question": "Question with 50% discount", "answer": "Answer about 50% discount offer", } - AppAnnotationService.insert_app_annotation_directly(annotation_with_percent, app.id) + AppAnnotationService.insert_app_annotation_directly( + annotation_with_percent, app.id, session=db_session_with_containers + ) annotation_with_underscore = { "question": "Question with test_data", "answer": "Answer about test_data value", } - AppAnnotationService.insert_app_annotation_directly(annotation_with_underscore, app.id) + AppAnnotationService.insert_app_annotation_directly( + annotation_with_underscore, app.id, session=db_session_with_containers + ) annotation_with_backslash = { "question": "Question with path\\to\\file", "answer": "Answer about path\\to\\file location", } - AppAnnotationService.insert_app_annotation_directly(annotation_with_backslash, app.id) + AppAnnotationService.insert_app_annotation_directly( + annotation_with_backslash, app.id, session=db_session_with_containers + ) # Create annotation that should NOT match (contains % but as part of different text) annotation_no_match = { "question": "Question with 100% different", "answer": "Answer about 100% different content", } - AppAnnotationService.insert_app_annotation_directly(annotation_no_match, app.id) + AppAnnotationService.insert_app_annotation_directly( + annotation_no_match, app.id, session=db_session_with_containers + ) # Test 1: Search with % character - should find exact match only annotation_list, total = AppAnnotationService.get_annotation_list_by_app_id( - app.id, page=1, limit=10, keyword="50%" + app.id, + page=1, + limit=10, + keyword="50%", + session=db_session_with_containers, ) assert total == 1 assert len(annotation_list) == 1 @@ -521,7 +559,11 @@ class TestAnnotationService: # Test 2: Search with _ character - should find exact match only annotation_list, total = AppAnnotationService.get_annotation_list_by_app_id( - app.id, page=1, limit=10, keyword="test_data" + app.id, + page=1, + limit=10, + keyword="test_data", + session=db_session_with_containers, ) assert total == 1 assert len(annotation_list) == 1 @@ -529,7 +571,11 @@ class TestAnnotationService: # Test 3: Search with \ character - should find exact match only annotation_list, total = AppAnnotationService.get_annotation_list_by_app_id( - app.id, page=1, limit=10, keyword="path\\to\\file" + app.id, + page=1, + limit=10, + keyword="path\\to\\file", + session=db_session_with_containers, ) assert total == 1 assert len(annotation_list) == 1 @@ -537,7 +583,11 @@ class TestAnnotationService: # Test 4: Search with % should NOT match 100% (verifies escaping works) annotation_list, total = AppAnnotationService.get_annotation_list_by_app_id( - app.id, page=1, limit=10, keyword="50%" + app.id, + page=1, + limit=10, + keyword="50%", + session=db_session_with_containers, ) # Should only find the 50% annotation, not the 100% one assert total == 1 @@ -557,7 +607,9 @@ class TestAnnotationService: # Try to get annotation list with non-existent app with pytest.raises(NotFound, match="App not found"): - AppAnnotationService.get_annotation_list_by_app_id(non_existent_app_id, page=1, limit=10, keyword="") + AppAnnotationService.get_annotation_list_by_app_id( + non_existent_app_id, page=1, limit=10, keyword="", session=db_session_with_containers + ) def test_delete_app_annotation_success( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -573,7 +625,9 @@ class TestAnnotationService: "question": fake.sentence(), "answer": fake.text(max_nb_chars=200), } - annotation = AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + annotation_args, app.id, session=db_session_with_containers + ) annotation_id = annotation.id # Delete the annotation @@ -728,7 +782,9 @@ class TestAnnotationService: "question": fake.sentence(), "answer": fake.text(max_nb_chars=200), } - annotation = AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + annotation_args, app.id, session=db_session_with_containers + ) # Add some hit histories for i in range(3): @@ -742,6 +798,7 @@ class TestAnnotationService: message_id=fake.uuid4(), from_source=ConversationFromSource.CONSOLE, score=0.8 + (i * 0.1), + session=db_session_with_containers, ) # Get hit histories @@ -749,6 +806,7 @@ class TestAnnotationService: self._annotation_ref(app, annotation.id), page=1, limit=10, + session=db_session_with_containers, ) # Verify results @@ -775,7 +833,9 @@ class TestAnnotationService: "question": fake.sentence(), "answer": fake.text(max_nb_chars=200), } - annotation = AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + annotation_args, app.id, session=db_session_with_containers + ) # Get initial hit count initial_hit_count = annotation.hit_count @@ -795,6 +855,7 @@ class TestAnnotationService: message_id=message_id, from_source=ConversationFromSource.CONSOLE, score=score, + session=db_session_with_containers, ) # Verify hit count was incremented @@ -834,10 +895,14 @@ class TestAnnotationService: "question": fake.sentence(), "answer": fake.text(max_nb_chars=200), } - created_annotation = AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + created_annotation = AppAnnotationService.insert_app_annotation_directly( + annotation_args, app.id, session=db_session_with_containers + ) # Get annotation by ID - retrieved_annotation = AppAnnotationService.get_annotation_by_id(created_annotation.id) + retrieved_annotation = AppAnnotationService.get_annotation_by_id( + created_annotation.id, session=db_session_with_containers + ) # Verify annotation was retrieved correctly assert retrieved_annotation is not None @@ -880,7 +945,9 @@ class TestAnnotationService: mock_pd.read_csv.return_value = mock_df # Batch import annotations - result = AppAnnotationService.batch_import_app_annotations(app.id, file_storage) + result = AppAnnotationService.batch_import_app_annotations( + app.id, file_storage, session=db_session_with_containers + ) # Verify result structure assert "job_id" in result @@ -920,7 +987,9 @@ class TestAnnotationService: mock_pd.read_csv.return_value = mock_df # Batch import annotations - result = AppAnnotationService.batch_import_app_annotations(app.id, file_storage) + result = AppAnnotationService.batch_import_app_annotations( + app.id, file_storage, session=db_session_with_containers + ) # Verify error result assert "error_msg" in result @@ -966,7 +1035,9 @@ class TestAnnotationService: ].get_features.return_value.annotation_quota_limit.size = 0 # Batch import annotations - result = AppAnnotationService.batch_import_app_annotations(app.id, file_storage) + result = AppAnnotationService.batch_import_app_annotations( + app.id, file_storage, session=db_session_with_containers + ) # Verify error result assert "error_msg" in result @@ -1008,7 +1079,7 @@ class TestAnnotationService: db_session_with_containers.commit() # Get annotation setting - result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id) + result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id, session=db_session_with_containers) # Verify result structure assert result["enabled"] is True @@ -1027,7 +1098,7 @@ class TestAnnotationService: app, account = self._create_test_app_and_account(db_session_with_containers, mock_external_service_dependencies) # Get annotation setting (no setting exists) - result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id) + result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id, session=db_session_with_containers) # Verify result structure assert result["enabled"] is False @@ -1072,7 +1143,9 @@ class TestAnnotationService: "score_threshold": 0.9, } - result = AppAnnotationService.update_app_annotation_setting(app.id, annotation_setting.id, update_args) + result = AppAnnotationService.update_app_annotation_setting( + app.id, annotation_setting.id, update_args, session=db_session_with_containers + ) # Verify result structure assert result["enabled"] is True @@ -1101,11 +1174,15 @@ class TestAnnotationService: "question": f"Question {i}: {fake.sentence()}", "answer": f"Answer {i}: {fake.text(max_nb_chars=200)}", } - annotation = AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + annotation_args, app.id, session=db_session_with_containers + ) annotations.append(annotation) # Export annotation list - exported_annotations = AppAnnotationService.export_annotation_list_by_app_id(app.id) + exported_annotations = AppAnnotationService.export_annotation_list_by_app_id( + app.id, session=db_session_with_containers + ) # Verify results assert len(exported_annotations) == 3 @@ -1132,7 +1209,9 @@ class TestAnnotationService: # Try to export annotation list with non-existent app with pytest.raises(NotFound, match="App not found"): - AppAnnotationService.export_annotation_list_by_app_id(non_existent_app_id) + AppAnnotationService.export_annotation_list_by_app_id( + non_existent_app_id, session=db_session_with_containers + ) def test_insert_app_annotation_directly_with_setting_success( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -1176,7 +1255,9 @@ class TestAnnotationService: } # Insert annotation directly - annotation = AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + annotation_args, app.id, session=db_session_with_containers + ) # Verify annotation was created correctly assert annotation.app_id == app.id @@ -1235,7 +1316,9 @@ class TestAnnotationService: "question": fake.sentence(), "answer": fake.text(max_nb_chars=200), } - annotation = AppAnnotationService.insert_app_annotation_directly(original_args, app.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + original_args, app.id, session=db_session_with_containers + ) # Reset mock to clear previous calls mock_external_service_dependencies["update_task"].delay.reset_mock() @@ -1312,7 +1395,9 @@ class TestAnnotationService: "question": fake.sentence(), "answer": fake.text(max_nb_chars=200), } - annotation = AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + annotation_args, app.id, session=db_session_with_containers + ) annotation_id = annotation.id # Reset mock to clear previous calls @@ -1382,7 +1467,9 @@ class TestAnnotationService: } # Insert annotation from message - annotation = AppAnnotationService.up_insert_app_annotation_from_message(annotation_args, app.id) + annotation = AppAnnotationService.up_insert_app_annotation_from_message( + annotation_args, app.id, session=db_session_with_containers + ) # Verify annotation was created correctly assert annotation.app_id == app.id diff --git a/api/tests/test_containers_integration_tests/services/test_api_based_extension_service.py b/api/tests/test_containers_integration_tests/services/test_api_based_extension_service.py index 1f88ce90621..de51f5077e6 100644 --- a/api/tests/test_containers_integration_tests/services/test_api_based_extension_service.py +++ b/api/tests/test_containers_integration_tests/services/test_api_based_extension_service.py @@ -82,7 +82,7 @@ class TestAPIBasedExtensionService: ) # Save extension - saved_extension = APIBasedExtensionService.save(db_session_with_containers, extension_data) + saved_extension = APIBasedExtensionService.save(extension_data, session=db_session_with_containers) # Verify extension was saved correctly assert saved_extension.id is not None @@ -120,21 +120,21 @@ class TestAPIBasedExtensionService: ) with pytest.raises(ValueError, match="name must not be empty"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) # Test empty api_endpoint extension_data.name = fake.company() extension_data.api_endpoint = "" with pytest.raises(ValueError, match="api_endpoint must not be empty"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) # Test empty api_key extension_data.api_endpoint = f"https://{fake.domain_name()}/api" extension_data.api_key = "" with pytest.raises(ValueError, match="api_key must not be empty"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) def test_get_all_by_tenant_id_success( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -158,11 +158,11 @@ class TestAPIBasedExtensionService: api_key=fake.password(length=20), ) - saved_extension = APIBasedExtensionService.save(db_session_with_containers, extension_data) + saved_extension = APIBasedExtensionService.save(extension_data, session=db_session_with_containers) extensions.append(saved_extension) # Get all extensions for tenant - extension_list = APIBasedExtensionService.get_all_by_tenant_id(db_session_with_containers, tenant.id) + extension_list = APIBasedExtensionService.get_all_by_tenant_id(tenant.id, session=db_session_with_containers) # Verify results assert len(extension_list) == 3 @@ -192,11 +192,11 @@ class TestAPIBasedExtensionService: api_key=fake.password(length=20), ) - created_extension = APIBasedExtensionService.save(db_session_with_containers, extension_data) + created_extension = APIBasedExtensionService.save(extension_data, session=db_session_with_containers) # Get extension by ID retrieved_extension = APIBasedExtensionService.get_with_tenant_id( - db_session_with_containers, tenant.id, created_extension.id + tenant.id, created_extension.id, session=db_session_with_containers ) # Verify extension was retrieved correctly @@ -223,7 +223,7 @@ class TestAPIBasedExtensionService: # Try to get non-existent extension with pytest.raises(ValueError, match="API based extension is not found"): APIBasedExtensionService.get_with_tenant_id( - db_session_with_containers, tenant.id, non_existent_extension_id + tenant.id, non_existent_extension_id, session=db_session_with_containers ) def test_delete_extension_success(self, db_session_with_containers: Session, mock_external_service_dependencies): @@ -243,11 +243,11 @@ class TestAPIBasedExtensionService: api_key=fake.password(length=20), ) - created_extension = APIBasedExtensionService.save(db_session_with_containers, extension_data) + created_extension = APIBasedExtensionService.save(extension_data, session=db_session_with_containers) extension_id = created_extension.id # Delete the extension - APIBasedExtensionService.delete(db_session_with_containers, created_extension) + APIBasedExtensionService.delete(created_extension, session=db_session_with_containers) # Verify extension was deleted @@ -275,7 +275,7 @@ class TestAPIBasedExtensionService: api_key=fake.password(length=20), ) - APIBasedExtensionService.save(db_session_with_containers, extension_data1) + APIBasedExtensionService.save(extension_data1, session=db_session_with_containers) # Try to create second extension with same name extension_data2 = APIBasedExtension( tenant_id=tenant.id, @@ -285,7 +285,7 @@ class TestAPIBasedExtensionService: ) with pytest.raises(ValueError, match="name must be unique, it is already existed"): - APIBasedExtensionService.save(db_session_with_containers, extension_data2) + APIBasedExtensionService.save(extension_data2, session=db_session_with_containers) def test_save_extension_update_existing( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -306,7 +306,7 @@ class TestAPIBasedExtensionService: api_key=fake.password(length=20), ) - created_extension = APIBasedExtensionService.save(db_session_with_containers, extension_data) + created_extension = APIBasedExtensionService.save(extension_data, session=db_session_with_containers) # Save original values for later comparison original_name = created_extension.name @@ -325,7 +325,7 @@ class TestAPIBasedExtensionService: created_extension.api_endpoint = new_endpoint created_extension.api_key = new_api_key - updated_extension = APIBasedExtensionService.save(db_session_with_containers, created_extension) + updated_extension = APIBasedExtensionService.save(created_extension, session=db_session_with_containers) # Verify extension was updated correctly assert updated_extension.id == created_extension.id @@ -342,7 +342,7 @@ class TestAPIBasedExtensionService: # Verify the update by retrieving the extension again retrieved_extension = APIBasedExtensionService.get_with_tenant_id( - db_session_with_containers, tenant.id, created_extension.id + tenant.id, created_extension.id, session=db_session_with_containers ) assert retrieved_extension.name == new_name assert retrieved_extension.api_endpoint == new_endpoint @@ -374,7 +374,7 @@ class TestAPIBasedExtensionService: # Try to save extension with connection error with pytest.raises(ValueError, match="connection error: request timeout"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) def test_save_extension_invalid_api_key_length( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -397,7 +397,7 @@ class TestAPIBasedExtensionService: # Try to save extension with short API key with pytest.raises(ValueError, match="api_key must be at least 5 characters"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) def test_save_extension_empty_fields(self, db_session_with_containers: Session, mock_external_service_dependencies): """ @@ -417,21 +417,21 @@ class TestAPIBasedExtensionService: ) with pytest.raises(ValueError, match="name must not be empty"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) # Test with None api_endpoint extension_data.name = fake.company() extension_data.api_endpoint = None with pytest.raises(ValueError, match="api_endpoint must not be empty"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) # Test with None api_key extension_data.api_endpoint = f"https://{fake.domain_name()}/api" extension_data.api_key = None with pytest.raises(ValueError, match="api_key must not be empty"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) def test_get_all_by_tenant_id_empty_list( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -445,7 +445,7 @@ class TestAPIBasedExtensionService: ) # Get all extensions for tenant (none exist) - extension_list = APIBasedExtensionService.get_all_by_tenant_id(db_session_with_containers, tenant.id) + extension_list = APIBasedExtensionService.get_all_by_tenant_id(tenant.id, session=db_session_with_containers) # Verify empty list is returned assert len(extension_list) == 0 @@ -475,7 +475,7 @@ class TestAPIBasedExtensionService: # Try to save extension with invalid ping response with pytest.raises(ValueError, match="{'result': 'invalid'}"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) def test_save_extension_missing_ping_result( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -501,7 +501,7 @@ class TestAPIBasedExtensionService: # Try to save extension with missing ping result with pytest.raises(ValueError, match="{'status': 'ok'}"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) def test_get_with_tenant_id_wrong_tenant( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -527,11 +527,13 @@ class TestAPIBasedExtensionService: api_key=fake.password(length=20), ) - created_extension = APIBasedExtensionService.save(db_session_with_containers, extension_data) + created_extension = APIBasedExtensionService.save(extension_data, session=db_session_with_containers) # Try to get extension with wrong tenant ID with pytest.raises(ValueError, match="API based extension is not found"): - APIBasedExtensionService.get_with_tenant_id(db_session_with_containers, tenant2.id, created_extension.id) + APIBasedExtensionService.get_with_tenant_id( + tenant2.id, created_extension.id, session=db_session_with_containers + ) def test_save_extension_api_key_exactly_four_chars_rejected( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -551,7 +553,7 @@ class TestAPIBasedExtensionService: ) with pytest.raises(ValueError, match="api_key must be at least 5 characters"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) def test_save_extension_api_key_exactly_five_chars_accepted( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -570,7 +572,7 @@ class TestAPIBasedExtensionService: api_key="12345", ) - saved = APIBasedExtensionService.save(db_session_with_containers, extension_data) + saved = APIBasedExtensionService.save(extension_data, session=db_session_with_containers) assert saved.id is not None def test_save_extension_requestor_constructor_error( @@ -593,7 +595,7 @@ class TestAPIBasedExtensionService: ) with pytest.raises(ValueError, match="connection error: bad config"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) def test_save_extension_network_exception( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -617,7 +619,7 @@ class TestAPIBasedExtensionService: ) with pytest.raises(ValueError, match="connection error: network failure"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) def test_save_extension_update_duplicate_name_rejected( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -630,28 +632,28 @@ class TestAPIBasedExtensionService: assert tenant is not None ext1 = APIBasedExtensionService.save( - db_session_with_containers, APIBasedExtension( tenant_id=tenant.id, name="Extension Alpha", api_endpoint=f"https://{fake.domain_name()}/api", api_key=fake.password(length=20), ), + session=db_session_with_containers, ) ext2 = APIBasedExtensionService.save( - db_session_with_containers, APIBasedExtension( tenant_id=tenant.id, name="Extension Beta", api_endpoint=f"https://{fake.domain_name()}/api", api_key=fake.password(length=20), ), + session=db_session_with_containers, ) # Try to rename ext2 to ext1's name ext2.name = "Extension Alpha" with pytest.raises(ValueError, match="name must be unique, it is already existed"): - APIBasedExtensionService.save(db_session_with_containers, ext2) + APIBasedExtensionService.save(ext2, session=db_session_with_containers) def test_get_all_returns_empty_for_different_tenant( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -667,15 +669,15 @@ class TestAPIBasedExtensionService: assert tenant1 is not None APIBasedExtensionService.save( - db_session_with_containers, APIBasedExtension( tenant_id=tenant1.id, name=fake.company(), api_endpoint=f"https://{fake.domain_name()}/api", api_key=fake.password(length=20), ), + session=db_session_with_containers, ) assert tenant2 is not None - result = APIBasedExtensionService.get_all_by_tenant_id(db_session_with_containers, tenant2.id) + result = APIBasedExtensionService.get_all_by_tenant_id(tenant2.id, session=db_session_with_containers) assert result == [] diff --git a/api/tests/test_containers_integration_tests/services/test_app_dsl_service.py b/api/tests/test_containers_integration_tests/services/test_app_dsl_service.py index cee08c4c33e..24c14637296 100644 --- a/api/tests/test_containers_integration_tests/services/test_app_dsl_service.py +++ b/api/tests/test_containers_integration_tests/services/test_app_dsl_service.py @@ -162,7 +162,7 @@ class TestAppDslService: api_rpm=10, ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) return app, account def _create_simple_yaml_content(self, app_name: str = "Test App", app_mode: str = "chat") -> str: @@ -841,7 +841,7 @@ class TestAppDslService: # ── Export ───────────────────────────────────────────────────────── - def test_export_dsl_delegates_by_mode(self, monkeypatch: pytest.MonkeyPatch): + def test_export_dsl_delegates_by_mode(self, monkeypatch: pytest.MonkeyPatch, db_session_with_containers: Session): workflow_calls: list[bool] = [] model_calls: list[bool] = [] monkeypatch.setattr( @@ -859,7 +859,7 @@ class TestAppDslService: mode=AppMode.WORKFLOW, icon_type="emoji", ) - AppDslService.export_dsl(workflow_app) + AppDslService.export_dsl(workflow_app, session=db_session_with_containers) assert workflow_calls == [True] chat_app = _app_stub( @@ -867,10 +867,12 @@ class TestAppDslService: icon_type="emoji", app_model_config=SimpleNamespace(to_dict=lambda: {"agent_mode": {"tools": []}}), ) - AppDslService.export_dsl(chat_app) + AppDslService.export_dsl(chat_app, session=db_session_with_containers) assert model_calls == [True] - def test_export_dsl_preserves_icon_and_icon_type(self, monkeypatch: pytest.MonkeyPatch): + def test_export_dsl_preserves_icon_and_icon_type( + self, monkeypatch: pytest.MonkeyPatch, db_session_with_containers: Session + ): monkeypatch.setattr( AppDslService, "_append_workflow_export_data", @@ -886,7 +888,7 @@ class TestAppDslService: description="App with emoji icon", use_icon_as_answer_icon=True, ) - yaml_output = AppDslService.export_dsl(emoji_app) + yaml_output = AppDslService.export_dsl(emoji_app, session=db_session_with_containers) data = yaml.safe_load(yaml_output) assert data["app"]["icon"] == "🎨" assert data["app"]["icon_type"] == "emoji" @@ -901,7 +903,7 @@ class TestAppDslService: description="App with image icon", use_icon_as_answer_icon=False, ) - yaml_output = AppDslService.export_dsl(image_app) + yaml_output = AppDslService.export_dsl(image_app, session=db_session_with_containers) data = yaml.safe_load(yaml_output) assert data["app"]["icon"] == "https://example.com/icon.png" assert data["app"]["icon_type"] == "image" @@ -936,7 +938,7 @@ class TestAppDslService: db_session_with_containers.add(model_config) db_session_with_containers.commit() - exported_dsl = AppDslService.export_dsl(app, include_secret=False) + exported_dsl = AppDslService.export_dsl(app, include_secret=False, session=db_session_with_containers) exported_data = yaml.safe_load(exported_dsl) assert exported_data["kind"] == "app" @@ -972,7 +974,7 @@ class TestAppDslService: "workflow_service" ].return_value.get_draft_workflow.return_value = mock_workflow - exported_dsl = AppDslService.export_dsl(app, include_secret=False) + exported_dsl = AppDslService.export_dsl(app, include_secret=False, session=db_session_with_containers) exported_data = yaml.safe_load(exported_dsl) assert exported_data["kind"] == "app" @@ -1006,7 +1008,7 @@ class TestAppDslService: workflow_id = str(uuid4()) - def mock_get_draft_workflow(app_model, wf_id=None): + def mock_get_draft_workflow(app_model, wf_id=None, **_kwargs): if wf_id == workflow_id: return mock_workflow return None @@ -1015,7 +1017,9 @@ class TestAppDslService: "workflow_service" ].return_value.get_draft_workflow.side_effect = mock_get_draft_workflow - exported_dsl = AppDslService.export_dsl(app, include_secret=False, workflow_id=workflow_id) + exported_dsl = AppDslService.export_dsl( + app, include_secret=False, workflow_id=workflow_id, session=db_session_with_containers + ) exported_data = yaml.safe_load(exported_dsl) assert exported_data["kind"] == "app" @@ -1034,11 +1038,15 @@ class TestAppDslService: WorkflowNotFoundError, match="Missing draft workflow configuration, please check.", ): - AppDslService.export_dsl(app, include_secret=False, workflow_id=str(uuid4())) + AppDslService.export_dsl( + app, include_secret=False, workflow_id=str(uuid4()), session=db_session_with_containers + ) # ── Workflow Export Data ─────────────────────────────────────────── - def test_append_workflow_export_data_filters_and_overrides(self, monkeypatch: pytest.MonkeyPatch): + def test_append_workflow_export_data_filters_and_overrides( + self, monkeypatch: pytest.MonkeyPatch, db_session_with_containers: Session + ): workflow_dict = { "graph": { "nodes": [ @@ -1123,6 +1131,7 @@ class TestAppDslService: app_model=_app_stub(), include_secret=False, workflow_id=None, + session=db_session_with_containers, ) nodes = export_data["workflow"]["graph"]["nodes"] @@ -1138,7 +1147,9 @@ class TestAppDslService: assert nodes[5]["data"]["subscription_id"] == "" assert export_data["dependencies"] == [{"tenant": _DEFAULT_TENANT_ID, "dep": "dep-1"}] - def test_append_workflow_export_data_missing_workflow_raises(self, monkeypatch: pytest.MonkeyPatch): + def test_append_workflow_export_data_missing_workflow_raises( + self, monkeypatch: pytest.MonkeyPatch, db_session_with_containers: Session + ): workflow_service = MagicMock() workflow_service.get_draft_workflow.return_value = None monkeypatch.setattr(app_dsl_service, "WorkflowService", lambda: workflow_service) @@ -1149,6 +1160,7 @@ class TestAppDslService: app_model=_app_stub(), include_secret=False, workflow_id=None, + session=db_session_with_containers, ) # ── Model Config Export Data ────────────────────────────────────── diff --git a/api/tests/test_containers_integration_tests/services/test_app_generate_service.py b/api/tests/test_containers_integration_tests/services/test_app_generate_service.py index 473111f364a..89cc7715d1c 100644 --- a/api/tests/test_containers_integration_tests/services/test_app_generate_service.py +++ b/api/tests/test_containers_integration_tests/services/test_app_generate_service.py @@ -187,7 +187,7 @@ class TestAppGenerateService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) return app, account @@ -234,12 +234,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -267,12 +267,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -298,12 +298,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -329,12 +329,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -362,12 +362,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -399,12 +399,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -431,12 +431,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.DEBUGGER, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -461,12 +461,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=db_session_with_containers, ) # Verify the result @@ -503,12 +503,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=end_user, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -535,12 +535,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -574,12 +574,12 @@ class TestAppGenerateService: # StatementError (from EnumText validation during autoflush) with pytest.raises((ValueError, sa.exc.StatementError)): AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) def test_generate_with_workflow_id_format_error( @@ -603,12 +603,12 @@ class TestAppGenerateService: # Execute the method under test and expect WorkflowIdFormatError with pytest.raises(WorkflowIdFormatError) as exc_info: AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify error message @@ -642,12 +642,12 @@ class TestAppGenerateService: # Execute the method under test and expect WorkflowNotFoundError with pytest.raises(WorkflowNotFoundError) as exc_info: AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify error message @@ -673,12 +673,12 @@ class TestAppGenerateService: # Execute the method under test and expect ValueError with pytest.raises(ValueError) as exc_info: AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.DEBUGGER, streaming=True, + session=db_session_with_containers, ) # Verify error message @@ -704,12 +704,12 @@ class TestAppGenerateService: # Execute the method under test and expect ValueError with pytest.raises(ValueError) as exc_info: AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify error message @@ -731,7 +731,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate_single_iteration( - app_model=app, user=account, node_id=node_id, args=args, streaming=True + app_model=app, + user=account, + node_id=node_id, + args=args, + streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -758,7 +763,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate_single_iteration( - app_model=app, user=account, node_id=node_id, args=args, streaming=True + app_model=app, + user=account, + node_id=node_id, + args=args, + streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -786,7 +796,12 @@ class TestAppGenerateService: # Execute the method under test and expect ValueError with pytest.raises(ValueError) as exc_info: AppGenerateService.generate_single_iteration( - app_model=app, user=account, node_id=node_id, args=args, streaming=True + app_model=app, + user=account, + node_id=node_id, + args=args, + streaming=True, + session=db_session_with_containers, ) # Verify error message @@ -808,7 +823,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate_single_loop( - app_model=app, user=account, node_id=node_id, args=args, streaming=True + app_model=app, + user=account, + node_id=node_id, + args=args, + streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -835,7 +855,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate_single_loop( - app_model=app, user=account, node_id=node_id, args=args, streaming=True + app_model=app, + user=account, + node_id=node_id, + args=args, + streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -861,7 +886,12 @@ class TestAppGenerateService: # Execute the method under test and expect ValueError with pytest.raises(ValueError) as exc_info: AppGenerateService.generate_single_loop( - app_model=app, user=account, node_id=node_id, args=args, streaming=True + app_model=app, + user=account, + node_id=node_id, + args=args, + streaming=True, + session=db_session_with_containers, ) # Verify error message @@ -1021,12 +1051,12 @@ class TestAppGenerateService: # Execute the method under test and expect exception with pytest.raises(Exception) as exc_info: AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify exception message @@ -1054,12 +1084,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -1094,12 +1124,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=invoke_from, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -1137,12 +1167,12 @@ class TestAppGenerateService: mock_exec_params.new.return_value = mock_payload result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify the result diff --git a/api/tests/test_containers_integration_tests/services/test_app_service.py b/api/tests/test_containers_integration_tests/services/test_app_service.py index f9df99c5594..8deaf6d462d 100644 --- a/api/tests/test_containers_integration_tests/services/test_app_service.py +++ b/api/tests/test_containers_integration_tests/services/test_app_service.py @@ -84,7 +84,7 @@ class TestAppService: # Create app app_service = AppService() - app = app_service.create_app(tenant.id, app_params, account) + app = app_service.create_app(tenant.id, app_params, account, session=db_session_with_containers) # Verify app was created correctly assert app.name == app_params.name @@ -144,7 +144,7 @@ class TestAppService: icon_background="#4ECDC4", ) - app = app_service.create_app(tenant.id, app_params, account) + app = app_service.create_app(tenant.id, app_params, account, session=db_session_with_containers) # Verify app mode was set correctly assert app.mode == mode @@ -183,7 +183,7 @@ class TestAppService: ) app_service = AppService() - created_app = app_service.create_app(tenant.id, app_params, account) + created_app = app_service.create_app(tenant.id, app_params, account, session=db_session_with_containers) # Get app using the service - needs current_user mock mock_current_user = create_autospec(Account, instance=True) @@ -234,7 +234,7 @@ class TestAppService: icon="📱", icon_background="#96CEB4", ) - app_service.create_app(tenant.id, app_params, account) + app_service.create_app(tenant.id, app_params, account, session=db_session_with_containers) # Get paginated apps params = AppListParams(page=1, limit=10, mode="chat") @@ -277,16 +277,19 @@ class TestAppService: tenant.id, CreateAppParams(name="Oldest Created", mode="chat", icon_type="emoji", icon="1"), account, + session=db_session_with_containers, ) newest_modified = app_service.create_app( tenant.id, CreateAppParams(name="Newest Modified", mode="chat", icon_type="emoji", icon="2"), account, + session=db_session_with_containers, ) newest_created = app_service.create_app( tenant.id, CreateAppParams(name="Newest Created", mode="chat", icon_type="emoji", icon="3"), account, + session=db_session_with_containers, ) timestamp_by_app_id = { @@ -362,15 +365,17 @@ class TestAppService: tenant.id, CreateAppParams(name="Starred App", mode="chat", icon_type="emoji", icon="1"), account, + session=db_session_with_containers, ) unstarred_app = app_service.create_app( tenant.id, CreateAppParams(name="Unstarred App", mode="chat", icon_type="emoji", icon="2"), account, + session=db_session_with_containers, ) - app_service.star_app(db_session_with_containers, app=starred_app, account_id=account.id) - app_service.star_app(db_session_with_containers, app=starred_app, account_id=account.id) + app_service.star_app(app=starred_app, account_id=account.id, session=db_session_with_containers) + app_service.star_app(app=starred_app, account_id=account.id, session=db_session_with_containers) db_session_with_containers.commit() star_count = db_session_with_containers.scalar( @@ -386,7 +391,7 @@ class TestAppService: assert starred_by_app_id[starred_app.id] is True assert starred_by_app_id[unstarred_app.id] is False - app_service.unstar_app(db_session_with_containers, app=starred_app, account_id=account.id) + app_service.unstar_app(app=starred_app, account_id=account.id, session=db_session_with_containers) db_session_with_containers.commit() paginated_apps = app_service.get_paginate_apps( @@ -422,26 +427,30 @@ class TestAppService: tenant.id, CreateAppParams(name="Oldest Created Starred App", mode="chat", icon_type="emoji", icon="1"), account, + session=db_session_with_containers, ) newest_modified_app = app_service.create_app( tenant.id, CreateAppParams(name="Newest Modified Starred App", mode="chat", icon_type="emoji", icon="2"), account, + session=db_session_with_containers, ) newest_created_app = app_service.create_app( tenant.id, CreateAppParams(name="Newest Created Starred App", mode="chat", icon_type="emoji", icon="3"), account, + session=db_session_with_containers, ) unstarred_app = app_service.create_app( tenant.id, CreateAppParams(name="Unstarred App", mode="chat", icon_type="emoji", icon="4"), account, + session=db_session_with_containers, ) - app_service.star_app(db_session_with_containers, app=oldest_created_app, account_id=account.id) - app_service.star_app(db_session_with_containers, app=newest_modified_app, account_id=account.id) - app_service.star_app(db_session_with_containers, app=newest_created_app, account_id=account.id) + app_service.star_app(app=oldest_created_app, account_id=account.id, session=db_session_with_containers) + app_service.star_app(app=newest_modified_app, account_id=account.id, session=db_session_with_containers) + app_service.star_app(app=newest_created_app, account_id=account.id, session=db_session_with_containers) timestamp_by_app_id = { oldest_created_app.id: (datetime(2026, 1, 1, 10, 0, 0), datetime(2026, 1, 1, 10, 0, 0)), @@ -535,8 +544,10 @@ class TestAppService: icon_background="#4ECDC4", ) - chat_app = app_service.create_app(tenant.id, chat_app_params, account) - completion_app = app_service.create_app(tenant.id, completion_app_params, account) + chat_app = app_service.create_app(tenant.id, chat_app_params, account, session=db_session_with_containers) + completion_app = app_service.create_app( + tenant.id, completion_app_params, account, session=db_session_with_containers + ) # Test filter by mode chat_apps = app_service.get_paginate_apps( @@ -599,7 +610,7 @@ class TestAppService: icon="💬", icon_background="#FF6B6B", ) - app_service.create_app(tenant.id, app_params, first_account) + app_service.create_app(tenant.id, app_params, first_account, session=db_session_with_containers) other_app_params = CreateAppParams( name="Second Creator App", description="Created by the second account", @@ -608,7 +619,7 @@ class TestAppService: icon="✍️", icon_background="#4ECDC4", ) - app_service.create_app(tenant.id, other_app_params, second_account) + app_service.create_app(tenant.id, other_app_params, second_account, session=db_session_with_containers) filtered_apps = app_service.get_paginate_apps( first_account.id, @@ -654,7 +665,7 @@ class TestAppService: icon="🏷️", icon_background="#FFEAA7", ) - app = app_service.create_app(tenant.id, app_params, account) + app = app_service.create_app(tenant.id, app_params, account, session=db_session_with_containers) # Mock TagService to return the app ID for tag filtering with patch("services.app_service.TagService.get_target_ids_by_tag_ids") as mock_tag_service: @@ -717,7 +728,7 @@ class TestAppService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_params, account) + app = app_service.create_app(tenant.id, app_params, account, session=db_session_with_containers) # Store original values original_name = app.name @@ -741,7 +752,7 @@ class TestAppService: mock_current_user.current_tenant_id = account.current_tenant_id with patch("services.app_service.current_user", mock_current_user): - updated_app = app_service.update_app(app, update_args) + updated_app = app_service.update_app(app, update_args, session=db_session_with_containers) # Verify updated fields assert updated_app.name == update_args["name"] @@ -788,6 +799,7 @@ class TestAppService: icon_background="#45B7D1", ), account, + session=db_session_with_containers, ) mock_current_user = create_autospec(Account, instance=True) @@ -805,6 +817,7 @@ class TestAppService: "icon_background": "#FF8C42", "use_icon_as_answer_icon": True, }, + session=db_session_with_containers, ) assert updated_app.icon_type == IconType.EMOJI @@ -841,6 +854,7 @@ class TestAppService: icon_background="#45B7D1", ), account, + session=db_session_with_containers, ) mock_current_user = create_autospec(Account, instance=True) @@ -859,6 +873,7 @@ class TestAppService: "icon_background": "#FF8C42", "use_icon_as_answer_icon": True, }, + session=db_session_with_containers, ) def test_update_app_name_success(self, db_session_with_containers: Session, mock_external_service_dependencies): @@ -892,7 +907,7 @@ class TestAppService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_params, account) + app = app_service.create_app(tenant.id, app_params, account, session=db_session_with_containers) # Store original name original_name = app.name @@ -904,7 +919,7 @@ class TestAppService: mock_current_user.current_tenant_id = account.current_tenant_id with patch("services.app_service.current_user", mock_current_user): - updated_app = app_service.update_app_name(app, new_name) + updated_app = app_service.update_app_name(app, new_name, session=db_session_with_containers) assert updated_app.name == new_name assert updated_app.updated_by == account.id @@ -946,7 +961,7 @@ class TestAppService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_params, account) + app = app_service.create_app(tenant.id, app_params, account, session=db_session_with_containers) # Store original values original_icon = app.icon @@ -961,7 +976,9 @@ class TestAppService: mock_current_user.current_tenant_id = account.current_tenant_id with patch("services.app_service.current_user", mock_current_user): - updated_app = app_service.update_app_icon(app, new_icon, new_icon_background, new_icon_type) + updated_app = app_service.update_app_icon( + app, new_icon, new_icon_background, new_icon_type, session=db_session_with_containers + ) assert updated_app.icon == new_icon assert updated_app.icon_background == new_icon_background @@ -1007,7 +1024,7 @@ class TestAppService: icon_background="#74B9FF", ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Store original site status original_site_status = app.enable_site @@ -1018,13 +1035,13 @@ class TestAppService: mock_current_user.current_tenant_id = account.current_tenant_id with patch("services.app_service.current_user", mock_current_user): - updated_app = app_service.update_app_site_status(app, False) + updated_app = app_service.update_app_site_status(app, False, session=db_session_with_containers) assert updated_app.enable_site is False assert updated_app.updated_by == account.id # Update site status back to enabled with patch("services.app_service.current_user", mock_current_user): - updated_app = app_service.update_app_site_status(updated_app, True) + updated_app = app_service.update_app_site_status(updated_app, True, session=db_session_with_containers) assert updated_app.enable_site is True assert updated_app.updated_by == account.id @@ -1067,7 +1084,7 @@ class TestAppService: icon_background="#A29BFE", ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Store original API status original_api_status = app.enable_api @@ -1078,13 +1095,13 @@ class TestAppService: mock_current_user.current_tenant_id = account.current_tenant_id with patch("services.app_service.current_user", mock_current_user): - updated_app = app_service.update_app_api_status(app, False) + updated_app = app_service.update_app_api_status(app, False, session=db_session_with_containers) assert updated_app.enable_api is False assert updated_app.updated_by == account.id # Update API status back to enabled with patch("services.app_service.current_user", mock_current_user): - updated_app = app_service.update_app_api_status(updated_app, True) + updated_app = app_service.update_app_api_status(updated_app, True, session=db_session_with_containers) assert updated_app.enable_api is True assert updated_app.updated_by == account.id @@ -1127,14 +1144,14 @@ class TestAppService: icon_background="#FD79A8", ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Store original values original_site_status = app.enable_site original_updated_at = app.updated_at # Update site status to the same value (no change) - updated_app = app_service.update_app_site_status(app, original_site_status) + updated_app = app_service.update_app_site_status(app, original_site_status, session=db_session_with_containers) # Verify app is returned unchanged assert updated_app.id == app.id @@ -1178,7 +1195,7 @@ class TestAppService: icon_background="#E17055", ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Store app ID for verification app_id = app.id @@ -1188,7 +1205,7 @@ class TestAppService: mock_delete_task.delay.return_value = None # Delete app - app_service.delete_app(app) + app_service.delete_app(app, session=db_session_with_containers) # Verify async deletion task was called mock_delete_task.delay.assert_called_once_with(tenant_id=tenant.id, app_id=app_id) @@ -1230,7 +1247,7 @@ class TestAppService: icon_background="#00B894", ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Store app ID for verification app_id = app.id @@ -1245,7 +1262,7 @@ class TestAppService: mock_delete_task.delay.return_value = None # Delete app - app_service.delete_app(app) + app_service.delete_app(app, session=db_session_with_containers) # Verify webapp auth cleanup was called mock_external_service_dependencies["enterprise_service"].WebAppAuth.cleanup_webapp.assert_called_once_with( @@ -1290,10 +1307,10 @@ class TestAppService: icon_background="#6C5CE7", ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Get app metadata - app_meta = app_service.get_app_meta(app) + app_meta = app_service.get_app_meta(app, session=db_session_with_containers) # Verify metadata contains expected fields assert "tool_icons" in app_meta @@ -1329,10 +1346,10 @@ class TestAppService: icon_background="#FDCB6E", ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Get app code by ID - app_code = AppService.get_app_code_by_id(app.id) + app_code = AppService.get_app_code_by_id(app.id, session=db_session_with_containers) # Verify app code was retrieved correctly # Note: Site would be created when App is created, site.code is auto-generated @@ -1369,7 +1386,7 @@ class TestAppService: icon_background="#E84393", ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Create a site for the app site = Site() @@ -1384,7 +1401,7 @@ class TestAppService: db_session_with_containers.commit() # Get app ID by code - app_id = AppService.get_app_id_by_code(site.code) + app_id = AppService.get_app_id_by_code(site.code, session=db_session_with_containers) # Verify app ID was retrieved correctly assert app_id == app.id @@ -1462,6 +1479,7 @@ class TestAppService: api_rpm=10, ), account, + session=db_session_with_containers, ) app_with_underscore = app_service.create_app( @@ -1477,6 +1495,7 @@ class TestAppService: api_rpm=10, ), account, + session=db_session_with_containers, ) app_with_backslash = app_service.create_app( @@ -1492,6 +1511,7 @@ class TestAppService: api_rpm=10, ), account, + session=db_session_with_containers, ) # Create app that should NOT match @@ -1508,6 +1528,7 @@ class TestAppService: api_rpm=10, ), account, + session=db_session_with_containers, ) # Test 1: Search with % character @@ -1560,7 +1581,7 @@ class TestAppService: from services.app_service import AppService with pytest.raises(ValueError, match="not found"): - AppService.get_app_code_by_id(str(uuid4())) + AppService.get_app_code_by_id(str(uuid4()), session=db_session_with_containers) def test_get_app_id_by_code_not_found( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -1569,7 +1590,7 @@ class TestAppService: from services.app_service import AppService with pytest.raises(ValueError, match="not found"): - AppService.get_app_id_by_code("nonexistent-code") + AppService.get_app_id_by_code("nonexistent-code", session=db_session_with_containers) def test_get_app_meta_returns_empty_when_workflow_missing( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -1582,7 +1603,7 @@ class TestAppService: app_service = AppService() workflow_app = SimpleNamespace(mode="workflow", workflow=None) - meta = app_service.get_app_meta(workflow_app) + meta = app_service.get_app_meta(workflow_app, session=db_session_with_containers) assert meta == {"tool_icons": {}} def test_get_app_meta_returns_empty_when_model_config_missing( @@ -1596,5 +1617,5 @@ class TestAppService: app_service = AppService() chat_app = SimpleNamespace(mode="chat", app_model_config=None) - meta = app_service.get_app_meta(chat_app) + meta = app_service.get_app_meta(chat_app, session=db_session_with_containers) assert meta == {"tool_icons": {}} diff --git a/api/tests/test_containers_integration_tests/services/test_billing_service.py b/api/tests/test_containers_integration_tests/services/test_billing_service.py index a3a4a0e6edd..777fb7721b4 100644 --- a/api/tests/test_containers_integration_tests/services/test_billing_service.py +++ b/api/tests/test_containers_integration_tests/services/test_billing_service.py @@ -417,7 +417,7 @@ class TestBillingServiceIsTenantOwnerOrAdmin: account, _ = self._create_account_with_tenant_role(db_session_with_containers, TenantAccountRole.EDITOR) with pytest.raises(ValueError, match="Only team owner or team admin can perform this action"): - BillingService.is_tenant_owner_or_admin(db_session_with_containers, account) + BillingService.is_tenant_owner_or_admin(account, session=db_session_with_containers) def test_is_tenant_owner_or_admin_dataset_operator_raises_error(self, db_session_with_containers: Session) -> None: """is_tenant_owner_or_admin raises ValueError for DATASET_OPERATOR role.""" @@ -426,4 +426,4 @@ class TestBillingServiceIsTenantOwnerOrAdmin: ) with pytest.raises(ValueError, match="Only team owner or team admin can perform this action"): - BillingService.is_tenant_owner_or_admin(db_session_with_containers, account) + BillingService.is_tenant_owner_or_admin(account, session=db_session_with_containers) diff --git a/api/tests/test_containers_integration_tests/services/test_conversation_service.py b/api/tests/test_containers_integration_tests/services/test_conversation_service.py index b19b6b9c984..19dd4d6cf70 100644 --- a/api/tests/test_containers_integration_tests/services/test_conversation_service.py +++ b/api/tests/test_containers_integration_tests/services/test_conversation_service.py @@ -350,6 +350,7 @@ class TestConversationServiceMessageCreation: conversation_id=conversation.id, first_id=None, # No starting point specified limit=10, + session=db_session_with_containers, ) # Assert - Verify the results @@ -395,6 +396,7 @@ class TestConversationServiceMessageCreation: conversation_id=conversation.id, first_id=first_message.id, limit=10, + session=db_session_with_containers, ) # Assert - Verify the results @@ -426,6 +428,7 @@ class TestConversationServiceMessageCreation: conversation_id=conversation.id, first_id=str(uuid4()), limit=10, + session=db_session_with_containers, ) def test_pagination_with_has_more_flag(self, db_session_with_containers: Session): @@ -461,6 +464,7 @@ class TestConversationServiceMessageCreation: conversation_id=conversation.id, first_id=None, limit=limit, + session=db_session_with_containers, ) # Assert @@ -498,7 +502,8 @@ class TestConversationServiceMessageCreation: conversation_id=conversation.id, first_id=None, limit=10, - order="asc", # Ascending order + order="asc", # Ascending order, + session=db_session_with_containers, ) # Assert @@ -547,7 +552,7 @@ class TestConversationServiceSummarization: mock_llm_generator.return_value = generated_name # Act - result = ConversationService.auto_generate_name(app_model, conversation) + result = ConversationService.auto_generate_name(app_model, conversation, session=db_session_with_containers) # Assert assert conversation.name == generated_name # Name updated on conversation object @@ -572,7 +577,7 @@ class TestConversationServiceSummarization: # Act & Assert with pytest.raises(MessageNotExistsError): - ConversationService.auto_generate_name(app_model, conversation) + ConversationService.auto_generate_name(app_model, conversation, session=db_session_with_containers) @patch("services.conversation_service.LLMGenerator.generate_conversation_name") def test_auto_generate_name_handles_llm_failure_gracefully( @@ -604,7 +609,7 @@ class TestConversationServiceSummarization: mock_llm_generator.side_effect = Exception("LLM service unavailable") # Act - result = ConversationService.auto_generate_name(app_model, conversation) + result = ConversationService.auto_generate_name(app_model, conversation, session=db_session_with_containers) # Assert assert conversation.name == original_name # Name remains unchanged @@ -637,6 +642,7 @@ class TestConversationServiceSummarization: user=user, name=new_name, auto_generate=False, + session=db_session_with_containers, ) # Assert @@ -671,6 +677,7 @@ class TestConversationServiceSummarization: user=user, name=None, auto_generate=True, + session=db_session_with_containers, ) # Assert @@ -719,7 +726,9 @@ class TestConversationServiceMessageAnnotation: args = {"message_id": message.id, "answer": "AI is artificial intelligence"} # Act - result = AppAnnotationService.up_insert_app_annotation_from_message(args, app_model.id) + result = AppAnnotationService.up_insert_app_annotation_from_message( + args, app_model.id, session=db_session_with_containers + ) # Assert assert result.message_id == message.id @@ -753,7 +762,9 @@ class TestConversationServiceMessageAnnotation: } # Act - result = AppAnnotationService.up_insert_app_annotation_from_message(args, app_model.id) + result = AppAnnotationService.up_insert_app_annotation_from_message( + args, app_model.id, session=db_session_with_containers + ) # Assert assert result.message_id is None @@ -802,7 +813,9 @@ class TestConversationServiceMessageAnnotation: args = {"message_id": message.id, "answer": "Updated annotation content"} # Act - result = AppAnnotationService.up_insert_app_annotation_from_message(args, app_model.id) + result = AppAnnotationService.up_insert_app_annotation_from_message( + args, app_model.id, session=db_session_with_containers + ) # Assert assert result.id == existing_annotation.id @@ -838,7 +851,11 @@ class TestConversationServiceMessageAnnotation: # Act result_items, result_total = AppAnnotationService.get_annotation_list_by_app_id( - app_id=app_model.id, page=1, limit=10, keyword="" + app_id=app_model.id, + page=1, + limit=10, + keyword="", + session=db_session_with_containers, ) # Assert @@ -886,7 +903,8 @@ class TestConversationServiceMessageAnnotation: app_id=app_model.id, page=1, limit=10, - keyword="machine", # Search keyword + keyword="machine", # Search keyword, + session=db_session_with_containers, ) # Assert @@ -914,7 +932,9 @@ class TestConversationServiceMessageAnnotation: } # Act - result = AppAnnotationService.insert_app_annotation_directly(args, app_model.id) + result = AppAnnotationService.insert_app_annotation_directly( + args, app_model.id, session=db_session_with_containers + ) # Assert assert result.question == args["question"] @@ -942,7 +962,9 @@ class TestConversationServiceExport: ) # Act - result = ConversationService.get_conversation(app_model=app_model, conversation_id=conversation.id, user=user) + result = ConversationService.get_conversation( + app_model=app_model, conversation_id=conversation.id, user=user, session=db_session_with_containers + ) # Assert assert result == conversation @@ -956,7 +978,12 @@ class TestConversationServiceExport: # Act & Assert with pytest.raises(ConversationNotExistsError): - ConversationService.get_conversation(app_model=app_model, conversation_id=str(uuid4()), user=user) + ConversationService.get_conversation( + app_model=app_model, + conversation_id=str(uuid4()), + user=user, + session=db_session_with_containers, + ) @patch("services.annotation_service.current_account_with_tenant") def test_export_annotation_list(self, mock_current_account, db_session_with_containers: Session): @@ -982,7 +1009,7 @@ class TestConversationServiceExport: mock_current_account.return_value = (account, app_model.tenant_id) # Act - result = AppAnnotationService.export_annotation_list_by_app_id(app_model.id) + result = AppAnnotationService.export_annotation_list_by_app_id(app_model.id, session=db_session_with_containers) # Assert assert len(result) == 10 @@ -1006,7 +1033,9 @@ class TestConversationServiceExport: ) # Act - result = MessageService.get_message(app_model=app_model, user=user, message_id=message.id) + result = MessageService.get_message( + app_model=app_model, user=user, message_id=message.id, session=db_session_with_containers + ) # Assert assert result == message @@ -1020,7 +1049,9 @@ class TestConversationServiceExport: # Act & Assert with pytest.raises(MessageNotExistsError): - MessageService.get_message(app_model=app_model, user=user, message_id=str(uuid4())) + MessageService.get_message( + app_model=app_model, user=user, message_id=str(uuid4()), session=db_session_with_containers + ) def test_get_conversation_for_end_user(self, db_session_with_containers: Session): """ @@ -1041,7 +1072,10 @@ class TestConversationServiceExport: # Act result = ConversationService.get_conversation( - app_model=app_model, conversation_id=conversation.id, user=end_user + app_model=app_model, + conversation_id=conversation.id, + user=end_user, + session=db_session_with_containers, ) # Assert @@ -1069,7 +1103,9 @@ class TestConversationServiceExport: conversation_id = conversation.id # Act - Delete the conversation - ConversationService.delete(app_model=app_model, conversation_id=conversation_id, user=user) + ConversationService.delete( + app_model=app_model, conversation_id=conversation_id, user=user, session=db_session_with_containers + ) # Assert - Verify two-step deletion process # Step 1: Immediate database deletion @@ -1104,6 +1140,7 @@ class TestConversationServiceExport: app_model=app_model, conversation_id=conversation.id, user=other_account, + session=db_session_with_containers, ) # Verify no deletion and no async cleanup trigger @@ -1129,9 +1166,14 @@ class TestConversationServiceExport: conversation_id = conversation.id # Act — force an error during the delete to exercise the rollback path - with patch("services.conversation_service.db.session.delete", side_effect=Exception("DB error")): + with patch.object(db_session_with_containers, "delete", side_effect=Exception("DB error")): with pytest.raises(Exception, match="DB error"): - ConversationService.delete(app_model=app_model, conversation_id=conversation_id, user=user) + ConversationService.delete( + app_model=app_model, + conversation_id=conversation_id, + user=user, + session=db_session_with_containers, + ) # Assert — async cleanup must NOT have been scheduled mock_delete_task.delay.assert_not_called() diff --git a/api/tests/test_containers_integration_tests/services/test_conversation_service_variables.py b/api/tests/test_containers_integration_tests/services/test_conversation_service_variables.py index 33d4563904e..9a725b06b64 100644 --- a/api/tests/test_containers_integration_tests/services/test_conversation_service_variables.py +++ b/api/tests/test_containers_integration_tests/services/test_conversation_service_variables.py @@ -6,10 +6,9 @@ from uuid import uuid4 import pytest from flask import Flask -from sqlalchemy.orm import Session, sessionmaker +from sqlalchemy.orm import Session from core.app.entities.app_invoke_entities import InvokeFrom -from extensions.ext_database import db from graphon.variables import FloatVariable, IntegerVariable, StringVariable from models.account import Account, Tenant, TenantAccountJoin from models.enums import ConversationFromSource, EndUserType @@ -152,13 +151,6 @@ class ConversationServiceVariableIntegrationFactory: @pytest.fixture def real_conversation_service_session_factory(flask_app_with_containers: Flask): del flask_app_with_containers - real_session_maker = sessionmaker(bind=db.engine, expire_on_commit=False) - - with ( - patch("services.conversation_service.session_factory.create_session", side_effect=lambda: real_session_maker()), - patch("services.conversation_service.session_factory.get_session_maker", return_value=real_session_maker), - ): - yield class TestConversationServiceVariables: @@ -193,6 +185,7 @@ class TestConversationServiceVariables: user=account, limit=10, last_id=None, + session=db_session_with_containers, ) assert [item["id"] for item in result.data] == [first_variable.id, second_variable.id] @@ -237,6 +230,7 @@ class TestConversationServiceVariables: user=account, limit=10, last_id=first_variable.id, + session=db_session_with_containers, ) assert [item["id"] for item in result.data] == [second_variable.id, third_variable.id] @@ -257,6 +251,7 @@ class TestConversationServiceVariables: user=account, limit=10, last_id=str(uuid4()), + session=db_session_with_containers, ) def test_get_conversational_variable_sets_has_more( @@ -282,6 +277,7 @@ class TestConversationServiceVariables: user=account, limit=2, last_id=None, + session=db_session_with_containers, ) assert len(result.data) == 2 @@ -309,6 +305,7 @@ class TestConversationServiceVariables: variable_id=existing.id, user=account, new_value="support", + session=db_session_with_containers, ) db_session_with_containers.expire_all() @@ -335,6 +332,7 @@ class TestConversationServiceVariables: variable_id=str(uuid4()), user=account, new_value="support", + session=db_session_with_containers, ) def test_update_conversation_variable_type_mismatch_raises_error( @@ -358,6 +356,7 @@ class TestConversationServiceVariables: variable_id=existing.id, user=account, new_value="wrong-type", + session=db_session_with_containers, ) def test_update_conversation_variable_integer_number_compatibility( @@ -380,6 +379,7 @@ class TestConversationServiceVariables: variable_id=existing.id, user=account, new_value=42, + session=db_session_with_containers, ) db_session_with_containers.expire_all() diff --git a/api/tests/test_containers_integration_tests/services/test_credit_pool_service.py b/api/tests/test_containers_integration_tests/services/test_credit_pool_service.py index de8e6ba612c..9cbe5252bbf 100644 --- a/api/tests/test_containers_integration_tests/services/test_credit_pool_service.py +++ b/api/tests/test_containers_integration_tests/services/test_credit_pool_service.py @@ -4,10 +4,8 @@ from unittest.mock import patch from uuid import uuid4 import pytest -from flask import has_app_context from sqlalchemy.orm import Session -from core.db.session_factory import session_factory from core.errors.error import QuotaExceededError from models import TenantCreditPool from models.enums import ProviderQuotaType @@ -35,11 +33,10 @@ class TestCreditPoolService: db_session.add(pool) db_session.commit() - @pytest.mark.usefixtures("db_session_with_containers") - def test_create_default_pool(self) -> None: + def test_create_default_pool(self, db_session_with_containers: Session) -> None: tenant_id = self._create_tenant_id() - pool = CreditPoolService.create_default_pool(tenant_id) + pool = CreditPoolService.create_default_pool(tenant_id, session=db_session_with_containers) assert isinstance(pool, TenantCreditPool) assert pool.tenant_id == tenant_id @@ -51,43 +48,46 @@ class TestCreditPoolService: tenant_id = self._create_tenant_id() self._create_pool(db_session_with_containers, tenant_id=tenant_id, quota_limit=10, quota_used=0) - result = CreditPoolService.get_pool(tenant_id=tenant_id, pool_type=ProviderQuotaType.TRIAL) + result = CreditPoolService.get_pool( + tenant_id=tenant_id, pool_type=ProviderQuotaType.TRIAL, session=db_session_with_containers + ) assert result is not None assert result.tenant_id == tenant_id assert result.pool_type == ProviderQuotaType.TRIAL - @pytest.mark.usefixtures("flask_app_with_containers") - def test_get_pool_uses_configured_session_factory_without_flask_app_context(self) -> None: + def test_get_pool_uses_provided_session(self, db_session_with_containers: Session) -> None: tenant_id = self._create_tenant_id() - session_maker = session_factory.get_session_maker() - with session_maker.begin() as session: - session.add( - TenantCreditPool( - tenant_id=tenant_id, - pool_type=ProviderQuotaType.TRIAL, - quota_limit=10, - quota_used=2, - ) + db_session_with_containers.add( + TenantCreditPool( + tenant_id=tenant_id, + pool_type=ProviderQuotaType.TRIAL, + quota_limit=10, + quota_used=2, ) + ) + db_session_with_containers.commit() - assert not has_app_context() - result = CreditPoolService.get_pool(tenant_id=tenant_id, pool_type=ProviderQuotaType.TRIAL) + result = CreditPoolService.get_pool( + tenant_id=tenant_id, pool_type=ProviderQuotaType.TRIAL, session=db_session_with_containers + ) assert result is not None assert result.tenant_id == tenant_id assert result.pool_type == ProviderQuotaType.TRIAL assert result.quota_used == 2 - @pytest.mark.usefixtures("flask_app_with_containers") - def test_get_pool_returns_none_when_not_exists(self) -> None: - result = CreditPoolService.get_pool(tenant_id=self._create_tenant_id(), pool_type=ProviderQuotaType.TRIAL) + def test_get_pool_returns_none_when_not_exists(self, db_session_with_containers: Session) -> None: + result = CreditPoolService.get_pool( + tenant_id=self._create_tenant_id(), pool_type=ProviderQuotaType.TRIAL, session=db_session_with_containers + ) assert result is None - @pytest.mark.usefixtures("flask_app_with_containers") - def test_check_credits_available_returns_false_when_no_pool(self) -> None: - result = CreditPoolService.check_credits_available(tenant_id=self._create_tenant_id(), credits_required=10) + def test_check_credits_available_returns_false_when_no_pool(self, db_session_with_containers: Session) -> None: + result = CreditPoolService.check_credits_available( + tenant_id=self._create_tenant_id(), credits_required=10, session=db_session_with_containers + ) assert result is False @@ -95,7 +95,9 @@ class TestCreditPoolService: tenant_id = self._create_tenant_id() self._create_pool(db_session_with_containers, tenant_id=tenant_id, quota_limit=10, quota_used=0) - result = CreditPoolService.check_credits_available(tenant_id=tenant_id, credits_required=10) + result = CreditPoolService.check_credits_available( + tenant_id=tenant_id, credits_required=10, session=db_session_with_containers + ) assert result is True @@ -103,14 +105,17 @@ class TestCreditPoolService: tenant_id = self._create_tenant_id() self._create_pool(db_session_with_containers, tenant_id=tenant_id, quota_limit=10, quota_used=10) - result = CreditPoolService.check_credits_available(tenant_id=tenant_id, credits_required=1) + result = CreditPoolService.check_credits_available( + tenant_id=tenant_id, credits_required=1, session=db_session_with_containers + ) assert result is False - @pytest.mark.usefixtures("flask_app_with_containers") - def test_check_and_deduct_credits_raises_when_no_pool(self) -> None: + def test_check_and_deduct_credits_raises_when_no_pool(self, db_session_with_containers: Session) -> None: with pytest.raises(QuotaExceededError, match="Credit pool not found"): - CreditPoolService.check_and_deduct_credits(tenant_id=self._create_tenant_id(), credits_required=1) + CreditPoolService.check_and_deduct_credits( + tenant_id=self._create_tenant_id(), credits_required=1, session=db_session_with_containers + ) def test_check_and_deduct_credits_returns_zero_for_non_positive_request( self, db_session_with_containers: Session @@ -118,10 +123,12 @@ class TestCreditPoolService: tenant_id = self._create_tenant_id() self._create_pool(db_session_with_containers, tenant_id=tenant_id, quota_limit=10, quota_used=2) - result = CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=0) + result = CreditPoolService.check_and_deduct_credits( + tenant_id=tenant_id, credits_required=0, session=db_session_with_containers + ) assert result == 0 - updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id) + updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id, session=db_session_with_containers) assert updated_pool is not None assert updated_pool.quota_used == 2 @@ -130,9 +137,11 @@ class TestCreditPoolService: self._create_pool(db_session_with_containers, tenant_id=tenant_id, quota_limit=10, quota_used=10) with pytest.raises(QuotaExceededError, match="No credits remaining"): - CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=1) + CreditPoolService.check_and_deduct_credits( + tenant_id=tenant_id, credits_required=1, session=db_session_with_containers + ) - updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id) + updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id, session=db_session_with_containers) assert updated_pool is not None assert updated_pool.quota_used == 10 @@ -141,10 +150,12 @@ class TestCreditPoolService: self._create_pool(db_session_with_containers, tenant_id=tenant_id, quota_limit=10, quota_used=2) credits_required = 3 - result = CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=credits_required) + result = CreditPoolService.check_and_deduct_credits( + tenant_id=tenant_id, credits_required=credits_required, session=db_session_with_containers + ) assert result == credits_required - pool = CreditPoolService.get_pool(tenant_id=tenant_id) + pool = CreditPoolService.get_pool(tenant_id=tenant_id, session=db_session_with_containers) assert pool is not None assert pool.quota_used == 5 @@ -155,9 +166,11 @@ class TestCreditPoolService: self._create_pool(db_session_with_containers, tenant_id=tenant_id, quota_limit=10, quota_used=9) with pytest.raises(QuotaExceededError, match="Insufficient credits remaining"): - CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=3) + CreditPoolService.check_and_deduct_credits( + tenant_id=tenant_id, credits_required=3, session=db_session_with_containers + ) - updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id) + updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id, session=db_session_with_containers) assert updated_pool is not None assert updated_pool.quota_used == 9 @@ -171,9 +184,11 @@ class TestCreditPoolService: patch.object(CreditPoolService, "_get_locked_pool", side_effect=RuntimeError("database unavailable")), pytest.raises(QuotaExceededError, match="Failed to deduct credits"), ): - CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=1) + CreditPoolService.check_and_deduct_credits( + tenant_id=tenant_id, credits_required=1, session=db_session_with_containers + ) - updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id) + updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id, session=db_session_with_containers) assert updated_pool is not None assert updated_pool.quota_used == 2 @@ -181,10 +196,12 @@ class TestCreditPoolService: tenant_id = self._create_tenant_id() self._create_pool(db_session_with_containers, tenant_id=tenant_id, quota_limit=10, quota_used=9) - result = CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=3) + result = CreditPoolService.deduct_credits_capped( + tenant_id=tenant_id, credits_required=3, session=db_session_with_containers + ) assert result == 1 - updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id) + updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id, session=db_session_with_containers) assert updated_pool is not None assert updated_pool.quota_used == 10 @@ -194,16 +211,19 @@ class TestCreditPoolService: tenant_id = self._create_tenant_id() self._create_pool(db_session_with_containers, tenant_id=tenant_id, quota_limit=10, quota_used=2) - result = CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=0) + result = CreditPoolService.deduct_credits_capped( + tenant_id=tenant_id, credits_required=0, session=db_session_with_containers + ) assert result == 0 - updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id) + updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id, session=db_session_with_containers) assert updated_pool is not None assert updated_pool.quota_used == 2 - @pytest.mark.usefixtures("flask_app_with_containers") - def test_deduct_credits_capped_returns_zero_when_no_pool(self) -> None: - result = CreditPoolService.deduct_credits_capped(tenant_id=self._create_tenant_id(), credits_required=1) + def test_deduct_credits_capped_returns_zero_when_no_pool(self, db_session_with_containers: Session) -> None: + result = CreditPoolService.deduct_credits_capped( + tenant_id=self._create_tenant_id(), credits_required=1, session=db_session_with_containers + ) assert result == 0 @@ -211,10 +231,12 @@ class TestCreditPoolService: tenant_id = self._create_tenant_id() self._create_pool(db_session_with_containers, tenant_id=tenant_id, quota_limit=10, quota_used=10) - result = CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=1) + result = CreditPoolService.deduct_credits_capped( + tenant_id=tenant_id, credits_required=1, session=db_session_with_containers + ) assert result == 0 - updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id) + updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id, session=db_session_with_containers) assert updated_pool is not None assert updated_pool.quota_used == 10 @@ -226,9 +248,11 @@ class TestCreditPoolService: patch.object(CreditPoolService, "_get_locked_pool", side_effect=RuntimeError("database unavailable")), pytest.raises(QuotaExceededError, match="Failed to deduct credits"), ): - CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=1) + CreditPoolService.deduct_credits_capped( + tenant_id=tenant_id, credits_required=1, session=db_session_with_containers + ) - updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id) + updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id, session=db_session_with_containers) assert updated_pool is not None assert updated_pool.quota_used == 2 @@ -240,8 +264,10 @@ class TestCreditPoolService: patch.object(CreditPoolService, "_get_locked_pool", side_effect=QuotaExceededError("quota unavailable")), pytest.raises(QuotaExceededError, match="quota unavailable"), ): - CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=1) + CreditPoolService.deduct_credits_capped( + tenant_id=tenant_id, credits_required=1, session=db_session_with_containers + ) - updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id) + updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id, session=db_session_with_containers) assert updated_pool is not None assert updated_pool.quota_used == 2 diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service.py b/api/tests/test_containers_integration_tests/services/test_dataset_service.py index 40c00267043..912e00b0b7d 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service.py @@ -602,10 +602,7 @@ class TestDatasetServiceUpdateAndDeleteDataset: # Act / Assert with pytest.raises(ValueError, match="Dataset name already exists"): DatasetService.update_dataset( - db_session_with_containers, - source_dataset.id, - {"name": "Existing Dataset"}, - account, + source_dataset.id, {"name": "Existing Dataset"}, account, session=db_session_with_containers ) def test_delete_dataset_with_documents_success(self, db_session_with_containers: Session): @@ -728,7 +725,7 @@ class TestDatasetServiceRetrievalConfiguration: } # Act - result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, account) + result = DatasetService.update_dataset(dataset.id, update_data, account, session=db_session_with_containers) # Assert db_session_with_containers.refresh(dataset) diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py b/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py index ced144e8d6e..6b32273624b 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py @@ -574,11 +574,7 @@ class TestDatasetPermissionServiceIntegration: with pytest.raises(NoPermissionError, match="does not have permission"): DatasetPermissionService.check_permission( - db_session_with_containers, - user, - dataset, - DatasetPermissionEnum.ALL_TEAM, - [], + user, dataset, DatasetPermissionEnum.ALL_TEAM, [], session=db_session_with_containers ) def test_check_permission_prevents_dataset_operator_from_changing_permission_mode( @@ -589,11 +585,7 @@ class TestDatasetPermissionServiceIntegration: with pytest.raises(NoPermissionError, match="cannot change the dataset permissions"): DatasetPermissionService.check_permission( - db_session_with_containers, - user, - dataset, - DatasetPermissionEnum.ONLY_ME, - [], + user, dataset, DatasetPermissionEnum.ONLY_ME, [], session=db_session_with_containers ) def test_check_permission_requires_partial_member_list_for_partial_members_mode( @@ -604,11 +596,7 @@ class TestDatasetPermissionServiceIntegration: with pytest.raises(ValueError, match="Partial member list is required"): DatasetPermissionService.check_permission( - db_session_with_containers, - user, - dataset, - DatasetPermissionEnum.PARTIAL_TEAM, - [], + user, dataset, DatasetPermissionEnum.PARTIAL_TEAM, [], session=db_session_with_containers ) def test_check_permission_rejects_dataset_operator_member_list_changes(self, db_session_with_containers: Session): @@ -618,11 +606,11 @@ class TestDatasetPermissionServiceIntegration: with patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=["user-1"]): with pytest.raises(ValueError, match="cannot change the dataset permissions"): DatasetPermissionService.check_permission( - db_session_with_containers, user, dataset, DatasetPermissionEnum.PARTIAL_TEAM, [{"user_id": "user-2"}], + session=db_session_with_containers, ) def test_check_permission_allows_dataset_operator_when_member_list_is_unchanged( @@ -633,11 +621,11 @@ class TestDatasetPermissionServiceIntegration: with patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=["user-1"]): DatasetPermissionService.check_permission( - db_session_with_containers, user, dataset, DatasetPermissionEnum.PARTIAL_TEAM, [{"user_id": "user-1"}], + session=db_session_with_containers, ) def test_clear_partial_member_list_deletes_permissions_and_commits(self, db_session_with_containers: Session): diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service_update_dataset.py b/api/tests/test_containers_integration_tests/services/test_dataset_service_update_dataset.py index f719a465dbd..d9fb23e8e33 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service_update_dataset.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service_update_dataset.py @@ -189,7 +189,7 @@ class TestDatasetServiceUpdateDataset: "external_knowledge_api_id": external_api.id, } - result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) db_session_with_containers.refresh(dataset) updated_binding = db_session_with_containers.query(ExternalKnowledgeBindings).filter_by(id=binding_id).first() @@ -221,7 +221,7 @@ class TestDatasetServiceUpdateDataset: update_data = {"name": "new_name", "external_knowledge_api_id": str(uuid4())} with pytest.raises(ValueError) as context: - DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) assert "External knowledge id is required" in str(context.value) db_session_with_containers.rollback() @@ -245,7 +245,7 @@ class TestDatasetServiceUpdateDataset: update_data = {"name": "new_name", "external_knowledge_id": "knowledge_id"} with pytest.raises(ValueError) as context: - DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) assert "External knowledge api id is required" in str(context.value) db_session_with_containers.rollback() @@ -272,7 +272,7 @@ class TestDatasetServiceUpdateDataset: } with pytest.raises(ValueError) as context: - DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) assert "External knowledge binding not found" in str(context.value) db_session_with_containers.rollback() @@ -303,7 +303,7 @@ class TestDatasetServiceUpdateDataset: "embedding_model": "text-embedding-ada-002", } - result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) db_session_with_containers.refresh(dataset) assert dataset.name == "new_name" @@ -338,7 +338,7 @@ class TestDatasetServiceUpdateDataset: "embedding_model": None, } - result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) db_session_with_containers.refresh(dataset) assert dataset.name == "new_name" @@ -371,7 +371,7 @@ class TestDatasetServiceUpdateDataset: } with patch("services.dataset_service.deal_dataset_vector_index_task") as mock_task: - result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) mock_task.delay.assert_called_once_with(dataset.id, "remove") db_session_with_containers.refresh(dataset) @@ -418,7 +418,7 @@ class TestDatasetServiceUpdateDataset: mock_model_manager.return_value.get_model_instance.return_value = embedding_model mock_get_binding.return_value = binding - result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) mock_model_manager.return_value.get_model_instance.assert_called_once_with( tenant_id=tenant.id, @@ -462,7 +462,7 @@ class TestDatasetServiceUpdateDataset: "retrieval_model": "new_model", } - result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) db_session_with_containers.refresh(dataset) assert dataset.name == "new_name" @@ -514,7 +514,7 @@ class TestDatasetServiceUpdateDataset: mock_model_manager.return_value.get_model_instance.return_value = embedding_model mock_get_binding.return_value = binding - result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) mock_model_manager.return_value.get_model_instance.assert_called_once_with( tenant_id=tenant.id, @@ -545,7 +545,7 @@ class TestDatasetServiceUpdateDataset: update_data = {"name": "new_name"} with pytest.raises(ValueError) as context: - DatasetService.update_dataset(db_session_with_containers, str(uuid4()), update_data, user) + DatasetService.update_dataset(str(uuid4()), update_data, user, session=db_session_with_containers) assert "Dataset not found" in str(context.value) @@ -568,7 +568,7 @@ class TestDatasetServiceUpdateDataset: update_data = {"name": "new_name"} with pytest.raises(NoPermissionError): - DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, outsider) + DatasetService.update_dataset(dataset.id, update_data, outsider, session=db_session_with_containers) def test_update_internal_dataset_embedding_model_error(self, db_session_with_containers: Session): """Test error when embedding model is not available.""" @@ -595,6 +595,6 @@ class TestDatasetServiceUpdateDataset: mock_model_manager.return_value.get_model_instance.side_effect = Exception("No Embedding Model available") with pytest.raises(Exception) as context: - DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) assert "No Embedding Model available".lower() in str(context.value).lower() diff --git a/api/tests/test_containers_integration_tests/services/test_file_service_zip_and_lookup.py b/api/tests/test_containers_integration_tests/services/test_file_service_zip_and_lookup.py index 5eb84f805aa..a5445a17297 100644 --- a/api/tests/test_containers_integration_tests/services/test_file_service_zip_and_lookup.py +++ b/api/tests/test_containers_integration_tests/services/test_file_service_zip_and_lookup.py @@ -69,7 +69,7 @@ def test_build_upload_files_zip_tempfile_sanitizes_and_dedupes_names(monkeypatch def test_get_upload_files_by_ids_returns_empty_when_no_ids(db_session_with_containers: Session) -> None: """Ensure empty input returns an empty mapping without hitting the database.""" - assert FileService.get_upload_files_by_ids(db_session_with_containers, str(uuid4()), []) == {} + assert FileService.get_upload_files_by_ids(str(uuid4()), [], session=db_session_with_containers) == {} def test_get_upload_files_by_ids_returns_id_keyed_mapping(db_session_with_containers: Session) -> None: @@ -78,7 +78,9 @@ def test_get_upload_files_by_ids_returns_id_keyed_mapping(db_session_with_contai file1 = _create_upload_file(db_session_with_containers, tenant_id=tenant_id, key="k1", name="file1.txt") file2 = _create_upload_file(db_session_with_containers, tenant_id=tenant_id, key="k2", name="file2.txt") - result = FileService.get_upload_files_by_ids(db_session_with_containers, tenant_id, [file1.id, file1.id, file2.id]) + result = FileService.get_upload_files_by_ids( + tenant_id, [file1.id, file1.id, file2.id], session=db_session_with_containers + ) assert set(result.keys()) == {file1.id, file2.id} assert result[file1.id].id == file1.id @@ -92,6 +94,6 @@ def test_get_upload_files_by_ids_filters_by_tenant(db_session_with_containers: S file_a = _create_upload_file(db_session_with_containers, tenant_id=tenant_a, key="ka", name="a.txt") _create_upload_file(db_session_with_containers, tenant_id=tenant_b, key="kb", name="b.txt") - result = FileService.get_upload_files_by_ids(db_session_with_containers, tenant_a, [file_a.id]) + result = FileService.get_upload_files_by_ids(tenant_a, [file_a.id], session=db_session_with_containers) assert set(result.keys()) == {file_a.id} diff --git a/api/tests/test_containers_integration_tests/services/test_hit_testing_service.py b/api/tests/test_containers_integration_tests/services/test_hit_testing_service.py index 4a73f98f50e..67b3a2d3e57 100644 --- a/api/tests/test_containers_integration_tests/services/test_hit_testing_service.py +++ b/api/tests/test_containers_integration_tests/services/test_hit_testing_service.py @@ -192,7 +192,7 @@ class TestHitTestingService: mock_format.return_value = [mock_record] response = _RetrieveResponse.model_validate( - HitTestingService.compact_retrieve_response(db_session_with_containers, query, [mock_doc]) + HitTestingService.compact_retrieve_response(query, [mock_doc], session=db_session_with_containers) ) assert response.query.content == query @@ -246,12 +246,12 @@ class TestHitTestingService: response = _RetrieveResponse.model_validate( HitTestingService.external_retrieve( - db_session_with_containers, dataset=dataset, query='test "query"', account=account, external_retrieval_model={"model": "test"}, metadata_filtering_conditions={"key": "val"}, + session=db_session_with_containers, ) ) @@ -276,7 +276,7 @@ class TestHitTestingService: account = MagicMock() response = _RetrieveResponse.model_validate( - HitTestingService.external_retrieve(db_session_with_containers, dataset, "test query", account) + HitTestingService.external_retrieve(dataset, "test query", account, session=db_session_with_containers) ) assert response.query.content == "test query" @@ -300,12 +300,12 @@ class TestHitTestingService: response = _RetrieveResponse.model_validate( HitTestingService.retrieve( - db_session_with_containers, dataset=dataset, query="test query", account=account, retrieval_model=None, external_retrieval_model=external_retrieval_model, + session=db_session_with_containers, ) ) @@ -343,12 +343,12 @@ class TestHitTestingService: mock_retrieve.return_value = retrieved_documents HitTestingService.retrieve( - db_session_with_containers, dataset=dataset, query="test query", account=account, retrieval_model=retrieval_model, external_retrieval_model=external_retrieval_model, + session=db_session_with_containers, ) mock_get_meta.assert_called_once() @@ -380,12 +380,12 @@ class TestHitTestingService: response = _RetrieveResponse.model_validate( HitTestingService.retrieve( - db_session_with_containers, dataset=dataset, query="test query", account=account, retrieval_model=retrieval_model, external_retrieval_model=external_retrieval_model, + session=db_session_with_containers, ) ) @@ -412,13 +412,13 @@ class TestHitTestingService: mock_retrieve.return_value = retrieved_documents HitTestingService.retrieve( - db_session_with_containers, dataset=dataset, query="test query", account=account, retrieval_model=retrieval_model, external_retrieval_model=external_retrieval_model, attachment_ids=attachment_ids, + session=db_session_with_containers, ) mock_retrieve.assert_called_once_with( @@ -472,12 +472,12 @@ class TestHitTestingService: mock_retrieve.return_value = retrieved_documents HitTestingService.retrieve( - db_session_with_containers, dataset=dataset, query="test query", account=account, retrieval_model=retrieval_model, external_retrieval_model=external_retrieval_model, + session=db_session_with_containers, ) mock_retrieve.assert_called_once() diff --git a/api/tests/test_containers_integration_tests/services/test_human_input_delivery_test.py b/api/tests/test_containers_integration_tests/services/test_human_input_delivery_test.py index c1188d3d0f9..84a0226ba17 100644 --- a/api/tests/test_containers_integration_tests/services/test_human_input_delivery_test.py +++ b/api/tests/test_containers_integration_tests/services/test_human_input_delivery_test.py @@ -119,6 +119,7 @@ def test_human_input_delivery_test_sends_email( account=account, node_id="human-node", delivery_method_id=str(delivery_method_id), + session=db_session_with_containers, ) assert send_mock.call_count == 1 @@ -145,6 +146,7 @@ def test_human_input_delivery_test_form_accepts_file_upload( account=account, node_id="human-node", delivery_method_id=str(delivery_method_id), + session=db_session_with_containers, ) form = db_session_with_containers.scalar( @@ -213,6 +215,7 @@ def test_human_input_delivery_test_form_accepts_remote_file_upload( account=account, node_id="human-node", delivery_method_id=str(delivery_method_id), + session=db_session_with_containers, ) form = db_session_with_containers.scalar( diff --git a/api/tests/test_containers_integration_tests/services/test_message_service.py b/api/tests/test_containers_integration_tests/services/test_message_service.py index f2d682be3bf..702812b96de 100644 --- a/api/tests/test_containers_integration_tests/services/test_message_service.py +++ b/api/tests/test_containers_integration_tests/services/test_message_service.py @@ -1,4 +1,4 @@ -from unittest.mock import patch +from unittest.mock import ANY, patch import pytest from faker import Faker @@ -117,7 +117,7 @@ class TestMessageService: # Create app app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Setup current_user mock self._mock_current_user(mock_external_service_dependencies, account.id, tenant.id) @@ -222,6 +222,7 @@ class TestMessageService: first_id=messages[2].id, # Use middle message as first_id limit=2, order="asc", + session=db_session_with_containers, ) # Verify results @@ -243,7 +244,12 @@ class TestMessageService: # Test pagination with no user result = MessageService.pagination_by_first_id( - app_model=app, user=None, conversation_id=fake.uuid4(), first_id=None, limit=10 + app_model=app, + user=None, + conversation_id=fake.uuid4(), + first_id=None, + limit=10, + session=db_session_with_containers, ) # Verify empty result @@ -262,7 +268,12 @@ class TestMessageService: # Test pagination with no conversation ID result = MessageService.pagination_by_first_id( - app_model=app, user=account, conversation_id="", first_id=None, limit=10 + app_model=app, + user=account, + conversation_id="", + first_id=None, + limit=10, + session=db_session_with_containers, ) # Verify empty result @@ -291,6 +302,7 @@ class TestMessageService: conversation_id=conversation.id, first_id=fake.uuid4(), # Non-existent message ID limit=10, + session=db_session_with_containers, ) def test_pagination_by_last_id_success( @@ -316,6 +328,7 @@ class TestMessageService: last_id=messages[2].id, # Use middle message as last_id limit=2, conversation_id=conversation.id, + session=db_session_with_containers, ) # Verify results @@ -345,7 +358,12 @@ class TestMessageService: # Test pagination with include_ids include_ids = [messages[0].id, messages[1].id, messages[2].id] result = MessageService.pagination_by_last_id( - app_model=app, user=account, last_id=messages[1].id, limit=2, include_ids=include_ids + app_model=app, + user=account, + last_id=messages[1].id, + limit=2, + include_ids=include_ids, + session=db_session_with_containers, ) # Verify results @@ -364,8 +382,10 @@ class TestMessageService: fake = Faker() app, account = self._create_test_app_and_account(db_session_with_containers, mock_external_service_dependencies) - # Test pagination with no user - result = MessageService.pagination_by_last_id(app_model=app, user=None, last_id=None, limit=10) + # Test pagination with no user, + result = MessageService.pagination_by_last_id( + app_model=app, user=None, last_id=None, limit=10, session=db_session_with_containers + ) # Verify empty result assert result.limit == 10 @@ -393,6 +413,7 @@ class TestMessageService: last_id=fake.uuid4(), # Non-existent message ID limit=10, conversation_id=conversation.id, + session=db_session_with_containers, ) def test_create_feedback_success(self, db_session_with_containers: Session, mock_external_service_dependencies): @@ -410,7 +431,12 @@ class TestMessageService: rating = FeedbackRating.LIKE content = fake.text(max_nb_chars=100) feedback = MessageService.create_feedback( - app_model=app, message_id=message.id, user=account, rating=rating, content=content + app_model=app, + message_id=message.id, + user=account, + rating=rating, + content=content, + session=db_session_with_containers, ) # Verify feedback was created correctly @@ -442,6 +468,7 @@ class TestMessageService: user=None, rating=FeedbackRating.LIKE, content=fake.text(max_nb_chars=100), + session=db_session_with_containers, ) def test_create_feedback_update_existing( @@ -461,14 +488,24 @@ class TestMessageService: initial_rating = FeedbackRating.LIKE initial_content = fake.text(max_nb_chars=100) feedback = MessageService.create_feedback( - app_model=app, message_id=message.id, user=account, rating=initial_rating, content=initial_content + app_model=app, + message_id=message.id, + user=account, + rating=initial_rating, + content=initial_content, + session=db_session_with_containers, ) # Update feedback updated_rating = FeedbackRating.DISLIKE updated_content = fake.text(max_nb_chars=100) updated_feedback = MessageService.create_feedback( - app_model=app, message_id=message.id, user=account, rating=updated_rating, content=updated_content + app_model=app, + message_id=message.id, + user=account, + rating=updated_rating, + content=updated_content, + session=db_session_with_containers, ) # Verify feedback was updated correctly @@ -498,10 +535,18 @@ class TestMessageService: user=account, rating=FeedbackRating.LIKE, content=fake.text(max_nb_chars=100), + session=db_session_with_containers, ) - # Delete feedback by setting rating to None - MessageService.create_feedback(app_model=app, message_id=message.id, user=account, rating=None, content=None) + # Delete feedback by setting rating to None, + MessageService.create_feedback( + app_model=app, + message_id=message.id, + user=account, + rating=None, + content=None, + session=db_session_with_containers, + ) # Verify feedback was deleted @@ -526,7 +571,12 @@ class TestMessageService: # Test creating feedback with no rating when no feedback exists with pytest.raises(ValueError, match="rating cannot be None when feedback not exists"): MessageService.create_feedback( - app_model=app, message_id=message.id, user=account, rating=None, content=None + app_model=app, + message_id=message.id, + user=account, + rating=None, + content=None, + session=db_session_with_containers, ) def test_get_all_messages_feedbacks_success( @@ -550,11 +600,12 @@ class TestMessageService: user=account, rating=FeedbackRating.LIKE if i % 2 == 0 else FeedbackRating.DISLIKE, content=f"Feedback {i}: {fake.text(max_nb_chars=50)}", + session=db_session_with_containers, ) feedbacks.append(feedback) - # Get all feedbacks - result = MessageService.get_all_messages_feedbacks(app, page=1, limit=10) + # Get all feedbacks, + result = MessageService.get_all_messages_feedbacks(app, page=1, limit=10, session=db_session_with_containers) # Verify results assert len(result) == 3 @@ -583,11 +634,16 @@ class TestMessageService: user=account, rating=FeedbackRating.LIKE, content=f"Feedback {i}", + session=db_session_with_containers, ) # Get feedbacks with pagination - result_page_1 = MessageService.get_all_messages_feedbacks(app, page=1, limit=3) - result_page_2 = MessageService.get_all_messages_feedbacks(app, page=2, limit=3) + result_page_1 = MessageService.get_all_messages_feedbacks( + app, page=1, limit=3, session=db_session_with_containers + ) + result_page_2 = MessageService.get_all_messages_feedbacks( + app, page=2, limit=3, session=db_session_with_containers + ) # Verify pagination results assert len(result_page_1) == 3 @@ -609,8 +665,10 @@ class TestMessageService: conversation = self._create_test_conversation(db_session_with_containers, app, account, fake) message = self._create_test_message(db_session_with_containers, app, conversation, account, fake) - # Get message - retrieved_message = MessageService.get_message(app_model=app, user=account, message_id=message.id) + # Get message, + retrieved_message = MessageService.get_message( + app_model=app, user=account, message_id=message.id, session=db_session_with_containers + ) # Verify message was retrieved correctly assert retrieved_message.id == message.id @@ -628,7 +686,9 @@ class TestMessageService: # Test getting non-existent message with pytest.raises(MessageNotExistsError): - MessageService.get_message(app_model=app, user=account, message_id=fake.uuid4()) + MessageService.get_message( + app_model=app, user=account, message_id=fake.uuid4(), session=db_session_with_containers + ) def test_get_message_wrong_user(self, db_session_with_containers: Session, mock_external_service_dependencies): """ @@ -657,7 +717,9 @@ class TestMessageService: # Test getting message with different user with pytest.raises(MessageNotExistsError): - MessageService.get_message(app_model=app, user=other_account, message_id=message.id) + MessageService.get_message( + app_model=app, user=other_account, message_id=message.id, session=db_session_with_containers + ) def test_get_suggested_questions_after_answer_success( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -682,7 +744,11 @@ class TestMessageService: from core.app.entities.app_invoke_entities import InvokeFrom result = MessageService.get_suggested_questions_after_answer( - app_model=app, user=account, message_id=message.id, invoke_from=InvokeFrom.SERVICE_API + app_model=app, + user=account, + message_id=message.id, + invoke_from=InvokeFrom.SERVICE_API, + session=db_session_with_containers, ) # Verify results @@ -714,7 +780,11 @@ class TestMessageService: with pytest.raises(ValueError, match="user cannot be None"): MessageService.get_suggested_questions_after_answer( - app_model=app, user=None, message_id=message.id, invoke_from=InvokeFrom.SERVICE_API + app_model=app, + user=None, + message_id=message.id, + invoke_from=InvokeFrom.SERVICE_API, + session=db_session_with_containers, ) def test_get_suggested_questions_after_answer_disabled( @@ -740,7 +810,11 @@ class TestMessageService: with pytest.raises(SuggestedQuestionsAfterAnswerDisabledError): MessageService.get_suggested_questions_after_answer( - app_model=app, user=account, message_id=message.id, invoke_from=InvokeFrom.SERVICE_API + app_model=app, + user=account, + message_id=message.id, + invoke_from=InvokeFrom.SERVICE_API, + session=db_session_with_containers, ) def test_get_suggested_questions_after_answer_no_workflow( @@ -763,7 +837,11 @@ class TestMessageService: from core.app.entities.app_invoke_entities import InvokeFrom result = MessageService.get_suggested_questions_after_answer( - app_model=app, user=account, message_id=message.id, invoke_from=InvokeFrom.SERVICE_API + app_model=app, + user=account, + message_id=message.id, + invoke_from=InvokeFrom.SERVICE_API, + session=db_session_with_containers, ) # Verify empty result @@ -792,7 +870,11 @@ class TestMessageService: from core.app.entities.app_invoke_entities import InvokeFrom result = MessageService.get_suggested_questions_after_answer( - app_model=app, user=account, message_id=message.id, invoke_from=InvokeFrom.DEBUGGER + app_model=app, + user=account, + message_id=message.id, + invoke_from=InvokeFrom.DEBUGGER, + session=db_session_with_containers, ) # Verify results @@ -800,7 +882,7 @@ class TestMessageService: # Verify draft workflow was used instead of published workflow mock_external_service_dependencies["workflow_service"].return_value.get_draft_workflow.assert_called_once_with( - app_model=app + app_model=app, session=ANY ) # Verify TraceQueueManager was called diff --git a/api/tests/test_containers_integration_tests/services/test_message_service_execution_extra_content.py b/api/tests/test_containers_integration_tests/services/test_message_service_execution_extra_content.py index 6a9046acd4a..9da20d96f06 100644 --- a/api/tests/test_containers_integration_tests/services/test_message_service_execution_extra_content.py +++ b/api/tests/test_containers_integration_tests/services/test_message_service_execution_extra_content.py @@ -20,6 +20,7 @@ def test_pagination_returns_extra_contents(db_session_with_containers: Session): conversation_id=fixture.conversation.id, first_id=None, limit=10, + session=db_session_with_containers, ) assert pagination.data @@ -59,6 +60,7 @@ def test_pagination_returns_waiting_human_input_extra_contents(db_session_with_c conversation_id=fixture.conversation.id, first_id=None, limit=10, + session=db_session_with_containers, ) assert pagination.data diff --git a/api/tests/test_containers_integration_tests/services/test_metadata_partial_update.py b/api/tests/test_containers_integration_tests/services/test_metadata_partial_update.py index fbdc265265d..a9399985307 100644 --- a/api/tests/test_containers_integration_tests/services/test_metadata_partial_update.py +++ b/api/tests/test_containers_integration_tests/services/test_metadata_partial_update.py @@ -95,7 +95,9 @@ class TestMetadataPartialUpdate: ) metadata_args = MetadataOperationData(operation_data=[operation]) - MetadataService.update_documents_metadata(db_session_with_containers, dataset, metadata_args, current_account) + MetadataService.update_documents_metadata( + dataset, metadata_args, current_account, session=db_session_with_containers + ) db_session_with_containers.expire_all() updated_doc = db_session_with_containers.get(Document, document.id) @@ -126,7 +128,9 @@ class TestMetadataPartialUpdate: ) metadata_args = MetadataOperationData(operation_data=[operation]) - MetadataService.update_documents_metadata(db_session_with_containers, dataset, metadata_args, current_account) + MetadataService.update_documents_metadata( + dataset, metadata_args, current_account, session=db_session_with_containers + ) db_session_with_containers.expire_all() updated_doc = db_session_with_containers.get(Document, document.id) @@ -168,7 +172,9 @@ class TestMetadataPartialUpdate: ) metadata_args = MetadataOperationData(operation_data=[operation]) - MetadataService.update_documents_metadata(db_session_with_containers, dataset, metadata_args, current_account) + MetadataService.update_documents_metadata( + dataset, metadata_args, current_account, session=db_session_with_containers + ) db_session_with_containers.expire_all() bindings = db_session_with_containers.scalars( @@ -205,5 +211,5 @@ class TestMetadataPartialUpdate: with patch.object(db_session_with_containers, "commit", side_effect=RuntimeError("database connection lost")): with pytest.raises(RuntimeError, match="database connection lost"): MetadataService.update_documents_metadata( - db_session_with_containers, dataset, metadata_args, current_account + dataset, metadata_args, current_account, session=db_session_with_containers ) diff --git a/api/tests/test_containers_integration_tests/services/test_metadata_service.py b/api/tests/test_containers_integration_tests/services/test_metadata_service.py index 7cc9fc7e696..00afe7f8467 100644 --- a/api/tests/test_containers_integration_tests/services/test_metadata_service.py +++ b/api/tests/test_containers_integration_tests/services/test_metadata_service.py @@ -184,7 +184,7 @@ class TestMetadataService: # Act: Execute the method under test result = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Assert: Verify the expected outcomes @@ -220,7 +220,9 @@ class TestMetadataService: # Act & Assert: Verify proper error handling with pytest.raises(ValueError, match="Metadata name cannot exceed 255 characters."): - MetadataService.create_metadata(db_session_with_containers, dataset.id, metadata_args, account, tenant.id) + MetadataService.create_metadata( + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers + ) def test_create_metadata_name_already_exists( self, db_session_with_containers: Session, mock_external_service_dependencies: MetadataServiceDeps @@ -238,7 +240,9 @@ class TestMetadataService: # Create first metadata first_metadata_args = MetadataArgs(type="string", name="duplicate_name") - MetadataService.create_metadata(db_session_with_containers, dataset.id, first_metadata_args, account, tenant.id) + MetadataService.create_metadata( + dataset.id, first_metadata_args, account, tenant.id, session=db_session_with_containers + ) # Try to create second metadata with same name second_metadata_args = MetadataArgs(type="number", name="duplicate_name") @@ -246,7 +250,7 @@ class TestMetadataService: # Act & Assert: Verify proper error handling with pytest.raises(ValueError, match="Metadata name already exists."): MetadataService.create_metadata( - db_session_with_containers, dataset.id, second_metadata_args, account, tenant.id + dataset.id, second_metadata_args, account, tenant.id, session=db_session_with_containers ) def test_create_metadata_name_conflicts_with_built_in_field( @@ -269,7 +273,9 @@ class TestMetadataService: # Act & Assert: Verify proper error handling with pytest.raises(ValueError, match="Metadata name already exists in Built-in fields."): - MetadataService.create_metadata(db_session_with_containers, dataset.id, metadata_args, account, tenant.id) + MetadataService.create_metadata( + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers + ) def test_update_metadata_name_success( self, db_session_with_containers: Session, mock_external_service_dependencies: MetadataServiceDeps @@ -288,13 +294,13 @@ class TestMetadataService: # Create metadata first metadata_args = MetadataArgs(type="string", name="old_name") metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Act: Execute the method under test new_name = "new_name" result = MetadataService.update_metadata_name( - db_session_with_containers, dataset.id, metadata.id, new_name, account, tenant.id + dataset.id, metadata.id, new_name, account, tenant.id, session=db_session_with_containers ) # Assert: Verify the expected outcomes @@ -325,7 +331,7 @@ class TestMetadataService: # Create metadata first metadata_args = MetadataArgs(type="string", name="old_name") metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Try to update with too long name @@ -334,7 +340,7 @@ class TestMetadataService: # Act & Assert: Verify proper error handling with pytest.raises(ValueError, match="Metadata name cannot exceed 255 characters."): MetadataService.update_metadata_name( - db_session_with_containers, dataset.id, metadata.id, long_name, account, tenant.id + dataset.id, metadata.id, long_name, account, tenant.id, session=db_session_with_containers ) def test_update_metadata_name_already_exists( @@ -354,18 +360,18 @@ class TestMetadataService: # Create two metadata entries first_metadata_args = MetadataArgs(type="string", name="first_metadata") first_metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, first_metadata_args, account, tenant.id + dataset.id, first_metadata_args, account, tenant.id, session=db_session_with_containers ) second_metadata_args = MetadataArgs(type="number", name="second_metadata") second_metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, second_metadata_args, account, tenant.id + dataset.id, second_metadata_args, account, tenant.id, session=db_session_with_containers ) # Try to update first metadata with second metadata's name with pytest.raises(ValueError, match="Metadata name already exists."): MetadataService.update_metadata_name( - db_session_with_containers, dataset.id, first_metadata.id, "second_metadata", account, tenant.id + dataset.id, first_metadata.id, "second_metadata", account, tenant.id, session=db_session_with_containers ) def test_update_metadata_name_conflicts_with_built_in_field( @@ -385,7 +391,7 @@ class TestMetadataService: # Create metadata first metadata_args = MetadataArgs(type="string", name="old_name") metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Try to update with built-in field name @@ -393,7 +399,7 @@ class TestMetadataService: with pytest.raises(ValueError, match="Metadata name already exists in Built-in fields."): MetadataService.update_metadata_name( - db_session_with_containers, dataset.id, metadata.id, built_in_field_name, account, tenant.id + dataset.id, metadata.id, built_in_field_name, account, tenant.id, session=db_session_with_containers ) def test_update_metadata_name_not_found( @@ -418,7 +424,7 @@ class TestMetadataService: # Act: Execute the method under test result = MetadataService.update_metadata_name( - db_session_with_containers, dataset.id, fake_metadata_id, new_name, account, tenant.id + dataset.id, fake_metadata_id, new_name, account, tenant.id, session=db_session_with_containers ) # Assert: Verify the method returns None when metadata is not found @@ -441,11 +447,11 @@ class TestMetadataService: # Create metadata first metadata_args = MetadataArgs(type="string", name="to_be_deleted") metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Act: Execute the method under test - result = MetadataService.delete_metadata(db_session_with_containers, dataset.id, metadata.id) + result = MetadataService.delete_metadata(dataset.id, metadata.id, session=db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -476,7 +482,7 @@ class TestMetadataService: fake_metadata_id = str(uuid.uuid4()) # Use valid UUID format # Act: Execute the method under test - result = MetadataService.delete_metadata(db_session_with_containers, dataset.id, fake_metadata_id) + result = MetadataService.delete_metadata(dataset.id, fake_metadata_id, session=db_session_with_containers) # Assert: Verify the method returns None when metadata is not found assert result is None @@ -501,7 +507,7 @@ class TestMetadataService: # Create metadata metadata_args = MetadataArgs(type="string", name="test_metadata") metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Create metadata binding @@ -522,7 +528,7 @@ class TestMetadataService: db_session_with_containers.commit() # Act: Execute the method under test - result = MetadataService.delete_metadata(db_session_with_containers, dataset.id, metadata.id) + result = MetadataService.delete_metadata(dataset.id, metadata.id, session=db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -587,7 +593,7 @@ class TestMetadataService: assert dataset.built_in_field_enabled is False # Act: Execute the method under test - MetadataService.enable_built_in_field(db_session_with_containers, dataset) + MetadataService.enable_built_in_field(dataset, session=db_session_with_containers) # Assert: Verify the expected outcomes @@ -623,7 +629,7 @@ class TestMetadataService: ]() # Act: Execute the method under test - MetadataService.enable_built_in_field(db_session_with_containers, dataset) + MetadataService.enable_built_in_field(dataset, session=db_session_with_containers) # Assert: Verify the method returns early without changes db_session_with_containers.refresh(dataset) @@ -649,7 +655,7 @@ class TestMetadataService: ]() # Act: Execute the method under test - MetadataService.enable_built_in_field(db_session_with_containers, dataset) + MetadataService.enable_built_in_field(dataset, session=db_session_with_containers) # Assert: Verify the expected outcomes @@ -696,7 +702,7 @@ class TestMetadataService: ] # Act: Execute the method under test - MetadataService.disable_built_in_field(db_session_with_containers, dataset) + MetadataService.disable_built_in_field(dataset, session=db_session_with_containers) # Assert: Verify the expected outcomes db_session_with_containers.refresh(dataset) @@ -728,7 +734,7 @@ class TestMetadataService: ]() # Act: Execute the method under test - MetadataService.disable_built_in_field(db_session_with_containers, dataset) + MetadataService.disable_built_in_field(dataset, session=db_session_with_containers) # Assert: Verify the method returns early without changes @@ -761,7 +767,7 @@ class TestMetadataService: ]() # Act: Execute the method under test - MetadataService.disable_built_in_field(db_session_with_containers, dataset) + MetadataService.disable_built_in_field(dataset, session=db_session_with_containers) # Assert: Verify the expected outcomes db_session_with_containers.refresh(dataset) @@ -787,7 +793,7 @@ class TestMetadataService: # Create metadata metadata_args = MetadataArgs(type="string", name="test_metadata") metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Mock DocumentService.get_document @@ -807,7 +813,7 @@ class TestMetadataService: operation_data = MetadataOperationData(operation_data=[operation]) # Act: Execute the method under test - MetadataService.update_documents_metadata(db_session_with_containers, dataset, operation_data, account) + MetadataService.update_documents_metadata(dataset, operation_data, account, session=db_session_with_containers) # Assert: Verify the expected outcomes @@ -853,7 +859,7 @@ class TestMetadataService: # Create metadata metadata_args = MetadataArgs(type="string", name="test_metadata") metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Mock DocumentService.get_document @@ -873,7 +879,7 @@ class TestMetadataService: operation_data = MetadataOperationData(operation_data=[operation]) # Act: Execute the method under test - MetadataService.update_documents_metadata(db_session_with_containers, dataset, operation_data, account) + MetadataService.update_documents_metadata(dataset, operation_data, account, session=db_session_with_containers) # Assert: Verify the expected outcomes # Verify document metadata was updated with both custom and built-in fields @@ -902,7 +908,7 @@ class TestMetadataService: # Create metadata metadata_args = MetadataArgs(type="string", name="test_metadata") metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Create metadata operation data @@ -924,7 +930,9 @@ class TestMetadataService: # Act & Assert: The method should raise ValueError("Document not found.") # because the exception is now re-raised after rollback with pytest.raises(ValueError, match="Document not found"): - MetadataService.update_documents_metadata(db_session_with_containers, dataset, operation_data, account) + MetadataService.update_documents_metadata( + dataset, operation_data, account, session=db_session_with_containers + ) def test_knowledge_base_metadata_lock_check_dataset_id( self, db_session_with_containers: Session, mock_external_service_dependencies: MetadataServiceDeps @@ -1021,7 +1029,7 @@ class TestMetadataService: # Create metadata metadata_args = MetadataArgs(type="string", name="test_metadata") metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Create document and metadata binding @@ -1041,7 +1049,7 @@ class TestMetadataService: db_session_with_containers.commit() # Act: Execute the method under test - result = MetadataService.get_dataset_metadatas(db_session_with_containers, dataset) + result = MetadataService.get_dataset_metadatas(dataset, session=db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -1082,11 +1090,11 @@ class TestMetadataService: # Create metadata metadata_args = MetadataArgs(type="string", name="test_metadata") metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Act: Execute the method under test - result = MetadataService.get_dataset_metadatas(db_session_with_containers, dataset) + result = MetadataService.get_dataset_metadatas(dataset, session=db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -1115,7 +1123,7 @@ class TestMetadataService: ) # Act: Execute the method under test - result = MetadataService.get_dataset_metadatas(db_session_with_containers, dataset) + result = MetadataService.get_dataset_metadatas(dataset, session=db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None diff --git a/api/tests/test_containers_integration_tests/services/test_model_load_balancing_service.py b/api/tests/test_containers_integration_tests/services/test_model_load_balancing_service.py index aca38391353..71d2c1c6800 100644 --- a/api/tests/test_containers_integration_tests/services/test_model_load_balancing_service.py +++ b/api/tests/test_containers_integration_tests/services/test_model_load_balancing_service.py @@ -339,7 +339,11 @@ class TestModelLoadBalancingService: # Act: Execute the method under test service = ModelLoadBalancingService() is_enabled, configs = service.get_load_balancing_configs( - tenant_id=tenant.id, provider="openai", model="gpt-3.5-turbo", model_type="llm" + tenant_id=tenant.id, + provider="openai", + model="gpt-3.5-turbo", + model_type="llm", + session=db_session_with_containers, ) # Assert: Verify the expected outcomes @@ -381,7 +385,11 @@ class TestModelLoadBalancingService: service = ModelLoadBalancingService() with pytest.raises(ValueError) as exc_info: service.get_load_balancing_configs( - tenant_id=tenant.id, provider="nonexistent_provider", model="gpt-3.5-turbo", model_type="llm" + tenant_id=tenant.id, + provider="nonexistent_provider", + model="gpt-3.5-turbo", + model_type="llm", + session=db_session_with_containers, ) # Verify correct error message @@ -443,7 +451,11 @@ class TestModelLoadBalancingService: # Act: Execute the method under test service = ModelLoadBalancingService() is_enabled, configs = service.get_load_balancing_configs( - tenant_id=tenant.id, provider="openai", model="gpt-3.5-turbo", model_type="llm" + tenant_id=tenant.id, + provider="openai", + model="gpt-3.5-turbo", + model_type="llm", + session=db_session_with_containers, ) # Assert: Verify the expected outcomes diff --git a/api/tests/test_containers_integration_tests/services/test_oauth_server_service.py b/api/tests/test_containers_integration_tests/services/test_oauth_server_service.py index 0969198ecf3..d397c62b6a8 100644 --- a/api/tests/test_containers_integration_tests/services/test_oauth_server_service.py +++ b/api/tests/test_containers_integration_tests/services/test_oauth_server_service.py @@ -4,7 +4,7 @@ from __future__ import annotations import uuid from typing import cast -from unittest.mock import ANY, MagicMock, patch +from unittest.mock import MagicMock, patch from uuid import uuid4 import pytest @@ -159,17 +159,19 @@ class TestOAuthServerServiceTokenOperations: def test_validate_access_token_returns_none_when_not_found(self, mock_redis): mock_redis.get.return_value = None + session = MagicMock() - result = OAuthServerService.validate_oauth_access_token("client-1", "missing-token") + result = OAuthServerService.validate_oauth_access_token("client-1", "missing-token", session) assert result is None def test_validate_access_token_loads_user_when_exists(self, mock_redis): mock_redis.get.return_value = b"user-88" expected_user = MagicMock() + session = MagicMock() with patch("services.oauth_server.AccountService.load_user", return_value=expected_user) as mock_load: - result = OAuthServerService.validate_oauth_access_token("client-1", "access-token") + result = OAuthServerService.validate_oauth_access_token("client-1", "access-token", session) assert result is expected_user - mock_load.assert_called_once_with("user-88", ANY) + mock_load.assert_called_once_with("user-88", session) diff --git a/api/tests/test_containers_integration_tests/services/test_ops_service.py b/api/tests/test_containers_integration_tests/services/test_ops_service.py index 9643fb61d44..b4b8521fb2e 100644 --- a/api/tests/test_containers_integration_tests/services/test_ops_service.py +++ b/api/tests/test_containers_integration_tests/services/test_ops_service.py @@ -67,6 +67,7 @@ class TestOpsService: icon_background="#FF6B6B", ), account, + session=db_session_with_containers, ) return app, account @@ -91,13 +92,13 @@ class TestOpsService: # ── get_tracing_app_config ───────────────────────────────────────── def test_get_tracing_app_config_no_config(self, db_session_with_containers: Session, mock_ops_trace_manager): - result = OpsService.get_tracing_app_config(str(uuid.uuid4()), "arize") + result = OpsService.get_tracing_app_config(str(uuid.uuid4()), "arize", db_session_with_containers) assert result is None def test_get_tracing_app_config_no_app(self, db_session_with_containers: Session, mock_ops_trace_manager): fake_app_id = str(uuid.uuid4()) self._insert_trace_config(db_session_with_containers, fake_app_id, "arize") - result = OpsService.get_tracing_app_config(fake_app_id, "arize") + result = OpsService.get_tracing_app_config(fake_app_id, "arize", db_session_with_containers) assert result is None def test_get_tracing_app_config_none_config( @@ -107,7 +108,7 @@ class TestOpsService: self._insert_trace_config(db_session_with_containers, app.id, "arize", tracing_config=None) with pytest.raises(ValueError, match="Tracing config cannot be None."): - OpsService.get_tracing_app_config(app.id, "arize") + OpsService.get_tracing_app_config(app.id, "arize", db_session_with_containers) @pytest.mark.parametrize( ("provider", "default_url"), @@ -135,7 +136,7 @@ class TestOpsService: app, _ = self._create_app(db_session_with_containers, mock_external_service_dependencies) self._insert_trace_config(db_session_with_containers, app.id, provider) - result = OpsService.get_tracing_app_config(app.id, provider) + result = OpsService.get_tracing_app_config(app.id, provider, db_session_with_containers) assert result is not None assert result["tracing_config"]["project_url"] == default_url @@ -155,7 +156,7 @@ class TestOpsService: app, _ = self._create_app(db_session_with_containers, mock_external_service_dependencies) self._insert_trace_config(db_session_with_containers, app.id, provider) - result = OpsService.get_tracing_app_config(app.id, provider) + result = OpsService.get_tracing_app_config(app.id, provider, db_session_with_containers) assert result is not None assert result["tracing_config"]["project_url"] == "success_url" @@ -171,7 +172,7 @@ class TestOpsService: app, _ = self._create_app(db_session_with_containers, mock_external_service_dependencies) self._insert_trace_config(db_session_with_containers, app.id, "langfuse") - result = OpsService.get_tracing_app_config(app.id, "langfuse") + result = OpsService.get_tracing_app_config(app.id, "langfuse", db_session_with_containers) assert result is not None assert result["tracing_config"]["project_url"] == "https://api.langfuse.com/project/key" @@ -187,7 +188,7 @@ class TestOpsService: app, _ = self._create_app(db_session_with_containers, mock_external_service_dependencies) self._insert_trace_config(db_session_with_containers, app.id, "langfuse") - result = OpsService.get_tracing_app_config(app.id, "langfuse") + result = OpsService.get_tracing_app_config(app.id, "langfuse", db_session_with_containers) assert result is not None assert result["tracing_config"]["project_url"] == "https://api.langfuse.com/" @@ -195,7 +196,9 @@ class TestOpsService: # ── create_tracing_app_config ────────────────────────────────────── def test_create_tracing_app_config_invalid_provider(self, db_session_with_containers: Session): - result = OpsService.create_tracing_app_config(str(uuid.uuid4()), "invalid_provider", {}) + result = OpsService.create_tracing_app_config( + str(uuid.uuid4()), "invalid_provider", {}, db_session_with_containers + ) assert result == {"error": "Invalid tracing provider: invalid_provider"} def test_create_tracing_app_config_invalid_credentials( @@ -203,7 +206,10 @@ class TestOpsService: ): mock_ops_trace_manager.check_trace_config_is_effective.return_value = False result = OpsService.create_tracing_app_config( - str(uuid.uuid4()), TracingProviderEnum.LANGFUSE, {"public_key": "p", "secret_key": "s"} + str(uuid.uuid4()), + TracingProviderEnum.LANGFUSE, + {"public_key": "p", "secret_key": "s"}, + db_session_with_containers, ) assert result == {"error": "Invalid Credentials"} @@ -228,7 +234,7 @@ class TestOpsService: app, _ = self._create_app(db_session_with_containers, mock_external_service_dependencies) self._insert_trace_config(db_session_with_containers, app.id, str(provider)) - result = OpsService.create_tracing_app_config(app.id, provider, config) + result = OpsService.create_tracing_app_config(app.id, provider, config, db_session_with_containers) assert result is None @@ -245,6 +251,7 @@ class TestOpsService: app.id, TracingProviderEnum.LANGFUSE, {"public_key": "p", "secret_key": "s", "host": "https://api.langfuse.com"}, + db_session_with_containers, ) assert result == {"result": "success"} @@ -258,13 +265,17 @@ class TestOpsService: app, _ = self._create_app(db_session_with_containers, mock_external_service_dependencies) self._insert_trace_config(db_session_with_containers, app.id, str(TracingProviderEnum.ARIZE)) - result = OpsService.create_tracing_app_config(app.id, TracingProviderEnum.ARIZE, {}) + result = OpsService.create_tracing_app_config( + app.id, TracingProviderEnum.ARIZE, {}, db_session_with_containers + ) assert result is None def test_create_tracing_app_config_no_app(self, db_session_with_containers: Session, mock_ops_trace_manager): mock_ops_trace_manager.check_trace_config_is_effective.return_value = True - result = OpsService.create_tracing_app_config(str(uuid.uuid4()), TracingProviderEnum.ARIZE, {}) + result = OpsService.create_tracing_app_config( + str(uuid.uuid4()), TracingProviderEnum.ARIZE, {}, db_session_with_containers + ) assert result is None def test_create_tracing_app_config_with_empty_other_keys( @@ -277,7 +288,9 @@ class TestOpsService: mock_otm.encrypt_tracing_config.return_value = {} app, _ = self._create_app(db_session_with_containers, mock_external_service_dependencies) - result = OpsService.create_tracing_app_config(app.id, TracingProviderEnum.ARIZE, {"project": ""}) + result = OpsService.create_tracing_app_config( + app.id, TracingProviderEnum.ARIZE, {"project": ""}, db_session_with_containers + ) assert result == {"result": "success"} @@ -290,7 +303,9 @@ class TestOpsService: mock_otm.encrypt_tracing_config.return_value = {"encrypted": "config"} app, _ = self._create_app(db_session_with_containers, mock_external_service_dependencies) - result = OpsService.create_tracing_app_config(app.id, TracingProviderEnum.ARIZE, {}) + result = OpsService.create_tracing_app_config( + app.id, TracingProviderEnum.ARIZE, {}, db_session_with_containers + ) assert result == {"result": "success"} @@ -298,17 +313,21 @@ class TestOpsService: def test_update_tracing_app_config_invalid_provider(self, db_session_with_containers: Session): with pytest.raises(ValueError, match="Invalid tracing provider: invalid_provider"): - OpsService.update_tracing_app_config(str(uuid.uuid4()), "invalid_provider", {}) + OpsService.update_tracing_app_config(str(uuid.uuid4()), "invalid_provider", {}, db_session_with_containers) def test_update_tracing_app_config_no_config(self, db_session_with_containers: Session, mock_ops_trace_manager): - result = OpsService.update_tracing_app_config(str(uuid.uuid4()), TracingProviderEnum.ARIZE, {}) + result = OpsService.update_tracing_app_config( + str(uuid.uuid4()), TracingProviderEnum.ARIZE, {}, db_session_with_containers + ) assert result is None def test_update_tracing_app_config_no_app(self, db_session_with_containers: Session, mock_ops_trace_manager): fake_app_id = str(uuid.uuid4()) self._insert_trace_config(db_session_with_containers, fake_app_id, str(TracingProviderEnum.ARIZE)) mock_ops_trace_manager.encrypt_tracing_config.return_value = {} - result = OpsService.update_tracing_app_config(fake_app_id, TracingProviderEnum.ARIZE, {}) + result = OpsService.update_tracing_app_config( + fake_app_id, TracingProviderEnum.ARIZE, {}, db_session_with_containers + ) assert result is None def test_update_tracing_app_config_invalid_credentials( @@ -323,7 +342,7 @@ class TestOpsService: self._insert_trace_config(db_session_with_containers, app.id, str(TracingProviderEnum.ARIZE)) with pytest.raises(ValueError, match="Invalid Credentials"): - OpsService.update_tracing_app_config(app.id, TracingProviderEnum.ARIZE, {}) + OpsService.update_tracing_app_config(app.id, TracingProviderEnum.ARIZE, {}, db_session_with_containers) def test_update_tracing_app_config_success( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -336,7 +355,9 @@ class TestOpsService: app, _ = self._create_app(db_session_with_containers, mock_external_service_dependencies) self._insert_trace_config(db_session_with_containers, app.id, str(TracingProviderEnum.ARIZE)) - result = OpsService.update_tracing_app_config(app.id, TracingProviderEnum.ARIZE, {}) + result = OpsService.update_tracing_app_config( + app.id, TracingProviderEnum.ARIZE, {}, db_session_with_containers + ) assert result is not None assert result["app_id"] == app.id @@ -344,7 +365,7 @@ class TestOpsService: # ── delete_tracing_app_config ────────────────────────────────────── def test_delete_tracing_app_config_no_config(self, db_session_with_containers: Session): - result = OpsService.delete_tracing_app_config(str(uuid.uuid4()), "arize") + result = OpsService.delete_tracing_app_config(str(uuid.uuid4()), "arize", db_session_with_containers) assert result is None def test_delete_tracing_app_config_success( @@ -353,7 +374,7 @@ class TestOpsService: app, _ = self._create_app(db_session_with_containers, mock_external_service_dependencies) self._insert_trace_config(db_session_with_containers, app.id, "arize") - result = OpsService.delete_tracing_app_config(app.id, "arize") + result = OpsService.delete_tracing_app_config(app.id, "arize", db_session_with_containers) assert result is True remaining = db_session_with_containers.scalar( diff --git a/api/tests/test_containers_integration_tests/services/test_recommended_app_service.py b/api/tests/test_containers_integration_tests/services/test_recommended_app_service.py index 9b8eec08ef4..f27132b0fe9 100644 --- a/api/tests/test_containers_integration_tests/services/test_recommended_app_service.py +++ b/api/tests/test_containers_integration_tests/services/test_recommended_app_service.py @@ -14,6 +14,8 @@ from models.model import AccountTrialAppRecord, TrialApp from services import recommended_app_service as service_module from services.recommended_app_service import RecommendedAppService +pytestmark = pytest.mark.usefixtures("db_session_with_containers") + class RecommendedAppPayload(TypedDict, total=False): id: str @@ -118,13 +120,13 @@ class TestRecommendedAppServiceGetApps: mock_factory = MagicMock(return_value=mock_instance) mock_factory_class.get_recommend_app_factory.return_value = mock_factory - result = RecommendedAppService.get_recommended_apps_and_categories(db.session, "en-US") + result = RecommendedAppService.get_recommended_apps_and_categories("en-US", session=db.session()) assert result == expected assert len(result["recommended_apps"]) == 2 assert len(result["categories"]) == 3 mock_factory_class.get_recommend_app_factory.assert_called_once_with("remote") - mock_instance.get_recommended_apps_and_categories.assert_called_once_with("en-US") + mock_instance.get_recommended_apps_and_categories.assert_called_once_with("en-US", session=db.session()) @patch("services.recommended_app_service.RecommendAppRetrievalFactory", autospec=True) @patch("services.recommended_app_service.dify_config") @@ -143,7 +145,7 @@ class TestRecommendedAppServiceGetApps: mock_builtin_instance.fetch_recommended_apps_from_builtin.return_value = builtin_response mock_factory_class.get_buildin_recommend_app_retrieval.return_value = mock_builtin_instance - result = RecommendedAppService.get_recommended_apps_and_categories(db.session, "zh-CN") + result = RecommendedAppService.get_recommended_apps_and_categories("zh-CN", session=db.session()) assert result == builtin_response assert result["recommended_apps"][0]["id"] == "builtin-1" @@ -164,7 +166,7 @@ class TestRecommendedAppServiceGetApps: mock_builtin_instance.fetch_recommended_apps_from_builtin.return_value = builtin_response mock_factory_class.get_buildin_recommend_app_retrieval.return_value = mock_builtin_instance - result = RecommendedAppService.get_recommended_apps_and_categories(db.session, "en-US") + result = RecommendedAppService.get_recommended_apps_and_categories("en-US", session=db.session()) assert result == builtin_response mock_builtin_instance.fetch_recommended_apps_from_builtin.assert_called_once() @@ -182,10 +184,10 @@ class TestRecommendedAppServiceGetApps: mock_instance.get_recommended_apps_and_categories.return_value = lang_response mock_factory_class.get_recommend_app_factory.return_value = MagicMock(return_value=mock_instance) - result = RecommendedAppService.get_recommended_apps_and_categories(db.session, language) + result = RecommendedAppService.get_recommended_apps_and_categories(language, session=db.session()) assert result["recommended_apps"][0]["id"] == f"app-{language}" - mock_instance.get_recommended_apps_and_categories.assert_called_with(language) + mock_instance.get_recommended_apps_and_categories.assert_called_with(language, session=db.session()) @patch("services.recommended_app_service.RecommendAppRetrievalFactory", autospec=True) @patch("services.recommended_app_service.dify_config") @@ -197,7 +199,7 @@ class TestRecommendedAppServiceGetApps: mock_instance.get_recommended_apps_and_categories.return_value = response mock_factory_class.get_recommend_app_factory.return_value = MagicMock(return_value=mock_instance) - RecommendedAppService.get_recommended_apps_and_categories(db.session, "en-US") + RecommendedAppService.get_recommended_apps_and_categories("en-US", session=db.session()) mock_factory_class.get_recommend_app_factory.assert_called_with(mode) @@ -237,10 +239,10 @@ class TestRecommendedAppServiceGetDetail: mock_instance.get_recommend_app_detail.return_value = expected mock_factory_class.get_recommend_app_factory.return_value = MagicMock(return_value=mock_instance) - result = RecommendedAppService.get_recommend_app_detail(db.session, app_id) + result = RecommendedAppService.get_recommend_app_detail(app_id, session=db.session()) assert result == expected - mock_instance.get_recommend_app_detail.assert_called_once_with(app_id) + mock_instance.get_recommend_app_detail.assert_called_once_with(app_id, session=db.session()) @patch("services.recommended_app_service.FeatureService", autospec=True) @patch("services.recommended_app_service.RecommendAppRetrievalFactory", autospec=True) @@ -256,10 +258,10 @@ class TestRecommendedAppServiceGetDetail: mock_instance.get_recommend_app_detail.return_value = detail mock_factory_class.get_recommend_app_factory.return_value = MagicMock(return_value=mock_instance) - result = RecommendedAppService.get_recommend_app_detail(db.session, "test-app") + result = RecommendedAppService.get_recommend_app_detail("test-app", session=db.session()) assert result is not None - mock_instance.get_recommend_app_detail.assert_called_with("test-app") + mock_instance.get_recommend_app_detail.assert_called_with("test-app", session=db.session()) mock_factory_class.get_recommend_app_factory.assert_called_with(mode) @@ -283,11 +285,11 @@ class TestRecommendedAppServiceGetLearnDifyApps: } mock_factory_class.get_recommend_app_factory.return_value = MagicMock(return_value=mock_instance) - result = RecommendedAppService.get_learn_dify_apps(db.session, "en-US") + result = RecommendedAppService.get_learn_dify_apps("en-US", session=db.session()) assert result == {"recommended_apps": [expected_app]} mock_factory_class.get_recommend_app_factory.assert_called_once_with("remote") - mock_instance.get_learn_dify_apps.assert_called_once_with("en-US") + mock_instance.get_learn_dify_apps.assert_called_once_with("en-US", session=db.session()) @patch("services.recommended_app_service.dify_config") def test_sets_can_trial_when_trial_feature_enabled( @@ -314,10 +316,10 @@ class TestRecommendedAppServiceGetLearnDifyApps: can_trial_mock = MagicMock(return_value=True) monkeypatch.setattr(RecommendedAppService, "_can_trial_app", can_trial_mock) - result = RecommendedAppService.get_learn_dify_apps(db.session, "en-US") + result = RecommendedAppService.get_learn_dify_apps("en-US", session=db.session()) assert result["recommended_apps"][0]["can_trial"] is True - can_trial_mock.assert_called_once_with(db.session, "app-1") + can_trial_mock.assert_called_once_with(db.session(), "app-1") # ── Integration tests: trial app features (real DB) ──────────────────── @@ -333,10 +335,10 @@ class TestRecommendedAppServiceTrialFeatures: MagicMock(return_value=SimpleNamespace(enable_trial_app=False)), ) - result = RecommendedAppService.get_recommended_apps_and_categories(db.session, "en-US") + result = RecommendedAppService.get_recommended_apps_and_categories("en-US", session=db.session()) assert result == expected - retrieval_instance.get_recommended_apps_and_categories.assert_called_once_with("en-US") + retrieval_instance.get_recommended_apps_and_categories.assert_called_once_with("en-US", session=db.session()) builtin_instance.fetch_recommended_apps_from_builtin.assert_not_called() def test_get_apps_should_enrich_can_trial_when_enabled( @@ -364,7 +366,7 @@ class TestRecommendedAppServiceTrialFeatures: MagicMock(return_value=SimpleNamespace(enable_trial_app=True)), ) - result = RecommendedAppService.get_recommended_apps_and_categories(db.session, "ja-JP") + result = RecommendedAppService.get_recommended_apps_and_categories("ja-JP", session=db.session()) builtin_instance.fetch_recommended_apps_from_builtin.assert_called_once_with("en-US") assert result["recommended_apps"][0]["can_trial"] is True @@ -400,7 +402,7 @@ class TestRecommendedAppServiceTrialFeatures: MagicMock(return_value=SimpleNamespace(enable_trial_app=True)), ) - result = RecommendedAppService.get_recommend_app_detail(db.session, app_id) + result = RecommendedAppService.get_recommend_app_detail(app_id, session=db.session()) assert result is not None detail_result = cast(RecommendedAppPayload, result) @@ -421,10 +423,10 @@ class TestRecommendedAppServiceTrialFeatures: mock_instance.get_recommend_app_detail.return_value = None mock_factory_class.get_recommend_app_factory.return_value = MagicMock(return_value=mock_instance) - result = RecommendedAppService.get_recommend_app_detail(db.session, "nonexistent") + result = RecommendedAppService.get_recommend_app_detail("nonexistent", session=db.session()) assert result is None - mock_instance.get_recommend_app_detail.assert_called_once_with("nonexistent") + mock_instance.get_recommend_app_detail.assert_called_once_with("nonexistent", session=db.session()) mock_feature_service.get_system_features.assert_not_called() def test_add_trial_app_record_increments_count_for_existing(self, db_session_with_containers: Session) -> None: @@ -434,7 +436,7 @@ class TestRecommendedAppServiceTrialFeatures: db_session_with_containers.add(AccountTrialAppRecord(app_id=app_id, account_id=account_id, count=3)) db_session_with_containers.commit() - RecommendedAppService.add_trial_app_record(db.session, app_id, account_id) + RecommendedAppService.add_trial_app_record(app_id, account_id, session=db.session()) db_session_with_containers.expire_all() record = db_session_with_containers.scalar( @@ -449,7 +451,7 @@ class TestRecommendedAppServiceTrialFeatures: app_id = str(uuid.uuid4()) account_id = str(uuid.uuid4()) - RecommendedAppService.add_trial_app_record(db.session, app_id, account_id) + RecommendedAppService.add_trial_app_record(app_id, account_id, session=db.session()) db_session_with_containers.expire_all() record = db_session_with_containers.scalar( diff --git a/api/tests/test_containers_integration_tests/services/test_saved_message_service.py b/api/tests/test_containers_integration_tests/services/test_saved_message_service.py index cfd1d4e86b4..92741ac56cb 100644 --- a/api/tests/test_containers_integration_tests/services/test_saved_message_service.py +++ b/api/tests/test_containers_integration_tests/services/test_saved_message_service.py @@ -1,4 +1,4 @@ -from unittest.mock import patch +from unittest.mock import ANY, patch import pytest from faker import Faker @@ -86,7 +86,7 @@ class TestSavedMessageService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) return app, account @@ -222,7 +222,7 @@ class TestSavedMessageService: # Act: Execute the method under test result = SavedMessageService.pagination_by_last_id( - db_session_with_containers, app_model=app, user=account, last_id=None, limit=10 + app_model=app, user=account, last_id=None, limit=10, session=db_session_with_containers ) # Assert: Verify the expected outcomes @@ -297,7 +297,7 @@ class TestSavedMessageService: # Act: Execute the method under test result = SavedMessageService.pagination_by_last_id( - db_session_with_containers, app_model=app, user=end_user, last_id="test_last_id", limit=5 + app_model=app, user=end_user, last_id="test_last_id", limit=5, session=db_session_with_containers ) # Assert: Verify the expected outcomes @@ -347,7 +347,7 @@ class TestSavedMessageService: mock_external_service_dependencies["message_service"].get_message.return_value = message # Act: Execute the method under test - SavedMessageService.save(db_session_with_containers, app_model=app, user=account, message_id=message.id) + SavedMessageService.save(app_model=app, user=account, message_id=message.id, session=db_session_with_containers) # Assert: Verify the expected outcomes # Check if saved message was created in database @@ -372,7 +372,7 @@ class TestSavedMessageService: # Verify MessageService.get_message was called mock_external_service_dependencies["message_service"].get_message.assert_called_once_with( - app_model=app, user=account, message_id=message.id + app_model=app, user=account, message_id=message.id, session=ANY ) # Verify database state @@ -397,7 +397,7 @@ class TestSavedMessageService: # Act & Assert: Verify proper error handling with pytest.raises(ValueError) as exc_info: SavedMessageService.pagination_by_last_id( - db_session_with_containers, app_model=app, user=None, last_id=None, limit=10 + app_model=app, user=None, last_id=None, limit=10, session=db_session_with_containers ) assert "User is required" in str(exc_info.value) @@ -417,7 +417,9 @@ class TestSavedMessageService: message = self._create_test_message(db_session_with_containers, app, account) # Act: Execute the method under test with None user - result = SavedMessageService.save(db_session_with_containers, app_model=app, user=None, message_id=message.id) + result = SavedMessageService.save( + app_model=app, user=None, message_id=message.id, session=db_session_with_containers + ) # Assert: Verify the expected outcomes assert result is None @@ -476,7 +478,9 @@ class TestSavedMessageService: ) # Act: Execute the method under test - SavedMessageService.delete(db_session_with_containers, app_model=app, user=account, message_id=message.id) + SavedMessageService.delete( + app_model=app, user=account, message_id=message.id, session=db_session_with_containers + ) # Assert: Verify the expected outcomes # Check if saved message was deleted from database @@ -506,7 +510,9 @@ class TestSavedMessageService: mock_external_service_dependencies["message_service"].get_message.return_value = message - SavedMessageService.save(db_session_with_containers, app_model=app, user=end_user, message_id=message.id) + SavedMessageService.save( + app_model=app, user=end_user, message_id=message.id, session=db_session_with_containers + ) saved = ( db_session_with_containers.query(SavedMessage) @@ -527,9 +533,9 @@ class TestSavedMessageService: mock_external_service_dependencies["message_service"].get_message.return_value = message # Save once - SavedMessageService.save(db_session_with_containers, app_model=app, user=account, message_id=message.id) + SavedMessageService.save(app_model=app, user=account, message_id=message.id, session=db_session_with_containers) # Save again - SavedMessageService.save(db_session_with_containers, app_model=app, user=account, message_id=message.id) + SavedMessageService.save(app_model=app, user=account, message_id=message.id, session=db_session_with_containers) count = ( db_session_with_containers.query(SavedMessage) @@ -552,7 +558,7 @@ class TestSavedMessageService: db_session_with_containers.add(saved) db_session_with_containers.commit() - SavedMessageService.delete(db_session_with_containers, app_model=app, user=None, message_id=message.id) + SavedMessageService.delete(app_model=app, user=None, message_id=message.id, session=db_session_with_containers) # Should still exist assert ( @@ -571,7 +577,9 @@ class TestSavedMessageService: # Should not raise — use a valid UUID that doesn't exist in DB from uuid import uuid4 - SavedMessageService.delete(db_session_with_containers, app_model=app, user=account, message_id=str(uuid4())) + SavedMessageService.delete( + app_model=app, user=account, message_id=str(uuid4()), session=db_session_with_containers + ) def test_delete_for_end_user(self, db_session_with_containers: Session, mock_external_service_dependencies): """Test deleting a saved message for an EndUser.""" @@ -585,7 +593,9 @@ class TestSavedMessageService: db_session_with_containers.add(saved) db_session_with_containers.commit() - SavedMessageService.delete(db_session_with_containers, app_model=app, user=end_user, message_id=message.id) + SavedMessageService.delete( + app_model=app, user=end_user, message_id=message.id, session=db_session_with_containers + ) assert ( db_session_with_containers.query(SavedMessage) @@ -615,7 +625,9 @@ class TestSavedMessageService: db_session_with_containers.commit() # Delete only account1's saved message - SavedMessageService.delete(db_session_with_containers, app_model=app, user=account1, message_id=message.id) + SavedMessageService.delete( + app_model=app, user=account1, message_id=message.id, session=db_session_with_containers + ) # Account's saved message should be gone assert ( diff --git a/api/tests/test_containers_integration_tests/services/test_tag_service.py b/api/tests/test_containers_integration_tests/services/test_tag_service.py index 748cca6c845..86b635ac23d 100644 --- a/api/tests/test_containers_integration_tests/services/test_tag_service.py +++ b/api/tests/test_containers_integration_tests/services/test_tag_service.py @@ -205,7 +205,7 @@ def test_get_tags_success(db_session_with_containers: Session, current_user_stub db_session_with_containers, tags=tags[:2], target_id=dataset.id, tenant_id=tenant.id, user_id=account.id ) - result = TagService.get_tags(db_session_with_containers, TagType.KNOWLEDGE, tenant.id) + result = TagService.get_tags(TagType.KNOWLEDGE, tenant.id, session=db_session_with_containers) assert result is not None assert len(result) == 3 @@ -235,7 +235,7 @@ def test_get_tags_with_keyword_filter(db_session_with_containers: Session, curre tags[2].name = "web_development" db_session_with_containers.flush() - result = TagService.get_tags(db_session_with_containers, TagType.APP, tenant.id, keyword="development") + result = TagService.get_tags(TagType.APP, tenant.id, keyword="development", session=db_session_with_containers) assert result is not None assert len(result) == 2 @@ -243,7 +243,9 @@ def test_get_tags_with_keyword_filter(db_session_with_containers: Session, curre for tag_result in result: assert "development" in tag_result.name.lower() - result_no_match = TagService.get_tags(db_session_with_containers, TagType.APP, tenant.id, keyword="nonexistent") + result_no_match = TagService.get_tags( + TagType.APP, tenant.id, keyword="nonexistent", session=db_session_with_containers + ) assert result_no_match == [] @@ -291,19 +293,19 @@ def test_get_tags_with_special_characters_in_keyword( db_session_with_containers.flush() - result = TagService.get_tags(db_session_with_containers, TagType.APP, tenant.id, keyword="50%") + result = TagService.get_tags(TagType.APP, tenant.id, keyword="50%", session=db_session_with_containers) assert len(result) == 1 assert result[0].name == "50% discount" - result = TagService.get_tags(db_session_with_containers, TagType.APP, tenant.id, keyword="test_data") + result = TagService.get_tags(TagType.APP, tenant.id, keyword="test_data", session=db_session_with_containers) assert len(result) == 1 assert result[0].name == "test_data_tag" - result = TagService.get_tags(db_session_with_containers, TagType.APP, tenant.id, keyword="path\\to\\tag") + result = TagService.get_tags(TagType.APP, tenant.id, keyword="path\\to\\tag", session=db_session_with_containers) assert len(result) == 1 assert result[0].name == "path\\to\\tag" - result = TagService.get_tags(db_session_with_containers, TagType.APP, tenant.id, keyword="50%") + result = TagService.get_tags(TagType.APP, tenant.id, keyword="50%", session=db_session_with_containers) assert len(result) == 1 assert all("50%" in item.name for item in result) @@ -312,7 +314,7 @@ def test_get_tags_empty_result(db_session_with_containers: Session, current_user account, tenant = _create_account_with_tenant(db_session_with_containers) _set_current_user(current_user_stub, account, tenant) - result = TagService.get_tags(db_session_with_containers, TagType.KNOWLEDGE, tenant.id) + result = TagService.get_tags(TagType.KNOWLEDGE, tenant.id, session=db_session_with_containers) assert result == [] diff --git a/api/tests/test_containers_integration_tests/services/test_web_conversation_service.py b/api/tests/test_containers_integration_tests/services/test_web_conversation_service.py index 664c1167994..ed063ceaccc 100644 --- a/api/tests/test_containers_integration_tests/services/test_web_conversation_service.py +++ b/api/tests/test_containers_integration_tests/services/test_web_conversation_service.py @@ -90,7 +90,7 @@ class TestWebConversationService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) return app, account @@ -312,7 +312,7 @@ class TestWebConversationService: conversation = self._create_test_conversation(db_session_with_containers, app, account, fake) # Pin the conversation - WebConversationService.pin(app, conversation.id, account) + WebConversationService.pin(app, conversation.id, account, db_session_with_containers) # Verify the conversation was pinned @@ -346,10 +346,10 @@ class TestWebConversationService: conversation = self._create_test_conversation(db_session_with_containers, app, account, fake) # Pin the conversation first time - WebConversationService.pin(app, conversation.id, account) + WebConversationService.pin(app, conversation.id, account, db_session_with_containers) # Pin the conversation again - WebConversationService.pin(app, conversation.id, account) + WebConversationService.pin(app, conversation.id, account, db_session_with_containers) # Verify only one pinned conversation record exists @@ -380,7 +380,7 @@ class TestWebConversationService: conversation = self._create_test_conversation(db_session_with_containers, app, end_user, fake) # Pin the conversation - WebConversationService.pin(app, conversation.id, end_user) + WebConversationService.pin(app, conversation.id, end_user, db_session_with_containers) # Verify the conversation was pinned @@ -412,7 +412,7 @@ class TestWebConversationService: conversation = self._create_test_conversation(db_session_with_containers, app, account, fake) # Pin the conversation first - WebConversationService.pin(app, conversation.id, account) + WebConversationService.pin(app, conversation.id, account, db_session_with_containers) # Verify it was pinned @@ -430,7 +430,7 @@ class TestWebConversationService: assert pinned_conversation is not None # Unpin the conversation - WebConversationService.unpin(app, conversation.id, account) + WebConversationService.unpin(app, conversation.id, account, db_session_with_containers) # Verify it was unpinned pinned_conversation = ( @@ -459,7 +459,7 @@ class TestWebConversationService: conversation = self._create_test_conversation(db_session_with_containers, app, account, fake) # Try to unpin a conversation that was never pinned - WebConversationService.unpin(app, conversation.id, account) + WebConversationService.unpin(app, conversation.id, account, db_session_with_containers) # Verify no pinned conversation record exists @@ -509,7 +509,7 @@ class TestWebConversationService: conversation = self._create_test_conversation(db_session_with_containers, app, account, fake) # Try to pin with None user - WebConversationService.pin(app, conversation.id, None) + WebConversationService.pin(app, conversation.id, None, db_session_with_containers) # Verify no pinned conversation was created @@ -537,7 +537,7 @@ class TestWebConversationService: conversation = self._create_test_conversation(db_session_with_containers, app, account, fake) # Pin the conversation first - WebConversationService.pin(app, conversation.id, account) + WebConversationService.pin(app, conversation.id, account, db_session_with_containers) # Verify it was pinned @@ -555,7 +555,7 @@ class TestWebConversationService: assert pinned_conversation is not None # Try to unpin with None user - WebConversationService.unpin(app, conversation.id, None) + WebConversationService.unpin(app, conversation.id, None, db_session_with_containers) # Verify the conversation is still pinned pinned_conversation = ( diff --git a/api/tests/test_containers_integration_tests/services/test_webapp_auth_service.py b/api/tests/test_containers_integration_tests/services/test_webapp_auth_service.py index 7825f502f77..52d1fde7927 100644 --- a/api/tests/test_containers_integration_tests/services/test_webapp_auth_service.py +++ b/api/tests/test_containers_integration_tests/services/test_webapp_auth_service.py @@ -1,6 +1,6 @@ import time import uuid -from unittest.mock import patch +from unittest.mock import ANY, patch import pytest from faker import Faker @@ -223,7 +223,7 @@ class TestWebAppAuthService: ) # Act: Execute authentication - result = WebAppAuthService.authenticate(account.email, password) + result = WebAppAuthService.authenticate(account.email, password, db_session_with_containers) # Assert: Verify successful authentication assert result is not None @@ -260,7 +260,7 @@ class TestWebAppAuthService: # Act & Assert: Verify proper error handling with pytest.raises(AccountNotFoundError): - WebAppAuthService.authenticate(non_existent_email, "any_password") + WebAppAuthService.authenticate(non_existent_email, "any_password", db_session_with_containers) def test_authenticate_account_banned(self, db_session_with_containers: Session, mock_external_service_dependencies): """ @@ -297,7 +297,7 @@ class TestWebAppAuthService: # Act & Assert: Verify proper error handling with pytest.raises(AccountLoginError) as exc_info: - WebAppAuthService.authenticate(account.email, password) + WebAppAuthService.authenticate(account.email, password, db_session_with_containers) assert "Account is banned." in str(exc_info.value) @@ -318,7 +318,7 @@ class TestWebAppAuthService: # Act & Assert: Verify proper error handling with wrong password with pytest.raises(AccountPasswordError) as exc_info: - WebAppAuthService.authenticate(account.email, "wrong_password") + WebAppAuthService.authenticate(account.email, "wrong_password", db_session_with_containers) assert "Invalid email or password." in str(exc_info.value) @@ -350,7 +350,7 @@ class TestWebAppAuthService: # Act & Assert: Verify proper error handling with pytest.raises(AccountPasswordError) as exc_info: - WebAppAuthService.authenticate(account.email, "any_password") + WebAppAuthService.authenticate(account.email, "any_password", db_session_with_containers) assert "Invalid email or password." in str(exc_info.value) @@ -403,7 +403,7 @@ class TestWebAppAuthService: ) # Act: Execute user retrieval - result = WebAppAuthService.get_user_through_email(account.email) + result = WebAppAuthService.get_user_through_email(account.email, db_session_with_containers) # Assert: Verify successful retrieval assert result is not None @@ -430,7 +430,7 @@ class TestWebAppAuthService: non_existent_email = f"nonexistent_{uuid.uuid4().hex}@example.com" # Act: Execute user retrieval - result = WebAppAuthService.get_user_through_email(non_existent_email) + result = WebAppAuthService.get_user_through_email(non_existent_email, db_session_with_containers) # Assert: Verify proper handling assert result is None @@ -463,7 +463,7 @@ class TestWebAppAuthService: # Act & Assert: Verify proper error handling with pytest.raises(Unauthorized) as exc_info: - WebAppAuthService.get_user_through_email(account.email) + WebAppAuthService.get_user_through_email(account.email, db_session_with_containers) assert "Account is banned." in str(exc_info.value) @@ -659,7 +659,7 @@ class TestWebAppAuthService: ) # Act: Execute end user creation - result = WebAppAuthService.create_end_user(site.code, "test@example.com") + result = WebAppAuthService.create_end_user(site.code, "test@example.com", db_session_with_containers) # Assert: Verify successful creation assert result is not None @@ -694,7 +694,7 @@ class TestWebAppAuthService: # Act & Assert: Verify proper error handling with pytest.raises(NotFound) as exc_info: - WebAppAuthService.create_end_user(non_existent_code, "test@example.com") + WebAppAuthService.create_end_user(non_existent_code, "test@example.com", db_session_with_containers) assert "Site not found." in str(exc_info.value) @@ -732,7 +732,7 @@ class TestWebAppAuthService: # Act & Assert: Verify proper error handling with pytest.raises(NotFound) as exc_info: - WebAppAuthService.create_end_user(site.code, "test@example.com") + WebAppAuthService.create_end_user(site.code, "test@example.com", db_session_with_containers) assert "App not found." in str(exc_info.value) @@ -750,7 +750,9 @@ class TestWebAppAuthService: # Arrange: Setup test with private access mode # Act: Execute permission check requirement test - result = WebAppAuthService.is_app_require_permission_check(access_mode="private") + result = WebAppAuthService.is_app_require_permission_check( + access_mode="private", session=db_session_with_containers + ) # Assert: Verify correct result assert result is True @@ -769,7 +771,9 @@ class TestWebAppAuthService: # Arrange: Setup test with public access mode # Act: Execute permission check requirement test - result = WebAppAuthService.is_app_require_permission_check(access_mode="public") + result = WebAppAuthService.is_app_require_permission_check( + access_mode="public", session=db_session_with_containers + ) # Assert: Verify correct result assert result is False @@ -789,13 +793,17 @@ class TestWebAppAuthService: mock_external_service_dependencies["app_service"].get_app_id_by_code.return_value = "mock_app_id" # Act: Execute permission check requirement test - result = WebAppAuthService.is_app_require_permission_check(app_code="mock_app_code") + result = WebAppAuthService.is_app_require_permission_check( + app_code="mock_app_code", session=db_session_with_containers + ) # Assert: Verify correct result assert result is True # Verify mock service was called correctly - mock_external_service_dependencies["app_service"].get_app_id_by_code.assert_called_once_with("mock_app_code") + mock_external_service_dependencies["app_service"].get_app_id_by_code.assert_called_once_with( + "mock_app_code", session=ANY + ) mock_external_service_dependencies[ "enterprise_service" ].WebAppAuth.get_app_access_mode_by_id.assert_called_once_with("mock_app_id") @@ -814,7 +822,7 @@ class TestWebAppAuthService: # Act & Assert: Verify proper error handling with pytest.raises(ValueError) as exc_info: - WebAppAuthService.is_app_require_permission_check() + WebAppAuthService.is_app_require_permission_check(session=db_session_with_containers) assert "Either app_code or app_id must be provided." in str(exc_info.value) @@ -832,7 +840,7 @@ class TestWebAppAuthService: # Arrange: Setup test with public access mode # Act: Execute authentication type determination - result = WebAppAuthService.get_app_auth_type(access_mode="public") + result = WebAppAuthService.get_app_auth_type(access_mode="public", session=db_session_with_containers) # Assert: Verify correct result assert result == WebAppAuthType.PUBLIC @@ -851,7 +859,7 @@ class TestWebAppAuthService: # Arrange: Setup test with private access mode # Act: Execute authentication type determination - result = WebAppAuthService.get_app_auth_type(access_mode="private") + result = WebAppAuthService.get_app_auth_type(access_mode="private", session=db_session_with_containers) # Assert: Verify correct result assert result == WebAppAuthType.INTERNAL @@ -875,7 +883,9 @@ class TestWebAppAuthService: ].WebAppAuth.get_app_access_mode_by_id.return_value = setting # Act: Execute authentication type determination - result: WebAppAuthType = WebAppAuthService.get_app_auth_type(app_code="mock_app_code") + result: WebAppAuthType = WebAppAuthService.get_app_auth_type( + app_code="mock_app_code", session=db_session_with_containers + ) # Assert: Verify correct result assert result == WebAppAuthType.EXTERNAL @@ -899,6 +909,6 @@ class TestWebAppAuthService: # Act & Assert: Verify proper error handling with pytest.raises(ValueError) as exc_info: - WebAppAuthService.get_app_auth_type() + WebAppAuthService.get_app_auth_type(session=db_session_with_containers) assert "Either app_code or access_mode must be provided." in str(exc_info.value) diff --git a/api/tests/test_containers_integration_tests/services/test_webhook_service_relationships.py b/api/tests/test_containers_integration_tests/services/test_webhook_service_relationships.py index c699d39dde1..902134e053d 100644 --- a/api/tests/test_containers_integration_tests/services/test_webhook_service_relationships.py +++ b/api/tests/test_containers_integration_tests/services/test_webhook_service_relationships.py @@ -350,9 +350,10 @@ class TestWebhookServiceTriggerExecutionWithContainers: quota_charge.commit.assert_called_once() mock_trigger.assert_called_once() trigger_args = mock_trigger.call_args.args - assert trigger_args[1] is end_user - assert trigger_args[2].workflow_id == workflow.id - assert trigger_args[2].root_node_id == webhook_trigger.node_id + assert trigger_args[0] is end_user + assert trigger_args[1].workflow_id == workflow.id + assert trigger_args[1].root_node_id == webhook_trigger.node_id + assert mock_trigger.call_args.kwargs["session"] is not None def test_trigger_workflow_execution_marks_tenant_rate_limited_when_quota_exceeded( self, db_session_with_containers: Session, flask_app_with_containers: Flask diff --git a/api/tests/test_containers_integration_tests/services/test_workflow_app_service.py b/api/tests/test_containers_integration_tests/services/test_workflow_app_service.py index cf76afb303c..f553b0f72a0 100644 --- a/api/tests/test_containers_integration_tests/services/test_workflow_app_service.py +++ b/api/tests/test_containers_integration_tests/services/test_workflow_app_service.py @@ -99,7 +99,7 @@ class TestWorkflowAppService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) return app, account @@ -164,7 +164,7 @@ class TestWorkflowAppService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) return app diff --git a/api/tests/test_containers_integration_tests/services/test_workflow_run_service.py b/api/tests/test_containers_integration_tests/services/test_workflow_run_service.py index 726c360d77e..7c528f06b10 100644 --- a/api/tests/test_containers_integration_tests/services/test_workflow_run_service.py +++ b/api/tests/test_containers_integration_tests/services/test_workflow_run_service.py @@ -92,7 +92,7 @@ class TestWorkflowRunService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) return app, account @@ -544,7 +544,7 @@ class TestWorkflowRunService: icon="🚀", icon_background="#4ECDC4", ) - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Create workflow run without node executions workflow_run = self._create_test_workflow_run(db_session_with_containers, app, account, "debugging") @@ -596,7 +596,7 @@ class TestWorkflowRunService: icon="🚀", icon_background="#4ECDC4", ) - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Use invalid workflow run ID invalid_workflow_run_id = str(uuid.uuid4()) @@ -648,7 +648,7 @@ class TestWorkflowRunService: icon="🚀", icon_background="#4ECDC4", ) - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Create workflow run workflow_run = self._create_test_workflow_run(db_session_with_containers, app, account, "debugging") diff --git a/api/tests/test_containers_integration_tests/services/test_workflow_service.py b/api/tests/test_containers_integration_tests/services/test_workflow_service.py index 349aac1be36..6531ed4fbb0 100644 --- a/api/tests/test_containers_integration_tests/services/test_workflow_service.py +++ b/api/tests/test_containers_integration_tests/services/test_workflow_service.py @@ -227,7 +227,7 @@ class TestWorkflowService: workflow_service = WorkflowService() # Act - result = workflow_service.is_workflow_exist(app) + result = workflow_service.is_workflow_exist(app, session=db_session_with_containers) # Assert assert result is True @@ -247,7 +247,7 @@ class TestWorkflowService: workflow_service = WorkflowService() # Act - result = workflow_service.is_workflow_exist(app) + result = workflow_service.is_workflow_exist(app, session=db_session_with_containers) # Assert assert result is False @@ -269,7 +269,7 @@ class TestWorkflowService: workflow_service = WorkflowService() # Act - result = workflow_service.get_draft_workflow(app) + result = workflow_service.get_draft_workflow(app, session=db_session_with_containers) # Assert assert result is not None @@ -293,7 +293,7 @@ class TestWorkflowService: workflow_service = WorkflowService() # Act - result = workflow_service.get_draft_workflow(app) + result = workflow_service.get_draft_workflow(app, session=db_session_with_containers) # Assert assert result is None @@ -320,7 +320,7 @@ class TestWorkflowService: workflow_service = WorkflowService() # Act - result = workflow_service.get_published_workflow_by_id(app, workflow.id) + result = workflow_service.get_published_workflow_by_id(app, workflow.id, session=db_session_with_containers) # Assert assert result is not None @@ -349,7 +349,7 @@ class TestWorkflowService: from services.errors.app import IsDraftWorkflowError with pytest.raises(IsDraftWorkflowError): - workflow_service.get_published_workflow_by_id(app, workflow.id) + workflow_service.get_published_workflow_by_id(app, workflow.id, session=db_session_with_containers) def test_get_published_workflow_by_id_not_found(self, db_session_with_containers: Session): """ @@ -366,7 +366,9 @@ class TestWorkflowService: workflow_service = WorkflowService() # Act - result = workflow_service.get_published_workflow_by_id(app, non_existent_workflow_id) + result = workflow_service.get_published_workflow_by_id( + app, non_existent_workflow_id, session=db_session_with_containers + ) # Assert assert result is None @@ -393,7 +395,7 @@ class TestWorkflowService: workflow_service = WorkflowService() # Act - result = workflow_service.get_published_workflow(app) + result = workflow_service.get_published_workflow(app, session=db_session_with_containers) # Assert assert result is not None @@ -416,7 +418,7 @@ class TestWorkflowService: workflow_service = WorkflowService() # Act - result = workflow_service.get_published_workflow(app) + result = workflow_service.get_published_workflow(app, session=db_session_with_containers) # Assert assert result is None @@ -714,6 +716,7 @@ class TestWorkflowService: account=account, environment_variables=environment_variables, conversation_variables=conversation_variables, + session=db_session_with_containers, ) # Assert @@ -778,6 +781,7 @@ class TestWorkflowService: account=account, environment_variables=environment_variables, conversation_variables=conversation_variables, + session=db_session_with_containers, ) # Assert @@ -838,6 +842,7 @@ class TestWorkflowService: account=account, environment_variables=environment_variables, conversation_variables=conversation_variables, + session=db_session_with_containers, ) def test_publish_workflow_success(self, db_session_with_containers: Session): @@ -979,9 +984,7 @@ class TestWorkflowService: workflow_service = WorkflowService() restored_workflow = workflow_service.restore_published_workflow_to_draft( - app_model=app, - workflow_id=published_workflow.id, - account=account, + app_model=app, workflow_id=published_workflow.id, account=account, session=db_session_with_containers ) db_session_with_containers.expire_all() @@ -1130,7 +1133,9 @@ class TestWorkflowService: } # Act - result = workflow_service.convert_to_workflow(app_model=app, account=account, args=conversion_args) + result = workflow_service.convert_to_workflow( + app_model=app, account=account, args=conversion_args, session=db_session_with_containers + ) # Assert assert result is not None @@ -1190,7 +1195,9 @@ class TestWorkflowService: } # Act - result = workflow_service.convert_to_workflow(app_model=app, account=account, args=conversion_args) + result = workflow_service.convert_to_workflow( + app_model=app, account=account, args=conversion_args, session=db_session_with_containers + ) # Assert assert result is not None @@ -1222,7 +1229,9 @@ class TestWorkflowService: # Act & Assert with pytest.raises(ValueError, match="Current App mode: workflow is not supported convert to workflow"): - workflow_service.convert_to_workflow(app_model=app, account=account, args=conversion_args) + workflow_service.convert_to_workflow( + app_model=app, account=account, args=conversion_args, session=db_session_with_containers + ) def test_validate_features_structure_advanced_chat(self, db_session_with_containers: Session): """ diff --git a/api/tests/test_containers_integration_tests/services/test_workspace_service.py b/api/tests/test_containers_integration_tests/services/test_workspace_service.py index 4e89d906f16..d7cbfa91ab7 100644 --- a/api/tests/test_containers_integration_tests/services/test_workspace_service.py +++ b/api/tests/test_containers_integration_tests/services/test_workspace_service.py @@ -104,7 +104,7 @@ class TestWorkspaceService: # Mock current_user for flask_login with patch("services.workspace_service.current_user", account): # Act: Execute the method under test - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -151,7 +151,7 @@ class TestWorkspaceService: # Mock current_user for flask_login with patch("services.workspace_service.current_user", account): # Act: Execute the method under test - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -206,7 +206,7 @@ class TestWorkspaceService: # Mock current_user for flask_login with patch("services.workspace_service.current_user", account): # Act: Execute the method under test - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -261,7 +261,7 @@ class TestWorkspaceService: # Mock current_user for flask_login with patch("services.workspace_service.current_user", account): # Act: Execute the method under test - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -291,7 +291,7 @@ class TestWorkspaceService: # Arrange: No test data needed for this test # Act: Execute the method under test with None tenant - result = WorkspaceService.get_tenant_info(None) + result = WorkspaceService.get_tenant_info(None, db_session_with_containers) # Assert: Verify the expected outcomes assert result is None @@ -341,7 +341,7 @@ class TestWorkspaceService: # Mock current_user for flask_login with patch("services.workspace_service.current_user", account): # Act: Execute the method under test - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -398,7 +398,7 @@ class TestWorkspaceService: # Mock current_user for flask_login with patch("services.workspace_service.current_user", account): # Act: Execute the method under test - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -448,7 +448,7 @@ class TestWorkspaceService: # Mock current_user for flask_login with patch("services.workspace_service.current_user", account): # Act: Execute the method under test - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -513,7 +513,7 @@ class TestWorkspaceService: # Mock current_user for flask_login with patch("services.workspace_service.current_user", account): # Act: Execute the method under test - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -553,7 +553,7 @@ class TestWorkspaceService: # No TenantAccountJoin created with patch("services.workspace_service.current_user", account): with pytest.raises(AssertionError, match="TenantAccountJoin not found"): - WorkspaceService.get_tenant_info(tenant) + WorkspaceService.get_tenant_info(tenant, db_session_with_containers) def test_get_tenant_info_should_set_replace_webapp_logo_to_none_when_flag_absent( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -572,7 +572,7 @@ class TestWorkspaceService: mock_external_service_dependencies["tenant_service"].has_roles.return_value = True with patch("services.workspace_service.current_user", account): - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) assert result is not None assert result["custom_config"]["replace_webapp_logo"] is None @@ -596,7 +596,7 @@ class TestWorkspaceService: mock_external_service_dependencies["tenant_service"].has_roles.return_value = True with patch("services.workspace_service.current_user", account): - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) assert result is not None assert result["custom_config"]["replace_webapp_logo"].startswith(custom_base) @@ -615,7 +615,7 @@ class TestWorkspaceService: mock_external_service_dependencies["tenant_service"].has_roles.return_value = False with patch("services.workspace_service.current_user", account): - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) assert result is not None assert "next_credit_reset_date" not in result @@ -642,7 +642,7 @@ class TestWorkspaceService: patch("services.workspace_service.current_user", account), patch("services.credit_pool_service.CreditPoolService.get_pool", return_value=None), ): - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) assert result is not None assert result["next_credit_reset_date"] == "2025-02-01" @@ -669,7 +669,7 @@ class TestWorkspaceService: patch("services.workspace_service.current_user", account), patch("services.credit_pool_service.CreditPoolService.get_pool", return_value=paid_pool), ): - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) assert result is not None assert result["trial_credits"] == 1000 @@ -697,7 +697,7 @@ class TestWorkspaceService: patch("services.workspace_service.current_user", account), patch("services.credit_pool_service.CreditPoolService.get_pool", side_effect=[paid_pool, None]), ): - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) assert result is not None assert result["trial_credits"] == -1 @@ -726,7 +726,7 @@ class TestWorkspaceService: patch("services.workspace_service.current_user", account), patch("services.credit_pool_service.CreditPoolService.get_pool", side_effect=[paid_pool, trial_pool]), ): - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) assert result is not None assert result["trial_credits"] == 100 @@ -754,7 +754,7 @@ class TestWorkspaceService: patch("services.workspace_service.current_user", account), patch("services.credit_pool_service.CreditPoolService.get_pool", side_effect=[None, trial_pool]), ): - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) assert result is not None assert result["trial_credits"] == 50 @@ -785,7 +785,7 @@ class TestWorkspaceService: patch("services.workspace_service.current_user", account), patch("services.credit_pool_service.CreditPoolService.get_pool", side_effect=[paid_pool, trial_pool]), ): - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) assert result is not None assert result["trial_credits"] == 200 @@ -811,7 +811,7 @@ class TestWorkspaceService: patch("services.workspace_service.current_user", account), patch("services.credit_pool_service.CreditPoolService.get_pool", side_effect=[None, None]), ): - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) assert result is not None assert "trial_credits" not in result diff --git a/api/tests/test_containers_integration_tests/services/tools/test_workflow_tools_manage_service.py b/api/tests/test_containers_integration_tests/services/tools/test_workflow_tools_manage_service.py index 6f342e63dc8..b12472c586c 100644 --- a/api/tests/test_containers_integration_tests/services/tools/test_workflow_tools_manage_service.py +++ b/api/tests/test_containers_integration_tests/services/tools/test_workflow_tools_manage_service.py @@ -107,7 +107,7 @@ class TestWorkflowToolManageService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Create workflow for the app workflow = WorkflowModel( diff --git a/api/tests/test_containers_integration_tests/services/workflow/test_workflow_converter.py b/api/tests/test_containers_integration_tests/services/workflow/test_workflow_converter.py index ce5c2bd162f..8cd9526f6a0 100644 --- a/api/tests/test_containers_integration_tests/services/workflow/test_workflow_converter.py +++ b/api/tests/test_containers_integration_tests/services/workflow/test_workflow_converter.py @@ -217,6 +217,7 @@ class TestWorkflowConverter: icon_type="emoji", icon="🚀", icon_background="#4CAF50", + session=db_session_with_containers, ) # Assert: Verify the expected outcomes @@ -291,6 +292,7 @@ class TestWorkflowConverter: icon_type="emoji", icon="🚀", icon_background="#4CAF50", + session=db_session_with_containers, ) # Verify database state remains unchanged @@ -325,6 +327,7 @@ class TestWorkflowConverter: app_model=app, app_model_config=app.app_model_config, account_id=account.id, + session=db_session_with_containers, ) # Assert: Verify the expected outcomes @@ -467,6 +470,7 @@ class TestWorkflowConverter: app_model=app, variables=variables, external_data_variables=external_data_variables, + session=db_session_with_containers, ) # Assert: Verify the expected outcomes @@ -569,7 +573,7 @@ class TestConvertToHttpRequestNodeVariants: """Tests for chatbot vs workflow differences in HTTP request node conversion.""" @staticmethod - def _setup(app_mode, default_variables): + def _setup(app_mode, default_variables, db_session_with_containers: Session): app_model = App( tenant_id="tenant_id", mode=app_mode, @@ -598,19 +602,20 @@ class TestConvertToHttpRequestNodeVariants: app_model=app_model, variables=default_variables, external_data_variables=ext_vars, + session=db_session_with_containers, ) return nodes - def test_chatbot_query_uses_sys_query(self, default_variables): - nodes = self._setup(AppMode.CHAT, default_variables) + def test_chatbot_query_uses_sys_query(self, default_variables, db_session_with_containers: Session): + nodes = self._setup(AppMode.CHAT, default_variables, db_session_with_containers) body = json.loads(nodes[0]["data"]["body"]["data"]) assert body["params"]["query"] == "{{#sys.query#}}" assert body["point"] == APIBasedExtensionPoint.APP_EXTERNAL_DATA_TOOL_QUERY assert nodes[1]["data"]["type"] == "code" - def test_workflow_query_is_empty(self, default_variables): - nodes = self._setup(AppMode.WORKFLOW, default_variables) + def test_workflow_query_is_empty(self, default_variables, db_session_with_containers: Session): + nodes = self._setup(AppMode.WORKFLOW, default_variables, db_session_with_containers) body = json.loads(nodes[0]["data"]["body"]["data"]) assert body["params"]["query"] == "" diff --git a/api/tests/test_containers_integration_tests/trigger/test_trigger_e2e.py b/api/tests/test_containers_integration_tests/trigger/test_trigger_e2e.py index 9c20118e278..b6865510adf 100644 --- a/api/tests/test_containers_integration_tests/trigger/test_trigger_e2e.py +++ b/api/tests/test_containers_integration_tests/trigger/test_trigger_e2e.py @@ -194,7 +194,7 @@ def test_webhook_trigger_creates_trigger_log( db_session_with_containers.add_all([webhook_trigger, app_trigger]) db_session_with_containers.commit() - def _fake_trigger_workflow_async(session: Session, user: Any, trigger_data: Any) -> SimpleNamespace: + def _fake_trigger_workflow_async(user: Any, trigger_data: Any, *, session: Session) -> SimpleNamespace: log = WorkflowTriggerLog( tenant_id=trigger_data.tenant_id, app_id=trigger_data.app_id, @@ -575,7 +575,7 @@ def test_schedule_trigger_creates_trigger_log( db_session_with_containers.commit() # Mock AsyncWorkflowService to create WorkflowTriggerLog - def _fake_trigger_workflow_async(session: Session, user: Any, trigger_data: Any) -> SimpleNamespace: + def _fake_trigger_workflow_async(user: Any, trigger_data: Any, *, session: Session) -> SimpleNamespace: log = WorkflowTriggerLog( tenant_id=trigger_data.tenant_id, app_id=trigger_data.app_id, diff --git a/api/tests/unit_tests/commands/test_data_migration_commands.py b/api/tests/unit_tests/commands/test_data_migration_commands.py index b7f92f3291a..84e39d19a2d 100644 --- a/api/tests/unit_tests/commands/test_data_migration_commands.py +++ b/api/tests/unit_tests/commands/test_data_migration_commands.py @@ -109,8 +109,8 @@ def test_export_command_uses_cli_owned_session(monkeypatch, tmp_path: Path): package = MigrationPackage.from_mapping({"metadata": {"version": "1", "source_scope": "single"}}) class FakeMigrationExportService: - def export(self, export_session, selection): - captured["session"] = export_session + def export(self, selection, *, session): + captured["session"] = session captured["selection"] = selection return ExportResult(package=package, report_items=[], report_context=ReportContext()) @@ -156,8 +156,8 @@ def test_import_command_uses_cli_owned_session(monkeypatch, tmp_path: Path): ) class FakeMigrationImportService: - def import_package(self, import_session, request): - captured["session"] = import_session + def import_package(self, request, *, session): + captured["session"] = session captured["request"] = request return ImportResult(report_items=[], report_context=ReportContext(target_tenant="target")) diff --git a/api/tests/unit_tests/controllers/common/test_app_access.py b/api/tests/unit_tests/controllers/common/test_app_access.py index d070cc6e0fc..60a576346a0 100644 --- a/api/tests/unit_tests/controllers/common/test_app_access.py +++ b/api/tests/unit_tests/controllers/common/test_app_access.py @@ -152,7 +152,7 @@ class TestResolveAppAccessFilter: self._patch_whitelist(monkeypatch, ResourceWhitelistResources(unrestricted=False, resource_ids=[])) monkeypatch.setattr( f"{_RBAC_MODULE}.RBACService.MyPermissions.get", - lambda tenant_id, account_id: _permissions(workspace_keys=["app.create_and_management"]), + lambda tenant_id, account_id, session: _permissions(workspace_keys=["app.create_and_management"]), ) flt = resolve_app_access_filter("tenant-1", "acc-1") diff --git a/api/tests/unit_tests/controllers/console/agent/test_agent_controllers.py b/api/tests/unit_tests/controllers/console/agent/test_agent_controllers.py index 2e5851349d4..8f294293c07 100644 --- a/api/tests/unit_tests/controllers/console/agent/test_agent_controllers.py +++ b/api/tests/unit_tests/controllers/console/agent/test_agent_controllers.py @@ -236,7 +236,7 @@ def test_agent_app_list_and_create_use_agent_route( items=[_app_detail_obj(id="app-list", bound_agent_id="agent-list")], ) - def create_app(self, tenant_id: str, params, current_user: object) -> object: + def create_app(self, tenant_id: str, params, current_user: object, *, session: object) -> object: captured["create"] = {"tenant_id": tenant_id, "params": params, "current_user": current_user} return _app_detail_obj(id="app-created", bound_agent_id="agent-created") @@ -392,7 +392,8 @@ def test_agent_app_create_omits_optional_role_as_empty_string( captured: dict[str, object] = {} class FakeAppService: - def create_app(self, tenant_id: str, params: object, account: object) -> object: + def create_app(self, tenant_id: str, params: object, account: object, *, session: object) -> object: + del session captured["create"] = {"tenant_id": tenant_id, "params": params, "account": account} return _app_detail_obj(id="app-created", bound_agent_id="agent-created") @@ -472,11 +473,11 @@ def test_agent_app_detail_update_delete_resolve_app_from_agent_id( captured["get_app"] = app_obj return app_obj - def update_app(self, app_obj: object, args: dict[str, object]) -> object: + def update_app(self, app_obj: object, args: dict[str, object], *, session: object) -> object: captured["update"] = {"app": app_obj, "args": args} return _app_detail_obj(id="app-1", name=args["name"], bound_agent_id=agent_id) - def delete_app(self, app_obj: object) -> None: + def delete_app(self, app_obj: object, *, session: object) -> None: captured["delete"] = app_obj monkeypatch.setattr(roster_controller, "AppService", FakeAppService) @@ -661,18 +662,26 @@ def test_agent_publish_and_build_draft_routes_call_composer_service( discard_agent_app_build_draft, ) + def assert_call_without_session(key: str, expected: dict[str, object]) -> None: + call = dict(captured[key]) # type: ignore[arg-type] + assert call.pop("session", None) is not None + assert call == expected + with app.test_request_context( "/console/api/agent/00000000-0000-0000-0000-000000000001/publish", json={"version_note": "publish v1"}, ): published = unwrap(AgentPublishApi.post)(AgentPublishApi(), "tenant-1", current_user, agent_id) assert published["active_config_snapshot_id"] == "version-1" - assert captured["publish"] == { - "tenant_id": "tenant-1", - "agent_id": agent_id, - "account_id": account_id, - "version_note": "publish v1", - } + assert_call_without_session( + "publish", + { + "tenant_id": "tenant-1", + "agent_id": agent_id, + "account_id": account_id, + "version_note": "publish v1", + }, + ) with app.test_request_context( "/console/api/agent/00000000-0000-0000-0000-000000000001/build-draft/checkout", @@ -682,17 +691,20 @@ def test_agent_publish_and_build_draft_routes_call_composer_service( AgentBuildDraftCheckoutApi(), "tenant-1", current_user, agent_id ) assert checked_out["draft"]["id"] == "build-draft-1" - assert captured["checkout"] == { - "tenant_id": "tenant-1", - "agent_id": agent_id, - "account_id": account_id, - "force": True, - } + assert_call_without_session( + "checkout", + { + "tenant_id": "tenant-1", + "agent_id": agent_id, + "account_id": account_id, + "force": True, + }, + ) with app.test_request_context("/console/api/agent/00000000-0000-0000-0000-000000000001/build-draft"): loaded = unwrap(AgentBuildDraftApi.get)(AgentBuildDraftApi(), "tenant-1", current_user, agent_id) assert loaded["draft"]["id"] == "build-draft-1" - assert captured["load"] == {"tenant_id": "tenant-1", "agent_id": agent_id, "account_id": account_id} + assert_call_without_session("load", {"tenant_id": "tenant-1", "agent_id": agent_id, "account_id": account_id}) with app.test_request_context( "/console/api/agent/00000000-0000-0000-0000-000000000001/build-draft", @@ -711,7 +723,7 @@ def test_agent_publish_and_build_draft_routes_call_composer_service( ): applied = unwrap(AgentBuildDraftApplyApi.post)(AgentBuildDraftApplyApi(), "tenant-1", current_user, agent_id) assert applied == {"result": "success", "draft": {"id": "draft-1"}} - assert captured["apply"] == {"tenant_id": "tenant-1", "agent_id": agent_id, "account_id": account_id} + assert_call_without_session("apply", {"tenant_id": "tenant-1", "agent_id": agent_id, "account_id": account_id}) with app.test_request_context( "/console/api/agent/00000000-0000-0000-0000-000000000001/build-draft", @@ -719,7 +731,7 @@ def test_agent_publish_and_build_draft_routes_call_composer_service( ): discarded = unwrap(AgentBuildDraftApi.delete)(AgentBuildDraftApi(), "tenant-1", current_user, agent_id) assert discarded == {"result": "success"} - assert captured["discard"] == {"tenant_id": "tenant-1", "agent_id": agent_id, "account_id": account_id} + assert_call_without_session("discard", {"tenant_id": "tenant-1", "agent_id": agent_id, "account_id": account_id}) def test_agent_api_access_uses_agent_id_and_returns_service_api_metadata( @@ -775,7 +787,7 @@ def test_agent_api_status_and_key_routes_resolve_backing_app( monkeypatch.setattr(roster_controller, "_agent_api_key_count", lambda app_id: 1) class FakeAppService: - def update_app_api_status(self, app_obj: object, enable_api: bool) -> object: + def update_app_api_status(self, app_obj: object, enable_api: bool, *, session: object) -> object: captured["enable"] = {"app": app_obj, "enable_api": enable_api} app_model.enable_api = enable_api return app_model @@ -890,7 +902,7 @@ def test_agent_app_update_allows_empty_role(app: Flask, monkeypatch: pytest.Monk def get_app(self, app_obj: object) -> object: return app_obj - def update_app(self, app_obj: object, args: dict[str, object]) -> object: + def update_app(self, app_obj: object, args: dict[str, object], *, session: object) -> object: captured["update"] = {"app": app_obj, "args": args} return _app_detail_obj(id="app-1", name=args["name"], bound_agent_id=agent_id) @@ -1292,6 +1304,7 @@ def test_workflow_composer_copy_from_roster(app: Flask, monkeypatch: pytest.Monk ) assert result["binding"]["binding_type"] == "inline_agent" + assert captured.pop("session") is not None assert captured == { "tenant_id": "tenant-1", "app_id": "app-1", @@ -1896,8 +1909,18 @@ def test_list_agent_chat_messages_uses_current_user_conversation( captured.update(kwargs) return conversation + class SessionProxy: + def __call__(self): + return session + + def scalar(self, stmt: object): + return session.scalar(stmt) + + def scalars(self, stmt: object): + return session.scalars(stmt) + monkeypatch.setattr(message_controller.ConversationService, "get_conversation", get_conversation) - monkeypatch.setattr(message_controller, "db", SimpleNamespace(session=session)) + monkeypatch.setattr(message_controller, "db", SimpleNamespace(session=SessionProxy())) monkeypatch.setattr(message_controller, "attach_message_extra_contents", lambda messages: None) monkeypatch.setattr(message_controller, "MessageInfiniteScrollPaginationResponse", FakeMessagePaginationResponse) @@ -1905,6 +1928,7 @@ def test_list_agent_chat_messages_uses_current_user_conversation( result = message_controller._list_chat_messages(app_model=app_model, current_user=current_user) assert result == {"data": [message_id], "limit": 20, "has_more": False} + assert captured.pop("session") is session assert captured == {"app_model": app_model, "conversation_id": conversation_id, "user": current_user} diff --git a/api/tests/unit_tests/controllers/console/app/test_agent_app_sandbox.py b/api/tests/unit_tests/controllers/console/app/test_agent_app_sandbox.py index 8086f578956..0ab8814f368 100644 --- a/api/tests/unit_tests/controllers/console/app/test_agent_app_sandbox.py +++ b/api/tests/unit_tests/controllers/console/app/test_agent_app_sandbox.py @@ -48,6 +48,7 @@ class _WorkflowService: node_id: str, node_execution_id: str | None, path: str, + session, ) -> SandboxListResponse: self.calls.append(("list", tenant_id, app_id, workflow_run_id, node_id, node_execution_id, path)) return SandboxListResponse(path=path, entries=[], truncated=False) @@ -61,6 +62,7 @@ class _WorkflowService: node_id: str, node_execution_id: str | None, path: str, + session, ) -> SandboxReadResponse: self.calls.append(("read", tenant_id, app_id, workflow_run_id, node_id, node_execution_id, path)) return SandboxReadResponse(path=path, size=5, truncated=False, binary=False, text="hello") @@ -74,6 +76,7 @@ class _WorkflowService: node_id: str, node_execution_id: str | None, path: str, + session, ) -> AgentSandboxUploadDownload: self.calls.append(("upload", tenant_id, app_id, workflow_run_id, node_id, node_execution_id, path)) return AgentSandboxUploadDownload(url="https://files.example/upload.txt") diff --git a/api/tests/unit_tests/controllers/console/app/test_annotation_api.py b/api/tests/unit_tests/controllers/console/app/test_annotation_api.py index 8a6094b94b8..cc95f7f8a94 100644 --- a/api/tests/unit_tests/controllers/console/app/test_annotation_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_annotation_api.py @@ -2,7 +2,7 @@ from __future__ import annotations from inspect import unwrap from types import SimpleNamespace -from unittest.mock import Mock, patch +from unittest.mock import ANY, Mock, patch import pytest from flask import Flask @@ -142,7 +142,7 @@ class TestConsoleAnnotationRefBoundaries: assert response == "" assert status == 204 - delete_mock.assert_called_once_with(AppRef("tenant-1", "app-1"), ["ann-1", "ann-2"]) + delete_mock.assert_called_once_with(AppRef("tenant-1", "app-1"), ["ann-1", "ann-2"], session=ANY) def test_update_uses_annotation_ref(self, app: Flask): api = annotation_module.AnnotationUpdateDeleteApi() @@ -216,4 +216,4 @@ class TestConsoleAnnotationRefBoundaries: response = handler(api, "app-1", "ann-1") assert response["total"] == 1 - hit_history_mock.assert_called_once_with(AnnotationRef("tenant-1", "app-1", "ann-1"), 2, 5) + hit_history_mock.assert_called_once_with(AnnotationRef("tenant-1", "app-1", "ann-1"), 2, 5, session=ANY) diff --git a/api/tests/unit_tests/controllers/console/app/test_annotation_security.py b/api/tests/unit_tests/controllers/console/app/test_annotation_security.py index bfa4048191f..6a22d8769bc 100644 --- a/api/tests/unit_tests/controllers/console/app/test_annotation_security.py +++ b/api/tests/unit_tests/controllers/console/app/test_annotation_security.py @@ -193,9 +193,7 @@ class TestAnnotationImportServiceValidation: @pytest.fixture def mock_db_session(self): - """Mock database session.""" - with patch("services.annotation_service.db.session") as mock: - yield mock + return MagicMock() def test_max_records_limit_enforced(self, mock_app, mock_db_session): """Test that files with too many records are rejected.""" @@ -214,7 +212,7 @@ class TestAnnotationImportServiceValidation: with patch("services.annotation_service.FeatureService") as mock_features: mock_features.get_features.return_value.billing.enabled = False - result = AppAnnotationService.batch_import_app_annotations("app_id", file) + result = AppAnnotationService.batch_import_app_annotations("app_id", file, session=mock_db_session) # Should return error about too many records assert "error_msg" in result @@ -231,7 +229,7 @@ class TestAnnotationImportServiceValidation: with patch("services.annotation_service.current_account_with_tenant") as mock_auth: mock_auth.return_value = (MagicMock(id="user_id"), "tenant_id") - result = AppAnnotationService.batch_import_app_annotations("app_id", file) + result = AppAnnotationService.batch_import_app_annotations("app_id", file, session=mock_db_session) # Should return error about insufficient records assert "error_msg" in result @@ -250,7 +248,7 @@ class TestAnnotationImportServiceValidation: ): mock_auth.return_value = (MagicMock(id="user_id"), "tenant_id") - result = AppAnnotationService.batch_import_app_annotations("app_id", file) + result = AppAnnotationService.batch_import_app_annotations("app_id", file, session=mock_db_session) assert "error_msg" in result assert "malformed" in result["error_msg"].lower() @@ -271,7 +269,9 @@ class TestAnnotationImportServiceValidation: with patch("services.annotation_service.batch_import_annotations_task") as mock_task: with patch("services.annotation_service.redis_client"): - result = AppAnnotationService.batch_import_app_annotations("app_id", file) + result = AppAnnotationService.batch_import_app_annotations( + "app_id", file, session=mock_db_session + ) # Should return success response assert "job_id" in result diff --git a/api/tests/unit_tests/controllers/console/app/test_app_response_models.py b/api/tests/unit_tests/controllers/console/app/test_app_response_models.py index e7784f8fd94..79f109fd383 100644 --- a/api/tests/unit_tests/controllers/console/app/test_app_response_models.py +++ b/api/tests/unit_tests/controllers/console/app/test_app_response_models.py @@ -6,7 +6,7 @@ from datetime import datetime from importlib import util from pathlib import Path from types import ModuleType, SimpleNamespace -from unittest.mock import MagicMock +from unittest.mock import ANY, MagicMock import pytest from flask import Flask @@ -500,7 +500,8 @@ def test_app_list_uses_injected_session_for_draft_workflows( ) session = MagicMock() session.execute.return_value.scalars.return_value.all.return_value = [workflow] - scoped_session = SimpleNamespace(execute=MagicMock(side_effect=AssertionError("db.session should not be used"))) + scoped_session = MagicMock() + scoped_session.execute.side_effect = AssertionError("db.session should not be used") monkeypatch.setattr( app_module, @@ -515,7 +516,7 @@ def test_app_list_uses_injected_session_for_draft_workflows( monkeypatch.setattr( app_module.enterprise_rbac_service.RBACService.MyPermissions, "get", - lambda tenant_id, account_id: app_module.enterprise_rbac_service.MyPermissionsResponse( + lambda tenant_id, account_id, session: app_module.enterprise_rbac_service.MyPermissionsResponse( app=app_module.enterprise_rbac_service.ResourcePermissionSnapshot( overrides=[ app_module.enterprise_rbac_service.ResourcePermissionKeys( @@ -563,12 +564,12 @@ def test_app_create_api_attaches_permission_keys(app, app_module): monkeypatch.setattr( app_module, "AppService", - lambda: SimpleNamespace(create_app=lambda tenant_id, params, user: app_obj), + lambda: SimpleNamespace(create_app=lambda tenant_id, params, user, session: app_obj), ) monkeypatch.setattr( app_module.enterprise_rbac_service.RBACService.AppPermissions, "batch_get", - lambda tenant_id, account_id, app_ids: {"app-new": ["app.acl.view_layout", "app.acl.edit"]}, + lambda tenant_id, account_id, app_ids, session: {"app-new": ["app.acl.view_layout", "app.acl.edit"]}, ) resp, status = method(app_module.AppListApi(), "tenant-1", SimpleNamespace(id="acct-1")) @@ -611,7 +612,7 @@ def test_app_list_api_attaches_permission_keys(app, app_module): monkeypatch.setattr( app_module.enterprise_rbac_service.RBACService.MyPermissions, "get", - lambda tenant_id, account_id: app_module.enterprise_rbac_service.MyPermissionsResponse( + lambda tenant_id, account_id, session: app_module.enterprise_rbac_service.MyPermissionsResponse( app=app_module.enterprise_rbac_service.ResourcePermissionSnapshot( default_permission_keys=["app.preview", "app.acl.view_layout"], overrides=[ @@ -655,7 +656,7 @@ def test_app_list_api_limits_to_apps_created_by_current_user_without_view_permis monkeypatch.setattr( app_module.enterprise_rbac_service.RBACService.MyPermissions, "get", - lambda tenant_id, account_id: app_module.enterprise_rbac_service.MyPermissionsResponse( + lambda tenant_id, account_id, session: app_module.enterprise_rbac_service.MyPermissionsResponse( workspace=app_module.enterprise_rbac_service.WorkspacePermissionSnapshot( permission_keys=["app.create_and_management"] ) @@ -698,7 +699,7 @@ def test_app_list_api_limits_to_preview_overrides_without_manage_own_permission( monkeypatch.setattr( app_module.enterprise_rbac_service.RBACService.MyPermissions, "get", - lambda tenant_id, account_id: app_module.enterprise_rbac_service.MyPermissionsResponse( + lambda tenant_id, account_id, session: app_module.enterprise_rbac_service.MyPermissionsResponse( app=app_module.enterprise_rbac_service.ResourcePermissionSnapshot( overrides=[ app_module.enterprise_rbac_service.ResourcePermissionKeys( @@ -754,7 +755,7 @@ def test_app_list_api_returns_no_apps_without_workspace_or_resource_view_permiss monkeypatch.setattr( app_module.enterprise_rbac_service.RBACService.MyPermissions, "get", - lambda tenant_id, account_id: app_module.enterprise_rbac_service.MyPermissionsResponse(), + lambda tenant_id, account_id, session: app_module.enterprise_rbac_service.MyPermissionsResponse(), ) monkeypatch.setattr( app_module.enterprise_rbac_service.RBACService.AppAccess, @@ -820,7 +821,7 @@ def test_app_detail_api_attaches_current_user_permission_keys(app, app_module): resp = method(app_module.AppApi(), "tenant-1", SimpleNamespace(id="acct-1"), app_model=app_obj) - get_permissions.assert_called_once_with("tenant-1", "acct-1", app_id="app-1") + get_permissions.assert_called_once_with("tenant-1", "acct-1", app_id="app-1", session=ANY) assert resp["permission_keys"] == ["app.acl.view_layout", "app.acl.edit", "app.acl.monitor"] @@ -861,7 +862,7 @@ def test_app_copy_api_attaches_permission_keys(app, app_module): "get_system_features", lambda: SimpleNamespace(webapp_auth=SimpleNamespace(enabled=False)), ) - monkeypatch.setattr(app_module, "db", SimpleNamespace(engine=object())) + monkeypatch.setattr(app_module, "db", SimpleNamespace(engine=object(), session=lambda: MagicMock())) monkeypatch.setattr( app_module, "Session", @@ -870,7 +871,7 @@ def test_app_copy_api_attaches_permission_keys(app, app_module): monkeypatch.setattr( app_module.enterprise_rbac_service.RBACService.AppPermissions, "batch_get", - lambda tenant_id, account_id, app_ids: {"app-new": ["app.acl.view_layout", "app.acl.edit"]}, + lambda tenant_id, account_id, app_ids, session: {"app-new": ["app.acl.view_layout", "app.acl.edit"]}, ) resp, status = method( diff --git a/api/tests/unit_tests/controllers/console/app/test_workflow.py b/api/tests/unit_tests/controllers/console/app/test_workflow.py index 2f971eaf74f..2d811deb916 100644 --- a/api/tests/unit_tests/controllers/console/app/test_workflow.py +++ b/api/tests/unit_tests/controllers/console/app/test_workflow.py @@ -621,7 +621,7 @@ def test_workflow_online_users_filters_inaccessible_workflow(app: Flask, monkeyp monkeypatch.setattr( workflow_module, "WorkflowService", - lambda: SimpleNamespace(get_accessible_app_ids=lambda app_ids, tenant_id: {app_id_1}), + lambda: SimpleNamespace(get_accessible_app_ids=lambda app_ids, tenant_id, session: {app_id_1}), ) monkeypatch.setattr(workflow_module.file_helpers, "get_signed_file_url", sign_avatar) @@ -703,7 +703,7 @@ def test_workflow_online_users_batches_redis_reads(app: Flask, monkeypatch: pyte monkeypatch.setattr( workflow_module, "WorkflowService", - lambda: SimpleNamespace(get_accessible_app_ids=lambda app_ids, tenant_id: set(app_ids)), + lambda: SimpleNamespace(get_accessible_app_ids=lambda app_ids, tenant_id, session: set(app_ids)), ) first_pipeline = Mock() diff --git a/api/tests/unit_tests/controllers/console/app/test_workflow_human_input_debug_api.py b/api/tests/unit_tests/controllers/console/app/test_workflow_human_input_debug_api.py index f04ab6d6e7c..956706eafb6 100644 --- a/api/tests/unit_tests/controllers/console/app/test_workflow_human_input_debug_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_workflow_human_input_debug_api.py @@ -2,7 +2,7 @@ from __future__ import annotations from dataclasses import dataclass from types import SimpleNamespace -from unittest.mock import MagicMock +from unittest.mock import ANY, MagicMock import pytest from flask import Flask @@ -94,6 +94,7 @@ def test_human_input_preview_delegates_to_service( account=account, node_id="node-42", inputs={"topic": "tech"}, + session=ANY, ) @@ -144,6 +145,7 @@ def test_human_input_submit_forwards_payload(app: Flask, monkeypatch: pytest.Mon form_inputs={"answer": "42"}, inputs={"#node-1.result#": "LLM output"}, action="approve", + session=ANY, ) @@ -193,6 +195,7 @@ def test_human_input_delivery_test_calls_service( node_id="node-7", delivery_method_id="delivery-123", inputs={}, + session=ANY, ) diff --git a/api/tests/unit_tests/controllers/console/app/test_workflow_node_output_inspector.py b/api/tests/unit_tests/controllers/console/app/test_workflow_node_output_inspector.py index e66ae5246bc..dfe35a89f57 100644 --- a/api/tests/unit_tests/controllers/console/app/test_workflow_node_output_inspector.py +++ b/api/tests/unit_tests/controllers/console/app/test_workflow_node_output_inspector.py @@ -25,7 +25,7 @@ from __future__ import annotations import json from collections.abc import Iterator from typing import Any -from unittest.mock import MagicMock +from unittest.mock import ANY, MagicMock from uuid import UUID import pytest @@ -382,7 +382,9 @@ def test_serve_snapshot_happy_path(patch_service, app_model, run_id): result = ctrl._serve_snapshot(app_model, run_id) assert isinstance(result, dict) assert result["workflow_run_id"] == "00000000-0000-0000-0000-0000000000aa" - patch_service.snapshot_workflow_run.assert_called_once_with(app_model=app_model, workflow_run_id=str(run_id)) + patch_service.snapshot_workflow_run.assert_called_once_with( + app_model=app_model, workflow_run_id=str(run_id), session=ANY + ) def test_serve_snapshot_translates_inspector_error_to_404(patch_service, app_model, run_id): @@ -399,7 +401,7 @@ def test_serve_node_detail_happy_path(patch_service, app_model, run_id): result = ctrl._serve_node_detail(app_model, run_id, "agent-1") assert result["node_id"] == "agent-1" patch_service.node_detail.assert_called_once_with( - app_model=app_model, workflow_run_id=str(run_id), node_id="agent-1" + app_model=app_model, workflow_run_id=str(run_id), node_id="agent-1", session=ANY ) @@ -431,6 +433,7 @@ def test_serve_output_preview_happy_path(patch_service, app_model, run_id): workflow_run_id=str(run_id), node_id="agent-1", output_name="text", + session=ANY, ) diff --git a/api/tests/unit_tests/controllers/console/auth/test_account_activation.py b/api/tests/unit_tests/controllers/console/auth/test_account_activation.py index ebae7de6c15..001ca0bf8fb 100644 --- a/api/tests/unit_tests/controllers/console/auth/test_account_activation.py +++ b/api/tests/unit_tests/controllers/console/auth/test_account_activation.py @@ -597,7 +597,7 @@ class TestActivateApi: assert response["result"] == "success" mock_create_tenant_member.assert_called_once_with( - mock_invitation["tenant"], mock_account, mock_db.session, role=TenantAccountRole.ADMIN + mock_invitation["tenant"], mock_account, mock_db.session(), role=TenantAccountRole.ADMIN ) mock_switch_tenant.assert_called_once_with(mock_account, mock_invitation["tenant"].id, session=ANY) mock_revoke_token.assert_called_once_with("workspace-123", "invitee@example.com", "valid_token") diff --git a/api/tests/unit_tests/controllers/console/auth/test_data_source_bearer_auth.py b/api/tests/unit_tests/controllers/console/auth/test_data_source_bearer_auth.py index 21d1932f820..b231826aeac 100644 --- a/api/tests/unit_tests/controllers/console/auth/test_data_source_bearer_auth.py +++ b/api/tests/unit_tests/controllers/console/auth/test_data_source_bearer_auth.py @@ -43,7 +43,7 @@ def test_list_data_source_auth_uses_injected_tenant_id() -> None: ): result = method(api, "tenant-1") - get_provider_auth_list.assert_called_once_with(ANY, "tenant-1") + get_provider_auth_list.assert_called_once_with("tenant-1", session=ANY) assert result["sources"][0]["id"] == "binding-1" assert result["sources"][0]["provider"] == "custom" @@ -65,7 +65,7 @@ def test_create_data_source_auth_binding_uses_injected_tenant_id() -> None: ): result, status = method(api, "tenant-1") - create_auth.assert_called_once_with(ANY, "tenant-1", payload) + create_auth.assert_called_once_with("tenant-1", payload, session=ANY) assert result == {"result": "success"} assert status == 200 @@ -82,6 +82,6 @@ def test_delete_data_source_auth_binding_uses_injected_tenant_id() -> None: ): result, status = method(api, "tenant-1", "binding-1") - delete_provider_auth.assert_called_once_with(ANY, "tenant-1", "binding-1") + delete_provider_auth.assert_called_once_with("tenant-1", "binding-1", session=ANY) assert result == "" assert status == 204 diff --git a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_datasource_auth.py b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_datasource_auth.py index 33aaa19e640..8f66ca5c993 100644 --- a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_datasource_auth.py +++ b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_datasource_auth.py @@ -1,6 +1,6 @@ import inspect from datetime import UTC, datetime -from unittest.mock import MagicMock, patch +from unittest.mock import ANY, MagicMock, patch import pytest from flask import Flask @@ -508,6 +508,7 @@ class TestDatasourceAuthDeleteApi: auth_id="cred-1", provider="notion", plugin_id="langgenius/notion_datasource", + session=ANY, ) def test_delete_missing_credential_id(self, app: Flask): diff --git a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py index 2a1970d3837..39a6fa65a06 100644 --- a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py +++ b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py @@ -64,10 +64,9 @@ class TestPipelineTemplateListApi: tenant_id = "tenant-1" service_calls: list[tuple[str, str, str]] = [] - def get_pipeline_templates( - session: Mock, template_type: str, language: str, current_tenant_id: str - ) -> dict[str, object]: - service_calls.append((template_type, language, current_tenant_id)) + def get_pipeline_templates(*, type: str, language: str, current_tenant_id: str, session) -> dict[str, object]: + del session + service_calls.append((type, language, current_tenant_id)) return {"pipeline_templates": [_template_item()]} with ( @@ -94,10 +93,9 @@ class TestPipelineTemplateListApi: tenant_id = "tenant-1" service_calls: list[tuple[str, str, str]] = [] - def get_pipeline_templates( - session: Mock, template_type: str, language: str, current_tenant_id: str - ) -> dict[str, object]: - service_calls.append((template_type, language, current_tenant_id)) + def get_pipeline_templates(*, type: str, language: str, current_tenant_id: str, session) -> dict[str, object]: + del session + service_calls.append((type, language, current_tenant_id)) return {"pipeline_templates": []} with ( @@ -117,16 +115,18 @@ class TestPipelineTemplateDetailApi: method = unwrap(api.get) service_calls: list[tuple[str, str]] = [] - class Service: - def get_pipeline_template_detail( - self, session: Mock, template_id: str, template_type: str - ) -> dict[str, object]: - service_calls.append((template_id, template_type)) - return _template_detail() + def get_pipeline_template_detail(template_id: str, type: str, *, session) -> dict[str, object]: + del session + service_calls.append((template_id, type)) + return _template_detail() with ( app.test_request_context("/rag/pipeline/templates/template-1?type=customized"), - patch.object(module, "RagPipelineService", Service), + patch.object( + module.RagPipelineService, + "get_pipeline_template_detail", + side_effect=get_pipeline_template_detail, + ), ): response, status = method(api, Mock(), "template-1") @@ -138,13 +138,16 @@ class TestPipelineTemplateDetailApi: api = PipelineTemplateDetailApi() method = unwrap(api.get) - class Service: - def get_pipeline_template_detail(self, session: Mock, template_id: str, template_type: str) -> None: - return None + def get_pipeline_template_detail(template_id: str, type: str, *, session) -> None: + del template_id, type, session with ( app.test_request_context("/rag/pipeline/templates/missing"), - patch.object(module, "RagPipelineService", Service), + patch.object( + module.RagPipelineService, + "get_pipeline_template_detail", + side_effect=get_pipeline_template_detail, + ), ): with pytest.raises(NotFound): method(api, Mock(), "missing") @@ -160,8 +163,14 @@ class TestCustomizedPipelineTemplateApi: service_calls: list[tuple[str, PipelineTemplateInfoEntity, Account, str]] = [] def update_template( - template_id: str, template_info: PipelineTemplateInfoEntity, current_user: Account, current_tenant_id: str + template_id: str, + template_info: PipelineTemplateInfoEntity, + current_user: Account, + current_tenant_id: str, + *, + session, ) -> None: + del session service_calls.append((template_id, template_info, current_user, current_tenant_id)) with ( @@ -198,8 +207,14 @@ class TestCustomizedPipelineTemplateApi: service_calls: list[tuple[str, PipelineTemplateInfoEntity, Account, str]] = [] def update_template( - template_id: str, template_info: PipelineTemplateInfoEntity, current_user: Account, current_tenant_id: str + template_id: str, + template_info: PipelineTemplateInfoEntity, + current_user: Account, + current_tenant_id: str, + *, + session, ) -> None: + del session service_calls.append((template_id, template_info, current_user, current_tenant_id)) with ( @@ -228,7 +243,8 @@ class TestCustomizedPipelineTemplateApi: tenant_id = "tenant-1" deleted_templates: list[tuple[str, str]] = [] - def delete_template(template_id: str, current_tenant_id: str) -> None: + def delete_template(template_id: str, current_tenant_id: str, *, session) -> None: + del session deleted_templates.append((template_id, current_tenant_id)) with ( @@ -325,9 +341,19 @@ class TestPublishCustomizedPipelineTemplateApi: service_calls: list[tuple[str, dict[str, object], Account, str]] = [] class Service: + def __init__(self, *args, **kwargs) -> None: + pass + def publish_customized_pipeline_template( - self, pipeline_id: str, data: dict[str, object], current_user: Account, current_tenant_id: str + self, + pipeline_id: str, + data: dict[str, object], + current_user: Account, + current_tenant_id: str, + *, + session, ) -> None: + del session service_calls.append((pipeline_id, data, current_user, current_tenant_id)) with ( @@ -352,9 +378,19 @@ class TestPublishCustomizedPipelineTemplateApi: service_calls: list[tuple[str, dict[str, object], Account, str]] = [] class Service: + def __init__(self, *args, **kwargs) -> None: + pass + def publish_customized_pipeline_template( - self, pipeline_id: str, data: dict[str, object], current_user: Account, current_tenant_id: str + self, + pipeline_id: str, + data: dict[str, object], + current_user: Account, + current_tenant_id: str, + *, + session, ) -> None: + del session service_calls.append((pipeline_id, data, current_user, current_tenant_id)) with ( diff --git a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py index 19dc90ed8a4..e344a4c8bab 100644 --- a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py +++ b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py @@ -52,7 +52,9 @@ def _pipeline() -> Pipeline: def test_draft_rag_pipeline_workflow_get_serializes_response_model(monkeypatch: pytest.MonkeyPatch) -> None: workflow = _make_workflow() monkeypatch.setattr( - module, "RagPipelineService", lambda: SimpleNamespace(get_draft_workflow=lambda **_kwargs: workflow) + module, + "RagPipelineService", + lambda *_args, **_kwargs: SimpleNamespace(get_draft_workflow=lambda **_kwargs: workflow), ) api = module.DraftRagPipelineApi() @@ -97,12 +99,12 @@ def test_published_rag_pipeline_workflows_serialize_items_before_session_closes( assert session_state["open"] is True return getattr(base_workflow, name) - monkeypatch.setattr(module, "db", SimpleNamespace(engine=object())) + monkeypatch.setattr(module, "db", SimpleNamespace(engine=object(), session=lambda: object())) monkeypatch.setattr(module, "sessionmaker", lambda *_args, **_kwargs: _SessionMaker()) monkeypatch.setattr( module, "RagPipelineService", - lambda: SimpleNamespace(get_all_published_workflow=lambda **_kwargs: ([_Workflow()], False)), + lambda *_args, **_kwargs: SimpleNamespace(get_all_published_workflow=lambda **_kwargs: ([_Workflow()], False)), ) with app.test_request_context( @@ -132,12 +134,12 @@ def test_rag_pipeline_workflow_patch_serializes_response_model(app: Flask, monke def begin(self): return _SessionContext() - monkeypatch.setattr(module, "db", SimpleNamespace(engine=object())) + monkeypatch.setattr(module, "db", SimpleNamespace(engine=object(), session=lambda: object())) monkeypatch.setattr(module, "sessionmaker", lambda *_args, **_kwargs: _SessionMaker()) monkeypatch.setattr( module, "RagPipelineService", - lambda: SimpleNamespace(update_workflow=lambda **_kwargs: workflow), + lambda *_args, **_kwargs: SimpleNamespace(update_workflow=lambda **_kwargs: workflow), ) payload: dict[str, object] = {"marked_name": "Updated release"} @@ -165,7 +167,7 @@ def test_default_rag_pipeline_block_configs_serializes_root_response(monkeypatch monkeypatch.setattr( module, "RagPipelineService", - lambda: SimpleNamespace(get_default_block_configs=lambda: block_configs), + lambda *_args, **_kwargs: SimpleNamespace(get_default_block_configs=lambda: block_configs), ) api = module.DefaultRagPipelineBlockConfigsApi() @@ -190,7 +192,7 @@ def test_draft_rag_pipeline_second_step_parameters_serializes_variables(app, mon monkeypatch.setattr( module, "RagPipelineService", - lambda: SimpleNamespace(get_second_step_parameters=lambda **_kwargs: variables), + lambda *_args, **_kwargs: SimpleNamespace(get_second_step_parameters=lambda **_kwargs: variables), ) api = module.DraftRagPipelineSecondStepApi() @@ -210,7 +212,7 @@ def test_rag_pipeline_recommended_plugins_serializes_known_envelope(app, monkeyp monkeypatch.setattr( module, "RagPipelineService", - lambda: SimpleNamespace(get_recommended_plugins=lambda *_args: recommended_plugins), + lambda *_args, **_kwargs: SimpleNamespace(get_recommended_plugins=lambda *_args: recommended_plugins), ) api = module.RagPipelineRecommendedPluginApi() diff --git a/api/tests/unit_tests/controllers/console/datasets/test_datasets.py b/api/tests/unit_tests/controllers/console/datasets/test_datasets.py index 53f4f139937..6913825d599 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_datasets.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_datasets.py @@ -3,7 +3,7 @@ import json from contextlib import ExitStack from inspect import unwrap from types import SimpleNamespace -from unittest.mock import MagicMock, PropertyMock, patch +from unittest.mock import ANY, MagicMock, PropertyMock, patch import pytest from flask import Flask @@ -63,6 +63,18 @@ def dataset_model_property_defaults(): for name, value in properties.items(): property_mock = stack.enter_context(patch.object(Dataset, name, new_callable=PropertyMock)) property_mock.return_value = value + stack.enter_context( + patch( + "controllers.console.datasets.datasets.enterprise_rbac_service.RBACService.MyPermissions.get", + return_value=enterprise_rbac_service.MyPermissionsResponse(), + ) + ) + stack.enter_context( + patch( + "controllers.console.datasets.datasets.enterprise_rbac_service.RBACService.DatasetPermissions.batch_get", + return_value={}, + ) + ) yield @@ -245,7 +257,7 @@ class TestDatasetList: ): resp, status = method(api, "tenant-1", current_user) - get_permissions.assert_called_once_with("tenant-1", current_user.id) + get_permissions.assert_called_once_with("tenant-1", current_user.id, session=ANY) assert status == 200 assert resp["data"][0]["permission_keys"] == ["dataset.acl.readonly", "dataset.acl.edit"] @@ -742,7 +754,7 @@ class TestDatasetApiGet: data, status = method(api, tenant_id, user, dataset_id) - get_permissions.assert_called_once_with(tenant_id, user.id, dataset_id=dataset_id) + get_permissions.assert_called_once_with(tenant_id, user.id, dataset_id=dataset_id, session=ANY) assert status == 200 assert data["permission_keys"] == ["dataset.acl.readonly", "dataset.acl.edit"] diff --git a/api/tests/unit_tests/controllers/console/datasets/test_external.py b/api/tests/unit_tests/controllers/console/datasets/test_external.py index 8ac40f03d3b..1cffc90ae23 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_external.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_external.py @@ -1,5 +1,6 @@ import inspect -from unittest.mock import MagicMock, PropertyMock, patch +from types import SimpleNamespace +from unittest.mock import ANY, MagicMock, PropertyMock, patch import pytest from flask import Flask @@ -7,6 +8,7 @@ from werkzeug.exceptions import Forbidden, NotFound import services from controllers.console import console_ns +from controllers.console.datasets import external as external_module from controllers.console.datasets.error import DatasetNameDuplicateError from controllers.console.datasets.external import ( BedrockRetrievalApi, @@ -142,7 +144,7 @@ class TestExternalApiUseCheckApi: assert status == 200 assert response == {"is_using": True, "count": 2} - mock_use_check.assert_called_once_with(session, "api-id", "tenant-1") + mock_use_check.assert_called_once_with("api-id", "tenant-1", session=ANY) class TestExternalDatasetCreateApi: @@ -186,6 +188,7 @@ class TestExternalDatasetCreateApi: "create_external_dataset", return_value=dataset, ), + patch.object(external_module, "db", SimpleNamespace(session=lambda: MagicMock())), ): _, status = method(api, MagicMock(), "tenant-1", current_user) diff --git a/api/tests/unit_tests/controllers/console/explore/test_recommended_app.py b/api/tests/unit_tests/controllers/console/explore/test_recommended_app.py index 8a2e14cce9b..4adeaaa90dd 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_recommended_app.py +++ b/api/tests/unit_tests/controllers/console/explore/test_recommended_app.py @@ -32,7 +32,7 @@ class TestRecommendedAppListApi: ): result = method(api, make_account("fr-FR")) - service_mock.assert_called_once_with(ANY, "en-US") + service_mock.assert_called_once_with("en-US", session=ANY) assert result == result_data def test_get_fallback_to_user_language(self, app: Flask): @@ -51,7 +51,7 @@ class TestRecommendedAppListApi: ): result = method(api, make_account("fr-FR")) - service_mock.assert_called_once_with(ANY, "fr-FR") + service_mock.assert_called_once_with("fr-FR", session=ANY) assert result == result_data def test_get_fallback_to_default_language(self, app: Flask): @@ -70,7 +70,7 @@ class TestRecommendedAppListApi: ): result = method(api, make_account(None)) - service_mock.assert_called_once_with(ANY, module.languages[0]) + service_mock.assert_called_once_with(module.languages[0], session=ANY) assert result == result_data @@ -91,7 +91,7 @@ class TestLearnDifyAppListApi: ): result = method(api, make_account("fr-FR")) - service_mock.assert_called_once_with(ANY, "en-US") + service_mock.assert_called_once_with("en-US", session=ANY) assert result == result_data def test_get_fallback_to_user_language(self, app: Flask): @@ -110,7 +110,7 @@ class TestLearnDifyAppListApi: ): result = method(api, make_account("fr-FR")) - service_mock.assert_called_once_with(ANY, "fr-FR") + service_mock.assert_called_once_with("fr-FR", session=ANY) assert result == result_data @@ -131,7 +131,7 @@ class TestRecommendedAppApi: ): result = method(api, "11111111-1111-1111-1111-111111111111") - service_mock.assert_called_once_with(ANY, "11111111-1111-1111-1111-111111111111") + service_mock.assert_called_once_with("11111111-1111-1111-1111-111111111111", session=ANY) assert result == result_data diff --git a/api/tests/unit_tests/controllers/console/explore/test_saved_message.py b/api/tests/unit_tests/controllers/console/explore/test_saved_message.py index ae05b8f6a0e..f210d0d5d04 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_saved_message.py +++ b/api/tests/unit_tests/controllers/console/explore/test_saved_message.py @@ -63,7 +63,7 @@ class TestSavedMessageListApi: result = method(api, current_user, installed_app) pagination_mock.assert_called_once() - assert pagination_mock.call_args.args[2] is current_user + assert pagination_mock.call_args.args[1] is current_user assert result["limit"] == 20 assert result["has_more"] is False assert len(result["data"]) == 2 @@ -96,7 +96,7 @@ class TestSavedMessageListApi: result = method(api, current_user, installed_app) save_mock.assert_called_once() - assert save_mock.call_args.args[2] is current_user + assert save_mock.call_args.args[1] is current_user assert result == {"result": "success"} def test_post_message_not_exists(self, app: Flask, payload_patch): @@ -136,7 +136,7 @@ class TestSavedMessageApi: result, status = method(api, current_user, installed_app, str(uuid4())) delete_mock.assert_called_once() - assert delete_mock.call_args.args[2] is current_user + assert delete_mock.call_args.args[1] is current_user assert status == 204 assert result == "" diff --git a/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow.py b/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow.py index 8785ce85109..98b538800ac 100644 --- a/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow.py +++ b/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow.py @@ -38,7 +38,10 @@ def _snippet(**overrides) -> CustomizedSnippet: @pytest.fixture(autouse=True) def _patch_snippet_service_factory(monkeypatch: pytest.MonkeyPatch) -> None: def factory(): - return snippet_workflow_module.SnippetService() + try: + return snippet_workflow_module.SnippetService(snippet_workflow_module._snippet_session_maker()) + except TypeError: + return snippet_workflow_module.SnippetService() monkeypatch.setattr(snippet_workflow_module, "_snippet_service", factory) monkeypatch.setattr(snippet_workflow_module, "_snippet_session_maker", Mock(return_value=Mock())) diff --git a/api/tests/unit_tests/controllers/console/tag/test_tags.py b/api/tests/unit_tests/controllers/console/tag/test_tags.py index 2da11afa1f7..8aaebeb124a 100644 --- a/api/tests/unit_tests/controllers/console/tag/test_tags.py +++ b/api/tests/unit_tests/controllers/console/tag/test_tags.py @@ -3,7 +3,7 @@ from unittest.mock import MagicMock, PropertyMock, patch import pytest from flask import Flask -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden import controllers.console.tag.tags as module @@ -22,7 +22,7 @@ from services.tag_service import UpdateTagPayload class SessionMatcher: def __eq__(self, other): - return isinstance(other, Session | scoped_session) + return isinstance(other, Session) def unwrap(func): @@ -131,7 +131,7 @@ class TestTagListApi: ): result, status = method(api, "tenant-1") - get_tags_mock.assert_called_once_with(SessionMatcher(), "snippet", "tenant-1", None) + get_tags_mock.assert_called_once_with("snippet", "tenant-1", None, session=SessionMatcher()) assert status == 200 assert result == [{"id": "1", "name": "snippet-tag", "type": "snippet", "binding_count": "1"}] @@ -224,7 +224,7 @@ class TestTagUpdateDeleteApi: update_payload, tag_id, session = update_tags_mock.call_args.args assert update_payload == UpdateTagPayload(name="updated") assert tag_id == "tag-1" - assert session == module.db.session + assert session == SessionMatcher() assert result["binding_count"] == "3" def test_patch_forbidden(self, app: Flask, readonly_user, payload_patch): @@ -250,7 +250,7 @@ class TestTagUpdateDeleteApi: ): result, status = method(api, "tag-1") - delete_mock.assert_called_once_with("tag-1", module.db.session) + delete_mock.assert_called_once_with("tag-1", SessionMatcher()) assert status == 204 def test_delete_snippet_tag_checks_type_in_current_tenant(self, app: Flask, admin_user): @@ -278,7 +278,7 @@ class TestTagUpdateDeleteApi: scene=module.RBACPermission.SNIPPETS_CREATE_AND_MODIFY, resource_required=False, ) - delete_mock.assert_called_once_with("tag-1", module.db.session) + delete_mock.assert_called_once_with("tag-1", SessionMatcher()) assert result == "" assert status == 204 diff --git a/api/tests/unit_tests/controllers/console/test_extension.py b/api/tests/unit_tests/controllers/console/test_extension.py index bab825ca6f0..8ea327dfdce 100644 --- a/api/tests/unit_tests/controllers/console/test_extension.py +++ b/api/tests/unit_tests/controllers/console/test_extension.py @@ -114,7 +114,7 @@ def test_api_based_extension_get_returns_tenant_extensions(app: Flask, monkeypat assert response[0]["name"] == "Weather API" assert response[0]["api_endpoint"] == extension.api_endpoint assert response[0]["api_key"].startswith(extension.api_key[:3]) - service_mock.assert_called_once_with(ANY, "tenant-123") + service_mock.assert_called_once_with("tenant-123", session=ANY) def test_api_based_extension_post_creates_extension(app: Flask, monkeypatch: pytest.MonkeyPatch): @@ -132,7 +132,7 @@ def test_api_based_extension_post_creates_extension(app: Flask, monkeypatch: pyt response, status = APIBasedExtensionAPI().post() args, _ = save_mock.call_args - created_extension: APIBasedExtension = args[1] + created_extension: APIBasedExtension = args[0] assert created_extension.tenant_id == "tenant-123" assert created_extension.name == payload["name"] assert created_extension.api_endpoint == payload["api_endpoint"] @@ -157,7 +157,7 @@ def test_api_based_extension_detail_get_fetches_extension(app: Flask, monkeypatc assert response["id"] == extension.id assert response["name"] == extension.name - service_mock.assert_called_once_with(ANY, "tenant-123", str(extension_id)) + service_mock.assert_called_once_with("tenant-123", str(extension_id), session=ANY) def test_api_based_extension_detail_post_keeps_hidden_api_key(app: Flask, monkeypatch: pytest.MonkeyPatch): @@ -187,7 +187,7 @@ def test_api_based_extension_detail_post_keeps_hidden_api_key(app: Flask, monkey assert existing_extension.name == payload["name"] assert existing_extension.api_endpoint == payload["api_endpoint"] assert existing_extension.api_key == "keep-me" - save_mock.assert_called_once_with(ANY, existing_extension) + save_mock.assert_called_once_with(existing_extension, session=ANY) assert response["name"] == payload["name"] assert response["api_key"] == _masked_api_key("keep-me") @@ -217,7 +217,7 @@ def test_api_based_extension_detail_post_updates_api_key_when_provided(app: Flas response = APIBasedExtensionDetailAPI().post(extension_id) assert existing_extension.api_key == "new-secret" - save_mock.assert_called_once_with(ANY, existing_extension) + save_mock.assert_called_once_with(existing_extension, session=ANY) assert response["name"] == payload["name"] assert response["api_key"] == _masked_api_key(payload["api_key"]) @@ -239,6 +239,6 @@ def test_api_based_extension_detail_delete_removes_extension(app: Flask, monkeyp ): response, status = APIBasedExtensionDetailAPI().delete(extension_id) - delete_mock.assert_called_once_with(ANY, existing_extension) + delete_mock.assert_called_once_with(existing_extension, session=ANY) assert status == 204 assert response == "" diff --git a/api/tests/unit_tests/controllers/console/test_workspace_account.py b/api/tests/unit_tests/controllers/console/test_workspace_account.py index 5f36e805baa..39a3b2485dd 100644 --- a/api/tests/unit_tests/controllers/console/test_workspace_account.py +++ b/api/tests/unit_tests/controllers/console/test_workspace_account.py @@ -692,7 +692,7 @@ def test_get_account_by_email_with_case_fallback_uses_lowercase_lookup(): second.scalar_one_or_none.return_value = expected_account mock_session.execute.side_effect = [first, second] - result = AccountService.get_account_by_email_with_case_fallback(mock_session, "Mixed@Test.com") + result = AccountService.get_account_by_email_with_case_fallback("Mixed@Test.com", session=mock_session) assert result is expected_account assert mock_session.execute.call_count == 2 diff --git a/api/tests/unit_tests/controllers/console/workspace/test_load_balancing_config.py b/api/tests/unit_tests/controllers/console/workspace/test_load_balancing_config.py index a1d08849ee3..7d034f90642 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_load_balancing_config.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_load_balancing_config.py @@ -6,7 +6,7 @@ import builtins import importlib import sys from types import SimpleNamespace -from unittest.mock import MagicMock +from unittest.mock import ANY, MagicMock import pytest from flask import Flask @@ -92,6 +92,7 @@ def test_validate_credentials_success(app: Flask, load_balancing_module, monkeyp model="gpt-4o", model_type=ModelType.LLM, credentials={"api_key": "sk-***"}, + session=ANY, ) @@ -143,5 +144,6 @@ def test_validate_credentials_with_config_id(app: Flask, load_balancing_module, model="gpt-4o", model_type=ModelType.LLM, credentials={"api_key": "sk-***"}, + session=ANY, config_id="cfg-1", ) diff --git a/api/tests/unit_tests/controllers/console/workspace/test_tool_providers.py b/api/tests/unit_tests/controllers/console/workspace/test_tool_providers.py index 2a576d1c920..5a578c42603 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_tool_providers.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_tool_providers.py @@ -7,7 +7,7 @@ import importlib from contextlib import ExitStack, contextmanager from inspect import unwrap from types import ModuleType, SimpleNamespace -from unittest.mock import MagicMock, patch +from unittest.mock import ANY, MagicMock, patch import pytest from flask import Flask @@ -186,6 +186,7 @@ def test_builtin_provider_credentials_get(app: Flask, controller_module, monkeyp service_mock.assert_called_once_with( tenant_id="tenant-cred", provider_name="demo", + session=ANY, user=user, include_credential_ids=None, ) @@ -210,6 +211,7 @@ def test_builtin_provider_credentials_get_reads_repeated_include_ids( service_mock.assert_called_once_with( tenant_id="tenant-cred", provider_name="demo", + session=ANY, user=user, include_credential_ids=["cred-1", "cred-2"], ) diff --git a/api/tests/unit_tests/controllers/inner_api/app/test_dsl.py b/api/tests/unit_tests/controllers/inner_api/app/test_dsl.py index 71381e6a2b4..ad84eed1f5e 100644 --- a/api/tests/unit_tests/controllers/inner_api/app/test_dsl.py +++ b/api/tests/unit_tests/controllers/inner_api/app/test_dsl.py @@ -6,7 +6,7 @@ in test_auth_wraps.py; handler tests use inspect.unwrap() to bypass them. """ import inspect -from unittest.mock import MagicMock, patch +from unittest.mock import ANY, MagicMock, patch import pytest from flask import Flask @@ -19,7 +19,7 @@ from controllers.inner_api.app.dsl import ( _get_active_account, ) from models.account import AccountStatus -from services.app_dsl_service import ImportStatus +from services.app_dsl_service import Import, ImportStatus class TestInnerAppDSLImportPayload: @@ -117,9 +117,7 @@ class TestEnterpriseAppDSLImport: mock_dsl_cls.return_value = self._mock_dsl yield - def _make_import_result(self, status: ImportStatus, **kwargs) -> "Import": - from services.app_dsl_service import Import - + def _make_import_result(self, status: ImportStatus, **kwargs) -> Import: result = Import( id="import-id", status=status, @@ -224,7 +222,7 @@ class TestEnterpriseAppDSLExport: body, status_code = result assert status_code == 200 assert body["data"] == "version: 0.6.0\nkind: app\n" - mock_dsl_cls.export_dsl.assert_called_once_with(app_model=mock_app, include_secret=False) + mock_dsl_cls.export_dsl.assert_called_once_with(app_model=mock_app, session=ANY, include_secret=False) @patch("controllers.inner_api.app.dsl.AppDslService") @patch("controllers.inner_api.app.dsl.db") @@ -239,7 +237,7 @@ class TestEnterpriseAppDSLExport: body, status_code = result assert status_code == 200 - mock_dsl_cls.export_dsl.assert_called_once_with(app_model=mock_app, include_secret=True) + mock_dsl_cls.export_dsl.assert_called_once_with(app_model=mock_app, session=ANY, include_secret=True) @patch("controllers.inner_api.app.dsl.db") def test_export_app_not_found_returns_404(self, mock_db, api_instance, app: Flask): diff --git a/api/tests/unit_tests/controllers/inner_api/plugin/test_agent_drive.py b/api/tests/unit_tests/controllers/inner_api/plugin/test_agent_drive.py index 8c38564b3d0..8289a575050 100644 --- a/api/tests/unit_tests/controllers/inner_api/plugin/test_agent_drive.py +++ b/api/tests/unit_tests/controllers/inner_api/plugin/test_agent_drive.py @@ -9,7 +9,7 @@ from __future__ import annotations import inspect from types import SimpleNamespace -from unittest.mock import patch +from unittest.mock import ANY, patch import pytest from flask import Flask @@ -33,7 +33,7 @@ def test_manifest_parses_query_and_returns_items(): result = raw(AgentDriveManifestApi(), "agent-agent-1") assert result == {"items": [{"key": "docs/a.txt"}]} svc.return_value.manifest.assert_called_once_with( - tenant_id="tenant-1", agent_id="agent-1", prefix="docs/", include_download_url=True + tenant_id="tenant-1", agent_id="agent-1", prefix="docs/", include_download_url=True, session=ANY ) @@ -85,7 +85,11 @@ def test_skills_requires_tenant_id_and_returns_items(): } ] } - assert svc.return_value.list_skills.call_args.kwargs == {"tenant_id": "tenant-1", "agent_id": "agent-1"} + assert svc.return_value.list_skills.call_args.kwargs == { + "tenant_id": "tenant-1", + "agent_id": "agent-1", + "session": ANY, + } def test_commit_parses_body_and_returns_items(): diff --git a/api/tests/unit_tests/controllers/inner_api/workspace/test_workspace.py b/api/tests/unit_tests/controllers/inner_api/workspace/test_workspace.py index a6626adc420..bda25bb2fa8 100644 --- a/api/tests/unit_tests/controllers/inner_api/workspace/test_workspace.py +++ b/api/tests/unit_tests/controllers/inner_api/workspace/test_workspace.py @@ -117,7 +117,7 @@ class TestEnterpriseWorkspace: assert result["tenant"]["name"] == "My Workspace" mock_tenant_svc.create_tenant.assert_called_once_with("My Workspace", is_from_dashboard=True, session=ANY) mock_tenant_svc.create_tenant_member.assert_called_once_with( - mock_tenant, mock_account, mock_db.session, role="owner" + mock_tenant, mock_account, mock_db.session(), role="owner" ) mock_event.send.assert_called_once_with(mock_tenant) diff --git a/api/tests/unit_tests/controllers/openapi/test_workspaces_members.py b/api/tests/unit_tests/controllers/openapi/test_workspaces_members.py index b78473fadda..86d26420253 100644 --- a/api/tests/unit_tests/controllers/openapi/test_workspaces_members.py +++ b/api/tests/unit_tests/controllers/openapi/test_workspaces_members.py @@ -151,8 +151,8 @@ def _tenant_service(**overrides) -> SimpleNamespace: "get_tenant_members": Mock(return_value=[]), "remove_member_from_tenant": Mock(), "update_member_role": Mock(), - "get_tenant_by_id": lambda session, tenant_id: session.get(None, tenant_id), - "find_workspace_for_account": lambda session, account_id, workspace_id: session.execute(None).first(), + "get_tenant_by_id": lambda tenant_id, *, session: session.get(None, tenant_id), + "find_workspace_for_account": lambda account_id, workspace_id, *, session: session.execute(None).first(), } methods.update(overrides) return SimpleNamespace(**methods) @@ -162,12 +162,18 @@ def _account_service(**overrides) -> SimpleNamespace: """AccountService double; ``get_account_by_id`` delegates to the injected session (see :func:`_tenant_service`).""" methods: dict = { - "get_account_by_id": lambda session, account_id: session.get(None, account_id), + "get_account_by_id": lambda account_id, *, session: session.get(None, account_id), } methods.update(overrides) return SimpleNamespace(**methods) +def _db_mock() -> MagicMock: + mock_db = MagicMock() + mock_db.session.return_value = mock_db.session + return mock_db + + # --------------------------------------------------------------------------- # Route registration # --------------------------------------------------------------------------- @@ -272,7 +278,7 @@ def test_switch_returns_workspace_detail_with_current_true( acct_id = uuid.uuid4() api = WorkspaceSwitchApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.return_value = _account(account_id=str(acct_id)) membership = SimpleNamespace(role=TenantAccountRole.OWNER, current=True) mock_db.session.execute.return_value.first.return_value = (_tenant(ws_id), membership) @@ -304,7 +310,7 @@ def test_switch_404s_when_service_raises_account_not_link_tenant( acct_id = uuid.uuid4() api = WorkspaceSwitchApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.return_value = _account(account_id=str(acct_id)) monkeypatch.setattr( @@ -339,7 +345,7 @@ def test_members_list_returns_normalized_rows(app: Flask, bypass_pipeline, monke role=TenantAccountRole.ADMIN, ) - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.return_value = _tenant(ws_id) monkeypatch.setattr( @@ -381,7 +387,7 @@ def test_members_list_paginates_with_query_params(app: Flask, bypass_pipeline, m for i in range(5) ] - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.return_value = _tenant(ws_id) monkeypatch.setattr( @@ -409,7 +415,7 @@ def test_members_list_rejects_unknown_query_param(app: Flask, bypass_pipeline, m acct_id = uuid.uuid4() api = WorkspaceMembersApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.return_value = _tenant(ws_id) monkeypatch.setattr(sys.modules["controllers.openapi.workspaces"], "db", mock_db) @@ -433,7 +439,7 @@ def test_invite_happy_path_returns_invite_url_and_member_id( invited = _account(account_id="new-1", email="new@example.com") - mock_db = MagicMock() + mock_db = _db_mock() # session.get is called twice: once for inviter Account, once for Tenant mock_db.session.get.side_effect = [_account(account_id=str(acct_id)), _tenant(ws_id)] @@ -514,7 +520,7 @@ def test_invite_blocked_by_saas_members_cap(app: Flask, bypass_pipeline, monkeyp acct_id = uuid.uuid4() api = WorkspaceMembersApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.side_effect = [_account(account_id=str(acct_id)), _tenant(ws_id)] invite_mock = Mock() @@ -552,7 +558,7 @@ def test_invite_blocked_by_ee_workspace_members_license(app: Flask, bypass_pipel acct_id = uuid.uuid4() api = WorkspaceMembersApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.side_effect = [_account(account_id=str(acct_id)), _tenant(ws_id)] invite_mock = Mock() @@ -592,7 +598,7 @@ def test_invite_ce_passes_when_both_caps_disabled(app: Flask, bypass_pipeline, m api = WorkspaceMembersApi() invited = _account(account_id="new-1", email="new@example.com") - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.side_effect = [_account(account_id=str(acct_id)), _tenant(ws_id)] monkeypatch.setattr( @@ -625,7 +631,7 @@ def test_invite_400_when_already_in_tenant(app: Flask, bypass_pipeline, monkeypa acct_id = uuid.uuid4() api = WorkspaceMembersApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.side_effect = [_account(account_id=str(acct_id)), _tenant(ws_id)] monkeypatch.setattr( @@ -656,7 +662,7 @@ def test_delete_member_happy_path(app: Flask, bypass_pipeline, monkeypatch: pyte acct_id = uuid.uuid4() api = WorkspaceMemberApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.side_effect = [ _account(account_id=str(acct_id)), # operator _tenant(ws_id), # tenant @@ -698,7 +704,7 @@ def test_delete_member_exception_mapping(app: Flask, bypass_pipeline, monkeypatc acct_id = uuid.uuid4() api = WorkspaceMemberApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.side_effect = [ _account(account_id=str(acct_id)), _tenant(ws_id), @@ -731,7 +737,7 @@ def test_delete_member_404_when_member_missing(app: Flask, bypass_pipeline, monk acct_id = uuid.uuid4() api = WorkspaceMemberApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.side_effect = [ _account(account_id=str(acct_id)), _tenant(ws_id), @@ -763,7 +769,7 @@ def test_update_role_happy_path(app: Flask, bypass_pipeline, monkeypatch: pytest acct_id = uuid.uuid4() api = WorkspaceMemberApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.side_effect = [ _account(account_id=str(acct_id)), _tenant(ws_id), @@ -809,7 +815,7 @@ def test_update_role_exception_mapping(app: Flask, bypass_pipeline, monkeypatch, acct_id = uuid.uuid4() api = WorkspaceMemberApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.side_effect = [ _account(account_id=str(acct_id)), _tenant(ws_id), @@ -851,7 +857,7 @@ def test_load_tenant_rejects_archived_workspace(app: Flask, bypass_pipeline, mon api = WorkspaceMembersApi() archived = SimpleNamespace(id=ws_id, name="WS", status="archive", created_at=datetime(2026, 5, 18, tzinfo=UTC)) - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.return_value = archived monkeypatch.setattr( @@ -878,7 +884,7 @@ def test_invite_400_when_register_error(app: Flask, bypass_pipeline, monkeypatch acct_id = uuid.uuid4() api = WorkspaceMembersApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.side_effect = [_account(account_id=str(acct_id)), _tenant(ws_id)] monkeypatch.setattr( diff --git a/api/tests/unit_tests/controllers/service_api/app/test_annotation.py b/api/tests/unit_tests/controllers/service_api/app/test_annotation.py index 810101fb0a5..1ff925cba7e 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_annotation.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_annotation.py @@ -15,7 +15,7 @@ Note: API endpoint tests for annotation controllers are complex due to: import uuid from inspect import unwrap from types import SimpleNamespace -from unittest.mock import Mock +from unittest.mock import ANY, Mock import pytest from flask import Flask @@ -264,7 +264,7 @@ class TestAnnotationListApi: assert response["page"] == 1 assert response["limit"] == 20 - get_mock.assert_called_once_with("app", 1, 20, "") + get_mock.assert_called_once_with("app", 1, 20, "", session=ANY) def test_get_accepts_valid_numeric_strings(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: annotation = SimpleNamespace(id="a1", question="q", content="a", created_at=0) @@ -281,7 +281,7 @@ class TestAnnotationListApi: assert response["total"] == 1 assert response["page"] == 2 assert response["limit"] == 5 - get_mock.assert_called_once_with("app", 2, 5, "refund") + get_mock.assert_called_once_with("app", 2, 5, "refund", session=ANY) @pytest.mark.parametrize("query_string", ["page=abc&limit=5", "page=1&limit=abc", "page=&limit=5", "limit=0"]) def test_get_rejects_invalid_explicit_pagination_value( diff --git a/api/tests/unit_tests/controllers/service_api/app/test_app.py b/api/tests/unit_tests/controllers/service_api/app/test_app.py index 30979a26980..04e9220ad55 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_app.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_app.py @@ -3,7 +3,7 @@ Unit tests for Service API App controllers """ import uuid -from unittest.mock import Mock, patch +from unittest.mock import ANY, Mock, patch import pytest from flask import Flask @@ -368,7 +368,7 @@ class TestAppMetaApi: response = api.get() # Assert - mock_service_instance.get_app_meta.assert_called_once_with(mock_app_model) + mock_service_instance.get_app_meta.assert_called_once_with(mock_app_model, session=ANY) assert response == {"tool_icons": {}, "AgentIcons": {}} diff --git a/api/tests/unit_tests/controllers/service_api/app/test_completion.py b/api/tests/unit_tests/controllers/service_api/app/test_completion.py index 9f2a2edeff8..65652594294 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_completion.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_completion.py @@ -252,7 +252,12 @@ class TestAppGenerateService: mock_generate.return_value = expected result = AppGenerateService.generate( - app_model=Mock(spec=App), user=Mock(spec=EndUser), args={"query": "Hi"}, invoke_from=Mock(), streaming=False + app_model=Mock(spec=App), + user=Mock(spec=EndUser), + args={"query": "Hi"}, + invoke_from=Mock(), + session=Mock(), + streaming=False, ) assert result == expected @@ -264,7 +269,12 @@ class TestAppGenerateService: with pytest.raises(services.errors.conversation.ConversationNotExistsError): AppGenerateService.generate( - app_model=Mock(spec=App), user=Mock(spec=EndUser), args={}, invoke_from=Mock(), streaming=False + app_model=Mock(spec=App), + user=Mock(spec=EndUser), + args={}, + invoke_from=Mock(), + session=Mock(), + streaming=False, ) @patch.object(AppGenerateService, "generate") @@ -274,7 +284,12 @@ class TestAppGenerateService: with pytest.raises(QuotaExceededError): AppGenerateService.generate( - app_model=Mock(spec=App), user=Mock(spec=EndUser), args={}, invoke_from=Mock(), streaming=False + app_model=Mock(spec=App), + user=Mock(spec=EndUser), + args={}, + invoke_from=Mock(), + session=Mock(), + streaming=False, ) @patch.object(AppGenerateService, "generate") @@ -284,7 +299,12 @@ class TestAppGenerateService: with pytest.raises(InvokeError): AppGenerateService.generate( - app_model=Mock(spec=App), user=Mock(spec=EndUser), args={}, invoke_from=Mock(), streaming=False + app_model=Mock(spec=App), + user=Mock(spec=EndUser), + args={}, + invoke_from=Mock(), + session=Mock(), + streaming=False, ) diff --git a/api/tests/unit_tests/controllers/service_api/app/test_conversation.py b/api/tests/unit_tests/controllers/service_api/app/test_conversation.py index 97873c631ae..3197812bc31 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_conversation.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_conversation.py @@ -475,6 +475,7 @@ class TestConversationService: user=Mock(spec=EndUser), name="New Name", auto_generate=False, + session=Mock(), ) assert result.name == "New Name" diff --git a/api/tests/unit_tests/controllers/service_api/app/test_message.py b/api/tests/unit_tests/controllers/service_api/app/test_message.py index d8d5c61bcb3..0400abe0e65 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_message.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_message.py @@ -266,6 +266,7 @@ class TestMessageService: conversation_id=str(uuid.uuid4()), first_id=None, limit=20, + session=Mock(), ) assert hasattr(result, "data") @@ -281,7 +282,12 @@ class TestMessageService: with pytest.raises(services.errors.conversation.ConversationNotExistsError): MessageService.pagination_by_first_id( - app_model=Mock(spec=App), user=Mock(spec=EndUser), conversation_id="invalid_id", first_id=None, limit=20 + app_model=Mock(spec=App), + user=Mock(spec=EndUser), + conversation_id="invalid_id", + first_id=None, + limit=20, + session=Mock(), ) @patch.object(MessageService, "pagination_by_first_id") @@ -296,6 +302,7 @@ class TestMessageService: conversation_id=str(uuid.uuid4()), first_id="invalid_first_id", limit=20, + session=Mock(), ) @patch.object(MessageService, "create_feedback") @@ -309,6 +316,7 @@ class TestMessageService: user=Mock(spec=EndUser), rating=FeedbackRating.LIKE, content="Great response!", + session=Mock(), ) mock_create_feedback.assert_called_once() @@ -325,6 +333,7 @@ class TestMessageService: user=Mock(spec=EndUser), rating=FeedbackRating.LIKE, content=None, + session=Mock(), ) @patch.object(MessageService, "get_all_messages_feedbacks") @@ -336,7 +345,7 @@ class TestMessageService: ] mock_get_feedbacks.return_value = mock_feedbacks - result = MessageService.get_all_messages_feedbacks(app_model=Mock(spec=App), page=1, limit=20) + result = MessageService.get_all_messages_feedbacks(app_model=Mock(spec=App), page=1, limit=20, session=Mock()) assert len(result) == 2 assert result[0]["rating"] == "like" @@ -348,7 +357,11 @@ class TestMessageService: mock_get_questions.return_value = mock_questions result = MessageService.get_suggested_questions_after_answer( - app_model=Mock(spec=App), user=Mock(spec=EndUser), message_id=str(uuid.uuid4()), invoke_from=Mock() + app_model=Mock(spec=App), + user=Mock(spec=EndUser), + message_id=str(uuid.uuid4()), + invoke_from=Mock(), + session=Mock(), ) assert len(result) == 3 @@ -361,7 +374,11 @@ class TestMessageService: with pytest.raises(SuggestedQuestionsAfterAnswerDisabledError): MessageService.get_suggested_questions_after_answer( - app_model=Mock(spec=App), user=Mock(spec=EndUser), message_id=str(uuid.uuid4()), invoke_from=Mock() + app_model=Mock(spec=App), + user=Mock(spec=EndUser), + message_id=str(uuid.uuid4()), + invoke_from=Mock(), + session=Mock(), ) @patch.object(MessageService, "get_suggested_questions_after_answer") @@ -371,7 +388,11 @@ class TestMessageService: with pytest.raises(MessageNotExistsError): MessageService.get_suggested_questions_after_answer( - app_model=Mock(spec=App), user=Mock(spec=EndUser), message_id="invalid_message_id", invoke_from=Mock() + app_model=Mock(spec=App), + user=Mock(spec=EndUser), + message_id="invalid_message_id", + invoke_from=Mock(), + session=Mock(), ) diff --git a/api/tests/unit_tests/controllers/service_api/app/test_workflow.py b/api/tests/unit_tests/controllers/service_api/app/test_workflow.py index 3cabfe43ddc..2115bb85526 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_workflow.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_workflow.py @@ -18,7 +18,7 @@ import uuid from datetime import UTC, datetime from inspect import unwrap from types import SimpleNamespace -from unittest.mock import Mock, patch +from unittest.mock import MagicMock, Mock, patch import pytest from flask import Flask @@ -167,26 +167,6 @@ class TestWorkflowLogQuery: query_max_limit = WorkflowLogQuery(limit=100) assert query_max_limit.limit == 100 - def test_query_rejects_page_below_minimum(self): - """Test query rejects page < 1.""" - with pytest.raises(ValueError): - WorkflowLogQuery(page=0) - - def test_query_rejects_page_above_maximum(self): - """Test query rejects page > 99999.""" - with pytest.raises(ValueError): - WorkflowLogQuery(page=100000) - - def test_query_rejects_limit_below_minimum(self): - """Test query rejects limit < 1.""" - with pytest.raises(ValueError): - WorkflowLogQuery(limit=0) - - def test_query_rejects_limit_above_maximum(self): - """Test query rejects limit > 100.""" - with pytest.raises(ValueError): - WorkflowLogQuery(limit=101) - def test_query_with_keyword_search(self): """Test query with keyword filter.""" query = WorkflowLogQuery(keyword="workflow execution") @@ -263,7 +243,7 @@ class TestAppGenerateServiceWorkflow: """Test AppGenerateService workflow integration.""" @patch.object(AppGenerateService, "generate") - def test_generate_accepts_workflow_args(self, mock_generate): + def test_generate_accepts_workflow_args(self, mock_generate: MagicMock): """Test generate accepts workflow-specific args.""" mock_generate.return_value = {"result": "success"} @@ -272,6 +252,7 @@ class TestAppGenerateServiceWorkflow: user=Mock(), args={"inputs": {"key": "value"}, "workflow_id": "workflow_123"}, invoke_from=Mock(), + session=MagicMock(), streaming=False, ) @@ -279,7 +260,7 @@ class TestAppGenerateServiceWorkflow: mock_generate.assert_called_once() @patch.object(AppGenerateService, "generate") - def test_generate_raises_workflow_not_found_error(self, mock_generate): + def test_generate_raises_workflow_not_found_error(self, mock_generate: MagicMock): """Test generate raises WorkflowNotFoundError.""" mock_generate.side_effect = WorkflowNotFoundError("Workflow not found") @@ -289,11 +270,12 @@ class TestAppGenerateServiceWorkflow: user=Mock(), args={"workflow_id": "invalid_id"}, invoke_from=Mock(), + session=MagicMock(), streaming=False, ) @patch.object(AppGenerateService, "generate") - def test_generate_raises_is_draft_workflow_error(self, mock_generate): + def test_generate_raises_is_draft_workflow_error(self, mock_generate: MagicMock): """Test generate raises IsDraftWorkflowError.""" mock_generate.side_effect = IsDraftWorkflowError("Workflow is draft") @@ -303,11 +285,12 @@ class TestAppGenerateServiceWorkflow: user=Mock(), args={"workflow_id": "draft_workflow"}, invoke_from=Mock(), + session=MagicMock(), streaming=False, ) @patch.object(AppGenerateService, "generate") - def test_generate_supports_streaming_mode(self, mock_generate): + def test_generate_supports_streaming_mode(self, mock_generate: MagicMock): """Test generate supports streaming response mode.""" mock_stream = Mock() mock_generate.return_value = mock_stream @@ -317,6 +300,7 @@ class TestAppGenerateServiceWorkflow: user=Mock(), args={"inputs": {}, "response_mode": "streaming"}, invoke_from=Mock(), + session=MagicMock(), streaming=True, ) diff --git a/api/tests/unit_tests/controllers/service_api/dataset/rag_pipeline/test_rag_pipeline_workflow.py b/api/tests/unit_tests/controllers/service_api/dataset/rag_pipeline/test_rag_pipeline_workflow.py index 43cc2450db5..406037e268d 100644 --- a/api/tests/unit_tests/controllers/service_api/dataset/rag_pipeline/test_rag_pipeline_workflow.py +++ b/api/tests/unit_tests/controllers/service_api/dataset/rag_pipeline/test_rag_pipeline_workflow.py @@ -549,16 +549,14 @@ class TestPipelineRunApiPost: new_callable=lambda: Mock(spec=Account), ) @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.RagPipelineService") - @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.db") @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.service_api_ns") - def test_post_success_streaming( - self, mock_ns, mock_db, mock_svc_cls, mock_current_user, mock_gen_svc, mock_helper, app - ): + def test_post_success_streaming(self, mock_ns, mock_svc_cls, mock_current_user, mock_gen_svc, mock_helper, app): """Test successful pipeline run with streaming response.""" tenant_id = str(uuid.uuid4()) dataset_id = str(uuid.uuid4()) - mock_db.session.scalar.return_value = Mock() + session = Mock() + session.scalar.return_value = Mock() mock_ns.payload = { "inputs": {"key": "val"}, @@ -579,27 +577,33 @@ class TestPipelineRunApiPost: with app.test_request_context("/datasets/test/pipeline/run", method="POST"): api = PipelineRunApi() - response = api.post(tenant_id=tenant_id, dataset_id=dataset_id) + response = api.post.__wrapped__(api, session, tenant_id=tenant_id, dataset_id=dataset_id) assert response == {"result": "ok"} + mock_svc_cls.assert_called_once_with(session) mock_gen_svc.generate.assert_called_once() - @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.db") - def test_post_not_found(self, mock_db, app: Flask): + def test_post_not_found(self, app: Flask): """Test NotFound when dataset check fails.""" - mock_db.session.scalar.return_value = None + session = Mock() + session.scalar.return_value = None with app.test_request_context("/datasets/test/pipeline/run", method="POST"): api = PipelineRunApi() with pytest.raises(NotFound): - api.post(tenant_id=str(uuid.uuid4()), dataset_id=str(uuid.uuid4())) + api.post.__wrapped__( + api, + session, + tenant_id=str(uuid.uuid4()), + dataset_id=str(uuid.uuid4()), + ) @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.current_user", new="not_account") - @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.db") @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.service_api_ns") - def test_post_forbidden_non_account_user(self, mock_ns, mock_db, app: Flask): + def test_post_forbidden_non_account_user(self, mock_ns, app: Flask): """Test Forbidden when current_user is not an Account.""" - mock_db.session.scalar.return_value = Mock() + session = Mock() + session.scalar.return_value = Mock() mock_ns.payload = { "inputs": {}, "datasource_type": "online_document", @@ -612,7 +616,12 @@ class TestPipelineRunApiPost: with app.test_request_context("/datasets/test/pipeline/run", method="POST"): api = PipelineRunApi() with pytest.raises(Forbidden): - api.post(tenant_id=str(uuid.uuid4()), dataset_id=str(uuid.uuid4())) + api.post.__wrapped__( + api, + session, + tenant_id=str(uuid.uuid4()), + dataset_id=str(uuid.uuid4()), + ) class TestFileUploadApiPost: diff --git a/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_segment.py b/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_segment.py index a95baf1b482..0b1ca8741a9 100644 --- a/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_segment.py +++ b/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_segment.py @@ -1193,7 +1193,7 @@ class TestDatasetSegmentApiDelete: # Assert assert response == ("", 204) - mock_seg_svc.delete_segment.assert_called_once_with(mock_segment, mock_doc, mock_dataset, mock_db.session) + mock_seg_svc.delete_segment.assert_called_once_with(mock_segment, mock_doc, mock_dataset, mock_db.session()) @patch("controllers.service_api.dataset.segment.SegmentService") @patch("controllers.service_api.dataset.segment.DocumentService") diff --git a/api/tests/unit_tests/controllers/service_api/dataset/test_document.py b/api/tests/unit_tests/controllers/service_api/dataset/test_document.py index dd2caf4f3fc..e83724f955f 100644 --- a/api/tests/unit_tests/controllers/service_api/dataset/test_document.py +++ b/api/tests/unit_tests/controllers/service_api/dataset/test_document.py @@ -783,7 +783,7 @@ class TestDocumentApiDelete: # Assert assert response == ("", 204) - mock_doc_svc.delete_document.assert_called_once_with(mock_document, mock_db.session) + mock_doc_svc.delete_document.assert_called_once_with(mock_document, mock_db.session()) @patch("controllers.service_api.dataset.document.DocumentService") @patch("controllers.service_api.dataset.document.db") diff --git a/api/tests/unit_tests/controllers/service_api/dataset/test_metadata.py b/api/tests/unit_tests/controllers/service_api/dataset/test_metadata.py index b77c783ae16..dd1322a6344 100644 --- a/api/tests/unit_tests/controllers/service_api/dataset/test_metadata.py +++ b/api/tests/unit_tests/controllers/service_api/dataset/test_metadata.py @@ -408,7 +408,7 @@ class TestDatasetMetadataBuiltInFieldAction: assert status == 200 assert response["result"] == "success" - mock_meta_svc.enable_built_in_field.assert_called_once_with(ANY, mock_dataset) + mock_meta_svc.enable_built_in_field.assert_called_once_with(mock_dataset, session=ANY) @patch("controllers.service_api.dataset.metadata.MetadataService") @patch("controllers.service_api.dataset.metadata.DatasetService") @@ -439,7 +439,7 @@ class TestDatasetMetadataBuiltInFieldAction: ) assert status == 200 - mock_meta_svc.disable_built_in_field.assert_called_once_with(ANY, mock_dataset) + mock_meta_svc.disable_built_in_field.assert_called_once_with(mock_dataset, session=ANY) @patch("controllers.service_api.dataset.metadata.DatasetService") def test_action_dataset_not_found( diff --git a/api/tests/unit_tests/controllers/web/test_app.py b/api/tests/unit_tests/controllers/web/test_app.py index 542ee111e1b..73f308dc749 100644 --- a/api/tests/unit_tests/controllers/web/test_app.py +++ b/api/tests/unit_tests/controllers/web/test_app.py @@ -3,7 +3,7 @@ from __future__ import annotations from types import SimpleNamespace -from unittest.mock import MagicMock, patch +from unittest.mock import ANY, MagicMock, patch import pytest from flask import Flask @@ -148,7 +148,7 @@ class TestAppAccessMode: with app.test_request_context("/webapp/access-mode?appCode=code1"): result = AppAccessMode().get() - mock_resolve.assert_called_once_with("code1") + mock_resolve.assert_called_once_with("code1", session=ANY) mock_access.assert_called_once_with("resolved-id") assert result == {"accessMode": "external"} diff --git a/api/tests/unit_tests/controllers/web/test_message_list.py b/api/tests/unit_tests/controllers/web/test_message_list.py index 2bb425cdba2..b5d74df65ef 100644 --- a/api/tests/unit_tests/controllers/web/test_message_list.py +++ b/api/tests/unit_tests/controllers/web/test_message_list.py @@ -6,7 +6,7 @@ import builtins import uuid from datetime import datetime from types import ModuleType, SimpleNamespace -from unittest.mock import patch +from unittest.mock import ANY, patch from uuid import uuid4 import pytest @@ -158,7 +158,7 @@ def test_message_list_mapping(app: Flask) -> None: ): response = MessageListApi().get(app_model, end_user) - mock_page.assert_called_once_with(app_model, end_user, conversation_id, None, 20) + mock_page.assert_called_once_with(app_model, end_user, conversation_id, None, 20, session=ANY) assert response["limit"] == 20 assert response["has_more"] is False assert len(response["data"]) == 1 diff --git a/api/tests/unit_tests/controllers/web/test_web_login.py b/api/tests/unit_tests/controllers/web/test_web_login.py index 984be6ddba9..a91d4253aa8 100644 --- a/api/tests/unit_tests/controllers/web/test_web_login.py +++ b/api/tests/unit_tests/controllers/web/test_web_login.py @@ -1,7 +1,7 @@ import base64 import logging from types import SimpleNamespace -from unittest.mock import MagicMock, patch +from unittest.mock import ANY, MagicMock, patch import pytest from flask import Flask @@ -66,7 +66,7 @@ class TestEmailCodeLoginSendEmailApi: response = EmailCodeLoginSendEmailApi().post() assert response == {"result": "success", "data": "token-123"} - mock_get_user.assert_called_once_with("User@Example.com") + mock_get_user.assert_called_once_with("User@Example.com", ANY) mock_send_email.assert_called_once_with(account=mock_account, language="en-US") @@ -96,7 +96,7 @@ class TestEmailCodeLoginApi: response = EmailCodeLoginApi().post() assert response == {"result": "success", "data": {"access_token": "new-access-token"}} - mock_get_user.assert_called_once_with("User@Example.com") + mock_get_user.assert_called_once_with("User@Example.com", ANY) mock_revoke_token.assert_called_once_with("token-123") mock_login.assert_called_once() mock_reset_login_rate.assert_called_once_with("user@example.com") diff --git a/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_generator.py b/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_generator.py index cda5178e30c..41e14af72de 100644 --- a/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_generator.py +++ b/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_generator.py @@ -147,7 +147,7 @@ class TestAdvancedChatAppGeneratorInternals: ) monkeypatch.setattr( "core.app.apps.advanced_chat.app_generator.db", - SimpleNamespace(engine=object(), session=SimpleNamespace(close=lambda: None)), + SimpleNamespace(engine=object(), session=lambda: SimpleNamespace(close=lambda: None)), ) monkeypatch.setattr( "core.app.apps.advanced_chat.app_generator.sessionmaker", lambda **kwargs: SimpleNamespace() diff --git a/api/tests/unit_tests/core/app/apps/test_advanced_chat_app_generator.py b/api/tests/unit_tests/core/app/apps/test_advanced_chat_app_generator.py index 9b89b108207..ef12f0be965 100644 --- a/api/tests/unit_tests/core/app/apps/test_advanced_chat_app_generator.py +++ b/api/tests/unit_tests/core/app/apps/test_advanced_chat_app_generator.py @@ -138,7 +138,7 @@ def test_generate_falls_back_to_new_conversation_when_conversation_missing(monke ) monkeypatch.setattr( "core.app.apps.advanced_chat.app_generator.db", - SimpleNamespace(engine=object()), + SimpleNamespace(engine=object(), session=lambda: MagicMock()), ) trace_manager = object.__new__(TraceQueueManager) monkeypatch.setattr( diff --git a/api/tests/unit_tests/core/app/test_llm_quota.py b/api/tests/unit_tests/core/app/test_llm_quota.py index 13bdf765358..ec6ac134443 100644 --- a/api/tests/unit_tests/core/app/test_llm_quota.py +++ b/api/tests/unit_tests/core/app/test_llm_quota.py @@ -1,7 +1,7 @@ from collections.abc import Generator from contextlib import contextmanager from types import SimpleNamespace -from unittest.mock import MagicMock, patch +from unittest.mock import ANY, MagicMock, patch import pytest from sqlalchemy import create_engine, select @@ -28,8 +28,19 @@ from models.provider import Provider, ProviderType @contextmanager def _patched_credit_pool_session_factory(engine: Engine) -> Generator[None, None, None]: session_maker = sessionmaker(bind=engine, expire_on_commit=False) - with patch("services.credit_pool_service.session_factory.get_session_maker", return_value=session_maker): - yield + sessions = [] + + def _session(): + session = session_maker() + sessions.append(session) + return session + + with patch("core.app.llm.quota.db", SimpleNamespace(session=_session)): + try: + yield + finally: + for session in sessions: + session.close() def test_ensure_llm_quota_available_for_model_raises_when_system_model_is_exhausted() -> None: @@ -122,6 +133,7 @@ def test_deduct_llm_quota_for_model_uses_identity_based_trial_billing() -> None: mock_deduct_credits.assert_called_once_with( tenant_id="tenant-id", credits_required=42, + session=ANY, ) @@ -241,6 +253,7 @@ def test_deduct_llm_quota_for_model_uses_credit_configuration() -> None: mock_deduct_credits.assert_called_once_with( tenant_id="tenant-id", credits_required=9, + session=ANY, ) @@ -276,6 +289,7 @@ def test_deduct_llm_quota_for_model_uses_single_charge_for_times_quota() -> None mock_deduct_credits.assert_called_once_with( tenant_id="tenant-id", credits_required=1, + session=ANY, ) @@ -313,6 +327,7 @@ def test_deduct_llm_quota_for_model_uses_paid_billing_pool() -> None: tenant_id="tenant-id", credits_required=5, pool_type="paid", + session=ANY, ) diff --git a/api/tests/unit_tests/core/llm_generator/test_llm_generator_missing.py b/api/tests/unit_tests/core/llm_generator/test_llm_generator_missing.py index ddb33f0758f..1c4a6e2db7c 100644 --- a/api/tests/unit_tests/core/llm_generator/test_llm_generator_missing.py +++ b/api/tests/unit_tests/core/llm_generator/test_llm_generator_missing.py @@ -149,12 +149,12 @@ class TestWorkflowServiceInterface: from core.llm_generator.llm_generator import WorkflowServiceInterface class MockService(WorkflowServiceInterface): - def get_draft_workflow(self, app_model, workflow_id=None): - return super().get_draft_workflow(app_model, workflow_id) + def get_draft_workflow(self, app_model, workflow_id=None, *, session): + return super().get_draft_workflow(app_model, workflow_id, session=session) def get_node_last_run(self, app_model, workflow, node_id): return super().get_node_last_run(app_model, workflow, node_id) service = MockService() - service.get_draft_workflow(None) + service.get_draft_workflow(None, session=None) service.get_node_last_run(None, None, "node") diff --git a/api/tests/unit_tests/core/rag/datasource/test_datasource_retrieval.py b/api/tests/unit_tests/core/rag/datasource/test_datasource_retrieval.py index 7c672570bfa..d8452d91e2c 100644 --- a/api/tests/unit_tests/core/rag/datasource/test_datasource_retrieval.py +++ b/api/tests/unit_tests/core/rag/datasource/test_datasource_retrieval.py @@ -227,13 +227,12 @@ class TestRetrievalServiceInternals: @patch("core.rag.datasource.retrieval_service.ExternalDatasetService.fetch_external_knowledge_retrieval") @patch("core.rag.datasource.retrieval_service.MetadataFilteringCondition.model_validate") - @patch("core.rag.datasource.retrieval_service.db.session.scalar") - def test_external_retrieve_with_metadata_conditions(self, mock_scalar, mock_validate, mock_fetch): - mock_scalar.return_value = SimpleNamespace(tenant_id="tenant-1") + def test_external_retrieve_with_metadata_conditions(self, mock_validate, mock_fetch): mock_validate.return_value = "validated-condition" expected_documents = [create_mock_document("external-doc", "external-1", 0.8, provider="external")] mock_fetch.return_value = expected_documents session = MagicMock() + session.scalar.return_value = SimpleNamespace(tenant_id="tenant-1") results = RetrievalService.external_retrieve( session=session, @@ -246,19 +245,19 @@ class TestRetrievalServiceInternals: assert results == expected_documents mock_validate.assert_called_once() mock_fetch.assert_called_once_with( - session, - "tenant-1", - "dataset-1", - "test query", - {"top_k": 3}, + tenant_id="tenant-1", + dataset_id="dataset-1", + query="test query", + external_retrieval_parameters={"top_k": 3}, metadata_condition="validated-condition", + session=session, ) - @patch("core.rag.datasource.retrieval_service.db.session.scalar") - def test_external_retrieve_returns_empty_when_dataset_not_found(self, mock_scalar): - mock_scalar.return_value = None + def test_external_retrieve_returns_empty_when_dataset_not_found(self): + session = MagicMock() + session.scalar.return_value = None - results = RetrievalService.external_retrieve(session=MagicMock(), dataset_id="missing", query="q") + results = RetrievalService.external_retrieve(session=session, dataset_id="missing", query="q") assert results == [] diff --git a/api/tests/unit_tests/core/rag/indexing/processor/test_paragraph_index_processor.py b/api/tests/unit_tests/core/rag/indexing/processor/test_paragraph_index_processor.py index 302ababb48f..f5761b5ba3d 100644 --- a/api/tests/unit_tests/core/rag/indexing/processor/test_paragraph_index_processor.py +++ b/api/tests/unit_tests/core/rag/indexing/processor/test_paragraph_index_processor.py @@ -209,7 +209,7 @@ class TestParagraphIndexProcessor: vector = mock_vector_cls.return_value processor.clean(dataset, ["node-1"], delete_summaries=True) - mock_summary.assert_called_once_with(dataset, ["seg-1"]) + mock_summary.assert_called_once_with(dataset=dataset, segment_ids=["seg-1"]) vector.delete_by_ids.assert_called_once_with(["node-1"]) def test_clean_economy_deletes_summaries_and_keywords( @@ -225,7 +225,7 @@ class TestParagraphIndexProcessor: ): processor.clean(dataset, None, delete_summaries=True) - mock_summary.assert_called_once_with(dataset, None) + mock_summary.assert_called_once_with(dataset=dataset, segment_ids=None) mock_keyword_cls.return_value.delete.assert_called_once() def test_clean_deletes_keywords_by_ids(self, processor: ParagraphIndexProcessor, dataset: Mock) -> None: diff --git a/api/tests/unit_tests/core/rag/indexing/processor/test_parent_child_index_processor.py b/api/tests/unit_tests/core/rag/indexing/processor/test_parent_child_index_processor.py index 7d339a7701f..672764e5336 100644 --- a/api/tests/unit_tests/core/rag/indexing/processor/test_parent_child_index_processor.py +++ b/api/tests/unit_tests/core/rag/indexing/processor/test_parent_child_index_processor.py @@ -278,7 +278,7 @@ class TestParentChildIndexProcessor: ): processor.clean(dataset, ["node-1"], delete_summaries=True, precomputed_child_node_ids=[]) - mock_summary.assert_called_once_with(dataset, ["seg-1"]) + mock_summary.assert_called_once_with(dataset=dataset, segment_ids=["seg-1"]) def test_clean_deletes_all_summaries_when_node_ids_missing( self, processor: ParentChildIndexProcessor, dataset: Mock @@ -291,7 +291,7 @@ class TestParentChildIndexProcessor: ): processor.clean(dataset, None, delete_summaries=True) - mock_summary.assert_called_once_with(dataset, None) + mock_summary.assert_called_once_with(dataset=dataset, segment_ids=None) def test_split_child_nodes_requires_subchunk_segmentation(self, processor: ParentChildIndexProcessor) -> None: rules = Rule(subchunk_segmentation=None) diff --git a/api/tests/unit_tests/core/rag/indexing/processor/test_qa_index_processor.py b/api/tests/unit_tests/core/rag/indexing/processor/test_qa_index_processor.py index 6e5a4fabbb0..5dde1623d2d 100644 --- a/api/tests/unit_tests/core/rag/indexing/processor/test_qa_index_processor.py +++ b/api/tests/unit_tests/core/rag/indexing/processor/test_qa_index_processor.py @@ -243,7 +243,7 @@ class TestQAIndexProcessor: vector = mock_vector_cls.return_value processor.clean(dataset, ["node-1"], delete_summaries=True) - mock_summary.assert_called_once_with(dataset, ["seg-1"]) + mock_summary.assert_called_once_with(dataset=dataset, segment_ids=["seg-1"]) vector.delete_by_ids.assert_called_once_with(["node-1"]) def test_clean_handles_dataset_wide_cleanup(self, processor: QAIndexProcessor, dataset: Mock) -> None: @@ -256,7 +256,7 @@ class TestQAIndexProcessor: vector = mock_vector_cls.return_value processor.clean(dataset, None, delete_summaries=True) - mock_summary.assert_called_once_with(dataset, None) + mock_summary.assert_called_once_with(dataset=dataset, segment_ids=None) vector.delete.assert_called_once() def test_index_adds_documents_and_vectors_for_high_quality( diff --git a/api/tests/unit_tests/events/test_app_event_signals.py b/api/tests/unit_tests/events/test_app_event_signals.py index 29582a50f6d..a6059fadbcf 100644 --- a/api/tests/unit_tests/events/test_app_event_signals.py +++ b/api/tests/unit_tests/events/test_app_event_signals.py @@ -44,7 +44,7 @@ def _make_collector(target: list): @pytest.mark.usefixtures("mock_db", "_mock_deps") class TestAppWasDeletedSignal: - def test_sends_signal(self, app_model): + def test_sends_signal(self, app_model, mock_db): from events.app_event import app_was_deleted from services.app_service import AppService @@ -52,7 +52,7 @@ class TestAppWasDeletedSignal: handler = _make_collector(received) app_was_deleted.connect(handler) try: - AppService().delete_app(app_model) + AppService().delete_app(app_model, session=mock_db.session) finally: app_was_deleted.disconnect(handler) @@ -71,7 +71,7 @@ class TestAppWasDeletedSignal: mock_db.session.delete.side_effect = lambda _: call_order.append("db_delete") try: - AppService().delete_app(app_model) + AppService().delete_app(app_model, session=mock_db.session) finally: app_was_deleted.disconnect(handler) @@ -80,7 +80,7 @@ class TestAppWasDeletedSignal: @pytest.mark.usefixtures("mock_db") class TestAppWasUpdatedSignal: - def test_update_app(self, app_model): + def test_update_app(self, app_model, mock_db): from events.app_event import app_was_updated from services.app_service import AppService @@ -101,13 +101,14 @@ class TestAppWasUpdatedSignal: "use_icon_as_answer_icon": False, "max_active_requests": 0, }, + session=mock_db.session, ) finally: app_was_updated.disconnect(handler) assert received == [app_model] - def test_update_app_name(self, app_model): + def test_update_app_name(self, app_model, mock_db): from events.app_event import app_was_updated from services.app_service import AppService @@ -117,13 +118,13 @@ class TestAppWasUpdatedSignal: with patch("services.app_service.current_user", MagicMock(id="user-1")): try: - AppService().update_app_name(app_model, "New Name") + AppService().update_app_name(app_model, "New Name", session=mock_db.session) finally: app_was_updated.disconnect(handler) assert received == [app_model] - def test_update_app_icon(self, app_model): + def test_update_app_icon(self, app_model, mock_db): from events.app_event import app_was_updated from services.app_service import AppService @@ -133,13 +134,13 @@ class TestAppWasUpdatedSignal: with patch("services.app_service.current_user", MagicMock(id="user-1")): try: - AppService().update_app_icon(app_model, "🎉", "#000") + AppService().update_app_icon(app_model, "🎉", "#000", session=mock_db.session) finally: app_was_updated.disconnect(handler) assert received == [app_model] - def test_update_app_site_status_sends_when_changed(self, app_model): + def test_update_app_site_status_sends_when_changed(self, app_model, mock_db): from events.app_event import app_was_updated from services.app_service import AppService @@ -150,13 +151,13 @@ class TestAppWasUpdatedSignal: with patch("services.app_service.current_user", MagicMock(id="user-1")): try: app_model.enable_site = False - AppService().update_app_site_status(app_model, True) + AppService().update_app_site_status(app_model, True, session=mock_db.session) finally: app_was_updated.disconnect(handler) assert received == [app_model] - def test_update_app_site_status_skips_when_unchanged(self, app_model): + def test_update_app_site_status_skips_when_unchanged(self, app_model, mock_db): from events.app_event import app_was_updated from services.app_service import AppService @@ -166,13 +167,13 @@ class TestAppWasUpdatedSignal: try: app_model.enable_site = True - AppService().update_app_site_status(app_model, True) + AppService().update_app_site_status(app_model, True, session=mock_db.session) finally: app_was_updated.disconnect(handler) assert received == [] - def test_update_app_api_status_sends_when_changed(self, app_model): + def test_update_app_api_status_sends_when_changed(self, app_model, mock_db): from events.app_event import app_was_updated from services.app_service import AppService @@ -183,13 +184,13 @@ class TestAppWasUpdatedSignal: with patch("services.app_service.current_user", MagicMock(id="user-1")): try: app_model.enable_api = False - AppService().update_app_api_status(app_model, True) + AppService().update_app_api_status(app_model, True, session=mock_db.session) finally: app_was_updated.disconnect(handler) assert received == [app_model] - def test_update_app_api_status_skips_when_unchanged(self, app_model): + def test_update_app_api_status_skips_when_unchanged(self, app_model, mock_db): from events.app_event import app_was_updated from services.app_service import AppService @@ -199,7 +200,7 @@ class TestAppWasUpdatedSignal: try: app_model.enable_api = True - AppService().update_app_api_status(app_model, True) + AppService().update_app_api_status(app_model, True, session=mock_db.session) finally: app_was_updated.disconnect(handler) diff --git a/api/tests/unit_tests/events/test_update_provider_when_message_created.py b/api/tests/unit_tests/events/test_update_provider_when_message_created.py index f9ac5d9678e..327c80323b4 100644 --- a/api/tests/unit_tests/events/test_update_provider_when_message_created.py +++ b/api/tests/unit_tests/events/test_update_provider_when_message_created.py @@ -1,7 +1,7 @@ from collections.abc import Generator from contextlib import contextmanager from types import SimpleNamespace -from unittest.mock import patch +from unittest.mock import ANY, patch from uuid import uuid4 import pytest @@ -19,8 +19,19 @@ from models.provider import ProviderType @contextmanager def _patched_credit_pool_session_factory(engine: Engine) -> Generator[None, None, None]: session_maker = sessionmaker(bind=engine, expire_on_commit=False) - with patch("services.credit_pool_service.session_factory.get_session_maker", return_value=session_maker): - yield + sessions = [] + + def _session(): + session = session_maker() + sessions.append(session) + return session + + with patch("events.event_handlers.update_provider_when_message_created.db", SimpleNamespace(session=_session)): + try: + yield + finally: + for session in sessions: + session.close() def test_message_created_trial_credit_accounting_does_not_raise_when_balance_is_insufficient() -> None: @@ -140,5 +151,6 @@ def test_capped_credit_pool_accounting_skips_exhaustion_warning_when_full_amount tenant_id="tenant-id", credits_required=3, pool_type="trial", + session=ANY, ) assert "Credit pool exhausted during message-created accounting" not in caplog.text diff --git a/api/tests/unit_tests/services/agent/test_agent_services.py b/api/tests/unit_tests/services/agent/test_agent_services.py index 6fee58bb29d..f896a761904 100644 --- a/api/tests/unit_tests/services/agent/test_agent_services.py +++ b/api/tests/unit_tests/services/agent/test_agent_services.py @@ -117,7 +117,9 @@ def test_load_workflow_composer_returns_empty_state(monkeypatch: pytest.MonkeyPa monkeypatch.setattr(AgentComposerService, "_get_draft_workflow", lambda **kwargs: SimpleNamespace(id="workflow-1")) monkeypatch.setattr(AgentComposerService, "_get_workflow_binding", lambda **kwargs: None) - result = AgentComposerService.load_workflow_composer(tenant_id="tenant-1", app_id="app-1", node_id="node-1") + result = AgentComposerService.load_workflow_composer( + tenant_id="tenant-1", app_id="app-1", node_id="node-1", session=composer_service.db.session + ) assert result["binding"] is None assert result["save_options"] == ["node_job_only", "save_to_roster"] @@ -155,7 +157,9 @@ def test_load_workflow_composer_serializes_existing_binding(monkeypatch: pytest. lambda **kwargs: {"agent": kwargs["agent"].id, "version": kwargs["version"].id}, ) - result = AgentComposerService.load_workflow_composer(tenant_id="tenant-1", app_id="app-1", node_id="node-1") + result = AgentComposerService.load_workflow_composer( + tenant_id="tenant-1", app_id="app-1", node_id="node-1", session=composer_service.db.session + ) assert result == {"agent": "agent-1", "version": "version-1"} @@ -190,6 +194,7 @@ def test_load_workflow_composer_uses_roster_preview_snapshot(monkeypatch: pytest app_id="app-1", node_id="node-1", snapshot_id="preview-version", + session=composer_service.db.session, ) assert result == {"binding_snapshot_id": "binding-version", "version": "preview-version"} @@ -232,6 +237,7 @@ def test_load_workflow_composer_uses_inline_preview_snapshot(monkeypatch: pytest app_id="app-1", node_id="node-1", snapshot_id="inline-preview-version", + session=composer_service.db.session, ) assert result == {"agent": "inline-agent-1", "version": "inline-preview-version"} @@ -258,6 +264,7 @@ def test_workflow_inline_debug_conversation_seed(monkeypatch: pytest.MonkeyPatch binding=binding, agent=agent, account_id="account-1", + session="session-1", ) assert debug_conversation_id == "debug-conversation-1" @@ -279,6 +286,7 @@ def test_workflow_inline_debug_conversation_seed_skips_non_inline(monkeypatch: p binding=SimpleNamespace(binding_type=WorkflowAgentBindingType.ROSTER_AGENT), agent=SimpleNamespace(id="agent-1", scope=AgentScope.ROSTER), account_id="account-1", + session="session-1", ) is None ) @@ -288,6 +296,7 @@ def test_workflow_inline_debug_conversation_seed_skips_non_inline(monkeypatch: p binding=SimpleNamespace(binding_type=WorkflowAgentBindingType.INLINE_AGENT), agent=SimpleNamespace(id="inline-agent-1", scope=AgentScope.WORKFLOW_ONLY), account_id=None, + session="session-1", ) is None ) @@ -303,6 +312,7 @@ def test_load_workflow_composer_rejects_preview_without_binding(monkeypatch: pyt app_id="app-1", node_id="node-1", snapshot_id="preview-version", + session=composer_service.db.session, ) @@ -361,7 +371,12 @@ def test_save_workflow_composer_dispatches_save_strategy(monkeypatch, strategy, ) result = AgentComposerService.save_workflow_composer( - tenant_id="tenant-1", app_id="app-1", node_id="node-1", account_id="account-1", payload=payload + tenant_id="tenant-1", + app_id="app-1", + node_id="node-1", + account_id="account-1", + payload=payload, + session=composer_service.db.session, ) assert result.pop("validation") == {"warnings": [], "knowledge_retrieval_placeholder": []} @@ -382,7 +397,12 @@ def test_save_workflow_composer_rejects_agent_app_variant(): with pytest.raises(ValueError): AgentComposerService.save_workflow_composer( - tenant_id="tenant-1", app_id="app-1", node_id="node-1", account_id="account-1", payload=payload + tenant_id="tenant-1", + app_id="app-1", + node_id="node-1", + account_id="account-1", + payload=payload, + session=composer_service.db.session, ) @@ -444,7 +464,11 @@ def test_save_agent_app_composer_creates_agent_when_missing(monkeypatch: pytest. ) result = AgentComposerService.save_agent_app_composer( - tenant_id="tenant-1", app_id="app-1", account_id="account-1", payload=payload + tenant_id="tenant-1", + app_id="app-1", + account_id="account-1", + payload=payload, + session=composer_service.db.session, ) assert result.pop("validation") == {"warnings": [], "knowledge_retrieval_placeholder": []} @@ -475,7 +499,9 @@ def test_load_agent_app_composer_exposes_draft_save_only(monkeypatch: pytest.Mon monkeypatch.setattr(AgentComposerService, "_serialize_version", lambda _version: None) monkeypatch.setattr(AgentComposerService, "_serialize_draft", lambda _draft: {"id": "draft-1"}) - result = AgentComposerService.load_agent_app_composer(tenant_id="tenant-1", app_id="app-1") + result = AgentComposerService.load_agent_app_composer( + tenant_id="tenant-1", app_id="app-1", session=composer_service.db.session + ) assert result["save_options"] == [ComposerSaveStrategy.SAVE_TO_CURRENT_VERSION.value] @@ -495,6 +521,7 @@ def test_save_agent_app_composer_rejects_version_save_strategy(): app_id="app-1", account_id="account-1", payload=payload, + session=composer_service.db.session, ) @@ -528,7 +555,11 @@ def test_save_agent_app_composer_updates_normal_draft(monkeypatch: pytest.Monkey ) result = AgentComposerService.save_agent_app_composer( - tenant_id="tenant-1", app_id="app-1", account_id="account-1", payload=payload + tenant_id="tenant-1", + app_id="app-1", + account_id="account-1", + payload=payload, + session=composer_service.db.session, ) assert result.pop("validation") == {"warnings": [], "knowledge_retrieval_placeholder": []} @@ -570,7 +601,7 @@ def test_save_agent_app_composer_keeps_published_when_draft_matches_active_snaps ) AgentComposerService.save_agent_app_composer( - tenant_id="tenant-1", app_id="app-1", account_id="account-1", payload=payload + tenant_id="tenant-1", app_id="app-1", account_id="account-1", payload=payload, session=fake_session ) assert agent.active_config_is_published is True @@ -617,6 +648,7 @@ def test_publish_agent_app_draft_rejects_missing_model(monkeypatch: pytest.Monke agent_id="agent-1", account_id="account-1", version_note="ship it", + session=fake_session, ) assert exc_info.value.error_code == "agent_model_not_configured" @@ -665,6 +697,7 @@ def test_publish_agent_app_draft_creates_published_snapshot(monkeypatch: pytest. agent_id="agent-1", account_id="account-1", version_note="ship it", + session=composer_service.db.session, ) assert result["result"] == "success" @@ -708,6 +741,7 @@ def test_agent_app_build_draft_checkout_and_apply_use_user_isolated_draft(monkey tenant_id="tenant-1", agent_id="agent-1", account_id="account-1", + session=composer_service.db.session, ) build_draft = fake_session.added[0] @@ -729,6 +763,7 @@ def test_agent_app_build_draft_checkout_and_apply_use_user_isolated_draft(monkey tenant_id="tenant-1", agent_id="agent-1", account_id="account-1", + session=composer_service.db.session, ) assert applied["result"] == "success" @@ -787,6 +822,7 @@ def test_agent_app_build_draft_apply_marks_unpublished_when_build_draft_differs( tenant_id="tenant-1", agent_id="agent-1", account_id="account-1", + session=fake_session, ) assert normal_draft.config_snapshot_dict == build_draft.config_snapshot_dict @@ -812,12 +848,21 @@ def test_agent_app_composer_candidates_and_impact(monkeypatch: pytest.MonkeyPatc monkeypatch.setattr(AgentComposerService, "_workspace_dify_tools", lambda **kwargs: []) workflow_candidates = AgentComposerService.get_workflow_candidates( - tenant_id="tenant-1", app_id="app-1", node_id="node-1", user_id="account-1" + tenant_id="tenant-1", + app_id="app-1", + node_id="node-1", + user_id="account-1", + session=composer_service.db.session, ) agent_app_candidates = AgentComposerService.get_agent_app_candidates( - tenant_id="tenant-1", agent_id="agent-1", user_id="account-1" + tenant_id="tenant-1", + agent_id="agent-1", + user_id="account-1", + session=composer_service.db.session, + ) + impact = AgentComposerService.calculate_impact( + tenant_id="tenant-1", current_snapshot_id="version-1", session=composer_service.db.session ) - impact = AgentComposerService.calculate_impact(tenant_id="tenant-1", current_snapshot_id="version-1") assert workflow_candidates["variant"] == "workflow" assert workflow_candidates["allowed_node_job_candidates"]["previous_node_outputs"] == [] @@ -854,7 +899,9 @@ def test_serialize_workflow_state_changes_lock_and_save_options(monkeypatch: pyt version = AgentConfigSnapshot(id="version-1", version=1, config_snapshot='{"prompt":{"system_prompt":"x"}}') monkeypatch.setattr(AgentComposerService, "calculate_impact", lambda **kwargs: {"workflow_node_count": 1}) - state = AgentComposerService._serialize_workflow_state(binding=binding, agent=agent, version=version) + state = AgentComposerService._serialize_workflow_state( + binding=binding, agent=agent, version=version, session=composer_service.db.session + ) assert state["soul_lock"]["locked"] is True assert state["agent"]["role"] == "Tender Analyst" @@ -893,7 +940,9 @@ def test_serialize_workflow_state_passes_user_declared_outputs_through_effective version = AgentConfigSnapshot(id="version-1", version=1, config_snapshot='{"prompt":{"system_prompt":"x"}}') monkeypatch.setattr(AgentComposerService, "calculate_impact", lambda **kwargs: {"workflow_node_count": 1}) - state = AgentComposerService._serialize_workflow_state(binding=binding, agent=agent, version=version) + state = AgentComposerService._serialize_workflow_state( + binding=binding, agent=agent, version=version, session=composer_service.db.session + ) # When the user has declared outputs, effective_declared_outputs is the same # list (no defaults injected). @@ -943,6 +992,7 @@ def test_serialize_workflow_state_includes_inline_debug_conversation_message_sta agent=agent, version=version, account_id="account-1", + session=composer_service.db.session, ) assert state["debug_conversation_id"] == "debug-conversation-1" @@ -1024,6 +1074,7 @@ def test_composer_save_helpers_create_and_rebind_agents(monkeypatch: pytest.Monk account_id="account-1", binding=existing_binding, payload=payload, + session=composer_service.db.session, ) inline_binding = AgentComposerService._save_node_job_only( tenant_id="tenant-1", @@ -1033,6 +1084,7 @@ def test_composer_save_helpers_create_and_rebind_agents(monkeypatch: pytest.Monk account_id="account-1", binding=None, payload=payload, + session=composer_service.db.session, ) new_agent_binding = AgentComposerService._save_as_new_agent( tenant_id="tenant-1", @@ -1042,6 +1094,7 @@ def test_composer_save_helpers_create_and_rebind_agents(monkeypatch: pytest.Monk account_id="account-1", binding=None, payload=payload, + session=composer_service.db.session, ) save_to_roster_binding = AgentComposerService._save_to_roster( tenant_id="tenant-1", @@ -1055,12 +1108,14 @@ def test_composer_save_helpers_create_and_rebind_agents(monkeypatch: pytest.Monk current_snapshot_id="inline-version-1", ), payload=payload, + session=composer_service.db.session, ) new_version_binding = AgentComposerService._save_as_new_version( tenant_id="tenant-1", account_id="account-1", binding=WorkflowAgentNodeBinding(agent_id="roster-agent-1", current_snapshot_id="source-version-1"), payload=payload, + session=composer_service.db.session, ) assert updated_binding.updated_by == "account-1" @@ -1085,6 +1140,7 @@ def test_composer_save_helpers_create_and_rebind_agents(monkeypatch: pytest.Monk "account_id": "account-1", "agent_soul": payload.agent_soul, "node_job": payload.node_job, + "session": composer_service.db.session, } ] @@ -1151,6 +1207,7 @@ def test_node_job_only_updates_inline_agent_soul(monkeypatch: pytest.MonkeyPatch account_id="account-1", binding=binding, payload=payload, + session=composer_service.db.session, ) assert updated_binding.current_snapshot_id == "inline-version-2" @@ -1203,6 +1260,7 @@ def test_node_job_only_switches_roster_binding_to_inline_agent(monkeypatch: pyte account_id="account-1", binding=binding, payload=payload, + session=composer_service.db.session, ) assert updated_binding is binding @@ -1252,6 +1310,7 @@ def test_node_job_only_rejects_start_from_scratch_with_existing_inline_binding_i account_id="account-1", binding=binding, payload=payload, + session=composer_service.db.session, ) @@ -1299,6 +1358,7 @@ def test_node_job_only_rejects_inline_binding_pointing_to_roster_agent(monkeypat account_id="account-1", binding=binding, payload=payload, + session=composer_service.db.session, ) @@ -1385,6 +1445,7 @@ def test_copy_workflow_composer_from_roster_creates_inline_agent_and_preserves_n account_id="account-1", source_agent_id="roster-agent-1", source_snapshot_id="roster-version-2", + session=composer_service.db.session, ) assert state["binding"]["binding_type"] == WorkflowAgentBindingType.INLINE_AGENT.value @@ -1445,6 +1506,7 @@ def test_copy_workflow_composer_from_roster_rejects_stale_source_snapshot(monkey account_id="account-1", source_agent_id="roster-agent-1", source_snapshot_id="roster-version-1", + session=composer_service.db.session, ) @@ -1495,6 +1557,7 @@ def test_copy_workflow_composer_from_roster_is_idempotent_when_already_inline(mo account_id="account-1", source_agent_id="roster-agent-1", idempotency_key="same-click", + session=composer_service.db.session, ) assert state == {"binding_type": WorkflowAgentBindingType.INLINE_AGENT.value} @@ -1573,6 +1636,7 @@ def test_copy_workflow_composer_from_roster_rejects_invalid_source_binding( node_id="node-1", account_id="account-1", source_agent_id="roster-agent-1", + session=composer_service.db.session, ) @@ -1629,6 +1693,7 @@ def test_copy_agent_drive_rows_copies_skill_prefix_and_files(monkeypatch: pytest account_id="account-1", agent_soul=agent_soul, node_job=node_job, + session=composer_service.db.session, ) copied = [row for row in fake_session.added if isinstance(row, AgentDriveFile)] @@ -1654,6 +1719,7 @@ def test_copy_agent_drive_rows_skips_when_no_referenced_drive_keys(monkeypatch: target_agent_id="inline-agent-1", account_id="account-1", agent_soul=agent_soul, + session=composer_service.db.session, ) assert fake_session.added == [] @@ -1680,6 +1746,7 @@ def test_copy_agent_drive_rows_skips_existing_target_keys(monkeypatch: pytest.Mo target_agent_id="inline-agent-1", account_id="account-1", agent_soul=agent_soul, + session=composer_service.db.session, ) assert [row for row in fake_session.added if isinstance(row, AgentDriveFile)] == [] @@ -1743,7 +1810,7 @@ def test_composer_create_agents_syncs_active_config_has_model(monkeypatch: pytes ) class FakeAppService: - def create_app(self, tenant_id, params, account): + def create_app(self, tenant_id, params, account, session): created_apps.append((tenant_id, params, account)) return SimpleNamespace(id="app-agent-1") @@ -1781,6 +1848,7 @@ def test_composer_create_agents_syncs_active_config_has_model(monkeypatch: pytes node_id="node-1", account_id="account-1", agent_soul=_agent_soul_with_model(), + session=composer_service.db.session, ) roster_agent = AgentComposerService._create_roster_agent_for_composer( tenant_id="tenant-1", @@ -1789,6 +1857,7 @@ def test_composer_create_agents_syncs_active_config_has_model(monkeypatch: pytes agent_soul=_agent_soul_with_model(), operation=AgentConfigRevisionOperation.CREATE_VERSION, version_note=None, + session=composer_service.db.session, ) assert workflow_agent.active_config_snapshot_id == "version-with-model" @@ -1810,14 +1879,14 @@ def test_composer_require_account(monkeypatch: pytest.MonkeyPatch): account = SimpleNamespace(id="account-1") monkeypatch.setattr(composer_service.db, "session", SimpleNamespace(get=lambda model, account_id: account)) - assert AgentComposerService._require_account(account_id="account-1") is account + assert AgentComposerService._require_account(account_id="account-1", session=composer_service.db.session) is account def test_composer_require_account_raises_when_missing(monkeypatch: pytest.MonkeyPatch): monkeypatch.setattr(composer_service.db, "session", SimpleNamespace(get=lambda model, account_id: None)) with pytest.raises(ValueError, match="Account not found"): - AgentComposerService._require_account(account_id="missing-account") + AgentComposerService._require_account(account_id="missing-account", session=composer_service.db.session) def test_composer_create_roster_agent_rolls_back_name_conflict(monkeypatch: pytest.MonkeyPatch): @@ -1825,7 +1894,7 @@ def test_composer_create_roster_agent_rolls_back_name_conflict(monkeypatch: pyte monkeypatch.setattr(composer_service.db, "session", fake_session) class FakeAppService: - def create_app(self, tenant_id, params, account): + def create_app(self, tenant_id, params, account, session): raise IntegrityError("insert apps", params, Exception("duplicate")) monkeypatch.setattr(composer_service, "AppService", FakeAppService) @@ -1839,6 +1908,7 @@ def test_composer_create_roster_agent_rolls_back_name_conflict(monkeypatch: pyte agent_soul=_agent_soul_with_model(), operation=AgentConfigRevisionOperation.CREATE_VERSION, version_note=None, + session=composer_service.db.session, ) assert fake_session.rollbacks == 1 @@ -1849,7 +1919,7 @@ def test_composer_create_roster_agent_raises_when_backing_agent_missing(monkeypa monkeypatch.setattr(composer_service.db, "session", fake_session) class FakeAppService: - def create_app(self, tenant_id, params, account): + def create_app(self, tenant_id, params, account, session): return SimpleNamespace(id="app-agent-1") class FakeAgentRosterService: @@ -1871,6 +1941,7 @@ def test_composer_create_roster_agent_raises_when_backing_agent_missing(monkeypa agent_soul=_agent_soul_with_model(), operation=AgentConfigRevisionOperation.CREATE_VERSION, version_note=None, + session=composer_service.db.session, ) @@ -1892,6 +1963,7 @@ def test_agent_app_draft_match_does_not_mark_create_version_as_published(monkeyp tenant_id="tenant-1", agent=agent, agent_soul=agent_soul, + session=fake_session, ) is False ) @@ -1915,6 +1987,7 @@ def test_agent_app_draft_match_marks_publish_visible_revision_as_published(monke tenant_id="tenant-1", agent=agent, agent_soul=agent_soul, + session=fake_session, ) is True ) @@ -1945,6 +2018,7 @@ def test_composer_version_helpers_and_lookup_errors(monkeypatch: pytest.MonkeyPa agent_soul=agent_soul, operation=AgentConfigRevisionOperation.SAVE_NEW_VERSION, version_note="note", + session=composer_service.db.session, ) updated_snapshot = AgentComposerService._update_current_version( current_snapshot=AgentConfigSnapshot( @@ -1958,21 +2032,40 @@ def test_composer_version_helpers_and_lookup_errors(monkeypatch: pytest.MonkeyPa agent_soul=agent_soul, operation=AgentConfigRevisionOperation.SAVE_CURRENT_VERSION, version_note="updated", + session=composer_service.db.session, + ) + workflow = AgentComposerService._get_draft_workflow( + tenant_id="tenant-1", app_id="app-1", session=composer_service.db.session ) - workflow = AgentComposerService._get_draft_workflow(tenant_id="tenant-1", app_id="app-1") with pytest.raises(ValueError): - AgentComposerService._get_draft_workflow(tenant_id="tenant-1", app_id="missing") - assert AgentComposerService._require_agent(tenant_id="tenant-1", agent_id="agent-1").id == "agent-1" - with pytest.raises(composer_service.AgentNotFoundError): - AgentComposerService._require_agent(tenant_id="tenant-1", agent_id=None) - assert AgentComposerService._get_agent_if_present(tenant_id="tenant-1", agent_id="agent-1") is None + AgentComposerService._get_draft_workflow( + tenant_id="tenant-1", app_id="missing", session=composer_service.db.session + ) assert ( - AgentComposerService._require_version(tenant_id="tenant-1", agent_id="agent-1", version_id="version-1").id + AgentComposerService._require_agent( + tenant_id="tenant-1", agent_id="agent-1", session=composer_service.db.session + ).id + == "agent-1" + ) + with pytest.raises(composer_service.AgentNotFoundError): + AgentComposerService._require_agent(tenant_id="tenant-1", agent_id=None, session=composer_service.db.session) + assert ( + AgentComposerService._get_agent_if_present( + tenant_id="tenant-1", agent_id="agent-1", session=composer_service.db.session + ) + is None + ) + assert ( + AgentComposerService._require_version( + tenant_id="tenant-1", agent_id="agent-1", version_id="version-1", session=composer_service.db.session + ).id == "version-1" ) with pytest.raises(composer_service.AgentVersionNotFoundError): - AgentComposerService._require_version(tenant_id="tenant-1", agent_id="agent-1", version_id="missing") + AgentComposerService._require_version( + tenant_id="tenant-1", agent_id="agent-1", version_id="missing", session=composer_service.db.session + ) assert version.version == 2 assert updated_snapshot.version == 3 @@ -2006,7 +2099,11 @@ def test_composer_current_version_and_error_paths(monkeypatch: pytest.MonkeyPatc ) result = AgentComposerService._save_to_current_version( - tenant_id="tenant-1", account_id="account-1", binding=binding, payload=payload + tenant_id="tenant-1", + account_id="account-1", + binding=binding, + payload=payload, + session=composer_service.db.session, ) assert result.updated_by == "account-1" @@ -2027,6 +2124,7 @@ def test_composer_current_version_and_error_paths(monkeypatch: pytest.MonkeyPatc "save_strategy": ComposerSaveStrategy.SAVE_AS_NEW_AGENT.value, } ), + session=composer_service.db.session, ) @@ -3172,7 +3270,7 @@ class TestAgentAppBackingAgent: captured: dict[str, object] = {} class FakeAppService: - def create_app(self, tenant_id: str, params, account: object) -> object: + def create_app(self, tenant_id: str, params, account: object, session: object) -> object: captured["tenant_id"] = tenant_id captured["params"] = params captured["account"] = account @@ -3241,7 +3339,7 @@ class TestAgentAppBackingAgent: captured: dict[str, object] = {} class FakeAppService: - def create_app(self, tenant_id: str, params, account: object) -> object: + def create_app(self, tenant_id: str, params, account: object, session: object) -> object: captured["params"] = params return target_app @@ -3303,7 +3401,7 @@ class TestAgentAppBackingAgent: monkeypatch.setattr(service, "_next_duplicate_agent_name", lambda **_: "Iris copy") class FakeAppService: - def create_app(self, tenant_id: str, params, account: object) -> object: + def create_app(self, tenant_id: str, params, account: object, session: object) -> object: return target_app access_mode_updates = [] @@ -4307,7 +4405,33 @@ def test_dataset_rows_filters_malformed_ids(monkeypatch: pytest.MonkeyPatch): assert captured == {} -def test_validate_knowledge_datasets_rejects_malformed_ids_without_dataset_lookup(monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize( + ("variant", "save_call"), + [ + ( + ComposerVariant.AGENT_APP, + lambda payload: AgentComposerService.save_agent_app_composer( + tenant_id="tenant-1", + app_id="app-1", + account_id="account-1", + payload=payload, + session=composer_service.db.session, + ), + ), + ( + ComposerVariant.WORKFLOW, + lambda payload: AgentComposerService.save_workflow_composer( + tenant_id="tenant-1", + app_id="app-1", + node_id="node-1", + account_id="account-1", + payload=payload, + session=composer_service.db.session, + ), + ), + ], +) +def test_composer_save_rejects_malformed_knowledge_dataset_ids(monkeypatch: pytest.MonkeyPatch, variant, save_call): captured = {"calls": 0} def fake_get_datasets_by_ids(ids, tenant_id): @@ -4342,7 +4466,35 @@ def test_validate_knowledge_datasets_rejects_malformed_ids_without_dataset_looku assert captured == {"calls": 0} -def test_validate_knowledge_datasets_rejects_missing_or_out_of_scope_datasets(monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize( + ("variant", "save_call"), + [ + ( + ComposerVariant.AGENT_APP, + lambda payload: AgentComposerService.save_agent_app_composer( + tenant_id="tenant-1", + app_id="app-1", + account_id="account-1", + payload=payload, + session=composer_service.db.session, + ), + ), + ( + ComposerVariant.WORKFLOW, + lambda payload: AgentComposerService.save_workflow_composer( + tenant_id="tenant-1", + app_id="app-1", + node_id="node-1", + account_id="account-1", + payload=payload, + session=composer_service.db.session, + ), + ), + ], +) +def test_composer_save_rejects_missing_or_out_of_scope_knowledge_datasets( + monkeypatch: pytest.MonkeyPatch, variant, save_call +): captured = {} missing_dataset_id = "550e8400-e29b-41d4-a716-446655440000" @@ -4431,6 +4583,7 @@ def test_save_agent_composer_allows_incomplete_knowledge_draft(monkeypatch: pyte agent_id="agent-1", account_id="account-1", payload=payload, + session=fake_session, ) assert result["loaded"] is True @@ -4519,6 +4672,7 @@ def test_drive_mention_findings_reports_missing_keys(monkeypatch: pytest.MonkeyP tenant_id="tenant-1", agent_id="agent-1", prompt=_drive_soul().prompt.system_prompt, + session=composer_service.db.session, ) assert [(f["code"], f["id"]) for f in findings] == [("mention_target_missing", "files/sample.pdf")] @@ -4534,6 +4688,7 @@ def test_drive_mention_findings_clean_when_all_keys_exist(monkeypatch: pytest.Mo tenant_id="tenant-1", agent_id="agent-1", prompt=_drive_soul().prompt.system_prompt, + session=composer_service.db.session, ) == [] ) @@ -4546,6 +4701,7 @@ def test_drive_mention_findings_skips_prompt_without_drive_mentions(monkeypatch: tenant_id="tenant-1", agent_id="agent-1", prompt=soul.prompt.system_prompt, + session=composer_service.db.session, ) assert findings == [] @@ -4565,7 +4721,10 @@ def test_collect_validation_findings_appends_drive_mention_findings_with_agent_c ) findings = AgentComposerService.collect_validation_findings( - tenant_id="tenant-1", payload=payload, agent_id="agent-1" + tenant_id="tenant-1", + payload=payload, + agent_id="agent-1", + session=composer_service.db.session, ) codes = {w["code"] for w in findings["warnings"]} @@ -4575,7 +4734,9 @@ def test_collect_validation_findings_appends_drive_mention_findings_with_agent_c "files/sample.pdf", } # without agent context the drive check is skipped entirely - findings_no_agent = AgentComposerService.collect_validation_findings(tenant_id="tenant-1", payload=payload) + findings_no_agent = AgentComposerService.collect_validation_findings( + tenant_id="tenant-1", payload=payload, session=composer_service.db.session + ) assert all(w["code"] != "mention_target_missing" for w in findings_no_agent["warnings"]) @@ -4588,7 +4749,12 @@ def test_resolve_bound_agent_id_queries_active_roster_agent(monkeypatch: pytest. import services.agent.composer_service as module monkeypatch.setattr(module.db, "session", SimpleNamespace(scalar=lambda stmt: "agent-9")) - assert AgentComposerService.resolve_bound_agent_id(tenant_id="t-1", app_id="app-1") == "agent-9" + assert ( + AgentComposerService.resolve_bound_agent_id( + tenant_id="t-1", app_id="app-1", session=composer_service.db.session + ) + == "agent-9" + ) def test_resolve_workflow_node_agent_id_degrades_without_workflow_or_binding(monkeypatch: pytest.MonkeyPatch): @@ -4598,20 +4764,35 @@ def test_resolve_workflow_node_agent_id_degrades_without_workflow_or_binding(mon raise ValueError("no draft workflow") monkeypatch.setattr(AgentComposerService, "_get_draft_workflow", classmethod(boom)) - assert AgentComposerService.resolve_workflow_node_agent_id(tenant_id="t", app_id="a", node_id="n") is None + assert ( + AgentComposerService.resolve_workflow_node_agent_id( + tenant_id="t", app_id="a", node_id="n", session=composer_service.db.session + ) + is None + ) monkeypatch.setattr( AgentComposerService, "_get_draft_workflow", classmethod(lambda cls, **kwargs: SimpleNamespace(id="wf-1")) ) monkeypatch.setattr(AgentComposerService, "_get_workflow_binding", classmethod(lambda cls, **kwargs: None)) - assert AgentComposerService.resolve_workflow_node_agent_id(tenant_id="t", app_id="a", node_id="n") is None + assert ( + AgentComposerService.resolve_workflow_node_agent_id( + tenant_id="t", app_id="a", node_id="n", session=composer_service.db.session + ) + is None + ) monkeypatch.setattr( AgentComposerService, "_get_workflow_binding", classmethod(lambda cls, **kwargs: SimpleNamespace(agent_id="agent-7")), ) - assert AgentComposerService.resolve_workflow_node_agent_id(tenant_id="t", app_id="a", node_id="n") == "agent-7" + assert ( + AgentComposerService.resolve_workflow_node_agent_id( + tenant_id="t", app_id="a", node_id="n", session=composer_service.db.session + ) + == "agent-7" + ) def test_save_workflow_composer_reports_drive_mentions_for_inline_node_job_only(monkeypatch: pytest.MonkeyPatch): @@ -4654,7 +4835,7 @@ def test_save_workflow_composer_reports_drive_mentions_for_inline_node_job_only( ) guarded: dict[str, str] = {} - def fake_collect(cls, *, tenant_id, payload, agent_id=None): + def fake_collect(cls, *, tenant_id, payload, agent_id=None, session=None): guarded["tenant_id"] = tenant_id guarded["agent_id"] = agent_id return {"warnings": [{"code": "mention_target_missing", "id": "files/sample.pdf"}]} @@ -4662,7 +4843,12 @@ def test_save_workflow_composer_reports_drive_mentions_for_inline_node_job_only( monkeypatch.setattr(AgentComposerService, "collect_validation_findings", classmethod(fake_collect)) result = AgentComposerService.save_workflow_composer( - tenant_id="t-1", app_id="app-1", node_id="n-1", account_id="acc-1", payload=payload + tenant_id="t-1", + app_id="app-1", + node_id="n-1", + account_id="acc-1", + payload=payload, + session=composer_service.db.session, ) assert result == { @@ -4712,14 +4898,19 @@ def test_save_workflow_composer_reports_drive_mentions_for_roster_node_job_only( ) captured: dict[str, str | None] = {} - def fake_collect(cls, *, tenant_id, payload, agent_id=None): + def fake_collect(cls, *, tenant_id, payload, agent_id=None, session=None): captured["agent_id"] = agent_id return {"warnings": []} monkeypatch.setattr(AgentComposerService, "collect_validation_findings", classmethod(fake_collect)) result = AgentComposerService.save_workflow_composer( - tenant_id="t-1", app_id="app-1", node_id="n-1", account_id="acc-1", payload=payload + tenant_id="t-1", + app_id="app-1", + node_id="n-1", + account_id="acc-1", + payload=payload, + session=composer_service.db.session, ) assert result == {"state": "ok", "validation": {"warnings": []}} diff --git a/api/tests/unit_tests/services/agent/test_skill_standardize_service.py b/api/tests/unit_tests/services/agent/test_skill_standardize_service.py index 074ac59bb1c..5b3ade55721 100644 --- a/api/tests/unit_tests/services/agent/test_skill_standardize_service.py +++ b/api/tests/unit_tests/services/agent/test_skill_standardize_service.py @@ -50,6 +50,7 @@ def test_standardize_creates_drive_owned_toolfiles_and_commits_archive_manifest( tenant_id="tenant-1", user_id="user-1", agent_id="agent-1", + session=MagicMock(), ) # ToolFiles: SKILL.md and the full archive. Archive members stay lazy. diff --git a/api/tests/unit_tests/services/agent/test_skill_tool_inference_service.py b/api/tests/unit_tests/services/agent/test_skill_tool_inference_service.py index 25678bb4f4d..cfb32d63a92 100644 --- a/api/tests/unit_tests/services/agent/test_skill_tool_inference_service.py +++ b/api/tests/unit_tests/services/agent/test_skill_tool_inference_service.py @@ -38,14 +38,17 @@ def test_infer_returns_suggestions_with_inferred_from(monkeypatch): ' "env_suggestions": [{"key": "OPENAI_API_KEY", "reason": "whisper call", "secret_likely": true}]}]}' ) with patch.object(SkillToolInferenceService, "_invoke", staticmethod(lambda **kwargs: raw)): - result = service.infer(tenant_id="t-1", agent_id="a-1", slug="audio-transcribe") + session = MagicMock() + result = service.infer(tenant_id="t-1", agent_id="a-1", slug="audio-transcribe", session=session) assert result["inferable"] is True tool = result["cli_tools"][0] assert tool["name"] == "ffmpeg" assert tool["inferred_from"] == "audio-transcribe" assert tool["env_suggestions"] == [{"key": "OPENAI_API_KEY", "reason": "whisper call", "secret_likely": True}] - drive.preview.assert_called_once_with(tenant_id="t-1", agent_id="a-1", key="audio-transcribe/SKILL.md") + drive.preview.assert_called_once_with( + tenant_id="t-1", agent_id="a-1", key="audio-transcribe/SKILL.md", session=session + ) def test_infer_threads_skill_md_into_the_prompt(monkeypatch): @@ -57,7 +60,7 @@ def test_infer_threads_skill_md_into_the_prompt(monkeypatch): return '{"inferable": false, "cli_tools": [], "reason": "none"}' with patch.object(SkillToolInferenceService, "_invoke", staticmethod(fake_invoke)): - service.infer(tenant_id="t-1", agent_id="a-1", slug="audio-transcribe") + service.infer(tenant_id="t-1", agent_id="a-1", slug="audio-transcribe", session=MagicMock()) assert "Files inside the skill package" not in captured["prompt"] assert "ffmpeg" in captured["prompt"] # SKILL.md body present @@ -67,7 +70,7 @@ def test_infer_not_inferable_passes_reason_through(monkeypatch): service, _ = _service() raw = '{"inferable": false, "cli_tools": [], "reason": "SKILL.md 未描述任何外部命令依赖"}' with patch.object(SkillToolInferenceService, "_invoke", staticmethod(lambda **kwargs: raw)): - result = service.infer(tenant_id="t-1", agent_id="a-1", slug="audio-transcribe") + result = service.infer(tenant_id="t-1", agent_id="a-1", slug="audio-transcribe", session=MagicMock()) assert result == {"inferable": False, "cli_tools": [], "reason": "SKILL.md 未描述任何外部命令依赖"} @@ -81,7 +84,7 @@ def test_infer_retries_once_then_422(monkeypatch): with patch.object(SkillToolInferenceService, "_invoke", staticmethod(bad_invoke)): with pytest.raises(SkillToolInferenceError) as exc_info: - service.infer(tenant_id="t-1", agent_id="a-1", slug="audio-transcribe") + service.infer(tenant_id="t-1", agent_id="a-1", slug="audio-transcribe", session=MagicMock()) assert len(calls) == 2 # one retry assert exc_info.value.code == "inference_failed" @@ -92,7 +95,7 @@ def test_infer_repairs_slightly_malformed_json(monkeypatch): service, _ = _service() raw = 'Here you go: {"inferable": true, "cli_tools": [], "reason": null,}' with patch.object(SkillToolInferenceService, "_invoke", staticmethod(lambda **kwargs: raw)): - result = service.infer(tenant_id="t-1", agent_id="a-1", slug="audio-transcribe") + result = service.infer(tenant_id="t-1", agent_id="a-1", slug="audio-transcribe", session=MagicMock()) assert result["inferable"] is True @@ -102,7 +105,7 @@ def test_missing_skill_maps_to_404(): service = SkillToolInferenceService(drive_service=drive) with pytest.raises(SkillToolInferenceError) as exc_info: - service.infer(tenant_id="t-1", agent_id="a-1", slug="ghost") + service.infer(tenant_id="t-1", agent_id="a-1", slug="ghost", session=MagicMock()) assert exc_info.value.code == "skill_not_found" assert exc_info.value.status_code == 404 @@ -110,7 +113,7 @@ def test_missing_skill_maps_to_404(): def test_binary_skill_md_maps_to_404(): service, _ = _service(preview={"key": "x/SKILL.md", "size": 1, "truncated": False, "binary": True, "text": None}) with pytest.raises(SkillToolInferenceError) as exc_info: - service.infer(tenant_id="t-1", agent_id="a-1", slug="x") + service.infer(tenant_id="t-1", agent_id="a-1", slug="x", session=MagicMock()) assert exc_info.value.code == "skill_not_found" @@ -160,5 +163,5 @@ def test_load_skill_md_passes_through_non_missing_drive_errors(): service = SkillToolInferenceService(drive_service=drive) with pytest.raises(SkillToolInferenceError) as exc_info: - service.infer(tenant_id="t-1", agent_id="a-1", slug="x") + service.infer(tenant_id="t-1", agent_id="a-1", slug="x", session=MagicMock()) assert exc_info.value.code == "agent_not_found" diff --git a/api/tests/unit_tests/services/data_migration/test_export_service.py b/api/tests/unit_tests/services/data_migration/test_export_service.py index f5480ff52af..5479de5ba22 100644 --- a/api/tests/unit_tests/services/data_migration/test_export_service.py +++ b/api/tests/unit_tests/services/data_migration/test_export_service.py @@ -1,3 +1,5 @@ +from unittest.mock import MagicMock + import pytest from services.data_migration.dependency_discovery_service import DiscoveredDependency @@ -126,13 +128,13 @@ def test_secret_free_mcp_dependencies_are_dependency_only(): report_items = [] service._export_mcp_tools( - object(), tenant_id="tenant-1", provider_ids=["mcp-1"], include_secrets=False, exported_mcp_tools=mcp_tools, dependencies=dependencies, report_items=report_items, + session=MagicMock(), ) assert mcp_tools == [] @@ -151,12 +153,14 @@ def test_secret_free_mcp_dependencies_are_dependency_only(): def test_get_mcp_provider_does_not_compare_non_uuid_identifier_to_uuid_id(): statements = [] - class StubSession: - def scalar(self, statement): - statements.append(str(statement)) + def capture_scalar(statement): + statements.append(str(statement)) + + session = MagicMock() + session.scalar.side_effect = capture_scalar with pytest.raises(MigrationDataError, match="MCP provider not found"): - MigrationExportService()._get_mcp_provider(StubSession(), "tenant-1", "my-test-mcp") + MigrationExportService()._get_mcp_provider("tenant-1", "my-test-mcp", session=session) assert len(statements) == 1 assert "tool_mcp_providers.id =" not in statements[0] diff --git a/api/tests/unit_tests/services/data_migration/test_import_service.py b/api/tests/unit_tests/services/data_migration/test_import_service.py index 2b11d575ed6..10460aba470 100644 --- a/api/tests/unit_tests/services/data_migration/test_import_service.py +++ b/api/tests/unit_tests/services/data_migration/test_import_service.py @@ -3,6 +3,7 @@ import yaml from models.tools import MCPToolProvider, WorkflowToolProvider from services.app_dsl_service import Import +from services.data_migration import import_service from services.data_migration.entities import ( ConflictStrategy, IdStrategy, @@ -92,8 +93,12 @@ def test_package_target_tenant_id_ignores_invalid_uuid(monkeypatch): return EmptyResult() + from services.data_migration import import_service + + monkeypatch.setattr(import_service.db, "session", StubSession()) + with pytest.raises(MigrationDataError, match="Target tenant not found"): - ImportTargetResolver().resolve(StubSession(), ImportRequest(package=package)) + ImportTargetResolver().resolve(ImportRequest(package=package), session=import_service.db.session) def test_options_override_replaces_package_defaults(): @@ -113,7 +118,7 @@ def test_options_override_replaces_package_defaults(): captured_options: list[ImportOptions] = [] class StubResolver(ImportTargetResolver): - def resolve(self, session, request: ImportRequest) -> ImportTarget: + def resolve(self, request: ImportRequest, session) -> ImportTarget: return ImportTarget( tenant_id="tenant-1", tenant_name="target", @@ -124,7 +129,6 @@ def test_options_override_replaces_package_defaults(): class CapturingImportService(MigrationImportService): def _import_workflows( self, - session, package: MigrationPackage, target: ImportTarget, options: ImportOptions, @@ -137,7 +141,8 @@ def test_options_override_replaces_package_defaults(): override = ImportOptions(create_app_api_token_on_import=False, conflict_strategy=ConflictStrategy.SKIP) CapturingImportService(target_resolver=StubResolver()).import_package( - object(), ImportRequest(package=package, options_override=override) + ImportRequest(package=package, options_override=override), + session=import_service.db.session, ) assert captured_options == [override] @@ -150,37 +155,53 @@ def test_only_preserve_id_strategy_reuses_source_app_id(): assert service._should_preserve_source_app_id(ImportOptions(id_strategy=IdStrategy.GENERATE_NEW_ID)) is False -def test_find_existing_app_ignores_invalid_uuid(): +def test_find_existing_app_ignores_invalid_uuid(monkeypatch): class StubSession: def scalar(self, statement): raise AssertionError("invalid UUID should not be queried against App.id") - assert MigrationImportService()._find_existing_app(StubSession(), "not-a-uuid", "tenant-1") is None + from services.data_migration import import_service + + monkeypatch.setattr(import_service.db, "session", StubSession()) + + assert ( + MigrationImportService()._find_existing_app("not-a-uuid", "tenant-1", session=import_service.db.session) is None + ) -def test_find_existing_workflow_tool_does_not_compare_invalid_uuid(): +def test_find_existing_workflow_tool_does_not_compare_invalid_uuid(monkeypatch): captured = [] class StubSession: def scalar(self, statement): captured.append(statement) + from services.data_migration import import_service + + monkeypatch.setattr(import_service.db, "session", StubSession()) + MigrationImportService()._find_existing_workflow_tool( - StubSession(), "tenant-1", "not-a-uuid", "tool-name", "app-id" + "tenant-1", "not-a-uuid", "tool-name", "app-id", session=import_service.db.session ) where_clause = str(captured[0].whereclause) assert f"{WorkflowToolProvider.__tablename__}.id" not in where_clause -def test_find_existing_mcp_tool_does_not_compare_invalid_uuid(): +def test_find_existing_mcp_tool_does_not_compare_invalid_uuid(monkeypatch): captured = [] class StubSession: def scalar(self, statement): captured.append(statement) - MigrationImportService()._find_existing_mcp_tool(StubSession(), "tenant-1", "my-test-mcp", "my-test-mcp") + from services.data_migration import import_service + + monkeypatch.setattr(import_service.db, "session", StubSession()) + + MigrationImportService()._find_existing_mcp_tool( + "tenant-1", "my-test-mcp", "my-test-mcp", session=import_service.db.session + ) where_clause = str(captured[0].whereclause) assert f"{MCPToolProvider.__tablename__}.id" not in where_clause @@ -211,16 +232,17 @@ def test_workflow_app_import_does_not_wrap_app_dsl_import_in_nested_transaction( from services.data_migration import import_service + monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr(import_service, "AppDslService", StubAppDslService) imported_app_id = MigrationImportService()._import_workflow_app( - session=StubSession(), account=object(), workflow_data={"name": "main_chatflow"}, dsl_content="app:\n mode: workflow\n", app_id="source-app-id", existing_app=None, options=ImportOptions(id_strategy=IdStrategy.PRESERVE_ID), + session=import_service.db.session, ) assert imported_app_id == "imported-app-id" @@ -317,19 +339,20 @@ def test_workflow_tool_import_publishes_referenced_app_before_create(monkeypatch return account class PublishingImportService(MigrationImportService): - def _find_existing_app(self, session, app_id, tenant_id): + def _find_existing_app(self, app_id, tenant_id, session): return object() - def _find_existing_workflow_tool(self, session, tenant_id, workflow_tool_id, tool_name, app_id): + def _find_existing_workflow_tool(self, tenant_id, workflow_tool_id, tool_name, app_id, session): if ("created", app_id) in events: return type("WorkflowToolProvider", (), {"id": workflow_tool_id or "created-workflow-tool-id"})() return None - def _ensure_workflow_app_is_published(self, session, target, account, app_id): + def _ensure_workflow_app_is_published(self, target, account, app_id, session): events.append(("published", app_id)) from services.data_migration import import_service + monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr( import_service.WorkflowToolManageService, "create_workflow_tool", @@ -337,7 +360,6 @@ def test_workflow_tool_import_publishes_referenced_app_before_create(monkeypatch ) PublishingImportService()._import_workflow_tools( - StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -360,6 +382,7 @@ def test_workflow_tool_import_publishes_referenced_app_before_create(monkeypatch {}, [], [], + session=import_service.db.session, ) assert events == [("published", "workflow-app-1"), ("created", "workflow-app-1")] @@ -384,17 +407,18 @@ def test_workflow_tool_import_id_follows_id_strategy(monkeypatch: pytest.MonkeyP return account class StrategyImportService(MigrationImportService): - def _find_existing_app(self, session, app_id, tenant_id): + def _find_existing_app(self, app_id, tenant_id, session): return object() - def _find_existing_workflow_tool(self, session, tenant_id, workflow_tool_id, tool_name, app_id): + def _find_existing_workflow_tool(self, tenant_id, workflow_tool_id, tool_name, app_id, session): return target_provider if created_kwargs else None - def _ensure_workflow_app_is_published(self, session, target, account, app_id): + def _ensure_workflow_app_is_published(self, target, account, app_id, session): return None from services.data_migration import import_service + monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr( import_service.WorkflowToolManageService, "create_workflow_tool", @@ -402,7 +426,6 @@ def test_workflow_tool_import_id_follows_id_strategy(monkeypatch: pytest.MonkeyP ) StrategyImportService()._import_workflow_tools( - StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -425,6 +448,7 @@ def test_workflow_tool_import_id_follows_id_strategy(monkeypatch: pytest.MonkeyP id_mapping, id_mapping_details, [], + session=import_service.db.session, ) assert created_kwargs[0]["import_id"] == expected_import_id @@ -449,17 +473,20 @@ def test_workflow_tool_skip_records_id_mapping(monkeypatch): return account class SkipImportService(MigrationImportService): - def _find_existing_app(self, session, app_id, tenant_id): + def _find_existing_app(self, app_id, tenant_id, session): return object() - def _find_existing_workflow_tool(self, session, tenant_id, workflow_tool_id, tool_name, app_id): + def _find_existing_workflow_tool(self, tenant_id, workflow_tool_id, tool_name, app_id, session): return existing_provider - def _ensure_workflow_app_is_published(self, session, target, account, app_id): + def _ensure_workflow_app_is_published(self, target, account, app_id, session): return None + from services.data_migration import import_service + + monkeypatch.setattr(import_service.db, "session", StubSession()) + SkipImportService()._import_workflow_tools( - StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -482,6 +509,7 @@ def test_workflow_tool_skip_records_id_mapping(monkeypatch): id_mapping, [], [], + session=import_service.db.session, ) assert id_mapping["source-workflow-tool-id"] == "existing-workflow-tool-id" @@ -495,22 +523,18 @@ def test_api_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str report_items = [] class ExistingApiImportService(MigrationImportService): - def _find_api_tool_provider(self, session, tenant_id, provider_name): - return target_provider - - class StubSession: - def scalar(self, statement): + def _find_api_tool_provider(self, tenant_id, provider_name, session): return target_provider from services.data_migration import import_service + monkeypatch.setattr(import_service.db.session, "scalar", lambda statement: target_provider) monkeypatch.setattr( import_service.ApiToolManageService, "parser_api_schema", lambda schema: {"schema_type": "openapi"} ) monkeypatch.setattr(import_service.ApiToolManageService, "update_api_tool_provider", lambda **kwargs: None) ExistingApiImportService()._import_api_tools( - StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -528,6 +552,7 @@ def test_api_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str id_mapping, id_mapping_details, {"weather": {"source-api-provider-id-from-dsl"}}, + session=import_service.db.session, ) assert id_mapping == { @@ -549,18 +574,18 @@ def test_api_tool_create_records_id_mapping(monkeypatch): return None class CreatedApiImportService(MigrationImportService): - def _find_api_tool_provider(self, session, tenant_id, provider_name): + def _find_api_tool_provider(self, tenant_id, provider_name, session): return target_provider from services.data_migration import import_service + monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr( import_service.ApiToolManageService, "parser_api_schema", lambda schema: {"schema_type": "openapi"} ) monkeypatch.setattr(import_service.ApiToolManageService, "create_api_tool_provider", lambda **kwargs: None) CreatedApiImportService()._import_api_tools( - StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -578,6 +603,7 @@ def test_api_tool_create_records_id_mapping(monkeypatch): id_mapping, [], {}, + session=import_service.db.session, ) assert id_mapping["source-api-provider-id"] == "target-api-provider-id" @@ -605,10 +631,10 @@ def test_mcp_tool_import_restores_exported_tool_list(monkeypatch): from services.data_migration import import_service + monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr(import_service, "MCPToolManageService", StubMCPToolManageService) MigrationImportService()._import_mcp_tools( - StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -634,6 +660,7 @@ def test_mcp_tool_import_restores_exported_tool_list(monkeypatch): report_items, {}, [], + session=import_service.db.session, ) assert provider.tools == '[{"name": "echo"}]' @@ -653,7 +680,7 @@ def test_mcp_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str return None class ExistingMCPImportService(MigrationImportService): - def _find_existing_mcp_tool(self, session, tenant_id, provider_id, server_identifier): + def _find_existing_mcp_tool(self, tenant_id, provider_id, server_identifier, session): return provider class StubMCPToolManageService: @@ -665,10 +692,10 @@ def test_mcp_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str from services.data_migration import import_service + monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr(import_service, "MCPToolManageService", StubMCPToolManageService) ExistingMCPImportService()._import_mcp_tools( - StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -694,6 +721,7 @@ def test_mcp_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str [], id_mapping, id_mapping_details, + session=import_service.db.session, ) assert id_mapping["source-mcp-provider-id"] == "target-mcp-provider-id" @@ -715,7 +743,7 @@ def test_mcp_tool_create_records_id_mapping(monkeypatch): return None class CreatedMCPImportService(MigrationImportService): - def _find_existing_mcp_tool(self, session, tenant_id, provider_id, server_identifier): + def _find_existing_mcp_tool(self, tenant_id, provider_id, server_identifier, session): return provider if provider_created else None class StubMCPToolManageService: @@ -728,10 +756,10 @@ def test_mcp_tool_create_records_id_mapping(monkeypatch): from services.data_migration import import_service + monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr(import_service, "MCPToolManageService", StubMCPToolManageService) CreatedMCPImportService()._import_mcp_tools( - StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -756,6 +784,7 @@ def test_mcp_tool_create_records_id_mapping(monkeypatch): [], id_mapping, [], + session=import_service.db.session, ) assert id_mapping["source-mcp-provider-id"] == "target-mcp-provider-id" @@ -800,12 +829,11 @@ def test_dependency_only_mcp_preflight_reports_missing_target_provider_with_work } ) - class StubSession: - def scalar(self, statement): - return None + from services.data_migration import import_service + + monkeypatch.setattr(import_service.db.session, "scalar", lambda statement: None) MigrationImportService()._preflight_dependency_only_mcp( - StubSession(), package, ImportTarget( tenant_id="tenant-1", @@ -814,6 +842,7 @@ def test_dependency_only_mcp_preflight_reports_missing_target_provider_with_work operator_email="owner@example.com", ), report_items, + session=import_service.db.session, ) assert report_items == [ @@ -828,18 +857,22 @@ def test_dependency_only_mcp_preflight_reports_missing_target_provider_with_work ] -def test_dependency_only_mcp_lookup_does_not_compare_non_uuid_identifier_to_uuid_id(): +def test_dependency_only_mcp_lookup_does_not_compare_non_uuid_identifier_to_uuid_id(monkeypatch): captured = [] class StubSession: def scalar(self, statement): captured.append(statement) + from services.data_migration import import_service + + monkeypatch.setattr(import_service.db, "session", StubSession()) + MigrationImportService()._find_dependency_only_mcp_provider( - StubSession(), "tenant-1", "my-test-mcp-server", "my-test-mcp", + session=import_service.db.session, ) where_clause = str(captured[0].whereclause) @@ -860,12 +893,11 @@ def test_dependency_only_mcp_preflight_reports_available_target_provider(monkeyp {"id": "target-provider-id", "name": "my-test-mcp", "server_identifier": "my-test-mcp-server"}, )() - class StubSession: - def scalar(self, statement): - return provider + from services.data_migration import import_service + + monkeypatch.setattr(import_service.db.session, "scalar", lambda statement: provider) MigrationImportService()._preflight_dependency_only_mcp( - StubSession(), package, ImportTarget( tenant_id="tenant-1", @@ -874,6 +906,7 @@ def test_dependency_only_mcp_preflight_reports_available_target_provider(monkeyp operator_email="owner@example.com", ), report_items, + session=import_service.db.session, ) assert report_items == [ @@ -891,7 +924,7 @@ def test_import_package_imports_workflow_tool_provider_apps_before_consumers(): events = [] class StubResolver(ImportTargetResolver): - def resolve(self, session, request): + def resolve(self, request, session): return ImportTarget( tenant_id="tenant-1", tenant_name="target", @@ -902,7 +935,6 @@ def test_import_package_imports_workflow_tool_provider_apps_before_consumers(): class OrderedImportService(MigrationImportService): def _import_api_tools( self, - session, package, target, options, @@ -910,12 +942,13 @@ def test_import_package_imports_workflow_tool_provider_apps_before_consumers(): id_mapping, id_mapping_details, source_provider_ids_by_name, + *, + session=None, ): events.append(("api_tools", "imported")) def _import_workflows( self, - session, package, target, options, @@ -926,6 +959,7 @@ def test_import_package_imports_workflow_tool_provider_apps_before_consumers(): imported_workflow_ids=None, only_app_ids=None, skip_app_ids=None, + session=None, ): only_app_ids = set(only_app_ids or []) skip_app_ids = set(skip_app_ids or []) @@ -941,11 +975,13 @@ def test_import_package_imports_workflow_tool_provider_apps_before_consumers(): imported_workflow_ids.add(app_id) def _import_workflow_tools( - self, session, package, target, options, id_mapping, id_mapping_details, report_items + self, package, target, options, id_mapping, id_mapping_details, report_items, *, session=None ): events.append(("workflow_tool", package.workflow_tools[0]["id"])) - def _import_mcp_tools(self, session, package, target, options, report_items, id_mapping, id_mapping_details): + def _import_mcp_tools( + self, package, target, options, report_items, id_mapping, id_mapping_details, *, session=None + ): events.append(("mcp_tools", "imported")) package = MigrationPackage.from_mapping( @@ -959,7 +995,9 @@ def test_import_package_imports_workflow_tool_provider_apps_before_consumers(): } ) - OrderedImportService(target_resolver=StubResolver()).import_package(object(), ImportRequest(package=package)) + OrderedImportService(target_resolver=StubResolver()).import_package( + ImportRequest(package=package), session=import_service.db.session + ) assert events == [ ("api_tools", "imported"), diff --git a/api/tests/unit_tests/services/enterprise/test_rbac_service.py b/api/tests/unit_tests/services/enterprise/test_rbac_service.py index fdf921265b2..85638b11fff 100644 --- a/api/tests/unit_tests/services/enterprise/test_rbac_service.py +++ b/api/tests/unit_tests/services/enterprise/test_rbac_service.py @@ -558,7 +558,7 @@ class TestMyPermissions: } with patch(f"{MODULE}.dify_config.RBAC_ENABLED", True): - out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1") + out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", session=MagicMock()) call = _call_args(mock_send) assert call.method == "GET" @@ -613,11 +613,8 @@ class TestMyPermissions: mock_session = MagicMock() mock_session.__enter__.return_value = mock_session mock_session.scalar.return_value = role - with ( - patch(f"{MODULE}.dify_config.RBAC_ENABLED", False), - patch(f"{MODULE}.session_factory.create_session", return_value=mock_session), - ): - out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1") + with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False): + out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", session=mock_session) mock_send.assert_not_called() assert out.workspace.permission_keys == workspace_keys @@ -655,11 +652,8 @@ class TestMyPermissions: mock_session = MagicMock() mock_session.__enter__.return_value = mock_session mock_session.scalar.return_value = role - with ( - patch(f"{MODULE}.dify_config.RBAC_ENABLED", False), - patch(f"{MODULE}.session_factory.create_session", return_value=mock_session), - ): - out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1") + with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False): + out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", session=mock_session) actual_snippet_keys = { permission_key for permission_key in out.workspace.permission_keys if permission_key.startswith("snippets.") @@ -672,11 +666,8 @@ class TestMyPermissions: mock_session = MagicMock() mock_session.__enter__.return_value = mock_session mock_session.scalar.return_value = None - with ( - patch(f"{MODULE}.dify_config.RBAC_ENABLED", False), - patch(f"{MODULE}.session_factory.create_session", return_value=mock_session), - ): - out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1") + with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False): + out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", session=mock_session) mock_send.assert_not_called() assert out.workspace.permission_keys == [] @@ -694,7 +685,7 @@ class TestMyPermissions: } with patch(f"{MODULE}.dify_config.RBAC_ENABLED", True): - out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", app_id="app-1") + out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", app_id="app-1", session=MagicMock()) call = _call_args(mock_send) assert call.method == "GET" @@ -716,7 +707,7 @@ class TestMemberRoles: ], } with patch(f"{MODULE}.dify_config.RBAC_ENABLED", True): - out = svc.RBACService.MemberRoles.get("tenant-1", "acct-1", "acct-2") + out = svc.RBACService.MemberRoles.get("tenant-1", "acct-1", "acct-2", session=MagicMock()) call = _call_args(mock_send) assert call.method == "GET" assert call.endpoint == "/rbac/members/rbac-roles" @@ -728,12 +719,8 @@ class TestMemberRoles: session = MagicMock() session.scalar.return_value = svc.TenantAccountRole.EDITOR - with ( - patch(f"{MODULE}.dify_config.RBAC_ENABLED", False), - patch(f"{MODULE}.session_factory.create_session") as create_session, - ): - create_session.return_value.__enter__.return_value = session - out = svc.RBACService.MemberRoles.get("tenant-1", "acct-1", "acct-2") + with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False): + out = svc.RBACService.MemberRoles.get("tenant-1", "acct-1", "acct-2", session=session) mock_send.assert_not_called() assert out.account_id == "acct-2" @@ -755,7 +742,11 @@ class TestMemberRoles: mock_send.return_value = {"account_id": "acct-2", "roles": []} with patch(f"{MODULE}.dify_config.RBAC_ENABLED", True): svc.RBACService.MemberRoles.replace( - "tenant-1", "acct-1", "acct-2", role_ids=["workspace.owner", "workspace.editor"] + "tenant-1", + "acct-1", + "acct-2", + role_ids=["workspace.owner", "workspace.editor"], + session=MagicMock(), ) call = _call_args(mock_send) assert call.method == "PUT" @@ -769,11 +760,10 @@ class TestMemberRoles: target_join = SimpleNamespace(role=svc.TenantAccountRole.NORMAL, account_id="acct-2") session.scalar.return_value = target_join - with ( - patch(f"{MODULE}.dify_config.RBAC_ENABLED", False), - patch(f"{MODULE}.session_factory.create_session", return_value=session), - ): - out = svc.RBACService.MemberRoles.replace("tenant-1", "acct-1", "acct-2", role_ids=["editor"]) + with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False): + out = svc.RBACService.MemberRoles.replace( + "tenant-1", "acct-1", "acct-2", role_ids=["editor"], session=session + ) mock_send.assert_not_called() session.commit.assert_called_once() @@ -789,11 +779,10 @@ class TestMemberRoles: owner_join = SimpleNamespace(role=svc.TenantAccountRole.OWNER, account_id="acct-owner") session.scalar.side_effect = [target_join, owner_join] - with ( - patch(f"{MODULE}.dify_config.RBAC_ENABLED", False), - patch(f"{MODULE}.session_factory.create_session", return_value=session), - ): - out = svc.RBACService.MemberRoles.replace("tenant-1", "acct-1", "acct-2", role_ids=["owner"]) + with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False): + out = svc.RBACService.MemberRoles.replace( + "tenant-1", "acct-1", "acct-2", role_ids=["owner"], session=session + ) mock_send.assert_not_called() session.commit.assert_called_once() @@ -832,7 +821,9 @@ class TestResourcePermissions: } with patch(f"{MODULE}.dify_config.RBAC_ENABLED", True): - out = svc.RBACService.AppPermissions.batch_get("tenant-1", "acct-1", ["app-1", "app-2"]) + out = svc.RBACService.AppPermissions.batch_get( + "tenant-1", "acct-1", ["app-1", "app-2"], session=MagicMock() + ) call = _call_args(mock_send) assert call.method == "POST" @@ -847,11 +838,10 @@ class TestResourcePermissions: mock_session = MagicMock() mock_session.__enter__.return_value = mock_session mock_session.scalar.return_value = "editor" - with ( - patch(f"{MODULE}.dify_config.RBAC_ENABLED", False), - patch(f"{MODULE}.session_factory.create_session", return_value=mock_session), - ): - out = svc.RBACService.AppPermissions.batch_get("tenant-1", "acct-1", ["app-1", "app-2"]) + with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False): + out = svc.RBACService.AppPermissions.batch_get( + "tenant-1", "acct-1", ["app-1", "app-2"], session=mock_session + ) mock_send.assert_not_called() assert out == { @@ -868,7 +858,9 @@ class TestResourcePermissions: } with patch(f"{MODULE}.dify_config.RBAC_ENABLED", True): - out = svc.RBACService.DatasetPermissions.batch_get("tenant-1", "acct-1", ["ds-1", "ds-2"]) + out = svc.RBACService.DatasetPermissions.batch_get( + "tenant-1", "acct-1", ["ds-1", "ds-2"], session=MagicMock() + ) call = _call_args(mock_send) assert call.method == "POST" @@ -883,11 +875,10 @@ class TestResourcePermissions: mock_session = MagicMock() mock_session.__enter__.return_value = mock_session mock_session.scalar.return_value = "dataset_operator" - with ( - patch(f"{MODULE}.dify_config.RBAC_ENABLED", False), - patch(f"{MODULE}.session_factory.create_session", return_value=mock_session), - ): - out = svc.RBACService.DatasetPermissions.batch_get("tenant-1", "acct-1", ["ds-1", "ds-2"]) + with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False): + out = svc.RBACService.DatasetPermissions.batch_get( + "tenant-1", "acct-1", ["ds-1", "ds-2"], session=mock_session + ) mock_send.assert_not_called() assert out == { diff --git a/api/tests/unit_tests/services/hit_service.py b/api/tests/unit_tests/services/hit_service.py index ae19daba898..ffeb158e37a 100644 --- a/api/tests/unit_tests/services/hit_service.py +++ b/api/tests/unit_tests/services/hit_service.py @@ -186,7 +186,7 @@ class TestHitTestingServiceRetrieve: # Act result = HitTestingService.retrieve( - mock_db_session, dataset, query, account, retrieval_model, external_retrieval_model + dataset, query, account, retrieval_model, external_retrieval_model, session=mock_db_session ) # Assert @@ -234,7 +234,7 @@ class TestHitTestingServiceRetrieve: # Act result = HitTestingService.retrieve( - mock_db_session, dataset, query, account, retrieval_model, external_retrieval_model + dataset, query, account, retrieval_model, external_retrieval_model, session=mock_db_session ) # Assert @@ -292,7 +292,7 @@ class TestHitTestingServiceRetrieve: # Act result = HitTestingService.retrieve( - mock_db_session, dataset, query, account, retrieval_model, external_retrieval_model + dataset, query, account, retrieval_model, external_retrieval_model, session=mock_db_session ) # Assert @@ -337,7 +337,7 @@ class TestHitTestingServiceRetrieve: # Act result = HitTestingService.retrieve( - mock_db_session, dataset, query, account, retrieval_model, external_retrieval_model + dataset, query, account, retrieval_model, external_retrieval_model, session=mock_db_session ) # Assert @@ -380,7 +380,7 @@ class TestHitTestingServiceRetrieve: # Act result = HitTestingService.retrieve( - mock_db_session, dataset, query, account, retrieval_model, external_retrieval_model + dataset, query, account, retrieval_model, external_retrieval_model, session=mock_db_session ) # Assert @@ -438,7 +438,12 @@ class TestHitTestingServiceExternalRetrieve: # Act result = HitTestingService.external_retrieve( - mock_db_session, dataset, query, account, external_retrieval_model, metadata_filtering_conditions + dataset, + query, + account, + external_retrieval_model, + metadata_filtering_conditions, + session=mock_db_session, ) # Assert @@ -469,7 +474,7 @@ class TestHitTestingServiceExternalRetrieve: # Act result = HitTestingService.external_retrieve( - mock_db_session, dataset, query, account, external_retrieval_model, metadata_filtering_conditions + dataset, query, account, external_retrieval_model, metadata_filtering_conditions, session=mock_db_session ) # Assert @@ -504,7 +509,12 @@ class TestHitTestingServiceExternalRetrieve: # Act result = HitTestingService.external_retrieve( - mock_db_session, dataset, query, account, external_retrieval_model, metadata_filtering_conditions + dataset, + query, + account, + external_retrieval_model, + metadata_filtering_conditions, + session=mock_db_session, ) # Assert @@ -538,7 +548,12 @@ class TestHitTestingServiceExternalRetrieve: # Act result = HitTestingService.external_retrieve( - mock_db_session, dataset, query, account, external_retrieval_model, metadata_filtering_conditions + dataset, + query, + account, + external_retrieval_model, + metadata_filtering_conditions, + session=mock_db_session, ) # Assert @@ -579,7 +594,7 @@ class TestHitTestingServiceCompactRetrieveResponse: mock_format.return_value = mock_records # Act - result = HitTestingService.compact_retrieve_response(MagicMock(), query, documents) + result = HitTestingService.compact_retrieve_response(query, documents, session=MagicMock()) # Assert assert result["query"]["content"] == query @@ -605,7 +620,7 @@ class TestHitTestingServiceCompactRetrieveResponse: mock_format.return_value = [] # Act - result = HitTestingService.compact_retrieve_response(MagicMock(), query, documents) + result = HitTestingService.compact_retrieve_response(query, documents, session=MagicMock()) # Assert assert result["query"]["content"] == query diff --git a/api/tests/unit_tests/services/plugin/test_plugin_auto_upgrade_service.py b/api/tests/unit_tests/services/plugin/test_plugin_auto_upgrade_service.py index 0e793fae7ef..e66bb3fff04 100644 --- a/api/tests/unit_tests/services/plugin/test_plugin_auto_upgrade_service.py +++ b/api/tests/unit_tests/services/plugin/test_plugin_auto_upgrade_service.py @@ -15,47 +15,40 @@ PLUGIN_CATEGORY = TenantPluginAutoUpgradeCategory.TOOL def _patched_session(): - """Patch session_factory.create_session() to return a mock session as context manager.""" + """Return a mock SQLAlchemy session for service calls.""" session = MagicMock() - session.__enter__ = MagicMock(return_value=session) - session.__exit__ = MagicMock(return_value=False) - mock_factory = MagicMock() - mock_factory.create_session.return_value = session - patcher = patch(f"{MODULE}.session_factory", mock_factory) - return patcher, session + return session class TestGetStrategy: def test_returns_strategy_when_found(self): - p1, session = _patched_session() + session = _patched_session() strategy = MagicMock() session.scalar.return_value = strategy - with p1: - from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService + from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - result = PluginAutoUpgradeService.get_strategy("t1", PLUGIN_CATEGORY) + result = PluginAutoUpgradeService.get_strategy("t1", PLUGIN_CATEGORY, session=session) assert result is strategy def test_returns_none_when_not_found(self): - p1, session = _patched_session() + session = _patched_session() session.scalar.return_value = None - with p1: - from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService + from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - result = PluginAutoUpgradeService.get_strategy("t1", PLUGIN_CATEGORY) + result = PluginAutoUpgradeService.get_strategy("t1", PLUGIN_CATEGORY, session=session) assert result is None class TestChangeStrategy: def test_creates_new_strategy(self): - p1, session = _patched_session() + session = _patched_session() session.scalar.return_value = None - with p1, patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls: + with patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls: strat_cls.return_value = MagicMock() from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService @@ -67,28 +60,29 @@ class TestChangeStrategy: [], [], category=PLUGIN_CATEGORY, + session=session, ) assert result is True session.add.assert_called_once() def test_updates_existing_strategy(self): - p1, session = _patched_session() + session = _patched_session() existing = MagicMock() session.scalar.return_value = existing - with p1: - from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService + from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - result = PluginAutoUpgradeService.change_strategy( - "t1", - TenantPluginAutoUpgradeStrategySetting.LATEST, - 5, - TenantPluginAutoUpgradeMode.PARTIAL, - ["p1"], - ["p2"], - category=PLUGIN_CATEGORY, - ) + result = PluginAutoUpgradeService.change_strategy( + "t1", + TenantPluginAutoUpgradeStrategySetting.LATEST, + 5, + TenantPluginAutoUpgradeMode.PARTIAL, + ["p1"], + ["p2"], + category=PLUGIN_CATEGORY, + session=session, + ) assert result is True assert existing.strategy_setting == TenantPluginAutoUpgradeStrategySetting.LATEST @@ -100,11 +94,10 @@ class TestChangeStrategy: class TestExcludePlugin: def test_creates_default_strategy_when_none_exists(self): - p1, session = _patched_session() + session = _patched_session() session.scalar.return_value = None with ( - p1, patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy"), ): @@ -114,74 +107,87 @@ class TestExcludePlugin: "t1", "plugin-1", PLUGIN_CATEGORY, + session=session, ) assert result is True session.add.assert_called_once() def test_appends_to_exclude_list_in_exclude_mode(self): - p1, session = _patched_session() + session = _patched_session() existing = MagicMock() existing.upgrade_mode = TenantPluginAutoUpgradeMode.EXCLUDE existing.exclude_plugins = ["p-existing"] session.scalar.return_value = existing - with p1, patch(f"{MODULE}.select"): + with patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls: + strat_cls.UpgradeMode.EXCLUDE = "exclude" + strat_cls.UpgradeMode.PARTIAL = "partial" + strat_cls.UpgradeMode.ALL = "all" from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - result = PluginAutoUpgradeService.exclude_plugin("t1", "p-new", PLUGIN_CATEGORY) + result = PluginAutoUpgradeService.exclude_plugin("t1", "p-new", PLUGIN_CATEGORY, session=session) assert result is True assert existing.exclude_plugins == ["p-existing", "p-new"] def test_removes_from_include_list_in_partial_mode(self): - p1, session = _patched_session() + session = _patched_session() existing = MagicMock() existing.upgrade_mode = TenantPluginAutoUpgradeMode.PARTIAL existing.include_plugins = ["p1", "p2"] session.scalar.return_value = existing - with p1, patch(f"{MODULE}.select"): + with patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls: + strat_cls.UpgradeMode.EXCLUDE = "exclude" + strat_cls.UpgradeMode.PARTIAL = "partial" + strat_cls.UpgradeMode.ALL = "all" from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - result = PluginAutoUpgradeService.exclude_plugin("t1", "p1", PLUGIN_CATEGORY) + result = PluginAutoUpgradeService.exclude_plugin("t1", "p1", PLUGIN_CATEGORY, session=session) assert result is True assert existing.include_plugins == ["p2"] def test_switches_to_exclude_mode_from_all(self): - p1, session = _patched_session() + session = _patched_session() existing = MagicMock() existing.upgrade_mode = TenantPluginAutoUpgradeMode.ALL session.scalar.return_value = existing - with p1, patch(f"{MODULE}.select"): + with patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls: + strat_cls.UpgradeMode.EXCLUDE = "exclude" + strat_cls.UpgradeMode.PARTIAL = "partial" + strat_cls.UpgradeMode.ALL = "all" from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - result = PluginAutoUpgradeService.exclude_plugin("t1", "p1", PLUGIN_CATEGORY) + result = PluginAutoUpgradeService.exclude_plugin("t1", "p1", PLUGIN_CATEGORY, session=session) assert result is True assert existing.upgrade_mode == TenantPluginAutoUpgradeMode.EXCLUDE assert existing.exclude_plugins == ["p1"] def test_no_duplicate_in_exclude_list(self): - p1, session = _patched_session() + session = _patched_session() existing = MagicMock() existing.upgrade_mode = TenantPluginAutoUpgradeMode.EXCLUDE existing.exclude_plugins = ["p1"] session.scalar.return_value = existing - with p1, patch(f"{MODULE}.select"): + with patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls: + strat_cls.UpgradeMode.EXCLUDE = "exclude" + strat_cls.UpgradeMode.PARTIAL = "partial" + strat_cls.UpgradeMode.ALL = "all" from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - PluginAutoUpgradeService.exclude_plugin("t1", "p1", PLUGIN_CATEGORY) + PluginAutoUpgradeService.exclude_plugin("t1", "p1", PLUGIN_CATEGORY, session=session) assert existing.exclude_plugins == ["p1"] class TestBackfillStrategyCategories: def test_creates_default_missing_categories_without_fetching_daemon(self): - p1, session = _patched_session() + session = _patched_session() tool_strategy = SimpleNamespace( category=TenantPluginAutoUpgradeCategory.TOOL, strategy_setting=TenantPluginAutoUpgradeStrategySetting.FIX_ONLY, @@ -193,10 +199,10 @@ class TestBackfillStrategyCategories: session.scalars.return_value.all.return_value = [tool_strategy] installer = MagicMock() - with p1, patch(f"{MODULE}.PluginInstaller", return_value=installer): + with patch(f"{MODULE}.PluginInstaller", return_value=installer): from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - result = PluginAutoUpgradeService.backfill_strategy_categories("t1") + result = PluginAutoUpgradeService.backfill_strategy_categories("t1", session=session) expected_time = PluginAutoUpgradeService.default_upgrade_time_of_day("t1") assert result.created_count == len(TenantPluginAutoUpgradeCategory) - 1 @@ -219,7 +225,7 @@ class TestBackfillStrategyCategories: assert 0 <= default_time < 24 * 60 * 60 def test_creates_missing_categories_and_splits_known_plugins(self, caplog: pytest.LogCaptureFixture): - p1, session = _patched_session() + session = _patched_session() tool_strategy = SimpleNamespace( category=TenantPluginAutoUpgradeCategory.TOOL, strategy_setting=TenantPluginAutoUpgradeStrategySetting.FIX_ONLY, @@ -252,13 +258,12 @@ class TestBackfillStrategyCategories: installer.list_plugins.return_value = installed_plugins with ( - p1, patch(f"{MODULE}.PluginInstaller", return_value=installer), caplog.at_level(logging.WARNING, logger=MODULE), ): from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - result = PluginAutoUpgradeService.backfill_strategy_categories("t1") + result = PluginAutoUpgradeService.backfill_strategy_categories("t1", session=session) assert result.created_count == len(TenantPluginAutoUpgradeCategory) - 2 assert result.normalized is True diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_built_in_retrieval.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_built_in_retrieval.py index 5bc41fdb5bd..a7bb8cfeed6 100644 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_built_in_retrieval.py +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_built_in_retrieval.py @@ -22,8 +22,9 @@ def test_get_pipeline_templates(mocker: MockerFixture) -> None: }, ) retrieval = BuiltInPipelineTemplateRetrieval() + session = mocker.Mock() - templates = retrieval.get_pipeline_templates(mocker.Mock(), "en-US") + templates = retrieval.get_pipeline_templates("en-US", session=session) assert templates == {"pipeline_templates": [{"id": "tpl-1"}]} @@ -39,8 +40,9 @@ def test_get_pipeline_template_detail(mocker: MockerFixture) -> None: }, ) retrieval = BuiltInPipelineTemplateRetrieval() + session = mocker.Mock() - detail = retrieval.get_pipeline_template_detail(mocker.Mock(), "tpl-1") + detail = retrieval.get_pipeline_template_detail("tpl-1", session=session) assert detail == {"id": "tpl-1", "name": "Template 1"} @@ -52,8 +54,9 @@ def test_get_pipeline_templates_missing_language_returns_empty_dict(mocker: Mock return_value={"pipeline_templates": {}}, ) retrieval = BuiltInPipelineTemplateRetrieval() + session = mocker.Mock() - result = retrieval.get_pipeline_templates(mocker.Mock(), "fr-FR") + result = retrieval.get_pipeline_templates("fr-FR", session=session) assert result == {} @@ -65,8 +68,9 @@ def test_get_pipeline_template_detail_returns_none_for_unknown_id(mocker: Mocker return_value={"pipeline_templates": {"tpl-1": {"id": "tpl-1"}}}, ) retrieval = BuiltInPipelineTemplateRetrieval() + session = mocker.Mock() - result = retrieval.get_pipeline_template_detail(mocker.Mock(), "nonexistent-id") + result = retrieval.get_pipeline_template_detail("nonexistent-id", session=session) assert result is None diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_customized_retrieval.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_customized_retrieval.py index b3ef79961d3..b3befeb41fd 100644 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_customized_retrieval.py +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_customized_retrieval.py @@ -21,7 +21,7 @@ def test_get_pipeline_templates(mocker: MockerFixture) -> None: session_mock.scalars.return_value = scalars_mock retrieval = CustomizedPipelineTemplateRetrieval() - result = retrieval.get_pipeline_templates(session_mock, "en-US", "tenant-id") + result = retrieval.get_pipeline_templates("en-US", "tenant-id", session=session_mock) assert retrieval.get_type() == PipelineTemplateType.CUSTOMIZED assert result == { @@ -51,7 +51,7 @@ def test_get_pipeline_template_detail_returns_detail(mocker: MockerFixture) -> N ) retrieval = CustomizedPipelineTemplateRetrieval() - detail = retrieval.get_pipeline_template_detail(session_mock, "tpl-1") + detail = retrieval.get_pipeline_template_detail("tpl-1", session=session_mock) assert detail == { "id": "tpl-1", @@ -70,6 +70,6 @@ def test_get_pipeline_template_detail_returns_none_when_not_found(mocker: Mocker session_mock.get.return_value = None retrieval = CustomizedPipelineTemplateRetrieval() - result = retrieval.get_pipeline_template_detail(session_mock, "missing") + result = retrieval.get_pipeline_template_detail("missing", session=session_mock) assert result is None diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_database_retrieval.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_database_retrieval.py index cae79175b1c..48ae26ce3aa 100644 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_database_retrieval.py +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_database_retrieval.py @@ -23,7 +23,7 @@ def test_get_pipeline_templates(mocker: MockerFixture) -> None: session_mock.scalars.return_value = scalars_mock retrieval = DatabasePipelineTemplateRetrieval() - result = retrieval.get_pipeline_templates(session_mock, "en-US") + result = retrieval.get_pipeline_templates("en-US", session=session_mock) assert retrieval.get_type() == PipelineTemplateType.DATABASE assert result == { @@ -54,7 +54,7 @@ def test_get_pipeline_template_detail_returns_detail(mocker: MockerFixture) -> N ) retrieval = DatabasePipelineTemplateRetrieval() - detail = retrieval.get_pipeline_template_detail(session_mock, "tpl-1") + detail = retrieval.get_pipeline_template_detail("tpl-1", session=session_mock) assert detail == { "id": "tpl-1", @@ -72,6 +72,6 @@ def test_get_pipeline_template_detail_returns_none_when_not_found(mocker: Mocker session_mock.get.return_value = None retrieval = DatabasePipelineTemplateRetrieval() - result = retrieval.get_pipeline_template_detail(session_mock, "missing") + result = retrieval.get_pipeline_template_detail("missing", session=session_mock) assert result is None diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_pipeline_template_base.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_pipeline_template_base.py index c8af1869732..17cd5db7ab3 100644 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_pipeline_template_base.py +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_pipeline_template_base.py @@ -4,11 +4,11 @@ from services.rag_pipeline.pipeline_template.pipeline_template_base import Pipel class DummyRetrieval(PipelineTemplateRetrievalBase): - def get_pipeline_templates(self, session: Mock, language: str, current_tenant_id: str | None = None) -> dict: - del session, current_tenant_id + def get_pipeline_templates(self, language: str, *, session) -> dict: + del session return {"language": language} - def get_pipeline_template_detail(self, session: Mock, template_id: str) -> dict | None: + def get_pipeline_template_detail(self, template_id: str, *, session) -> dict | None: del session return {"id": template_id} @@ -20,6 +20,6 @@ def test_pipeline_template_retrieval_base_concrete_implementation() -> None: retrieval = DummyRetrieval() session = Mock() - assert retrieval.get_pipeline_templates(session, "en-US") == {"language": "en-US"} - assert retrieval.get_pipeline_template_detail(session, "tpl-1") == {"id": "tpl-1"} + assert retrieval.get_pipeline_templates("en-US", session=session) == {"language": "en-US"} + assert retrieval.get_pipeline_template_detail("tpl-1", session=session) == {"id": "tpl-1"} assert retrieval.get_type() == "dummy" diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_remote_retrieval.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_remote_retrieval.py index 8f55b4b1c2f..78e46d272c2 100644 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_remote_retrieval.py +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_remote_retrieval.py @@ -20,12 +20,12 @@ def test_get_pipeline_templates_fallbacks_to_database_on_error(mocker: MockerFix retrieval = RemotePipelineTemplateRetrieval() session = mocker.Mock() - result = retrieval.get_pipeline_templates(session, "en-US") + result = retrieval.get_pipeline_templates("en-US", session=session) assert retrieval.get_type() == PipelineTemplateType.REMOTE assert result == {"pipeline_templates": [{"id": "db-1"}]} fetch_mock.assert_called_once_with("en-US") - fallback_mock.assert_called_once_with(session, "en-US") + fallback_mock.assert_called_once_with("en-US", session=session) def test_get_pipeline_template_detail_fallbacks_to_database_on_error(mocker: MockerFixture) -> None: @@ -42,11 +42,11 @@ def test_get_pipeline_template_detail_fallbacks_to_database_on_error(mocker: Moc retrieval = RemotePipelineTemplateRetrieval() session = mocker.Mock() - result = retrieval.get_pipeline_template_detail(session, "tpl-1") + result = retrieval.get_pipeline_template_detail("tpl-1", session=session) assert result == {"id": "db-1"} fetch_mock.assert_called_once_with("tpl-1") - fallback_mock.assert_called_once_with(session, "tpl-1") + fallback_mock.assert_called_once_with("tpl-1", session=session) def test_fetch_pipeline_templates_from_dify_official(mocker: MockerFixture) -> None: diff --git a/api/tests/unit_tests/services/rag_pipeline/test_pipeline_generate_service.py b/api/tests/unit_tests/services/rag_pipeline/test_pipeline_generate_service.py index 0ae2ba97f1a..c8992585653 100644 --- a/api/tests/unit_tests/services/rag_pipeline/test_pipeline_generate_service.py +++ b/api/tests/unit_tests/services/rag_pipeline/test_pipeline_generate_service.py @@ -1,30 +1,15 @@ -from collections.abc import Iterator from types import SimpleNamespace from typing import cast -from uuid import uuid4 import pytest from pytest_mock import MockerFixture -from sqlalchemy import create_engine, func, select -from sqlalchemy.orm import Session, sessionmaker from core.app.entities.app_invoke_entities import InvokeFrom -from models.dataset import Document, Pipeline -from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus +from models.dataset import Pipeline from models.model import Account, App, EndUser from services.rag_pipeline.pipeline_generate_service import PipelineGenerateService -@pytest.fixture -def document_session() -> Iterator[Session]: - engine = create_engine("sqlite:///:memory:") - Document.__table__.create(engine) - session_factory = sessionmaker(bind=engine, expire_on_commit=False) - with session_factory() as session: - yield session - engine.dispose() - - def test_get_max_active_requests_uses_smallest_non_zero_limit(mocker: MockerFixture) -> None: mocker.patch("services.rag_pipeline.pipeline_generate_service.dify_config.APP_DEFAULT_ACTIVE_REQUESTS", 5) mocker.patch("services.rag_pipeline.pipeline_generate_service.dify_config.APP_MAX_ACTIVE_REQUESTS", 3) @@ -62,12 +47,13 @@ def test_get_workflow(mocker: MockerFixture, invoke_from, workflow, expected_err rag_pipeline_service.get_published_workflow.return_value = workflow pipeline = cast(Pipeline, SimpleNamespace(id="pipeline-1")) + session = mocker.Mock() if expected_error: with pytest.raises(ValueError, match=expected_error): - PipelineGenerateService._get_workflow(pipeline, invoke_from) + PipelineGenerateService._get_workflow(pipeline, invoke_from, session) else: - result = PipelineGenerateService._get_workflow(pipeline, invoke_from) + result = PipelineGenerateService._get_workflow(pipeline, invoke_from, session) assert result == workflow @@ -75,10 +61,10 @@ def test_generate_updates_document_status_and_returns_event_stream(mocker: Mocke pipeline = cast(Pipeline, SimpleNamespace(id="pipeline-1")) user = cast(Account | EndUser, SimpleNamespace(id="user-1")) args = {"original_document_id": "doc-1", "query": "hello"} + session_mock = mocker.Mock() mocker.patch.object(PipelineGenerateService, "_get_workflow", return_value=SimpleNamespace(id="wf-1")) update_status_mock = mocker.patch.object(PipelineGenerateService, "update_document_status") - session = mocker.Mock() generator_cls = mocker.patch("services.rag_pipeline.pipeline_generate_service.PipelineGenerator") generator_instance = generator_cls.return_value @@ -86,49 +72,39 @@ def test_generate_updates_document_status_and_returns_event_stream(mocker: Mocke generator_cls.convert_to_event_stream.return_value = "stream-events" result = PipelineGenerateService.generate( - session=session, pipeline=pipeline, user=user, args=args, invoke_from=InvokeFrom.WEB_APP, streaming=True, + session=session_mock, ) assert result == "stream-events" - update_status_mock.assert_called_once_with("doc-1", session) + update_status_mock.assert_called_once_with("doc-1", session=session_mock) -def test_update_document_status_updates_existing_document(document_session: Session) -> None: - session = document_session - document_id = str(uuid4()) - document = Document( - id=document_id, - tenant_id=str(uuid4()), - dataset_id=str(uuid4()), - position=1, - data_source_type=DataSourceType.UPLOAD_FILE, - batch="batch-1", - name="Doc", - created_from=DocumentCreatedFrom.WEB, - created_by=str(uuid4()), - indexing_status=IndexingStatus.COMPLETED, - ) - session.add(document) - session.commit() +def test_update_document_status_updates_existing_document(mocker: MockerFixture) -> None: + document = SimpleNamespace(indexing_status="completed") - PipelineGenerateService.update_document_status(document_id, session) + session_mock = mocker.Mock() + session_mock.get.return_value = document + add_mock = session_mock.add - updated_document = session.get(Document, document_id) - assert updated_document is not None - assert updated_document.indexing_status == IndexingStatus.WAITING + PipelineGenerateService.update_document_status("doc-1", session=session_mock) + + assert document.indexing_status == "waiting" + add_mock.assert_called_once_with(document) -def test_update_document_status_skips_when_document_missing(document_session: Session) -> None: - session = document_session +def test_update_document_status_skips_when_document_missing(mocker: MockerFixture) -> None: + session_mock = mocker.Mock() + session_mock.get.return_value = None + add_mock = session_mock.add - PipelineGenerateService.update_document_status(str(uuid4()), session) + PipelineGenerateService.update_document_status("missing", session=session_mock) - assert session.scalar(select(func.count()).select_from(Document)) == 0 + add_mock.assert_not_called() # --- generate_single_iteration --- @@ -144,8 +120,9 @@ def test_generate_single_iteration_delegates(mocker: MockerFixture) -> None: pipeline = cast(Pipeline, SimpleNamespace(id="p1")) user = cast(Account, SimpleNamespace(id="u1")) + session = mocker.Mock() - result = PipelineGenerateService.generate_single_iteration(pipeline, user, "node-1", {"key": "val"}) + result = PipelineGenerateService.generate_single_iteration(pipeline, user, "node-1", {"key": "val"}, session) assert result == "stream-iter" generator_instance.single_iteration_generate.assert_called_once() @@ -164,8 +141,9 @@ def test_generate_single_loop_delegates(mocker: MockerFixture) -> None: pipeline = cast(Pipeline, SimpleNamespace(id="p1")) user = cast(Account, SimpleNamespace(id="u1")) + session = mocker.Mock() - result = PipelineGenerateService.generate_single_loop(pipeline, user, "node-1", {"key": "val"}) + result = PipelineGenerateService.generate_single_loop(pipeline, user, "node-1", {"key": "val"}, session) assert result == "stream-loop" generator_instance.single_loop_generate.assert_called_once() diff --git a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_service.py b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_service.py index 37141c97c83..0d74b3abf9c 100644 --- a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_service.py +++ b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_service.py @@ -49,9 +49,8 @@ def rag_pipeline_service(mocker: MockerFixture) -> RagPipelineServiceTestContext ) session = mocker.Mock() session_maker = _make_mock_session_maker(mocker, session) - mocker.patch("services.rag_pipeline.rag_pipeline.session_factory.get_session_maker", return_value=session_maker) mocker.patch("services.rag_pipeline.rag_pipeline.db", SimpleNamespace(engine=mocker.Mock())) - service = RagPipelineService(session_maker=session_maker) + service = RagPipelineService(session=session, session_maker=session_maker) return RagPipelineServiceTestContext(service=service, session=session, session_maker=session_maker) @@ -156,10 +155,10 @@ def test_get_pipeline_templates_fallbacks_to_builtin_for_non_english_empty_resul builtin_retrieval.fetch_pipeline_templates_from_builtin.return_value = {"pipeline_templates": [{"id": "builtin-1"}]} factory_mock.get_built_in_pipeline_template_retrieval.return_value = builtin_retrieval - result = RagPipelineService.get_pipeline_templates(session, type="built-in", language="ja-JP") + result = RagPipelineService.get_pipeline_templates(type="built-in", language="ja-JP", session=session) assert result == {"pipeline_templates": [{"id": "builtin-1"}]} - remote_retrieval.get_pipeline_templates.assert_called_once_with(session, "ja-JP", None) + remote_retrieval.get_pipeline_templates.assert_called_once_with("ja-JP", None, session=session) builtin_retrieval.fetch_pipeline_templates_from_builtin.assert_called_once_with("en-US") @@ -171,11 +170,11 @@ def test_get_pipeline_templates_customized_mode_uses_customized_factory(mocker: factory_mock = mocker.patch("services.rag_pipeline.rag_pipeline.PipelineTemplateRetrievalFactory") factory_mock.get_pipeline_template_factory.return_value.return_value = retrieval - result = RagPipelineService.get_pipeline_templates(session, type="customized", language="en-US") + result = RagPipelineService.get_pipeline_templates(type="customized", language="en-US", session=session) assert result == {"pipeline_templates": [{"id": "custom-1"}]} factory_mock.get_pipeline_template_factory.assert_called_with("customized") - retrieval.get_pipeline_templates.assert_called_once_with(session, "en-US", None) + retrieval.get_pipeline_templates.assert_called_once_with("en-US", None, session=session) @pytest.mark.parametrize("template_type", ["built-in", "customized"]) @@ -188,12 +187,12 @@ def test_get_pipeline_template_detail_uses_expected_mode(mocker: MockerFixture, factory_mock = mocker.patch("services.rag_pipeline.rag_pipeline.PipelineTemplateRetrievalFactory") factory_mock.get_pipeline_template_factory.return_value.return_value = retrieval - result = RagPipelineService.get_pipeline_template_detail(session, "tpl-1", type=template_type) + result = RagPipelineService.get_pipeline_template_detail("tpl-1", type=template_type, session=session) assert result == {"id": "tpl-1"} expected_mode = "remote" if template_type == "built-in" else "customized" factory_mock.get_pipeline_template_factory.assert_called_with(expected_mode) - retrieval.get_pipeline_template_detail.assert_called_once_with(session, "tpl-1") + retrieval.get_pipeline_template_detail.assert_called_once_with("tpl-1", session=session) def test_get_published_workflow_returns_none_when_pipeline_has_no_workflow_id( @@ -845,13 +844,14 @@ def test_publish_customized_pipeline_template_success( # 2. Run test args = {"name": "New Template", "description": "Desc", "icon_info": {"icon": "star"}, "tags": ["tag1"]} - rag_pipeline_service.service.publish_customized_pipeline_template("p1", args, account, "t1") + rag_pipeline_service.service.publish_customized_pipeline_template("p1", args, account, "t1", session=session) # 3. Assertions # Verify a new template was added to session or similar? # Since we can't easily check the session inside the context manager with Mock, # we just check that no error was raised and DSL was exported. - mock_dsl_service.export_rag_pipeline_dsl.assert_called_once() + pipeline.retrieve_dataset.assert_called_once_with(session=session) + mock_dsl_service.export_rag_pipeline_dsl.assert_called_once_with(pipeline=pipeline, include_secret=True) # --- get_datasource_plugins --- @@ -863,7 +863,7 @@ def test_get_datasource_plugins_success( # 1. Setup mocks dataset = _make_dataset() - pipeline = _make_pipeline() + pipeline = _make_pipeline(workflow_id="wf-1") workflow = mocker.Mock() workflow.graph_dict = { @@ -996,7 +996,7 @@ def test_set_datasource_variables_success( # --- Utility Methods --- -def test_get_draft_workflow_success(mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext) -> None: +def test_get_draft_workflow_success(rag_pipeline_service: RagPipelineServiceTestContext) -> None: # 1. Setup mocks pipeline = _make_pipeline() @@ -1011,9 +1011,7 @@ def test_get_draft_workflow_success(mocker: MockerFixture, rag_pipeline_service: assert result == workflow -def test_get_published_workflow_success( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext -) -> None: +def test_get_published_workflow_success(rag_pipeline_service: RagPipelineServiceTestContext) -> None: # 1. Setup mocks pipeline = _make_pipeline(workflow_id="wf-pub") @@ -1406,10 +1404,7 @@ def test_get_node_last_run_delegates_to_repository( ) -> None: repo = mocker.Mock() repo.get_node_last_execution.return_value = "node-exec" - mocker.patch( - "services.rag_pipeline.rag_pipeline.DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository", - return_value=repo, - ) + rag_pipeline_service.service._node_execution_service_repo = repo pipeline = _make_pipeline() workflow = _make_workflow(workflow_id="wf1") @@ -1785,21 +1780,25 @@ def test_run_datasource_node_preview_raises_for_unsupported_provider( def test_publish_customized_pipeline_template_raises_for_missing_pipeline( mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - rag_pipeline_service.session.get.return_value = None + session = mocker.Mock() + session.get.return_value = None with pytest.raises(ValueError, match="Pipeline not found"): - rag_pipeline_service.service.publish_customized_pipeline_template("p1", {}, _make_account(), "t1") + rag_pipeline_service.service.publish_customized_pipeline_template( + "p1", {}, _make_account(), "t1", session=session + ) def test_publish_customized_pipeline_template_raises_for_missing_workflow_id( mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: pipeline = _make_pipeline(workflow_id=None) - rag_pipeline_service.session.get.return_value = pipeline + session = mocker.Mock() + session.get.return_value = pipeline with pytest.raises(ValueError, match="Pipeline workflow not found"): rag_pipeline_service.service.publish_customized_pipeline_template( - "p1", {"name": "template-name"}, _make_account(), "t1" + "p1", {"name": "template-name"}, _make_account(), "t1", session=session ) @@ -1824,10 +1823,8 @@ def test_get_pipeline_raises_when_pipeline_missing( def test_init_uses_default_sessionmaker_when_none(mocker: MockerFixture) -> None: default_session_maker = mocker.Mock() - mocker.patch( - "services.rag_pipeline.rag_pipeline.session_factory.get_session_maker", - return_value=default_session_maker, - ) + mocker.patch("services.rag_pipeline.rag_pipeline.sessionmaker", return_value=default_session_maker) + mocker.patch("services.rag_pipeline.rag_pipeline.db", SimpleNamespace(engine=mocker.Mock(), session=mocker.Mock())) create_exec_repo = mocker.patch( "services.rag_pipeline.rag_pipeline.DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository" ) @@ -1835,7 +1832,7 @@ def test_init_uses_default_sessionmaker_when_none(mocker: MockerFixture) -> None "services.rag_pipeline.rag_pipeline.DifyAPIRepositoryFactory.create_api_workflow_run_repository" ) - RagPipelineService(session_maker=None) + RagPipelineService(session=mocker.Mock(), session_maker=None) create_exec_repo.assert_called_once_with(default_session_maker) create_run_repo.assert_called_once_with(default_session_maker) @@ -1849,11 +1846,12 @@ def test_get_pipeline_templates_builtin_en_us_no_fallback(mocker: MockerFixture) factory = mocker.patch("services.rag_pipeline.rag_pipeline.PipelineTemplateRetrievalFactory") factory.get_pipeline_template_factory.return_value.return_value = retrieval builtin = factory.get_built_in_pipeline_template_retrieval.return_value + session = mocker.Mock() - result = RagPipelineService.get_pipeline_templates(session, type="built-in", language="en-US") + result = RagPipelineService.get_pipeline_templates(type="built-in", language="en-US", session=session) assert result == {"pipeline_templates": []} - retrieval.get_pipeline_templates.assert_called_once_with(session, "en-US", None) + retrieval.get_pipeline_templates.assert_called_once_with("en-US", None, session=session) builtin.fetch_pipeline_templates_from_builtin.assert_not_called() @@ -1861,14 +1859,14 @@ def test_update_customized_pipeline_template_commits_when_name_empty(mocker: Moc template = _make_customized_template() session = mocker.Mock() session.scalar.return_value = template - session_maker = _make_mock_session_maker(mocker, session) - mocker.patch("services.rag_pipeline.rag_pipeline.session_factory.get_session_maker", return_value=session_maker) info = PipelineTemplateInfoEntity(name="", description="updated", icon_info=IconInfo(icon="i")) - result = RagPipelineService.update_customized_pipeline_template("tpl-1", info, _make_account(), "t1") + result = RagPipelineService.update_customized_pipeline_template( + "tpl-1", info, _make_account(), "t1", session=session + ) assert result.description == "updated" - session_maker.begin.assert_called_once() + session.commit.assert_called_once() def test_get_all_published_workflow_without_filters_has_no_more( @@ -2102,10 +2100,13 @@ def test_publish_customized_pipeline_template_raises_when_workflow_missing( mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: pipeline = _make_pipeline(workflow_id="wf-1") - rag_pipeline_service.session.get.side_effect = [pipeline, None] + session = mocker.Mock() + session.get.side_effect = [pipeline, None] with pytest.raises(ValueError, match="Workflow not found"): - rag_pipeline_service.service.publish_customized_pipeline_template("p1", {}, _make_account(), "t1") + rag_pipeline_service.service.publish_customized_pipeline_template( + "p1", {}, _make_account(), "t1", session=session + ) def test_publish_customized_pipeline_template_raises_when_dataset_missing( @@ -2113,11 +2114,14 @@ def test_publish_customized_pipeline_template_raises_when_dataset_missing( ) -> None: pipeline = _make_pipeline(workflow_id="wf-1") workflow = _make_workflow(workflow_id="wf-1") + session = rag_pipeline_service.session + session.get.side_effect = [pipeline, workflow] pipeline.retrieve_dataset = mocker.Mock(return_value=None) - rag_pipeline_service.session.get.side_effect = [pipeline, workflow] with pytest.raises(ValueError, match="Dataset not found"): - rag_pipeline_service.service.publish_customized_pipeline_template("p1", {}, _make_account(), "t1") + rag_pipeline_service.service.publish_customized_pipeline_template( + "p1", {}, _make_account(), "t1", session=session + ) def test_get_recommended_plugins_skips_manifest_when_missing( @@ -2165,7 +2169,7 @@ def test_get_datasource_plugins_returns_empty_for_non_datasource_nodes( mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: dataset = _make_dataset() - pipeline = _make_pipeline() + pipeline = _make_pipeline(workflow_id="wf-1") workflow = SimpleNamespace( graph_dict={"nodes": [{"id": "n1", "data": {"type": "start"}}]}, rag_pipeline_variables=[] ) @@ -2360,7 +2364,7 @@ def test_get_datasource_plugins_extracts_user_inputs_and_credentials( mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: dataset = _make_dataset() - pipeline = _make_pipeline() + pipeline = _make_pipeline(workflow_id="wf-1") workflow = SimpleNamespace( graph_dict={ "nodes": [ diff --git a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_transform_service.py b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_transform_service.py index cee8e55f8cc..4ee1a5831a0 100644 --- a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_transform_service.py +++ b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_transform_service.py @@ -1,31 +1,16 @@ import logging -from collections.abc import Iterator from datetime import UTC, datetime from types import SimpleNamespace from typing import cast import pytest from pytest_mock import MockerFixture -from sqlalchemy import create_engine, select -from sqlalchemy.orm import Session, sessionmaker -from models.dataset import Dataset, Pipeline -from models.enums import DatasetRuntimeMode +from models.dataset import Dataset from services.entities.knowledge_entities.rag_pipeline_entities import KnowledgeConfiguration from services.rag_pipeline.rag_pipeline_transform_service import RagPipelineTransformService -@pytest.fixture -def pipeline_session() -> Iterator[Session]: - engine = create_engine("sqlite:///:memory:") - Dataset.__table__.create(engine) - Pipeline.__table__.create(engine) - session_factory = sessionmaker(bind=engine, expire_on_commit=False) - with session_factory() as session: - yield session - engine.dispose() - - @pytest.mark.parametrize( ("doc_form", "datasource_type", "indexing_technique"), [ @@ -107,37 +92,47 @@ def test_deal_dependencies_installs_missing_marketplace_plugins(mocker: MockerFi install_mock.assert_called_once_with("tenant-1", ["missing-plugin:1.0.0"]) -def test_transform_to_empty_pipeline_updates_dataset_and_flushes( - mocker: MockerFixture, pipeline_session: Session -) -> None: +def test_transform_to_empty_pipeline_updates_dataset_and_commits(mocker: MockerFixture) -> None: service = RagPipelineTransformService() mocker.patch( "services.rag_pipeline.rag_pipeline_transform_service.current_user", SimpleNamespace(id="user-1"), ) - session = pipeline_session - dataset = Dataset( + class FakePipeline: + def __init__(self, **kwargs): + self.id = "pipeline-1" + self.tenant_id = kwargs["tenant_id"] + self.name = kwargs["name"] + self.description = kwargs["description"] + self.created_by = kwargs["created_by"] + + mocker.patch("services.rag_pipeline.rag_pipeline_transform_service.Pipeline", FakePipeline) + session_mock = mocker.Mock() + add_mock = session_mock.add + flush_mock = session_mock.flush + commit_mock = session_mock.commit + + dataset = SimpleNamespace( + id="dataset-1", tenant_id="tenant-1", name="Dataset", description="desc", - created_by="user-1", + pipeline_id=None, + runtime_mode="general", + updated_by=None, + updated_at=None, ) - session.add(dataset) - session.commit() - flush_spy = mocker.spy(session, "flush") - commit_spy = mocker.spy(session, "commit") - result = service._transform_to_empty_pipeline(dataset, session) + result = service._transform_to_empty_pipeline(cast(Dataset, dataset), session=session_mock) - assert flush_spy.call_count == 2 - commit_spy.assert_not_called() - pipeline = session.scalar(select(Pipeline).where(Pipeline.id == dataset.pipeline_id)) - assert pipeline is not None - assert result == {"pipeline_id": pipeline.id, "dataset_id": dataset.id, "status": "success"} - assert dataset.pipeline_id == pipeline.id - assert dataset.runtime_mode == DatasetRuntimeMode.RAG_PIPELINE + assert result == {"pipeline_id": "pipeline-1", "dataset_id": "dataset-1", "status": "success"} + assert dataset.pipeline_id == "pipeline-1" + assert dataset.runtime_mode == "rag_pipeline" assert dataset.updated_by == "user-1" + add_mock.assert_called() + flush_mock.assert_called_once() + commit_mock.assert_called_once() # --- transform_dataset --- @@ -373,6 +368,7 @@ def test_transform_dataset_full_flow(mocker: MockerFixture) -> None: mocker.patch.object(service, "_deal_dependencies") mocker.patch.object(service, "_deal_document_data") + session_mock.commit = mocker.Mock() # Mock current_user to have the same tenant_id as dataset mock_current_user = SimpleNamespace(current_tenant_id="t1") @@ -386,8 +382,6 @@ def test_transform_dataset_full_flow(mocker: MockerFixture) -> None: assert result["pipeline_id"] == "p-new" assert dataset.runtime_mode == "rag_pipeline" assert dataset.chunk_structure == "text_model" - session_mock.flush.assert_called_once_with() - session_mock.commit.assert_not_called() def test_transform_dataset_raises_for_unsupported_doc_form_after_pipeline_create(mocker: MockerFixture) -> None: @@ -439,11 +433,12 @@ def test_transform_dataset_raises_when_transform_yaml_missing_workflow(mocker: M service.transform_dataset("d1", session_mock) -def test_create_pipeline_raises_when_workflow_data_missing(pipeline_session: Session) -> None: +def test_create_pipeline_raises_when_workflow_data_missing(mocker: MockerFixture) -> None: service = RagPipelineTransformService() + session = mocker.Mock() with pytest.raises(ValueError, match="Missing workflow data for rag pipeline"): - service._create_pipeline({"rag_pipeline": {"name": "N"}}, pipeline_session) + service._create_pipeline({"rag_pipeline": {"name": "N"}}, session=session) def test_deal_document_data_upload_file_with_existing_file(mocker: MockerFixture) -> None: diff --git a/api/tests/unit_tests/services/recommend_app/test_buildin_retrieval.py b/api/tests/unit_tests/services/recommend_app/test_buildin_retrieval.py index c86aaf1db22..cafac0656d1 100644 --- a/api/tests/unit_tests/services/recommend_app/test_buildin_retrieval.py +++ b/api/tests/unit_tests/services/recommend_app/test_buildin_retrieval.py @@ -39,7 +39,7 @@ class TestBuildInRecommendAppRetrieval: return_value={"apps": []}, ) as mock_fetch: retrieval = BuildInRecommendAppRetrieval() - result = retrieval.get_recommended_apps_and_categories("en-US") + result = retrieval.get_recommended_apps_and_categories("en-US", session=MagicMock()) mock_fetch.assert_called_once_with("en-US") assert result == {"apps": []} @@ -47,11 +47,12 @@ class TestBuildInRecommendAppRetrieval: def test_get_learn_dify_apps_delegates_to_database(self, mock_database_retrieval): expected = {"recommended_apps": [{"id": "learn-dify-app"}]} mock_database_retrieval.fetch_learn_dify_apps_from_db.return_value = expected + session = MagicMock() - result = BuildInRecommendAppRetrieval().get_learn_dify_apps("en-US") + result = BuildInRecommendAppRetrieval().get_learn_dify_apps("en-US", session=session) assert result == expected - mock_database_retrieval.fetch_learn_dify_apps_from_db.assert_called_once_with("en-US") + mock_database_retrieval.fetch_learn_dify_apps_from_db.assert_called_once_with("en-US", session=session) def test_get_recommend_app_detail_delegates(self): with patch.object( @@ -60,7 +61,7 @@ class TestBuildInRecommendAppRetrieval: return_value={"id": "app-1"}, ) as mock_fetch: retrieval = BuildInRecommendAppRetrieval() - result = retrieval.get_recommend_app_detail("app-1") + result = retrieval.get_recommend_app_detail("app-1", session=MagicMock()) mock_fetch.assert_called_once_with("app-1") assert result == {"id": "app-1"} diff --git a/api/tests/unit_tests/services/recommend_app/test_remote_retrieval.py b/api/tests/unit_tests/services/recommend_app/test_remote_retrieval.py index 55165deec25..9575aa9f52e 100644 --- a/api/tests/unit_tests/services/recommend_app/test_remote_retrieval.py +++ b/api/tests/unit_tests/services/recommend_app/test_remote_retrieval.py @@ -17,7 +17,7 @@ class TestRemoteRecommendAppRetrieval: return_value={"id": "app-1"}, ) def test_get_recommend_app_detail_success(self, mock_fetch): - result = RemoteRecommendAppRetrieval().get_recommend_app_detail("app-1") + result = RemoteRecommendAppRetrieval().get_recommend_app_detail("app-1", session=MagicMock()) assert result == {"id": "app-1"} mock_fetch.assert_called_once_with("app-1") @@ -32,7 +32,7 @@ class TestRemoteRecommendAppRetrieval: side_effect=ConnectionError("timeout"), ) def test_get_recommend_app_detail_falls_back_on_error(self, mock_fetch, mock_builtin): - result = RemoteRecommendAppRetrieval().get_recommend_app_detail("app-1") + result = RemoteRecommendAppRetrieval().get_recommend_app_detail("app-1", session=MagicMock()) assert result == {"id": "fallback"} mock_builtin.assert_called_once_with("app-1") @@ -42,7 +42,7 @@ class TestRemoteRecommendAppRetrieval: return_value={"recommended_apps": [], "categories": []}, ) def test_get_recommended_apps_success(self, mock_fetch): - result = RemoteRecommendAppRetrieval().get_recommended_apps_and_categories("en-US") + result = RemoteRecommendAppRetrieval().get_recommended_apps_and_categories("en-US", session=MagicMock()) assert result == {"recommended_apps": [], "categories": []} @patch( @@ -56,7 +56,7 @@ class TestRemoteRecommendAppRetrieval: side_effect=ValueError("server error"), ) def test_get_recommended_apps_falls_back_on_error(self, mock_fetch, mock_builtin): - result = RemoteRecommendAppRetrieval().get_recommended_apps_and_categories("en-US") + result = RemoteRecommendAppRetrieval().get_recommended_apps_and_categories("en-US", session=MagicMock()) assert result == {"recommended_apps": [{"id": "builtin"}]} @patch.object( @@ -65,7 +65,7 @@ class TestRemoteRecommendAppRetrieval: return_value={"recommended_apps": [{"id": "learn-dify-app"}]}, ) def test_get_learn_dify_apps_success(self, mock_fetch): - result = RemoteRecommendAppRetrieval().get_learn_dify_apps("en-US") + result = RemoteRecommendAppRetrieval().get_learn_dify_apps("en-US", session=MagicMock()) assert result == {"recommended_apps": [{"id": "learn-dify-app"}]} mock_fetch.assert_called_once_with("en-US") @@ -80,10 +80,12 @@ class TestRemoteRecommendAppRetrieval: side_effect=ValueError("server error"), ) def test_get_learn_dify_apps_falls_back_to_database_on_error(self, mock_fetch, mock_database): - result = RemoteRecommendAppRetrieval().get_learn_dify_apps("en-US") + session = MagicMock() + + result = RemoteRecommendAppRetrieval().get_learn_dify_apps("en-US", session=session) assert result == {"recommended_apps": [{"id": "db-fallback"}]} - mock_database.assert_called_once_with("en-US") + mock_database.assert_called_once_with("en-US", session=session) class TestFetchFromDifyOfficial: diff --git a/api/tests/unit_tests/services/test_account_service.py b/api/tests/unit_tests/services/test_account_service.py index 233191ca0b7..b73fa112003 100644 --- a/api/tests/unit_tests/services/test_account_service.py +++ b/api/tests/unit_tests/services/test_account_service.py @@ -681,7 +681,7 @@ class TestTenantService: mock_session = MagicMock() mock_session.execute.return_value.scalar_one_or_none.return_value = TenantAccountRole.ADMIN - role = TenantService.get_account_role_in_tenant(mock_session, "account-1", "tenant-1") + role = TenantService.get_account_role_in_tenant("account-1", "tenant-1", session=mock_session) assert role == TenantAccountRole.ADMIN @@ -690,7 +690,7 @@ class TestTenantService: mock_session = MagicMock() mock_session.execute.return_value.scalar_one_or_none.return_value = None - role = TenantService.get_account_role_in_tenant(mock_session, "account-1", "tenant-1") + role = TenantService.get_account_role_in_tenant("account-1", "tenant-1", session=mock_session) assert role is None @@ -699,7 +699,7 @@ class TestTenantService: without ever touching the session.""" mock_session = MagicMock() - assert TenantService.get_account_role_in_tenant(mock_session, None, "tenant-1") is None + assert TenantService.get_account_role_in_tenant(None, "tenant-1", session=mock_session) is None mock_session.execute.assert_not_called() def test_get_account_role_in_tenant_query_is_scoped(self): @@ -711,7 +711,7 @@ class TestTenantService: mock_session = MagicMock() mock_session.execute.return_value.scalar_one_or_none.return_value = TenantAccountRole.NORMAL - TenantService.get_account_role_in_tenant(mock_session, account_id, tenant_id) + TenantService.get_account_role_in_tenant(account_id, tenant_id, session=mock_session) stmt = mock_session.execute.call_args.args[0] compiled = str(stmt.compile(compile_kwargs={"literal_binds": True})) @@ -760,11 +760,7 @@ class TestTenantService: mock_tenant_instance.name = "Test User's Workspace" mock_tenant_class.return_value = mock_tenant_instance - # Mock the db import in CreditPoolService to avoid database connection - with patch("services.credit_pool_service.db") as mock_credit_pool_db: - mock_credit_pool_db.session.add = MagicMock() - mock_credit_pool_db.session.commit = MagicMock() - + with patch("services.credit_pool_service.CreditPoolService.create_default_pool"): # Execute test TenantService.create_owner_tenant_if_not_exist( mock_account, session=mock_db_dependencies["db"].session @@ -1052,6 +1048,7 @@ class TestTenantService: account_id="user-rbac", member_account_id="user-rbac", role_ids=["rbac-owner-id"], + session=mock_db_dependencies["db"].session, ) def test_admin_can_update_admin_member_role(self): @@ -1191,7 +1188,11 @@ class TestTenantService: with pytest.raises(NoPermissionError): TenantService.check_member_permission( - mock_tenant, mock_operator, mock_member, "remove", session=MagicMock() + mock_tenant, + mock_operator, + mock_member, + "remove", + session=mock_db_dependencies["db"].session, ) def test_rbac_member_can_remove_non_owner_member(self): @@ -1265,7 +1266,9 @@ class TestTenantService: ), patch("services.account_service.RBACService.Roles", mock_rbac_roles), ): - owner_account_id = AccountService.get_rbac_workspace_owner_account_id("tenant-1", "acct-1") + owner_account_id = AccountService.get_rbac_workspace_owner_account_id( + "tenant-1", "acct-1", session=MagicMock() + ) assert owner_account_id == "owner-account" call = mock_rbac_roles.members.call_args @@ -1912,7 +1915,9 @@ class TestRegisterService: is_setup=True, session=mock_db_dependencies["db"].session, ) - mock_lookup.assert_called_once_with(mock_db_dependencies["db"].session, "newuser@example.com") + mock_lookup.assert_called_once_with( + "newuser@example.com", session=mock_db_dependencies["db"].session + ) def test_invite_new_member_normalizes_new_account_email( self, mock_db_dependencies, mock_redis_dependencies, mock_task_dependencies @@ -1958,7 +1963,7 @@ class TestRegisterService: is_setup=True, session=mock_db_dependencies["db"].session, ) - mock_lookup.assert_called_once_with(mock_db_dependencies["db"].session, mixed_email) + mock_lookup.assert_called_once_with(mixed_email, session=mock_db_dependencies["db"].session) mock_check_permission.assert_called_once_with( mock_tenant, mock_inviter, @@ -2025,7 +2030,7 @@ class TestRegisterService: mock_tenant, mock_existing_account, "normal", requires_setup=True ) mock_task_dependencies.delay.assert_called_once() - mock_lookup.assert_called_once_with(mock_db_dependencies["db"].session, "existing@example.com") + mock_lookup.assert_called_once_with("existing@example.com", session=mock_db_dependencies["db"].session) def test_invite_existing_active_account_requires_acceptance_before_joining( self, mock_db_dependencies, mock_redis_dependencies, mock_task_dependencies @@ -2171,6 +2176,7 @@ class TestRegisterService: account_id=mock_inviter.id, member_account_id=mock_new_account.id, role_ids=["rbac-role-id-123"], + session=mock_db_dependencies["db"].session, ) def test_invite_new_member_rbac_enabled_existing_account( @@ -2220,6 +2226,7 @@ class TestRegisterService: account_id=mock_inviter.id, member_account_id=mock_existing_account.id, role_ids=["rbac-role-id-456"], + session=mock_db_dependencies["db"].session, ) def test_invite_new_member_rbac_enabled_existing_active_account_adds_role_before_signin_response( @@ -2268,6 +2275,7 @@ class TestRegisterService: account_id=mock_inviter.id, member_account_id=mock_existing_account.id, role_ids=["rbac-role-id-456"], + session=mock_db_dependencies["db"].session, ) mock_task_dependencies.delay.assert_not_called() @@ -2615,7 +2623,7 @@ class TestSessionInjectedGetters: sentinel_account = MagicMock(spec=Account) mock_session.get.return_value = sentinel_account - result = AccountService.get_account_by_id(mock_session, "user-123") + result = AccountService.get_account_by_id("user-123", session=mock_session) assert result is sentinel_account mock_session.get.assert_called_once_with(Account, "user-123") @@ -2625,7 +2633,7 @@ class TestSessionInjectedGetters: mock_session = MagicMock() mock_session.get.return_value = None - assert AccountService.get_account_by_id(mock_session, "missing") is None + assert AccountService.get_account_by_id("missing", session=mock_session) is None @pytest.mark.parametrize("sqlite_session", [(Account,)], indirect=True) def test_get_account_by_email_returns_scalar_or_none(self, sqlite_session: Session): @@ -2637,9 +2645,9 @@ class TestSessionInjectedGetters: sqlite_session.add(account) sqlite_session.commit() - assert AccountService.get_account_by_email(sqlite_session, "alice@example.com") == account - assert AccountService.get_account_by_email(sqlite_session, "ALICE@example.com") is None - assert AccountService.get_account_by_email(sqlite_session, "ghost@example.com") is None + assert AccountService.get_account_by_email("alice@example.com", session=sqlite_session) == account + assert AccountService.get_account_by_email("ALICE@example.com", session=sqlite_session) is None + assert AccountService.get_account_by_email("ghost@example.com", session=sqlite_session) is None def test_account_belongs_to_tenant_short_circuits_on_falsy_account_id(self): """SSO bearers with no ``account_id`` (and any other falsy id) @@ -2648,22 +2656,22 @@ class TestSessionInjectedGetters: """ mock_session = MagicMock() - assert TenantService.account_belongs_to_tenant(mock_session, None, "tenant-1") is False - assert TenantService.account_belongs_to_tenant(mock_session, "", "tenant-1") is False + assert TenantService.account_belongs_to_tenant(None, "tenant-1", session=mock_session) is False + assert TenantService.account_belongs_to_tenant("", "tenant-1", session=mock_session) is False mock_session.execute.assert_not_called() def test_account_belongs_to_tenant_true_when_join_row_exists(self): mock_session = MagicMock() mock_session.execute.return_value.scalar_one_or_none.return_value = "join-id" - assert TenantService.account_belongs_to_tenant(mock_session, "user-1", "tenant-1") is True + assert TenantService.account_belongs_to_tenant("user-1", "tenant-1", session=mock_session) is True mock_session.execute.assert_called_once() def test_account_belongs_to_tenant_false_when_no_join(self): mock_session = MagicMock() mock_session.execute.return_value.scalar_one_or_none.return_value = None - assert TenantService.account_belongs_to_tenant(mock_session, "user-1", "tenant-1") is False + assert TenantService.account_belongs_to_tenant("user-1", "tenant-1", session=mock_session) is False def test_get_account_memberships_returns_join_tenant_pairs(self): """Returns whatever ``session.query(...).join(...).filter(...).all()`` @@ -2674,7 +2682,7 @@ class TestSessionInjectedGetters: rows = [(MagicMock(), MagicMock()), (MagicMock(), MagicMock())] mock_session.query.return_value.join.return_value.filter.return_value.all.return_value = rows - out = TenantService.get_account_memberships(mock_session, "user-123") + out = TenantService.get_account_memberships("user-123", session=mock_session) assert out == rows # No fall-through to the global db.session proxy. @@ -2688,7 +2696,7 @@ class TestSessionInjectedGetters: rows = [(MagicMock(), MagicMock())] mock_session.execute.return_value.all.return_value = rows - out = TenantService.get_workspaces_for_account(mock_session, "user-123") + out = TenantService.get_workspaces_for_account("user-123", session=mock_session) assert out == rows assert mock_session.execute.called @@ -2704,20 +2712,20 @@ class TestSessionInjectedGetters: sentinel = MagicMock(spec=Tenant) mock_session.get.return_value = sentinel - assert TenantService.get_tenant_by_id(mock_session, "tenant-1") is sentinel + assert TenantService.get_tenant_by_id("tenant-1", session=mock_session) is sentinel mock_session.get.assert_called_once_with(Tenant, "tenant-1") def test_get_tenant_by_id_returns_none_when_missing(self): mock_session = MagicMock() mock_session.get.return_value = None - assert TenantService.get_tenant_by_id(mock_session, "missing") is None + assert TenantService.get_tenant_by_id("missing", session=mock_session) is None def test_get_tenants_by_ids_short_circuits_on_empty_input(self): """Empty id list must not emit ``WHERE id IN ()``.""" mock_session = MagicMock() - assert TenantService.get_tenants_by_ids(mock_session, []) == [] + assert TenantService.get_tenants_by_ids([], session=mock_session) == [] mock_session.execute.assert_not_called() def test_get_tenants_by_ids_returns_scalars(self): @@ -2725,7 +2733,7 @@ class TestSessionInjectedGetters: tenants = [MagicMock(), MagicMock()] mock_session.execute.return_value.scalars.return_value.all.return_value = tenants - assert TenantService.get_tenants_by_ids(mock_session, ["t1", "t2"]) == tenants + assert TenantService.get_tenants_by_ids(["t1", "t2"], session=mock_session) == tenants mock_session.execute.assert_called_once() def test_get_tenant_name_returns_scalar_or_none(self): @@ -2736,10 +2744,10 @@ class TestSessionInjectedGetters: mock_session = MagicMock() mock_session.execute.return_value.scalar_one_or_none.return_value = "Acme Inc." - assert TenantService.get_tenant_name(mock_session, "tenant-1") == "Acme Inc." + assert TenantService.get_tenant_name("tenant-1", session=mock_session) == "Acme Inc." mock_session.execute.return_value.scalar_one_or_none.return_value = None - assert TenantService.get_tenant_name(mock_session, "missing") is None + assert TenantService.get_tenant_name("missing", session=mock_session) is None def test_find_workspace_for_account_returns_first_row_or_none(self): """Per-id read returns ``session.execute(...).first()`` directly; @@ -2750,7 +2758,7 @@ class TestSessionInjectedGetters: sentinel_row = (MagicMock(), MagicMock()) mock_session.execute.return_value.first.return_value = sentinel_row - assert TenantService.find_workspace_for_account(mock_session, "user-123", "ws-1") is sentinel_row + assert TenantService.find_workspace_for_account("user-123", "ws-1", session=mock_session) is sentinel_row mock_session.execute.return_value.first.return_value = None - assert TenantService.find_workspace_for_account(mock_session, "user-123", "ws-1") is None + assert TenantService.find_workspace_for_account("user-123", "ws-1", session=mock_session) is None diff --git a/api/tests/unit_tests/services/test_agent_app_sandbox_service.py b/api/tests/unit_tests/services/test_agent_app_sandbox_service.py index a9ed82413bb..c36980f5829 100644 --- a/api/tests/unit_tests/services/test_agent_app_sandbox_service.py +++ b/api/tests/unit_tests/services/test_agent_app_sandbox_service.py @@ -258,6 +258,7 @@ def test_workflow_sandbox_service_resolves_locator_and_returns_download_url( node_id="node-1", node_execution_id="node-exec-1", path="report.txt", + session=session_factory.create_session(), ) assert result.url == "https://files.example/report.txt?token=1&as_attachment=true" @@ -376,6 +377,7 @@ def test_workflow_sandbox_service_filters_by_node_execution_id() -> None: node_id="node-1", node_execution_id="node-exec-2", path="out.txt", + session=session_factory.create_session(), ) assert result.text == "hello" @@ -409,6 +411,7 @@ def test_workflow_sandbox_service_uses_latest_active_session_when_execution_id_o node_id="node-1", node_execution_id=None, path=".", + session=session_factory.create_session(), ) assert result.path == "." @@ -428,6 +431,7 @@ def test_workflow_sandbox_service_raises_when_no_active_session() -> None: node_id="node-1", node_execution_id=None, path=".", + session=session_factory.create_session(), ) assert exc_info.value.code == "no_active_session" @@ -447,6 +451,7 @@ def test_workflow_sandbox_service_raises_when_runtime_specs_missing() -> None: node_id="node-1", node_execution_id=None, path=".", + session=session_factory.create_session(), ) assert exc_info.value.code == "no_sandbox" diff --git a/api/tests/unit_tests/services/test_agent_drive_service.py b/api/tests/unit_tests/services/test_agent_drive_service.py index ad72e142522..6371b197325 100644 --- a/api/tests/unit_tests/services/test_agent_drive_service.py +++ b/api/tests/unit_tests/services/test_agent_drive_service.py @@ -125,6 +125,7 @@ def _commit(key: str, tool_file_id: str, *, owned: bool = True): value_owned_by_drive=owned, ) ], + session=session_factory.create_session(), ) @@ -132,15 +133,25 @@ def test_commit_then_manifest_lists_the_entry(): tf = _seed_tool_file() _commit("data/report.txt", tf) - items = AgentDriveService().manifest(tenant_id=TENANT, agent_id=AGENT) + items = AgentDriveService().manifest(tenant_id=TENANT, agent_id=AGENT, session=session_factory.create_session()) assert [i["key"] for i in items] == ["data/report.txt"] assert items[0]["file_kind"] == "tool_file" assert items[0]["file_id"] == tf assert items[0]["mime_type"] == "text/plain" # prefix filter - assert AgentDriveService().manifest(tenant_id=TENANT, agent_id=AGENT, prefix="data/") != [] - assert AgentDriveService().manifest(tenant_id=TENANT, agent_id=AGENT, prefix="other/") == [] + assert ( + AgentDriveService().manifest( + tenant_id=TENANT, agent_id=AGENT, prefix="data/", session=session_factory.create_session() + ) + != [] + ) + assert ( + AgentDriveService().manifest( + tenant_id=TENANT, agent_id=AGENT, prefix="other/", session=session_factory.create_session() + ) + == [] + ) def test_commit_skill_row_persists_metadata_and_lists_catalog() -> None: @@ -157,6 +168,7 @@ def test_commit_skill_row_persists_metadata_and_lists_catalog() -> None: skill_metadata=DriveSkillMetadata(name="Tender Analyzer", description="Parses RFPs."), ) ], + session=session_factory.create_session(), ) with session_factory.create_session() as session: @@ -165,7 +177,7 @@ def test_commit_skill_row_persists_metadata_and_lists_catalog() -> None: assert row.is_skill is True assert row.skill_metadata == '{"description":"Parses RFPs.","name":"Tender Analyzer"}' - skills = AgentDriveService().list_skills(tenant_id=TENANT, agent_id=AGENT) + skills = AgentDriveService().list_skills(tenant_id=TENANT, agent_id=AGENT, session=session_factory.create_session()) assert len(skills) == 1 assert skills[0]["path"] == "tender-analyzer" assert skills[0]["skill_md_key"] == "tender-analyzer/SKILL.md" @@ -191,6 +203,7 @@ def test_commit_rejects_skill_row_without_skill_metadata() -> None: is_skill=True, ) ], + session=session_factory.create_session(), ) assert exc_info.value.code == "invalid_skill_metadata" @@ -220,7 +233,7 @@ def test_list_skills_raises_controlled_error_for_invalid_stored_metadata(raw_met session.commit() with pytest.raises(AgentDriveError) as exc_info: - AgentDriveService().list_skills(tenant_id=TENANT, agent_id=AGENT) + AgentDriveService().list_skills(tenant_id=TENANT, agent_id=AGENT, session=session_factory.create_session()) assert exc_info.value.code == "invalid_skill_metadata" @@ -239,6 +252,7 @@ def test_commit_rejects_non_skill_row_with_skill_metadata() -> None: skill_metadata=DriveSkillMetadata(name="Bad", description=""), ) ], + session=session_factory.create_session(), ) @@ -257,6 +271,7 @@ def test_commit_rejects_non_canonical_skill_key() -> None: skill_metadata=DriveSkillMetadata(name="Tender Analyzer", description=""), ) ], + session=session_factory.create_session(), ) @@ -282,6 +297,7 @@ def test_commit_rejects_agent_from_another_tenant(): value_owned_by_drive=True, ) ], + session=session_factory.create_session(), ) assert exc_info.value.status_code == 404 assert exc_info.value.code == "agent_not_found" @@ -311,24 +327,27 @@ def test_batch_failure_does_not_delete_old_storage_before_commit(): _commit("doc.txt", tf1, owned=True) with patch("services.agent_drive_service.storage") as storage_mock: - with pytest.raises(AgentDriveError): - AgentDriveService().commit( - tenant_id=TENANT, - user_id=USER, - agent_id=AGENT, - items=[ - DriveCommitItem( - key="doc.txt", - file_ref={"kind": "tool_file", "id": tf2}, - value_owned_by_drive=True, - ), - DriveCommitItem( - key="bad.txt", - file_ref={"kind": "tool_file", "id": "44444444-4444-4444-4444-444444444444"}, - value_owned_by_drive=True, - ), - ], - ) + with session_factory.create_session() as session: + with pytest.raises(AgentDriveError): + AgentDriveService().commit( + tenant_id=TENANT, + user_id=USER, + agent_id=AGENT, + items=[ + DriveCommitItem( + key="doc.txt", + file_ref={"kind": "tool_file", "id": tf2}, + value_owned_by_drive=True, + ), + DriveCommitItem( + key="bad.txt", + file_ref={"kind": "tool_file", "id": "44444444-4444-4444-4444-444444444444"}, + value_owned_by_drive=True, + ), + ], + session=session, + ) + session.rollback() storage_mock.delete.assert_not_called() with session_factory.create_session() as session: @@ -389,6 +408,7 @@ def test_recommit_same_skill_value_updates_metadata_without_cleaning_backing_fil skill_metadata=DriveSkillMetadata(name="Tender Analyzer", description="v1"), ) ], + session=session_factory.create_session(), ) with patch("services.agent_drive_service.storage") as storage_mock: @@ -405,6 +425,7 @@ def test_recommit_same_skill_value_updates_metadata_without_cleaning_backing_fil skill_metadata=DriveSkillMetadata(name="Tender Analyzer v2", description="v2"), ) ], + session=session_factory.create_session(), ) storage_mock.delete.assert_not_called() @@ -449,6 +470,7 @@ def _commit_upload(key: str, upload_file_id: str, *, owned: bool = True): value_owned_by_drive=owned, ) ], + session=session_factory.create_session(), ) @@ -456,7 +478,7 @@ def test_commit_upload_file_source_and_manifest(): uf = _seed_upload_file() _commit_upload("docs/u.txt", uf) - items = AgentDriveService().manifest(tenant_id=TENANT, agent_id=AGENT) + items = AgentDriveService().manifest(tenant_id=TENANT, agent_id=AGENT, session=session_factory.create_session()) assert items[0]["file_kind"] == "upload_file" assert items[0]["file_id"] == uf assert items[0]["mime_type"] == "text/plain" @@ -492,7 +514,9 @@ def test_manifest_includes_internal_download_url(): patch("core.app.workflow.file_runtime.DifyWorkflowFileRuntime") as runtime_cls, ): runtime_cls.return_value.resolve_file_url.return_value = "http://internal/files/x?sign=1" - items = AgentDriveService().manifest(tenant_id=TENANT, agent_id=AGENT, include_download_url=True) + items = AgentDriveService().manifest( + tenant_id=TENANT, agent_id=AGENT, include_download_url=True, session=session_factory.create_session() + ) assert items[0]["download_url"] == "http://internal/files/x?sign=1" # drive-owned resolution: internal URL (for_external=False) @@ -507,7 +531,9 @@ def test_manifest_download_url_none_when_unresolvable(): "services.agent_drive_service.file_factory.build_from_mapping", side_effect=ValueError("not found"), ): - items = AgentDriveService().manifest(tenant_id=TENANT, agent_id=AGENT, include_download_url=True) + items = AgentDriveService().manifest( + tenant_id=TENANT, agent_id=AGENT, include_download_url=True, session=session_factory.create_session() + ) assert items[0]["download_url"] is None @@ -524,6 +550,7 @@ def test_delete_by_key_cleans_drive_owned_value(): user_id=USER, agent_id=AGENT, items=[DriveCommitItem(key="files/doomed.txt", file_ref=None)], + session=session_factory.create_session(), ) storage_mock.delete.assert_called_once() @@ -560,6 +587,7 @@ def test_commit_null_batch_removes_multiple_skill_keys(): DriveCommitItem(key="tender-analyzer/SKILL.md", file_ref=None), DriveCommitItem(key="tender-analyzer/.DIFY-SKILL-FULL.zip", file_ref=None), ], + session=session_factory.create_session(), ) assert sorted(item["key"] for item in removed) == [ @@ -581,6 +609,7 @@ def test_commit_null_is_idempotent_for_missing_keys(): user_id=USER, agent_id=AGENT, items=[DriveCommitItem(key="files/never-there.txt", file_ref=None)], + session=session_factory.create_session(), ) assert removed == [{"key": "files/never-there.txt", "removed": True, "noop": True}] @@ -595,6 +624,7 @@ def test_commit_null_keeps_shared_value_records(): user_id=USER, agent_id=AGENT, items=[DriveCommitItem(key="files/shared.txt", file_ref=None)], + session=session_factory.create_session(), ) storage_mock.delete.assert_not_called() @@ -639,7 +669,9 @@ def test_preview_returns_text_with_truncation_flags(): with patch("services.agent_drive_service.storage") as storage_mock: storage_mock.load_stream.return_value = iter([b"# PDF Toolkit\nUse responsibly.\n"]) - result = AgentDriveService().preview(tenant_id=TENANT, agent_id=AGENT, key="pdf-toolkit/SKILL.md") + result = AgentDriveService().preview( + tenant_id=TENANT, agent_id=AGENT, key="pdf-toolkit/SKILL.md", session=session_factory.create_session() + ) assert result == { "key": "pdf-toolkit/SKILL.md", @@ -656,13 +688,17 @@ def test_preview_marks_binary_and_oversized_content(): with patch("services.agent_drive_service.storage") as storage_mock: storage_mock.load_stream.return_value = iter([b"\x00\x01\x02"]) - binary = AgentDriveService().preview(tenant_id=TENANT, agent_id=AGENT, key="files/blob.bin") + binary = AgentDriveService().preview( + tenant_id=TENANT, agent_id=AGENT, key="files/blob.bin", session=session_factory.create_session() + ) assert binary["binary"] is True assert binary["text"] is None with patch("services.agent_drive_service.storage") as storage_mock: storage_mock.load_stream.return_value = iter([b"x" * (AgentDriveService.PREVIEW_MAX_BYTES + 10)]) - oversized = AgentDriveService().preview(tenant_id=TENANT, agent_id=AGENT, key="files/blob.bin") + oversized = AgentDriveService().preview( + tenant_id=TENANT, agent_id=AGENT, key="files/blob.bin", session=session_factory.create_session() + ) assert oversized["truncated"] is True assert oversized["binary"] is False assert len(oversized["text"]) == AgentDriveService.PREVIEW_MAX_BYTES @@ -670,7 +706,9 @@ def test_preview_marks_binary_and_oversized_content(): def test_preview_unknown_key_is_404(): with pytest.raises(AgentDriveError) as exc_info: - AgentDriveService().preview(tenant_id=TENANT, agent_id=AGENT, key="ghost/SKILL.md") + AgentDriveService().preview( + tenant_id=TENANT, agent_id=AGENT, key="ghost/SKILL.md", session=session_factory.create_session() + ) assert exc_info.value.code == "drive_key_not_found" assert exc_info.value.status_code == 404 @@ -678,7 +716,10 @@ def test_preview_unknown_key_is_404(): def test_preview_rejects_cross_tenant_agent(): with pytest.raises(AgentDriveError) as exc_info: AgentDriveService().preview( - tenant_id="99999999-9999-9999-9999-999999999999", agent_id=AGENT, key="pdf-toolkit/SKILL.md" + tenant_id="99999999-9999-9999-9999-999999999999", + agent_id=AGENT, + key="pdf-toolkit/SKILL.md", + session=session_factory.create_session(), ) assert exc_info.value.code == "agent_not_found" @@ -688,7 +729,12 @@ def test_download_url_signs_external_audience(): _commit("pdf-toolkit/.DIFY-SKILL-FULL.zip", tf) with patch.object(AgentDriveService, "_resolve_download_url", return_value="https://signed.example/x") as resolver: - url = AgentDriveService().download_url(tenant_id=TENANT, agent_id=AGENT, key="pdf-toolkit/.DIFY-SKILL-FULL.zip") + url = AgentDriveService().download_url( + tenant_id=TENANT, + agent_id=AGENT, + key="pdf-toolkit/.DIFY-SKILL-FULL.zip", + session=session_factory.create_session(), + ) assert url == "https://signed.example/x" # console downloads are for browsers: external signing, never the internal URL @@ -702,7 +748,9 @@ def test_upload_file_download_url_uses_attachment_filename(): with patch("core.app.workflow.file_runtime.DifyWorkflowFileRuntime") as runtime_cls: runtime_cls.return_value.resolve_upload_file_url.return_value = "https://files.example/report.pdf" - url = AgentDriveService().download_url(tenant_id=TENANT, agent_id=AGENT, key="files/report.pdf") + url = AgentDriveService().download_url( + tenant_id=TENANT, agent_id=AGENT, key="files/report.pdf", session=session_factory.create_session() + ) assert url == "https://files.example/report.pdf" assert runtime_cls.return_value.resolve_upload_file_url.call_args.kwargs["for_external"] is True @@ -712,7 +760,7 @@ def test_upload_file_download_url_uses_attachment_filename(): def test_manifest_items_carry_created_at_for_inspector(): tf = _seed_tool_file() _commit("files/x.txt", tf) - items = AgentDriveService().manifest(tenant_id=TENANT, agent_id=AGENT) + items = AgentDriveService().manifest(tenant_id=TENANT, agent_id=AGENT, session=session_factory.create_session()) assert items[0]["created_at"] is None or isinstance(items[0]["created_at"], int) @@ -744,13 +792,14 @@ def _commit_skill(*, manifest_files: list[str] | None = None) -> None: value_owned_by_drive=True, ), ], + session=session_factory.create_session(), ) def test_list_skills_uses_canonical_skill_rows(): _commit_skill(manifest_files=["SKILL.md", "scripts/run.py"]) - skills = AgentDriveService().list_skills(tenant_id=TENANT, agent_id=AGENT) + skills = AgentDriveService().list_skills(tenant_id=TENANT, agent_id=AGENT, session=session_factory.create_session()) created_at = skills[0].pop("created_at") assert skills == [ @@ -773,7 +822,9 @@ def test_inspect_skill_returns_manifest_files_and_file_tree(): with patch("services.agent_drive_service.storage") as storage_mock: storage_mock.load_stream.return_value = iter([b"# PDF Toolkit\n"]) - result = AgentDriveService().inspect_skill(tenant_id=TENANT, agent_id=AGENT, skill_path="pdf-toolkit") + result = AgentDriveService().inspect_skill( + tenant_id=TENANT, agent_id=AGENT, skill_path="pdf-toolkit", session=session_factory.create_session() + ) assert result["source"] == "skill_md" assert result["warnings"] == [] @@ -792,7 +843,9 @@ def test_inspect_skill_falls_back_to_drive_keys_when_manifest_missing(): with patch("services.agent_drive_service.storage") as storage_mock: storage_mock.load_stream.return_value = iter([b"# PDF Toolkit\n"]) - result = AgentDriveService().inspect_skill(tenant_id=TENANT, agent_id=AGENT, skill_path="pdf-toolkit") + result = AgentDriveService().inspect_skill( + tenant_id=TENANT, agent_id=AGENT, skill_path="pdf-toolkit", session=session_factory.create_session() + ) assert result["warnings"] == ["manifest_files_unavailable"] assert [file["path"] for file in result["files"]] == ["SKILL.md"] @@ -808,6 +861,7 @@ def test_preview_skill_archive_member_from_manifest_without_drive_row(): tenant_id=TENANT, agent_id=AGENT, key="pdf-toolkit/references/guide.md", + session=session_factory.create_session(), ) assert result == { @@ -831,6 +885,7 @@ def test_download_url_signs_skill_archive_member_from_manifest_without_drive_row tenant_id=TENANT, agent_id=AGENT, key="pdf-toolkit/references/guide.md", + session=session_factory.create_session(), ) assert url == "https://signed.example/member" @@ -856,5 +911,6 @@ def test_skill_metadata_rejects_non_canonical_rows(): skill_metadata=DriveSkillMetadata(name="Bad"), ) ], + session=session_factory.create_session(), ) assert exc_info.value.code == "invalid_skill_key" diff --git a/api/tests/unit_tests/services/test_agent_tool_inner_service.py b/api/tests/unit_tests/services/test_agent_tool_inner_service.py index 93e20d267da..61049d29e9e 100644 --- a/api/tests/unit_tests/services/test_agent_tool_inner_service.py +++ b/api/tests/unit_tests/services/test_agent_tool_inner_service.py @@ -70,7 +70,7 @@ def test_invoke_uses_agent_tool_runtime_and_returns_observation() -> None: side_effect=lambda messages, **_kwargs: messages, ), ): - response = AgentToolInnerService().invoke(session, _request()) + response = AgentToolInnerService().invoke(_request(), session=session) assert response.observation == "ok" assert response.metadata == { @@ -89,7 +89,7 @@ def test_invoke_raises_app_not_found_when_session_has_no_app() -> None: session.get.return_value = None with pytest.raises(AgentToolInnerServiceError) as exc_info: - AgentToolInnerService().invoke(session, _request()) + AgentToolInnerService().invoke(_request(), session=session) assert exc_info.value.error_code == "app_not_found" assert exc_info.value.status_code == 404 @@ -102,7 +102,7 @@ def test_invoke_raises_app_tenant_mismatch_when_app_belongs_to_other_tenant() -> session.get.return_value = fake_app with pytest.raises(AgentToolInnerServiceError) as exc_info: - AgentToolInnerService().invoke(session, _request()) + AgentToolInnerService().invoke(_request(), session=session) assert exc_info.value.error_code == "app_tenant_mismatch" assert exc_info.value.status_code == 403 @@ -120,7 +120,7 @@ def test_invoke_maps_tool_runtime_app_not_found_value_error_to_specific_error_co patch("services.agent_tool_inner_service.ToolEngine.generic_invoke", side_effect=ValueError("app not found")), ): with pytest.raises(AgentToolInnerServiceError) as exc_info: - AgentToolInnerService().invoke(session, _request()) + AgentToolInnerService().invoke(_request(), session=session) assert exc_info.value.error_code == "app_not_found" assert exc_info.value.status_code == 404 @@ -141,7 +141,7 @@ def test_invoke_maps_tool_invoke_error_without_private_tool_engine_helper() -> N ), ): with pytest.raises(AgentToolInnerServiceError) as exc_info: - AgentToolInnerService().invoke(session, _request()) + AgentToolInnerService().invoke(_request(), session=session) assert exc_info.value.error_code == "agent_tool_invoke_failed" @@ -161,6 +161,6 @@ def test_invoke_maps_runtime_lookup_errors_to_service_error_codes(error: Excepti with patch("services.agent_tool_inner_service.ToolManager.get_agent_tool_runtime", side_effect=error): with pytest.raises(AgentToolInnerServiceError) as exc_info: - AgentToolInnerService().invoke(session, _request()) + AgentToolInnerService().invoke(_request(), session=session) assert exc_info.value.error_code == expected_code diff --git a/api/tests/unit_tests/services/test_annotation_service.py b/api/tests/unit_tests/services/test_annotation_service.py index 2975c4df14c..79bbb5873ac 100644 --- a/api/tests/unit_tests/services/test_annotation_service.py +++ b/api/tests/unit_tests/services/test_annotation_service.py @@ -103,7 +103,7 @@ class TestAppAnnotationServiceUpInsert: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.up_insert_app_annotation_from_message(args, "app-1") + AppAnnotationService.up_insert_app_annotation_from_message(args, "app-1", session=mock_db.session) def test_up_insert_app_annotation_from_message_should_raise_value_error_when_answer_missing(self) -> None: """Test missing answer and content raises ValueError.""" @@ -121,7 +121,7 @@ class TestAppAnnotationServiceUpInsert: # Act & Assert with pytest.raises(ValueError): - AppAnnotationService.up_insert_app_annotation_from_message(args, app.id) + AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, session=mock_db.session) def test_up_insert_app_annotation_from_message_should_raise_not_found_when_message_missing(self) -> None: """Test missing message raises NotFound.""" @@ -139,7 +139,7 @@ class TestAppAnnotationServiceUpInsert: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.up_insert_app_annotation_from_message(args, app.id) + AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, session=mock_db.session) def test_up_insert_app_annotation_from_message_should_update_existing_annotation_when_found(self) -> None: """Test existing annotation is updated and indexed.""" @@ -161,7 +161,7 @@ class TestAppAnnotationServiceUpInsert: mock_db.session.scalar.side_effect = [app, message, setting] # Act - result = AppAnnotationService.up_insert_app_annotation_from_message(args, app.id) + result = AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, session=mock_db.session) # Assert assert result == annotation @@ -199,7 +199,7 @@ class TestAppAnnotationServiceUpInsert: mock_db.session.scalar.side_effect = [app, message, None] # Act - result = AppAnnotationService.up_insert_app_annotation_from_message(args, app.id) + result = AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, session=mock_db.session) # Assert assert result == annotation_instance @@ -231,7 +231,7 @@ class TestAppAnnotationServiceUpInsert: # Act & Assert with pytest.raises(ValueError): - AppAnnotationService.up_insert_app_annotation_from_message(args, app.id) + AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, session=mock_db.session) def test_up_insert_app_annotation_from_message_should_create_annotation_when_message_missing(self) -> None: """Test annotation is created when message_id is not provided.""" @@ -252,7 +252,7 @@ class TestAppAnnotationServiceUpInsert: mock_db.session.scalar.side_effect = [app, setting] # Act - result = AppAnnotationService.up_insert_app_annotation_from_message(args, app.id) + result = AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, session=mock_db.session) # Assert assert result == annotation_instance @@ -383,7 +383,7 @@ class TestAppAnnotationServiceListAndExport: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.get_annotation_list_by_app_id("app-1", 1, 10, "") + AppAnnotationService.get_annotation_list_by_app_id("app-1", 1, 10, "", session=mock_db.session) def test_get_annotation_list_by_app_id_should_return_items_with_keyword(self) -> None: """Test keyword search returns items and total.""" @@ -402,7 +402,9 @@ class TestAppAnnotationServiceListAndExport: mock_paginate.return_value = pagination # Act - items, total = AppAnnotationService.get_annotation_list_by_app_id(app.id, 1, 10, "keyword") + items, total = AppAnnotationService.get_annotation_list_by_app_id( + app.id, 1, 10, "keyword", session=mock_db.session + ) # Assert assert items == ["a1"] @@ -424,7 +426,9 @@ class TestAppAnnotationServiceListAndExport: mock_paginate.return_value = pagination # Act - items, total = AppAnnotationService.get_annotation_list_by_app_id(app.id, 1, 10, "") + items, total = AppAnnotationService.get_annotation_list_by_app_id( + app.id, 1, 10, "", session=mock_db.session + ) # Assert assert items == ["a1", "a2"] @@ -451,7 +455,7 @@ class TestAppAnnotationServiceListAndExport: mock_db.session.scalars.return_value.all.return_value = [annotation1, annotation2] # Act - result = AppAnnotationService.export_annotation_list_by_app_id(app.id) + result = AppAnnotationService.export_annotation_list_by_app_id(app.id, session=mock_db.session) # Assert assert result == [annotation1, annotation2] @@ -473,7 +477,7 @@ class TestAppAnnotationServiceListAndExport: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.export_annotation_list_by_app_id("app-1") + AppAnnotationService.export_annotation_list_by_app_id("app-1", session=mock_db.session) class TestAppAnnotationServiceDirectManipulation: @@ -493,7 +497,7 @@ class TestAppAnnotationServiceDirectManipulation: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.insert_app_annotation_directly(args, "app-1") + AppAnnotationService.insert_app_annotation_directly(args, "app-1", session=mock_db.session) def test_insert_app_annotation_directly_should_raise_value_error_when_question_missing(self) -> None: """Test missing question raises ValueError.""" @@ -510,7 +514,7 @@ class TestAppAnnotationServiceDirectManipulation: # Act & Assert with pytest.raises(ValueError): - AppAnnotationService.insert_app_annotation_directly(args, app.id) + AppAnnotationService.insert_app_annotation_directly(args, app.id, session=mock_db.session) def test_insert_app_annotation_directly_should_create_annotation_and_index(self) -> None: """Test insert creates annotation and triggers index task.""" @@ -531,7 +535,7 @@ class TestAppAnnotationServiceDirectManipulation: mock_db.session.scalar.side_effect = [app, setting] # Act - result = AppAnnotationService.insert_app_annotation_directly(args, app.id) + result = AppAnnotationService.insert_app_annotation_directly(args, app.id, session=mock_db.session) # Assert assert result == annotation_instance @@ -696,7 +700,9 @@ class TestAppAnnotationServiceDirectManipulation: mock_db.session.execute.return_value.all.return_value = [] # Act - result = AppAnnotationService.delete_app_annotations_in_batch(_make_app_ref(app), ["ann-1"]) + result = AppAnnotationService.delete_app_annotations_in_batch( + _make_app_ref(app), ["ann-1"], session=mock_db.session + ) # Assert assert result == {"deleted_count": 0} @@ -723,7 +729,9 @@ class TestAppAnnotationServiceDirectManipulation: mock_db.session.execute.side_effect = [execute_result_multi, MagicMock(), execute_result_delete] # Act - result = AppAnnotationService.delete_app_annotations_in_batch(_make_app_ref(app), ["ann-1", "ann-2"]) + result = AppAnnotationService.delete_app_annotations_in_batch( + _make_app_ref(app), ["ann-1", "ann-2"], session=mock_db.session + ) # Assert assert result == {"deleted_count": 2} @@ -755,7 +763,7 @@ class TestAppAnnotationServiceBatchImport: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.batch_import_app_annotations("app-1", file) + AppAnnotationService.batch_import_app_annotations("app-1", file, session=mock_db.session) def test_batch_import_app_annotations_should_return_error_when_columns_invalid(self) -> None: """Test invalid column count returns error message.""" @@ -777,7 +785,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -801,7 +809,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -829,7 +837,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -855,7 +863,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -885,7 +893,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -911,7 +919,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -937,7 +945,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -963,7 +971,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -994,7 +1002,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -1027,7 +1035,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert assert result == {"job_id": "uuid-3", "job_status": "waiting", "record_count": 1} @@ -1067,7 +1075,7 @@ class TestAppAnnotationServiceBatchImport: # Act with caplog.at_level(logging.DEBUG): - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert assert result["error_msg"] == "An error occurred while processing the file: boom" @@ -1090,7 +1098,9 @@ class TestAppAnnotationServiceHitHistoryAndSettings: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.get_annotation_hit_histories(_make_annotation_ref(app, "ann-1"), 1, 10) + AppAnnotationService.get_annotation_hit_histories( + _make_annotation_ref(app, "ann-1"), 1, 10, session=mock_db.session + ) def test_get_annotation_hit_histories_should_return_items_and_total(self) -> None: """Test hit histories pagination returns items and total.""" @@ -1114,6 +1124,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: _make_annotation_ref(app, annotation.id), 1, 10, + session=mock_db.session, ) # Assert @@ -1129,7 +1140,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: mock_db.session.get.return_value = None # Act - result = AppAnnotationService.get_annotation_by_id("ann-1") + result = AppAnnotationService.get_annotation_by_id("ann-1", session=mock_db.session) # Assert assert result is None @@ -1142,7 +1153,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: mock_db.session.get.return_value = annotation # Act - result = AppAnnotationService.get_annotation_by_id("ann-1") + result = AppAnnotationService.get_annotation_by_id("ann-1", session=mock_db.session) # Assert assert result == annotation @@ -1165,6 +1176,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: message_id="msg-1", from_source="chat", score=0.8, + session=mock_db.session, ) # Assert @@ -1187,7 +1199,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: mock_db.session.scalar.side_effect = [app, setting] # Act - result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id) + result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id, session=mock_db.session) # Assert assert result["enabled"] is True @@ -1208,7 +1220,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.get_app_annotation_setting_by_app_id("app-1") + AppAnnotationService.get_app_annotation_setting_by_app_id("app-1", session=mock_db.session) def test_get_app_annotation_setting_by_app_id_should_return_empty_embedding_model_when_no_detail(self) -> None: """Test setting without detail returns empty embedding model.""" @@ -1224,7 +1236,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: mock_db.session.scalar.side_effect = [app, setting] # Act - result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id) + result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id, session=mock_db.session) # Assert assert result["enabled"] is True @@ -1243,7 +1255,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: mock_db.session.scalar.side_effect = [app, None] # Act - result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id) + result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id, session=mock_db.session) # Assert assert result == {"enabled": False} @@ -1265,7 +1277,9 @@ class TestAppAnnotationServiceHitHistoryAndSettings: mock_db.session.scalar.side_effect = [app, setting] # Act - result = AppAnnotationService.update_app_annotation_setting(app.id, setting.id, args) + result = AppAnnotationService.update_app_annotation_setting( + app.id, setting.id, args, session=mock_db.session + ) # Assert assert result["enabled"] is True @@ -1292,7 +1306,9 @@ class TestAppAnnotationServiceHitHistoryAndSettings: mock_db.session.scalar.side_effect = [app, setting] # Act - result = AppAnnotationService.update_app_annotation_setting(app.id, setting.id, args) + result = AppAnnotationService.update_app_annotation_setting( + app.id, setting.id, args, session=mock_db.session + ) # Assert assert result["enabled"] is True @@ -1312,7 +1328,9 @@ class TestAppAnnotationServiceHitHistoryAndSettings: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.update_app_annotation_setting("app-1", "setting-1", {"score_threshold": 0.5}) + AppAnnotationService.update_app_annotation_setting( + "app-1", "setting-1", {"score_threshold": 0.5}, session=mock_db.session + ) def test_update_app_annotation_setting_should_raise_not_found_when_setting_missing(self) -> None: """Test update raises NotFound when setting is missing.""" @@ -1328,7 +1346,9 @@ class TestAppAnnotationServiceHitHistoryAndSettings: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.update_app_annotation_setting(app.id, "setting-1", {"score_threshold": 0.5}) + AppAnnotationService.update_app_annotation_setting( + app.id, "setting-1", {"score_threshold": 0.5}, session=mock_db.session + ) class TestAppAnnotationServiceClearAll: @@ -1361,7 +1381,7 @@ class TestAppAnnotationServiceClearAll: mock_db.session.scalars.side_effect = [annotations_scalars, histories_scalars_1, histories_scalars_2] # Act - result = AppAnnotationService.clear_all_annotations(app.id) + result = AppAnnotationService.clear_all_annotations(app.id, session=mock_db.session) # Assert assert result == {"result": "success"} @@ -1385,4 +1405,4 @@ class TestAppAnnotationServiceClearAll: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.clear_all_annotations("app-1") + AppAnnotationService.clear_all_annotations("app-1", session=mock_db.session) diff --git a/api/tests/unit_tests/services/test_app_generate_service.py b/api/tests/unit_tests/services/test_app_generate_service.py index 865410fddf1..22c7514a522 100644 --- a/api/tests/unit_tests/services/test_app_generate_service.py +++ b/api/tests/unit_tests/services/test_app_generate_service.py @@ -235,12 +235,12 @@ class TestGenerate: side_effect=lambda x: x, ) result = AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) assert result == {"result": "ok"} gen_spy.assert_called_once() @@ -256,12 +256,12 @@ class TestGenerate: side_effect=lambda x: x, ) result = AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.AGENT_CHAT), user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) assert result == {"result": "agent"} gen_spy.assert_called_once() @@ -278,12 +278,12 @@ class TestGenerate: ) app = _make_app(AppMode.CHAT, is_agent=True) result = AppGenerateService.generate( - MagicMock(), app_model=app, user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) assert result == {"result": "agent-via-flag"} gen_spy.assert_called_once() @@ -300,12 +300,12 @@ class TestGenerate: ) app = _make_app(AppMode.CHAT, is_agent=False) result = AppGenerateService.generate( - MagicMock(), app_model=app, user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) assert result == {"result": "chat"} gen_spy.assert_called_once() @@ -342,12 +342,12 @@ class TestGenerate: ) result = AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.ADVANCED_CHAT), user=_make_user(), args={"workflow_id": None, "query": "hi", "inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) assert result == {"result": "advanced-blocking"} call_kwargs = gen_spy.call_args.kwargs @@ -375,12 +375,12 @@ class TestGenerate: ) result = AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.ADVANCED_CHAT), user=_make_user(), args={"workflow_id": None, "query": "hi", "inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=MagicMock(), ) # In streaming mode it should go through retrieve_events, not generate gen_instance.retrieve_events.assert_called_once() @@ -401,12 +401,12 @@ class TestGenerate: ) result = AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.WORKFLOW), user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) assert result == {"result": "workflow-blocking"} call_kwargs = gen_spy.call_args.kwargs @@ -435,12 +435,12 @@ class TestGenerate: ) result = AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.WORKFLOW), user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=MagicMock(), ) retrieve_spy.assert_called_once() # The inner on_subscribe closure was invoked by _build_streaming_task_on_subscribe @@ -451,12 +451,12 @@ class TestGenerate: app = _make_app("invalid-mode", is_agent=False) with pytest.raises(ValueError, match="Invalid app mode"): AppGenerateService.generate( - MagicMock(), app_model=app, user=_make_user(), args={}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) @@ -489,12 +489,12 @@ class TestGenerateBilling: ) AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) reserve_mock.assert_called_once_with(QuotaType.WORKFLOW, "tenant-id") quota_charge.commit.assert_called_once() @@ -513,12 +513,12 @@ class TestGenerateBilling: with pytest.raises(InvokeRateLimitError): AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) def test_exception_refunds_quota_and_exits_rate_limit(self, mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch): @@ -539,12 +539,12 @@ class TestGenerateBilling: with pytest.raises(RuntimeError, match="boom"): AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) quota_charge.refund.assert_called_once() @@ -571,12 +571,12 @@ class TestGenerateBilling: ) AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) # exit is called in finally block for non-streaming assert exit_calls == ["dummy-request-id"] @@ -669,12 +669,12 @@ class TestGenerateBilling: with pytest.raises(RuntimeError, match="boom"): AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) quota_charge.refund.assert_called_once() @@ -701,12 +701,12 @@ class TestGenerateBilling: with pytest.raises(RuntimeError, match="boom"): AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=MagicMock(), ) quota_charge.refund.assert_called_once() @@ -723,7 +723,7 @@ class TestGetWorkflow: ws.get_draft_workflow.return_value = draft_wf mocker.patch("services.app_generate_service.WorkflowService", return_value=ws) - result = AppGenerateService._get_workflow(_make_app(AppMode.WORKFLOW), InvokeFrom.DEBUGGER) + result = AppGenerateService._get_workflow(_make_app(AppMode.WORKFLOW), InvokeFrom.DEBUGGER, session=MagicMock()) assert result is draft_wf ws.get_draft_workflow.assert_called_once() @@ -733,7 +733,7 @@ class TestGetWorkflow: mocker.patch("services.app_generate_service.WorkflowService", return_value=ws) with pytest.raises(ValueError, match="Workflow not initialized"): - AppGenerateService._get_workflow(_make_app(AppMode.WORKFLOW), InvokeFrom.DEBUGGER) + AppGenerateService._get_workflow(_make_app(AppMode.WORKFLOW), InvokeFrom.DEBUGGER, session=MagicMock()) def test_non_debugger_fetches_published(self, mocker: MockerFixture): pub_wf = _make_workflow() @@ -741,7 +741,9 @@ class TestGetWorkflow: ws.get_published_workflow.return_value = pub_wf mocker.patch("services.app_generate_service.WorkflowService", return_value=ws) - result = AppGenerateService._get_workflow(_make_app(AppMode.WORKFLOW), InvokeFrom.SERVICE_API) + result = AppGenerateService._get_workflow( + _make_app(AppMode.WORKFLOW), InvokeFrom.SERVICE_API, session=MagicMock() + ) assert result is pub_wf ws.get_published_workflow.assert_called_once() @@ -751,7 +753,7 @@ class TestGetWorkflow: mocker.patch("services.app_generate_service.WorkflowService", return_value=ws) with pytest.raises(ValueError, match="Workflow not published"): - AppGenerateService._get_workflow(_make_app(AppMode.WORKFLOW), InvokeFrom.SERVICE_API) + AppGenerateService._get_workflow(_make_app(AppMode.WORKFLOW), InvokeFrom.SERVICE_API, session=MagicMock()) def test_specific_workflow_id_valid_uuid(self, mocker: MockerFixture): valid_uuid = str(uuid.uuid4()) @@ -761,7 +763,10 @@ class TestGetWorkflow: mocker.patch("services.app_generate_service.WorkflowService", return_value=ws) result = AppGenerateService._get_workflow( - _make_app(AppMode.WORKFLOW), InvokeFrom.SERVICE_API, workflow_id=valid_uuid + _make_app(AppMode.WORKFLOW), + InvokeFrom.SERVICE_API, + workflow_id=valid_uuid, + session=MagicMock(), ) assert result is specific_wf ws.get_published_workflow_by_id.assert_called_once() @@ -772,7 +777,10 @@ class TestGetWorkflow: with pytest.raises(WorkflowIdFormatError): AppGenerateService._get_workflow( - _make_app(AppMode.WORKFLOW), InvokeFrom.SERVICE_API, workflow_id="not-a-uuid" + _make_app(AppMode.WORKFLOW), + InvokeFrom.SERVICE_API, + workflow_id="not-a-uuid", + session=MagicMock(), ) def test_specific_workflow_id_not_found(self, mocker: MockerFixture): @@ -783,7 +791,10 @@ class TestGetWorkflow: with pytest.raises(WorkflowNotFoundError): AppGenerateService._get_workflow( - _make_app(AppMode.WORKFLOW), InvokeFrom.SERVICE_API, workflow_id=valid_uuid + _make_app(AppMode.WORKFLOW), + InvokeFrom.SERVICE_API, + workflow_id=valid_uuid, + session=MagicMock(), ) @@ -804,7 +815,11 @@ class TestGenerateSingleIteration: ) app = _make_app(AppMode.ADVANCED_CHAT) result = AppGenerateService.generate_single_iteration( - app_model=app, user=_make_user(), node_id="n1", args={"k": "v"} + app_model=app, + user=_make_user(), + node_id="n1", + args={"k": "v"}, + session=MagicMock(), ) iter_spy.assert_called_once() assert result == {"event": "iteration"} @@ -822,7 +837,11 @@ class TestGenerateSingleIteration: ) app = _make_app(AppMode.WORKFLOW) result = AppGenerateService.generate_single_iteration( - app_model=app, user=_make_user(), node_id="n1", args={"k": "v"} + app_model=app, + user=_make_user(), + node_id="n1", + args={"k": "v"}, + session=MagicMock(), ) iter_spy.assert_called_once() assert result == {"event": "wf-iteration"} @@ -830,7 +849,9 @@ class TestGenerateSingleIteration: def test_invalid_mode_raises(self, mocker: MockerFixture): app = _make_app(AppMode.CHAT) with pytest.raises(ValueError, match="Invalid app mode"): - AppGenerateService.generate_single_iteration(app_model=app, user=_make_user(), node_id="n1", args={}) + AppGenerateService.generate_single_iteration( + app_model=app, user=_make_user(), node_id="n1", args={}, session=MagicMock() + ) # --------------------------------------------------------------------------- @@ -850,7 +871,11 @@ class TestGenerateSingleLoop: ) app = _make_app(AppMode.ADVANCED_CHAT) result = AppGenerateService.generate_single_loop( - app_model=app, user=_make_user(), node_id="n1", args=MagicMock() + app_model=app, + user=_make_user(), + node_id="n1", + args=MagicMock(), + session=MagicMock(), ) loop_spy.assert_called_once() assert result == {"event": "loop"} @@ -868,7 +893,11 @@ class TestGenerateSingleLoop: ) app = _make_app(AppMode.WORKFLOW) result = AppGenerateService.generate_single_loop( - app_model=app, user=_make_user(), node_id="n1", args=MagicMock() + app_model=app, + user=_make_user(), + node_id="n1", + args=MagicMock(), + session=MagicMock(), ) loop_spy.assert_called_once() assert result == {"event": "wf-loop"} @@ -876,7 +905,9 @@ class TestGenerateSingleLoop: def test_invalid_mode_raises(self, mocker: MockerFixture): app = _make_app(AppMode.COMPLETION) with pytest.raises(ValueError, match="Invalid app mode"): - AppGenerateService.generate_single_loop(app_model=app, user=_make_user(), node_id="n1", args=MagicMock()) + AppGenerateService.generate_single_loop( + app_model=app, user=_make_user(), node_id="n1", args=MagicMock(), session=MagicMock() + ) # --------------------------------------------------------------------------- @@ -888,16 +919,18 @@ class TestGenerateMoreLikeThis: "services.app_generate_service.CompletionAppGenerator.generate_more_like_this", return_value={"result": "similar"}, ) + session = MagicMock() result = AppGenerateService.generate_more_like_this( - MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), message_id="msg-1", invoke_from=InvokeFrom.SERVICE_API, + session=session, streaming=True, ) assert result == {"result": "similar"} gen_spy.assert_called_once() + assert gen_spy.call_args.kwargs["session"] is session assert gen_spy.call_args.kwargs["stream"] is True diff --git a/api/tests/unit_tests/services/test_app_service.py b/api/tests/unit_tests/services/test_app_service.py index c57fb6ed775..36914679d3e 100644 --- a/api/tests/unit_tests/services/test_app_service.py +++ b/api/tests/unit_tests/services/test_app_service.py @@ -29,14 +29,14 @@ class TestOpenapiVisibilityHelpers: sentinel_app.status = "archived" # explicitly NOT "normal" mock_session.get.return_value = sentinel_app - assert AppService.get_app_by_id(mock_session, "app-uuid") is sentinel_app + assert AppService.get_app_by_id("app-uuid", session=mock_session) is sentinel_app mock_session.get.assert_called_once_with(App, "app-uuid") def test_get_app_by_id_returns_none_when_missing(self): mock_session = MagicMock() mock_session.get.return_value = None - assert AppService.get_app_by_id(mock_session, "missing") is None + assert AppService.get_app_by_id("missing", session=mock_session) is None def test_get_visible_app_by_id_returns_app_when_visible(self): mock_session = MagicMock() @@ -45,7 +45,7 @@ class TestOpenapiVisibilityHelpers: mock_session.get.return_value = app with patch("services.app_service.is_openapi_visible", return_value=True): - assert AppService.get_visible_app_by_id(mock_session, "app-uuid") is app + assert AppService.get_visible_app_by_id("app-uuid", session=mock_session) is app mock_session.get.assert_called_once_with(App, "app-uuid") @@ -53,7 +53,7 @@ class TestOpenapiVisibilityHelpers: mock_session = MagicMock() mock_session.get.return_value = None - assert AppService.get_visible_app_by_id(mock_session, "missing") is None + assert AppService.get_visible_app_by_id("missing", session=mock_session) is None def test_get_visible_app_by_id_returns_none_when_status_not_normal(self): """Soft-deleted/archived rows must not surface on the openapi @@ -65,7 +65,7 @@ class TestOpenapiVisibilityHelpers: mock_session.get.return_value = app with patch("services.app_service.is_openapi_visible", return_value=True): - assert AppService.get_visible_app_by_id(mock_session, "app-uuid") is None + assert AppService.get_visible_app_by_id("app-uuid", session=mock_session) is None def test_get_visible_app_by_id_returns_none_when_visibility_gate_rejects(self): """``is_openapi_visible`` is the per-row counterpart to @@ -78,7 +78,7 @@ class TestOpenapiVisibilityHelpers: mock_session.get.return_value = app with patch("services.app_service.is_openapi_visible", return_value=False): - assert AppService.get_visible_app_by_id(mock_session, "app-uuid") is None + assert AppService.get_visible_app_by_id("app-uuid", session=mock_session) is None def test_find_visible_apps_by_name_returns_scalars_through_visibility_gate(self): """Tenant-scoped name lookup. The helper passes the SELECT through @@ -90,7 +90,7 @@ class TestOpenapiVisibilityHelpers: mock_session.execute.return_value.scalars.return_value = iter(rows) with patch("services.app_service.apply_openapi_gate", side_effect=lambda q: q) as gate: - out = AppService.find_visible_apps_by_name(mock_session, name="my-app", tenant_id="tenant-1") + out = AppService.find_visible_apps_by_name(name="my-app", tenant_id="tenant-1", session=mock_session) assert out == rows # Visibility gate must wrap the SELECT exactly once. @@ -102,7 +102,7 @@ class TestOpenapiVisibilityHelpers: mock_session.execute.return_value.scalars.return_value = iter([]) with patch("services.app_service.apply_openapi_gate", side_effect=lambda q: q): - out = AppService.find_visible_apps_by_name(mock_session, name="nope", tenant_id="tenant-1") + out = AppService.find_visible_apps_by_name(name="nope", tenant_id="tenant-1", session=mock_session) assert out == [] @@ -113,7 +113,7 @@ class TestOpenapiVisibilityHelpers: """ mock_session = MagicMock() - assert AppService.find_visible_apps_by_ids(mock_session, []) == [] + assert AppService.find_visible_apps_by_ids([], session=mock_session) == [] mock_session.execute.assert_not_called() def test_find_visible_apps_by_ids_passes_through_visibility_gate(self): @@ -127,7 +127,7 @@ class TestOpenapiVisibilityHelpers: mock_session.execute.return_value.scalars.return_value.all.return_value = rows with patch("services.app_service.apply_openapi_gate", side_effect=lambda q: q) as gate: - out = AppService.find_visible_apps_by_ids(mock_session, ["a", "b"]) + out = AppService.find_visible_apps_by_ids(["a", "b"], session=mock_session) assert out == rows gate.assert_called_once() @@ -208,6 +208,7 @@ class TestAgentAppType: "use_icon_as_answer_icon": False, "max_active_requests": 0, }, + session=mock_db.session, ) assert updated_app.name == "Iris" @@ -266,6 +267,7 @@ class TestAgentAppType: "use_icon_as_answer_icon": False, "max_active_requests": 0, }, + session=mock_db.session, ) assert backing_agent.role == "research assistant" @@ -317,6 +319,7 @@ class TestAgentAppType: "use_icon_as_answer_icon": False, "max_active_requests": 0, }, + session=mock_db.session, ) assert backing_agent.role == "" @@ -370,6 +373,7 @@ class TestAgentAppType: "use_icon_as_answer_icon": False, "max_active_requests": 0, }, + session=mock_db.session, ) mock_db.session.rollback.assert_called_once() @@ -392,7 +396,7 @@ class TestAgentAppType: patch("services.app_service.remove_app_and_related_data_task"), ): mock_db.session.scalar.return_value = backing_agent - AppService().delete_app(app) # type: ignore[arg-type] + AppService().delete_app(app, session=mock_db.session) # type: ignore[arg-type] assert backing_agent.status == AgentStatus.ARCHIVED assert backing_agent.archived_by == "account-2" diff --git a/api/tests/unit_tests/services/test_async_workflow_service.py b/api/tests/unit_tests/services/test_async_workflow_service.py index 1b9cc8a2ff6..567066845bf 100644 --- a/api/tests/unit_tests/services/test_async_workflow_service.py +++ b/api/tests/unit_tests/services/test_async_workflow_service.py @@ -331,7 +331,7 @@ class TestAsyncWorkflowService: assert trigger_log.triggered_at is not None repo.update.assert_called_once_with(trigger_log) session.commit.assert_called_once() - called_trigger_data = mock_trigger_workflow_async.call_args[0][2] + called_trigger_data = mock_trigger_workflow_async.call_args.args[1] assert isinstance(called_trigger_data, TriggerData) assert called_trigger_data.app_id == "app-123" @@ -465,11 +465,16 @@ class TestAsyncWorkflowServiceGetWorkflow: workflow_service.get_published_workflow_by_id.return_value = workflow # Act - result = AsyncWorkflowService._get_workflow(workflow_service, app_model, workflow_id="workflow-123") + session = MagicMock() + result = AsyncWorkflowService._get_workflow( + workflow_service, app_model, workflow_id="workflow-123", session=session + ) # Assert assert result == workflow - workflow_service.get_published_workflow_by_id.assert_called_once_with(app_model, "workflow-123", session=None) + workflow_service.get_published_workflow_by_id.assert_called_once_with( + app_model, "workflow-123", session=session + ) workflow_service.get_published_workflow.assert_not_called() def test_should_raise_when_specific_workflow_id_not_found(self): @@ -481,7 +486,9 @@ class TestAsyncWorkflowServiceGetWorkflow: # Act / Assert with pytest.raises(WorkflowNotFoundError, match="Published workflow not found: workflow-404"): - AsyncWorkflowService._get_workflow(workflow_service, app_model, workflow_id="workflow-404") + AsyncWorkflowService._get_workflow( + workflow_service, app_model, workflow_id="workflow-404", session=MagicMock() + ) def test_should_return_default_published_workflow_when_workflow_id_not_provided(self): """Test _get_workflow returns default published workflow when no id is provided.""" @@ -493,11 +500,12 @@ class TestAsyncWorkflowServiceGetWorkflow: workflow_service.get_published_workflow.return_value = workflow # Act - result = AsyncWorkflowService._get_workflow(workflow_service, app_model) + session = MagicMock() + result = AsyncWorkflowService._get_workflow(workflow_service, app_model, session=session) # Assert assert result == workflow - workflow_service.get_published_workflow.assert_called_once_with(app_model, session=None) + workflow_service.get_published_workflow.assert_called_once_with(app_model, session=session) workflow_service.get_published_workflow_by_id.assert_not_called() def test_should_raise_when_default_published_workflow_not_found(self): @@ -510,4 +518,4 @@ class TestAsyncWorkflowServiceGetWorkflow: # Act / Assert with pytest.raises(WorkflowNotFoundError, match="No published workflow found for app: app-123"): - AsyncWorkflowService._get_workflow(workflow_service, app_model) + AsyncWorkflowService._get_workflow(workflow_service, app_model, session=MagicMock()) diff --git a/api/tests/unit_tests/services/test_billing_service.py b/api/tests/unit_tests/services/test_billing_service.py index e5610545aa9..dc691176114 100644 --- a/api/tests/unit_tests/services/test_billing_service.py +++ b/api/tests/unit_tests/services/test_billing_service.py @@ -1115,7 +1115,7 @@ class TestBillingServiceAccountManagement: mock_db_session.scalar.return_value = mock_join # Act - should not raise exception - BillingService.is_tenant_owner_or_admin(mock_db_session, current_user) + BillingService.is_tenant_owner_or_admin(current_user, session=mock_db_session) mock_db_session.scalar.assert_called_once() def test_is_tenant_owner_or_admin_admin(self, mock_db_session): @@ -1131,7 +1131,7 @@ class TestBillingServiceAccountManagement: mock_db_session.scalar.return_value = mock_join # Act - should not raise exception - BillingService.is_tenant_owner_or_admin(mock_db_session, current_user) + BillingService.is_tenant_owner_or_admin(current_user, session=mock_db_session) mock_db_session.scalar.assert_called_once() def test_is_tenant_owner_or_admin_normal_user_raises_error(self, mock_db_session): @@ -1148,7 +1148,7 @@ class TestBillingServiceAccountManagement: # Act & Assert with pytest.raises(ValueError) as exc_info: - BillingService.is_tenant_owner_or_admin(mock_db_session, current_user) + BillingService.is_tenant_owner_or_admin(current_user, session=mock_db_session) assert "Only team owner or team admin can perform this action" in str(exc_info.value) mock_db_session.scalar.assert_called_once() @@ -1163,7 +1163,7 @@ class TestBillingServiceAccountManagement: # Act & Assert with pytest.raises(ValueError) as exc_info: - BillingService.is_tenant_owner_or_admin(mock_db_session, current_user) + BillingService.is_tenant_owner_or_admin(current_user, session=mock_db_session) assert "Tenant account join not found" in str(exc_info.value) mock_db_session.scalar.assert_called_once() diff --git a/api/tests/unit_tests/services/test_conversation_service.py b/api/tests/unit_tests/services/test_conversation_service.py index 2c7f13b79f3..e6f7b48f651 100644 --- a/api/tests/unit_tests/services/test_conversation_service.py +++ b/api/tests/unit_tests/services/test_conversation_service.py @@ -330,12 +330,9 @@ class TestConversationServiceHelpers: class TestConversationServiceConversationalVariable: """Test conversational variable operations.""" - @patch("services.conversation_service.session_factory") @patch("services.conversation_service.ConversationService.get_conversation") @patch("services.conversation_service.dify_config") - def test_get_conversational_variable_with_name_filter_mysql( - self, mock_config, mock_get_conversation, mock_session_factory - ): + def test_get_conversational_variable_with_name_filter_mysql(self, mock_config, mock_get_conversation): """ Test variable filtering by name for MySQL databases. @@ -351,7 +348,6 @@ class TestConversationServiceConversationalVariable: # Mock session mock_session = MagicMock() - mock_session_factory.create_session.return_value.__enter__.return_value = mock_session mock_session.scalars.return_value.all.return_value = [] # Act @@ -362,6 +358,7 @@ class TestConversationServiceConversationalVariable: limit=10, last_id=None, variable_name="test_var", + session=mock_session, ) # Assert - JSON filter should be applied diff --git a/api/tests/unit_tests/services/test_credential_permission_service.py b/api/tests/unit_tests/services/test_credential_permission_service.py index e467e9c8c5e..cdcf4a6b00f 100644 --- a/api/tests/unit_tests/services/test_credential_permission_service.py +++ b/api/tests/unit_tests/services/test_credential_permission_service.py @@ -40,7 +40,7 @@ class TestGetPartialMemberList: session = MagicMock() session.scalars.return_value.all.return_value = [] result = CredentialPermissionService.get_partial_member_list( - session, credential_id, CredentialType.TRIGGER_SUBSCRIPTION + credential_id, CredentialType.TRIGGER_SUBSCRIPTION, session=session ) assert result == [] session.scalars.assert_called_once() @@ -49,7 +49,7 @@ class TestGetPartialMemberList: session = MagicMock() session.scalars.return_value.all.return_value = [user_id, other_user_id] result = CredentialPermissionService.get_partial_member_list( - session, credential_id, CredentialType.TRIGGER_SUBSCRIPTION + credential_id, CredentialType.TRIGGER_SUBSCRIPTION, session=session ) assert set(result) == {user_id, other_user_id} session.scalars.assert_called_once() diff --git a/api/tests/unit_tests/services/test_credit_pool_service.py b/api/tests/unit_tests/services/test_credit_pool_service.py index 5e589804c3d..f31d067525a 100644 --- a/api/tests/unit_tests/services/test_credit_pool_service.py +++ b/api/tests/unit_tests/services/test_credit_pool_service.py @@ -1,5 +1,3 @@ -from collections.abc import Generator -from contextlib import contextmanager from types import SimpleNamespace from unittest.mock import MagicMock, patch from uuid import uuid4 @@ -7,7 +5,7 @@ from uuid import uuid4 import pytest from sqlalchemy import create_engine, select from sqlalchemy.engine import Engine -from sqlalchemy.orm import sessionmaker +from sqlalchemy.orm import Session, sessionmaker from core.errors.error import QuotaExceededError from models import TenantCreditPool @@ -38,11 +36,8 @@ def _create_engine_with_pool(*, quota_limit: int, quota_used: int) -> tuple[Engi return engine, tenant_id, pool_id -@contextmanager -def _patched_session_factory(engine: Engine) -> Generator[None, None, None]: - session_maker = sessionmaker(bind=engine, expire_on_commit=False) - with patch("services.credit_pool_service.session_factory.get_session_maker", return_value=session_maker): - yield +def _make_session(engine: Engine) -> Session: + return sessionmaker(bind=engine, expire_on_commit=False)() def _get_quota_used(*, engine: Engine, pool_id: str) -> int | None: @@ -50,25 +45,17 @@ def _get_quota_used(*, engine: Engine, pool_id: str) -> int | None: return connection.scalar(select(TenantCreditPool.quota_used).where(TenantCreditPool.id == pool_id)) -def _make_session_maker(session: MagicMock) -> MagicMock: - session_maker = MagicMock() - transaction = session_maker.begin.return_value - transaction.__enter__.return_value = session - transaction.__exit__.return_value = None - return session_maker - - def _make_redis_lock() -> MagicMock: lock = MagicMock() lock.acquire.return_value = True return lock -def test_get_pool_uses_configured_session_factory_without_flask_app_context() -> None: +def test_get_pool_uses_provided_session() -> None: engine, tenant_id, _ = _create_engine_with_pool(quota_limit=10, quota_used=2) - with _patched_session_factory(engine): - pool = CreditPoolService.get_pool(tenant_id=tenant_id, pool_type=ProviderQuotaType.TRIAL) + with _make_session(engine) as session: + pool = CreditPoolService.get_pool(tenant_id=tenant_id, pool_type=ProviderQuotaType.TRIAL, session=session) assert pool is not None assert pool.tenant_id == tenant_id @@ -78,36 +65,34 @@ def test_get_pool_uses_configured_session_factory_without_flask_app_context() -> def test_check_and_deduct_credits_deducts_exact_amount_when_sufficient() -> None: engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=2) - with _patched_session_factory(engine): - deducted_credits = CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=3) + with _make_session(engine) as session: + deducted_credits = CreditPoolService.check_and_deduct_credits( + tenant_id=tenant_id, credits_required=3, session=session + ) assert deducted_credits == 3 assert _get_quota_used(engine=engine, pool_id=pool_id) == 5 def test_check_and_deduct_credits_returns_zero_for_non_positive_request() -> None: - assert CreditPoolService.check_and_deduct_credits(tenant_id=str(uuid4()), credits_required=0) == 0 + assert ( + CreditPoolService.check_and_deduct_credits(tenant_id=str(uuid4()), credits_required=0, session=MagicMock()) == 0 + ) def test_check_and_deduct_credits_raises_when_pool_is_missing() -> None: engine = create_engine("sqlite:///:memory:") TenantCreditPool.__table__.create(engine) - with ( - _patched_session_factory(engine), - pytest.raises(QuotaExceededError, match="Credit pool not found"), - ): - CreditPoolService.check_and_deduct_credits(tenant_id=str(uuid4()), credits_required=1) + with _make_session(engine) as session, pytest.raises(QuotaExceededError, match="Credit pool not found"): + CreditPoolService.check_and_deduct_credits(tenant_id=str(uuid4()), credits_required=1, session=session) def test_check_and_deduct_credits_raises_when_pool_is_empty() -> None: engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=10) - with ( - _patched_session_factory(engine), - pytest.raises(QuotaExceededError, match="No credits remaining"), - ): - CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=1) + with _make_session(engine) as session, pytest.raises(QuotaExceededError, match="No credits remaining"): + CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=1, session=session) assert _get_quota_used(engine=engine, pool_id=pool_id) == 10 @@ -115,11 +100,8 @@ def test_check_and_deduct_credits_raises_when_pool_is_empty() -> None: def test_check_and_deduct_credits_raises_without_partial_deduction_when_insufficient() -> None: engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=9) - with ( - _patched_session_factory(engine), - pytest.raises(QuotaExceededError, match="Insufficient credits remaining"), - ): - CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=3) + with _make_session(engine) as session, pytest.raises(QuotaExceededError, match="Insufficient credits remaining"): + CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=3, session=session) assert _get_quota_used(engine=engine, pool_id=pool_id) == 9 @@ -128,25 +110,27 @@ def test_check_and_deduct_credits_wraps_unexpected_deduction_errors() -> None: engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=2) with ( - _patched_session_factory(engine), + _make_session(engine) as session, patch.object(CreditPoolService, "_get_locked_pool", side_effect=RuntimeError("database unavailable")), pytest.raises(QuotaExceededError, match="Failed to deduct credits"), ): - CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=1) + CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=1, session=session) assert _get_quota_used(engine=engine, pool_id=pool_id) == 2 def test_deduct_credits_capped_returns_zero_for_non_positive_request() -> None: - assert CreditPoolService.deduct_credits_capped(tenant_id=str(uuid4()), credits_required=0) == 0 + assert CreditPoolService.deduct_credits_capped(tenant_id=str(uuid4()), credits_required=0, session=MagicMock()) == 0 def test_deduct_credits_capped_returns_zero_when_pool_is_missing() -> None: engine = create_engine("sqlite:///:memory:") TenantCreditPool.__table__.create(engine) - with _patched_session_factory(engine): - deducted_credits = CreditPoolService.deduct_credits_capped(tenant_id=str(uuid4()), credits_required=1) + with _make_session(engine) as session: + deducted_credits = CreditPoolService.deduct_credits_capped( + tenant_id=str(uuid4()), credits_required=1, session=session + ) assert deducted_credits == 0 @@ -154,8 +138,10 @@ def test_deduct_credits_capped_returns_zero_when_pool_is_missing() -> None: def test_deduct_credits_capped_returns_zero_when_pool_is_empty() -> None: engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=10) - with _patched_session_factory(engine): - deducted_credits = CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=1) + with _make_session(engine) as session: + deducted_credits = CreditPoolService.deduct_credits_capped( + tenant_id=tenant_id, credits_required=1, session=session + ) assert deducted_credits == 0 assert _get_quota_used(engine=engine, pool_id=pool_id) == 10 @@ -164,8 +150,10 @@ def test_deduct_credits_capped_returns_zero_when_pool_is_empty() -> None: def test_deduct_credits_capped_deducts_only_remaining_balance_when_insufficient() -> None: engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=9) - with _patched_session_factory(engine): - deducted_credits = CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=3) + with _make_session(engine) as session: + deducted_credits = CreditPoolService.deduct_credits_capped( + tenant_id=tenant_id, credits_required=3, session=session + ) assert deducted_credits == 1 assert _get_quota_used(engine=engine, pool_id=pool_id) == 10 @@ -175,11 +163,11 @@ def test_deduct_credits_capped_wraps_unexpected_deduction_errors() -> None: engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=2) with ( - _patched_session_factory(engine), + _make_session(engine) as session, patch.object(CreditPoolService, "_get_locked_pool", side_effect=RuntimeError("database unavailable")), pytest.raises(QuotaExceededError, match="Failed to deduct credits"), ): - CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=1) + CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=1, session=session) assert _get_quota_used(engine=engine, pool_id=pool_id) == 2 @@ -188,11 +176,11 @@ def test_deduct_credits_capped_reraises_quota_exceeded_errors() -> None: engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=2) with ( - _patched_session_factory(engine), + _make_session(engine) as session, patch.object(CreditPoolService, "_get_locked_pool", side_effect=QuotaExceededError("quota unavailable")), pytest.raises(QuotaExceededError, match="quota unavailable"), ): - CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=1) + CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=1, session=session) assert _get_quota_used(engine=engine, pool_id=pool_id) == 2 @@ -200,19 +188,18 @@ def test_deduct_credits_capped_reraises_quota_exceeded_errors() -> None: def test_check_and_deduct_credits_uses_tenant_redis_lock_before_db_deduction() -> None: tenant_id = "tenant-1" session = MagicMock() - session_maker = _make_session_maker(session) pool = SimpleNamespace(remaining_credits=10, quota_used=2) redis_lock = _make_redis_lock() with ( patch("services.credit_pool_service.redis_client.lock", return_value=redis_lock) as lock, - patch("services.credit_pool_service.session_factory.get_session_maker", return_value=session_maker), patch.object(CreditPoolService, "_get_locked_pool", return_value=pool) as get_locked_pool, ): result = CreditPoolService.check_and_deduct_credits( tenant_id=tenant_id, credits_required=3, pool_type=ProviderQuotaType.TRIAL, + session=session, ) assert result == 3 @@ -230,19 +217,18 @@ def test_check_and_deduct_credits_uses_tenant_redis_lock_before_db_deduction() - def test_deduct_credits_capped_uses_tenant_redis_lock_before_db_deduction() -> None: tenant_id = "tenant-1" session = MagicMock() - session_maker = _make_session_maker(session) pool = SimpleNamespace(remaining_credits=2, quota_used=8) redis_lock = _make_redis_lock() with ( patch("services.credit_pool_service.redis_client.lock", return_value=redis_lock) as lock, - patch("services.credit_pool_service.session_factory.get_session_maker", return_value=session_maker), patch.object(CreditPoolService, "_get_locked_pool", return_value=pool) as get_locked_pool, ): result = CreditPoolService.deduct_credits_capped( tenant_id=tenant_id, credits_required=5, pool_type=ProviderQuotaType.PAID, + session=session, ) assert result == 2 @@ -266,38 +252,35 @@ def test_deduct_credits_capped_uses_tenant_redis_lock_before_db_deduction() -> N ) def test_non_positive_credit_request_skips_tenant_redis_lock(deduct_method) -> None: with patch("services.credit_pool_service.redis_client.lock") as lock: - result = deduct_method(tenant_id="tenant-1", credits_required=0) + result = deduct_method(tenant_id="tenant-1", credits_required=0, session=MagicMock()) assert result == 0 lock.assert_not_called() def test_check_and_deduct_credits_wraps_redis_lock_errors_without_querying_db() -> None: - session_maker = MagicMock() + session = MagicMock() with ( patch("services.credit_pool_service.redis_client.lock", side_effect=RuntimeError("redis unavailable")), - patch("services.credit_pool_service.session_factory.get_session_maker", return_value=session_maker), pytest.raises(QuotaExceededError, match="Failed to deduct credits"), ): - CreditPoolService.check_and_deduct_credits(tenant_id="tenant-1", credits_required=1) + CreditPoolService.check_and_deduct_credits(tenant_id="tenant-1", credits_required=1, session=session) - session_maker.begin.assert_not_called() + session.scalar.assert_not_called() def test_deduct_credits_capped_ignores_release_errors_after_successful_deduction() -> None: session = MagicMock() - session_maker = _make_session_maker(session) pool = SimpleNamespace(remaining_credits=3, quota_used=7) redis_lock = _make_redis_lock() redis_lock.release.side_effect = RuntimeError("release failed") with ( patch("services.credit_pool_service.redis_client.lock", return_value=redis_lock), - patch("services.credit_pool_service.session_factory.get_session_maker", return_value=session_maker), patch.object(CreditPoolService, "_get_locked_pool", return_value=pool), ): - result = CreditPoolService.deduct_credits_capped(tenant_id="tenant-1", credits_required=2) + result = CreditPoolService.deduct_credits_capped(tenant_id="tenant-1", credits_required=2, session=session) assert result == 2 assert pool.quota_used == 9 diff --git a/api/tests/unit_tests/services/test_dataset_service_dataset.py b/api/tests/unit_tests/services/test_dataset_service_dataset.py index 02d965f4bd2..dcb6250a8f0 100644 --- a/api/tests/unit_tests/services/test_dataset_service_dataset.py +++ b/api/tests/unit_tests/services/test_dataset_service_dataset.py @@ -344,7 +344,9 @@ class TestDatasetServiceCreationAndUpdate: mock_db.session.scalar.return_value = object() with pytest.raises(DatasetNameDuplicateError, match="Dataset with name Dataset already exists"): - DatasetService.create_empty_dataset(mock_db.session, "tenant-1", "Dataset", None, "economy", account) + DatasetService.create_empty_dataset( + "tenant-1", "Dataset", None, "economy", account, session=mock_db.session + ) def test_create_empty_dataset_uses_default_embedding_model_for_high_quality_dataset(self): account = SimpleNamespace(id="user-1") @@ -512,7 +514,7 @@ class TestDatasetServiceCreationAndUpdate: session = MagicMock() with patch.object(DatasetService, "get_dataset", return_value=None): with pytest.raises(ValueError, match="Dataset not found"): - DatasetService.update_dataset(session, "dataset-1", {}, SimpleNamespace(id="user-1")) + DatasetService.update_dataset("dataset-1", {}, SimpleNamespace(id="user-1"), session=session) def test_update_dataset_raises_when_new_name_conflicts(self): dataset = DatasetServiceUnitDataFactory.create_dataset_mock(dataset_id="dataset-1", tenant_id="tenant-1") @@ -524,10 +526,7 @@ class TestDatasetServiceCreationAndUpdate: ): with pytest.raises(ValueError, match="Dataset name already exists"): DatasetService.update_dataset( - MagicMock(), - "dataset-1", - {"name": "New Dataset"}, - SimpleNamespace(id="user-1"), + "dataset-1", {"name": "New Dataset"}, SimpleNamespace(id="user-1"), session=MagicMock() ) def test_update_dataset_routes_external_datasets_to_external_helper(self): @@ -541,7 +540,7 @@ class TestDatasetServiceCreationAndUpdate: patch.object(DatasetService, "_update_external_dataset", return_value="updated") as update_external, ): session = MagicMock() - result = DatasetService.update_dataset(session, "dataset-1", {"name": dataset.name}, user) + result = DatasetService.update_dataset("dataset-1", {"name": dataset.name}, user, session=session) assert result == "updated" check_permission.assert_called_once() @@ -560,7 +559,7 @@ class TestDatasetServiceCreationAndUpdate: patch.object(DatasetService, "_update_internal_dataset", return_value="updated") as update_internal, ): session = MagicMock() - result = DatasetService.update_dataset(session, "dataset-1", {"name": dataset.name}, user) + result = DatasetService.update_dataset("dataset-1", {"name": dataset.name}, user, session=session) assert result == "updated" check_permission.assert_called_once() @@ -612,7 +611,7 @@ class TestDatasetServiceCreationAndUpdate: assert dataset.permission == DatasetPermissionEnum.PARTIAL_TEAM assert dataset.updated_by == "user-1" assert dataset.updated_at is now - get_external_knowledge_api.assert_called_once_with(mock_db.session, "api-1", dataset.tenant_id) + get_external_knowledge_api.assert_called_once_with("api-1", dataset.tenant_id, session=mock_db.session) update_binding.assert_called_once_with("dataset-1", "knowledge-1", "api-1", mock_db.session) mock_db.session.add.assert_called_once_with(dataset) mock_db.session.commit.assert_called_once() @@ -652,7 +651,7 @@ class TestDatasetServiceCreationAndUpdate: mock_db.session, ) - get_external_knowledge_api.assert_called_once_with(mock_db.session, "foreign-api", dataset.tenant_id) + get_external_knowledge_api.assert_called_once_with("foreign-api", dataset.tenant_id, session=mock_db.session) update_binding.assert_not_called() mock_db.session.commit.assert_not_called() @@ -1165,7 +1164,7 @@ class TestDatasetServiceRagPipelineSettings: with patch("services.dataset_service.current_user", SimpleNamespace(current_tenant_id=None)): with pytest.raises(ValueError, match="Current user or current tenant not found"): - DatasetService.update_rag_pipeline_dataset_settings(session, dataset, knowledge_configuration) + DatasetService.update_rag_pipeline_dataset_settings(dataset, knowledge_configuration, session=session) def test_update_rag_pipeline_dataset_settings_without_published_high_quality_updates_embedding_settings(self): session = MagicMock() @@ -1185,7 +1184,7 @@ class TestDatasetServiceRagPipelineSettings: ): model_manager_cls.for_tenant.return_value.get_model_instance.return_value = embedding_model - DatasetService.update_rag_pipeline_dataset_settings(session, dataset, knowledge_configuration) + DatasetService.update_rag_pipeline_dataset_settings(dataset, knowledge_configuration, session=session) assert dataset.chunk_structure == "paragraph" assert dataset.indexing_technique == "high_quality" @@ -1211,7 +1210,7 @@ class TestDatasetServiceRagPipelineSettings: ) with patch("services.dataset_service.current_user", SimpleNamespace(current_tenant_id="tenant-1")): - DatasetService.update_rag_pipeline_dataset_settings(session, dataset, knowledge_configuration) + DatasetService.update_rag_pipeline_dataset_settings(dataset, knowledge_configuration, session=session) assert dataset.indexing_technique == "economy" assert dataset.keyword_number == 12 @@ -1228,10 +1227,7 @@ class TestDatasetServiceRagPipelineSettings: with patch("services.dataset_service.current_user", SimpleNamespace(current_tenant_id="tenant-1")): with pytest.raises(ValueError, match="Chunk structure is not allowed to be updated"): DatasetService.update_rag_pipeline_dataset_settings( - session, - dataset, - knowledge_configuration, - has_published=True, + dataset, knowledge_configuration, has_published=True, session=session ) def test_update_rag_pipeline_dataset_settings_with_published_rejects_switch_to_economy(self): @@ -1252,10 +1248,7 @@ class TestDatasetServiceRagPipelineSettings: match="Knowledge base indexing technique is not allowed to be updated to economy", ): DatasetService.update_rag_pipeline_dataset_settings( - session, - dataset, - knowledge_configuration, - has_published=True, + dataset, knowledge_configuration, has_published=True, session=session ) def test_update_rag_pipeline_dataset_settings_with_published_adds_high_quality_index(self): @@ -1280,10 +1273,7 @@ class TestDatasetServiceRagPipelineSettings: model_manager_cls.for_tenant.return_value.get_model_instance.return_value = embedding_model DatasetService.update_rag_pipeline_dataset_settings( - session, - dataset, - knowledge_configuration, - has_published=True, + dataset, knowledge_configuration, has_published=True, session=session ) assert dataset.indexing_technique == "high_quality" @@ -1326,10 +1316,7 @@ class TestDatasetServiceRagPipelineSettings: ) DatasetService.update_rag_pipeline_dataset_settings( - session, - dataset, - knowledge_configuration, - has_published=True, + dataset, knowledge_configuration, has_published=True, session=session ) assert dataset.embedding_model_provider == "provider-two" @@ -1364,10 +1351,7 @@ class TestDatasetServiceRagPipelineSettings: ) DatasetService.update_rag_pipeline_dataset_settings( - session, - dataset, - knowledge_configuration, - has_published=True, + dataset, knowledge_configuration, has_published=True, session=session ) assert dataset.embedding_model_provider == "provider" @@ -1396,10 +1380,7 @@ class TestDatasetServiceRagPipelineSettings: patch("services.dataset_service.deal_dataset_index_update_task") as update_task, ): DatasetService.update_rag_pipeline_dataset_settings( - session, - dataset, - knowledge_configuration, - has_published=True, + dataset, knowledge_configuration, has_published=True, session=session ) assert dataset.keyword_number == 9 @@ -1457,7 +1438,7 @@ class TestDatasetPermissionService: session = MagicMock() with pytest.raises(NoPermissionError, match="does not have permission"): - DatasetPermissionService.check_permission(session, user, dataset, "all_team", []) + DatasetPermissionService.check_permission(user, dataset, "all_team", [], session=session) def test_check_permission_prevents_dataset_operator_from_changing_permission_mode(self): user = SimpleNamespace(is_dataset_editor=True, is_dataset_operator=True) @@ -1465,7 +1446,7 @@ class TestDatasetPermissionService: session = MagicMock() with pytest.raises(NoPermissionError, match="cannot change the dataset permissions"): - DatasetPermissionService.check_permission(session, user, dataset, "only_me", []) + DatasetPermissionService.check_permission(user, dataset, "only_me", [], session=session) def test_check_permission_requires_partial_member_list_for_partial_members_mode(self): user = SimpleNamespace(is_dataset_editor=True, is_dataset_operator=True) @@ -1473,7 +1454,7 @@ class TestDatasetPermissionService: session = MagicMock() with pytest.raises(ValueError, match="Partial member list is required"): - DatasetPermissionService.check_permission(session, user, dataset, "partial_members", []) + DatasetPermissionService.check_permission(user, dataset, "partial_members", [], session=session) def test_check_permission_rejects_dataset_operator_member_list_changes(self): user = SimpleNamespace(is_dataset_editor=True, is_dataset_operator=True) @@ -1485,11 +1466,7 @@ class TestDatasetPermissionService: with patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=["user-1"]): with pytest.raises(ValueError, match="cannot change the dataset permissions"): DatasetPermissionService.check_permission( - session, - user, - dataset, - "partial_members", - [{"user_id": "user-2"}], + user, dataset, "partial_members", [{"user_id": "user-2"}], session=session ) def test_check_permission_allows_dataset_operator_when_member_list_is_unchanged(self): @@ -1501,11 +1478,7 @@ class TestDatasetPermissionService: with patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=["user-1"]): DatasetPermissionService.check_permission( - session, - user, - dataset, - "partial_members", - [{"user_id": "user-1"}], + user, dataset, "partial_members", [{"user_id": "user-1"}], session=session ) def test_clear_partial_member_list_rolls_back_on_exception(self): diff --git a/api/tests/unit_tests/services/test_dataset_service_document.py b/api/tests/unit_tests/services/test_dataset_service_document.py index 02661fbe1f3..44619a29e83 100644 --- a/api/tests/unit_tests/services/test_dataset_service_document.py +++ b/api/tests/unit_tests/services/test_dataset_service_document.py @@ -1183,7 +1183,7 @@ class TestDocumentServiceTenantAndUpdateEdges: with patch("services.dataset_service.db") as mock_db: mock_db.session.scalar.return_value = 12 - result = DocumentService.get_tenant_documents_count(mock_db.session) + result = DocumentService.get_tenant_documents_count(session=mock_db.session) assert result == 12 diff --git a/api/tests/unit_tests/services/test_dataset_service_segment.py b/api/tests/unit_tests/services/test_dataset_service_segment.py index 34f3f947f96..c94093d59b7 100644 --- a/api/tests/unit_tests/services/test_dataset_service_segment.py +++ b/api/tests/unit_tests/services/test_dataset_service_segment.py @@ -306,13 +306,13 @@ class TestSegmentServiceQueries: def test_get_child_chunk_by_segment_ref_uses_full_ownership_chain(self): child_chunk = _make_child_chunk() segment_ref = _make_segment_ref() + session = MagicMock() + session.scalar.return_value = child_chunk - with patch("services.dataset_service.db") as mock_db: - mock_db.session.scalar.return_value = child_chunk - result = SegmentService.get_child_chunk_by_segment_ref("child-a", segment_ref) + result = SegmentService.get_child_chunk_by_segment_ref("child-a", segment_ref, session) assert result is child_chunk - stmt = mock_db.session.scalar.call_args.args[0] + stmt = session.scalar.call_args.args[0] sql = str(stmt.compile(compile_kwargs={"literal_binds": True})) assert "child_chunks.id = 'child-a'" in sql assert "child_chunks.tenant_id = 'tenant-1'" in sql @@ -381,13 +381,13 @@ class TestSegmentServiceQueries: ) segment.id = "segment-1" segment_ref = _make_segment_ref() + session = MagicMock() + session.scalar.return_value = segment - with patch("services.dataset_service.db") as mock_db: - mock_db.session.scalar.return_value = segment - result = SegmentService.get_segment_by_ref(segment_ref) + result = SegmentService.get_segment_by_ref(segment_ref, session) assert result is segment - stmt = mock_db.session.scalar.call_args.args[0] + stmt = session.scalar.call_args.args[0] sql = str(stmt.compile(compile_kwargs={"literal_binds": True})) assert "document_segments.id = 'segment-1'" in sql assert "document_segments.tenant_id = 'tenant-1'" in sql @@ -566,7 +566,7 @@ class TestSegmentServiceMutations: assert all(segment.error == "vector failed" for segment in result) assert document.word_count == 5 + sum(len(item["content"]) + len(item["answer"]) for item in segments) vector_service.create_segments_vector.assert_called_once_with( - [["k1"], None], result, dataset, document.doc_form + [["k1"], None], result, dataset, document.doc_form, mock_db.session ) mock_db.session.commit.assert_called_once() @@ -641,7 +641,7 @@ class TestSegmentServiceMutations: assert result is refreshed_segment assert segment.keywords == ["new"] vector_service.update_segment_vector.assert_called_once_with(["new"], segment, dataset) - vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset) + vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, mock_db.session) def test_update_segment_regenerates_child_chunks_and_updates_manual_summary(self, account_context): segment = _make_segment(content="same content", word_count=len("same content")) @@ -684,10 +684,11 @@ class TestSegmentServiceMutations: dataset, embedding_model_instance, processing_rule, + mock_db.session, True, ) - update_summary.assert_called_once_with(segment, dataset, "new summary") - vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset) + update_summary.assert_called_once_with(segment, dataset, "new summary", session=mock_db.session) + vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, mock_db.session) def test_update_segment_auto_regenerates_summary_after_content_change(self, account_context): segment = _make_segment(content="old", word_count=3) @@ -725,8 +726,8 @@ class TestSegmentServiceMutations: assert segment.tokens == 9 assert document.word_count == 18 vector_service.update_segment_vector.assert_called_once_with(["kw-1"], segment, dataset) - generate_summary.assert_called_once_with(segment, dataset, {"enable": True}) - vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset) + generate_summary.assert_called_once_with(segment, dataset, {"enable": True}, session=mock_db.session) + vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, mock_db.session) def test_update_segment_regenerates_summary_when_manual_summary_is_unchanged(self, account_context): segment = _make_segment(content="old", word_count=3) @@ -760,9 +761,9 @@ class TestSegmentServiceMutations: result = SegmentService.update_segment(args, segment, document, dataset, mock_db.session) assert result is refreshed_segment - generate_summary.assert_called_once_with(segment, dataset, {"enable": True}) + generate_summary.assert_called_once_with(segment, dataset, {"enable": True}, session=mock_db.session) update_summary.assert_not_called() - vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset) + vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, mock_db.session) def test_delete_segment_removes_index_and_updates_document_word_count(self): segment = _make_segment(word_count=4, index_node_id="parent-node") @@ -972,7 +973,7 @@ class TestSegmentServiceAdditionalRegenerationBranches: assert segment.word_count == len("question") + len("new answer") assert document.word_count == 20 + (len("question") + len("new answer") - 8) vector_service.update_segment_vector.assert_not_called() - vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset) + vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, mock_db.session) def test_update_segment_content_change_uses_answer_when_counting_tokens_for_qa_segments(self, account_context): segment = _make_segment(content="old", word_count=3) @@ -1009,7 +1010,7 @@ class TestSegmentServiceAdditionalRegenerationBranches: assert segment.tokens == 21 assert segment.word_count == len("new question") + len("new answer") vector_service.update_segment_vector.assert_called_once_with(["kw-1"], segment, dataset) - vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset) + vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, mock_db.session) def test_update_segment_content_change_parent_child_uses_default_embedding_and_ignores_summary_failures( self, account_context @@ -1063,10 +1064,11 @@ class TestSegmentServiceAdditionalRegenerationBranches: dataset, embedding_model_instance, processing_rule, + mock_db.session, True, ) - update_summary.assert_called_once_with(segment, dataset, "new summary") - vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset) + update_summary.assert_called_once_with(segment, dataset, "new summary", session=mock_db.session) + vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, mock_db.session) def test_update_segment_same_content_parent_child_marks_segment_error_for_non_high_quality_dataset( self, account_context diff --git a/api/tests/unit_tests/services/test_datasource_provider_service.py b/api/tests/unit_tests/services/test_datasource_provider_service.py index f374a294825..bd6891d846a 100644 --- a/api/tests/unit_tests/services/test_datasource_provider_service.py +++ b/api/tests/unit_tests/services/test_datasource_provider_service.py @@ -177,11 +177,11 @@ class TestDatasourceProviderService: def test_should_return_true_when_tenant_oauth_params_enabled(self, service, mock_db_session): mock_db_session.scalar.return_value = 1 - assert service.is_tenant_oauth_params_enabled("t1", make_id()) is True + assert service.is_tenant_oauth_params_enabled("t1", make_id(), session=mock_db_session) is True def test_should_return_false_when_tenant_oauth_params_disabled(self, service, mock_db_session): mock_db_session.scalar.return_value = 0 - assert service.is_tenant_oauth_params_enabled("t1", make_id()) is False + assert service.is_tenant_oauth_params_enabled("t1", make_id(), session=mock_db_session) is False # ----------------------------------------------------------------------- # remove_oauth_custom_client_params (lines 55-61) @@ -453,7 +453,7 @@ class TestDatasourceProviderService: tenant_params.client_params = {"k": "v"} mock_db_session.scalar.return_value = tenant_params with patch.object(service, "get_oauth_encrypter", return_value=(self._enc, None)): - result = service.get_tenant_oauth_client("t1", make_id(), mask=True) + result = service.get_tenant_oauth_client("t1", make_id(), mask=True, session=mock_db_session) assert result == {"k": "mask"} def test_should_return_decrypted_credentials_when_mask_is_false(self, service, mock_db_session): @@ -461,12 +461,12 @@ class TestDatasourceProviderService: tenant_params.client_params = {"k": "v"} mock_db_session.scalar.return_value = tenant_params with patch.object(service, "get_oauth_encrypter", return_value=(self._enc, None)): - result = service.get_tenant_oauth_client("t1", make_id(), mask=False) + result = service.get_tenant_oauth_client("t1", make_id(), mask=False, session=mock_db_session) assert result == {"k": "dec"} def test_should_return_none_when_no_tenant_oauth_config_exists(self, service, mock_db_session): mock_db_session.scalar.return_value = None - assert service.get_tenant_oauth_client("t1", make_id()) is None + assert service.get_tenant_oauth_client("t1", make_id(), session=mock_db_session) is None # ----------------------------------------------------------------------- # get_oauth_client (lines 423-457) @@ -657,7 +657,7 @@ class TestDatasourceProviderService: def test_should_return_empty_list_when_no_credentials_stored(self, service, mock_db_session): mock_db_session.scalars.return_value.all.return_value = [] - assert service.list_datasource_credentials("t1", "prov", "org/plug") == [] + assert service.list_datasource_credentials("t1", "prov", "org/plug", session=mock_db_session) == [] def test_should_return_masked_credentials_list_when_credentials_exist(self, service, mock_db_session): p = MagicMock(spec=DatasourceProvider) @@ -666,7 +666,7 @@ class TestDatasourceProviderService: p.is_default = False mock_db_session.scalars.return_value.all.return_value = [p] with patch.object(service, "extract_secret_variables", return_value=["sk"]): - result = service.list_datasource_credentials("t1", "prov", "org/plug") + result = service.list_datasource_credentials("t1", "prov", "org/plug", session=mock_db_session) assert len(result) == 1 # ----------------------------------------------------------------------- @@ -682,7 +682,9 @@ class TestDatasourceProviderService: mock_mgr.return_value.fetch_installed_datasource_providers.return_value = [ds] cred = {"credential": {"k": "v"}, "is_default": True} with patch.object(service, "list_datasource_credentials", return_value=[cred]): - results = service.get_all_datasource_credentials("t1") + session = MagicMock() + session.scalar.return_value = 0 + results = service.get_all_datasource_credentials("t1", session=session) assert len(results) == 1 def test_should_include_oauth_schema_for_hardcoded_plugin_ids(self, service, mock_db_session): @@ -707,7 +709,7 @@ class TestDatasourceProviderService: patch.object(service, "is_tenant_oauth_params_enabled", return_value=False), patch.object(service, "is_system_oauth_params_exist", return_value=False), ): - results = service.get_all_datasource_credentials("t1") + results = service.get_all_datasource_credentials("t1", session=mock_db_session) assert len(results) == 1 assert results[0]["oauth_schema"] is not None @@ -717,7 +719,7 @@ class TestDatasourceProviderService: def test_should_return_empty_list_when_no_real_credentials_exist(self, service, mock_db_session): mock_db_session.scalars.return_value.all.return_value = [] - assert service.get_real_datasource_credentials("t1", "prov", "org/plug") == [] + assert service.get_real_datasource_credentials("t1", "prov", "org/plug", session=mock_db_session) == [] def test_should_return_decrypted_credential_list_when_credentials_exist(self, service, mock_db_session): p = MagicMock(spec=DatasourceProvider) @@ -725,7 +727,7 @@ class TestDatasourceProviderService: p.encrypted_credentials = {"sk": "v"} mock_db_session.scalars.return_value.all.return_value = [p] with patch.object(service, "extract_secret_variables", return_value=["sk"]): - result = service.get_real_datasource_credentials("t1", "prov", "org/plug") + result = service.get_real_datasource_credentials("t1", "prov", "org/plug", session=mock_db_session) assert len(result) == 1 # ----------------------------------------------------------------------- @@ -788,11 +790,11 @@ class TestDatasourceProviderService: def test_should_delete_provider_and_commit_when_found(self, service, mock_db_session): p = MagicMock(spec=DatasourceProvider) mock_db_session.scalar.return_value = p - service.remove_datasource_credentials("t1", "id", "prov", "org/plug") + service.remove_datasource_credentials("t1", "id", "prov", "org/plug", session=mock_db_session) mock_db_session.delete.assert_called_once_with(p) def test_should_do_nothing_when_credential_not_found_on_remove(self, service, mock_db_session): """No error raised; no delete called when record doesn't exist (lines 994 branch).""" mock_db_session.scalar.return_value = None - service.remove_datasource_credentials("t1", "id", "prov", "org/plug") + service.remove_datasource_credentials("t1", "id", "prov", "org/plug", session=mock_db_session) mock_db_session.delete.assert_not_called() diff --git a/api/tests/unit_tests/services/test_external_dataset_service.py b/api/tests/unit_tests/services/test_external_dataset_service.py index dbb4627759c..9dff74f8dd5 100644 --- a/api/tests/unit_tests/services/test_external_dataset_service.py +++ b/api/tests/unit_tests/services/test_external_dataset_service.py @@ -145,7 +145,7 @@ class TestExternalDatasetServiceGetAPIs: @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_success_basic( - self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate_query, factory: ExternalDatasetServiceTestDataFactory ): """Test successful retrieval of external knowledge APIs with pagination.""" # Arrange @@ -158,7 +158,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis mock_pagination.total = 5 - mock_paginate.return_value = mock_pagination + mock_paginate_query.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -170,11 +170,11 @@ class TestExternalDatasetServiceGetAPIs: assert result_total == 5 assert result_items[0].id == "api-0" assert result_items[4].id == "api-4" - mock_paginate.assert_called_once() + mock_paginate_query.assert_called_once() @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_with_search_filter( - self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate_query, factory: ExternalDatasetServiceTestDataFactory ): """Test retrieval with search filter.""" # Arrange @@ -186,7 +186,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis mock_pagination.total = 1 - mock_paginate.return_value = mock_pagination + mock_paginate_query.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -200,14 +200,14 @@ class TestExternalDatasetServiceGetAPIs: @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_empty_results( - self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate_query, factory: ExternalDatasetServiceTestDataFactory ): """Test retrieval with no results.""" # Arrange mock_pagination = MagicMock() mock_pagination.items = [] mock_pagination.total = 0 - mock_paginate.return_value = mock_pagination + mock_paginate_query.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -220,7 +220,7 @@ class TestExternalDatasetServiceGetAPIs: @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_large_result_set( - self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate_query, factory: ExternalDatasetServiceTestDataFactory ): """Test retrieval with large result set.""" # Arrange @@ -229,7 +229,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis[:10] mock_pagination.total = 100 - mock_paginate.return_value = mock_pagination + mock_paginate_query.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -242,7 +242,7 @@ class TestExternalDatasetServiceGetAPIs: @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_pagination_last_page( - self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate_query, factory: ExternalDatasetServiceTestDataFactory ): """Test last page pagination with partial results.""" # Arrange @@ -251,7 +251,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis mock_pagination.total = 100 - mock_paginate.return_value = mock_pagination + mock_paginate_query.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -264,7 +264,7 @@ class TestExternalDatasetServiceGetAPIs: @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_case_insensitive_search( - self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate_query, factory: ExternalDatasetServiceTestDataFactory ): """Test case-insensitive search functionality.""" # Arrange @@ -276,7 +276,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis mock_pagination.total = 2 - mock_paginate.return_value = mock_pagination + mock_paginate_query.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -289,7 +289,7 @@ class TestExternalDatasetServiceGetAPIs: @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_special_characters_search( - self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate_query, factory: ExternalDatasetServiceTestDataFactory ): """Test search with special characters.""" # Arrange @@ -298,7 +298,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis mock_pagination.total = 1 - mock_paginate.return_value = mock_pagination + mock_paginate_query.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -310,7 +310,7 @@ class TestExternalDatasetServiceGetAPIs: @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_max_per_page_limit( - self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate_query, factory: ExternalDatasetServiceTestDataFactory ): """Test that max_per_page limit is enforced.""" # Arrange @@ -319,7 +319,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis mock_pagination.total = 1000 - mock_paginate.return_value = mock_pagination + mock_paginate_query.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -327,12 +327,12 @@ class TestExternalDatasetServiceGetAPIs: ) # Assert - call_args = mock_paginate.call_args + call_args = mock_paginate_query.call_args assert call_args.kwargs["max_per_page"] == 100 @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_ordered_by_created_at_desc( - self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate_query, factory: ExternalDatasetServiceTestDataFactory ): """Test that results are ordered by created_at descending.""" # Arrange @@ -344,7 +344,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis[::-1] # Reversed to simulate DESC order mock_pagination.total = 5 - mock_paginate.return_value = mock_pagination + mock_paginate_query.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -437,13 +437,13 @@ class TestExternalDatasetServiceValidateAPIList: class TestExternalDatasetServiceCreateAPI: """Test create_external_knowledge_api operations.""" + @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") def test_create_external_knowledge_api_success_full( - self, mock_check, factory: ExternalDatasetServiceTestDataFactory + self, mock_check, mock_db, factory: ExternalDatasetServiceTestDataFactory ): """Test successful creation with all fields.""" # Arrange - mock_session = MagicMock() tenant_id = "tenant-123" user_id = "user-123" args = { @@ -453,7 +453,7 @@ class TestExternalDatasetServiceCreateAPI: } # Act - result = ExternalDatasetService.create_external_knowledge_api(tenant_id, user_id, args, mock_session) + result = ExternalDatasetService.create_external_knowledge_api(tenant_id, user_id, args, session=mock_db.session) # Assert assert result.name == "Test API" @@ -462,55 +462,63 @@ class TestExternalDatasetServiceCreateAPI: assert result.created_by == user_id assert result.updated_by == user_id mock_check.assert_called_once_with(args["settings"]) - mock_session.add.assert_called_once() - mock_session.commit.assert_called_once() + mock_db.session.add.assert_called_once() + mock_db.session.commit.assert_called_once() + @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") def test_create_external_knowledge_api_minimal_fields( - self, mock_check, factory: ExternalDatasetServiceTestDataFactory + self, mock_check, mock_db, factory: ExternalDatasetServiceTestDataFactory ): """Test creation with minimal required fields.""" # Arrange - mock_session = MagicMock() args = { "name": "Minimal API", "settings": {"endpoint": "https://api.example.com", "api_key": "key"}, } # Act - result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_session) + result = ExternalDatasetService.create_external_knowledge_api( + "tenant-123", "user-123", args, session=mock_db.session + ) # Assert assert result.name == "Minimal API" assert result.description == "" - def test_create_external_knowledge_api_missing_settings(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_create_external_knowledge_api_missing_settings( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test creation fails when settings are missing.""" # Arrange - mock_session = MagicMock() args = {"name": "Test API", "description": "Test"} # Act & Assert with pytest.raises(ValueError, match="settings is required"): - ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_session) + ExternalDatasetService.create_external_knowledge_api( + "tenant-123", "user-123", args, session=mock_db.session + ) - def test_create_external_knowledge_api_none_settings(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_create_external_knowledge_api_none_settings(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test creation fails when settings are explicitly None.""" # Arrange - mock_session = MagicMock() args = {"name": "Test API", "settings": None} # Act & Assert with pytest.raises(ValueError, match="settings is required"): - ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_session) + ExternalDatasetService.create_external_knowledge_api( + "tenant-123", "user-123", args, session=mock_db.session + ) + @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") def test_create_external_knowledge_api_settings_json_serialization( - self, mock_check, factory: ExternalDatasetServiceTestDataFactory + self, mock_check, mock_db, factory: ExternalDatasetServiceTestDataFactory ): """Test that settings are properly JSON serialized.""" # Arrange - mock_session = MagicMock() settings = { "endpoint": "https://api.example.com", "api_key": "test-key", @@ -519,20 +527,22 @@ class TestExternalDatasetServiceCreateAPI: args = {"name": "Test API", "settings": settings} # Act - result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_session) + result = ExternalDatasetService.create_external_knowledge_api( + "tenant-123", "user-123", args, session=mock_db.session + ) # Assert assert isinstance(result.settings, str) parsed_settings = json.loads(result.settings) assert parsed_settings == settings + @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") def test_create_external_knowledge_api_unicode_handling( - self, mock_check, factory: ExternalDatasetServiceTestDataFactory + self, mock_check, mock_db, factory: ExternalDatasetServiceTestDataFactory ): """Test proper handling of Unicode characters in name and description.""" # Arrange - mock_session = MagicMock() args = { "name": "测试API", "description": "テストの説明", @@ -540,19 +550,21 @@ class TestExternalDatasetServiceCreateAPI: } # Act - result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_session) + result = ExternalDatasetService.create_external_knowledge_api( + "tenant-123", "user-123", args, session=mock_db.session + ) # Assert assert result.name == "测试API" assert result.description == "テストの説明" + @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") def test_create_external_knowledge_api_long_description( - self, mock_check, factory: ExternalDatasetServiceTestDataFactory + self, mock_check, mock_db, factory: ExternalDatasetServiceTestDataFactory ): """Test creation with very long description.""" # Arrange - mock_session = MagicMock() long_description = "A" * 1000 args = { "name": "Test API", @@ -561,7 +573,9 @@ class TestExternalDatasetServiceCreateAPI: } # Act - result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_session) + result = ExternalDatasetService.create_external_knowledge_api( + "tenant-123", "user-123", args, session=mock_db.session + ) # Assert assert result.description == long_description @@ -824,43 +838,43 @@ class TestExternalDatasetServiceCheckEndpoint: class TestExternalDatasetServiceGetAPI: """Test get_external_knowledge_api operations.""" - def test_get_external_knowledge_api_success(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_get_external_knowledge_api_success(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test successful retrieval of external knowledge API.""" # Arrange - mock_session = MagicMock() api_id = "api-123" expected_api = factory.create_external_knowledge_api_mock(api_id=api_id) - mock_session.scalar.return_value = expected_api + mock_db.session.scalar.return_value = expected_api # Act tenant_id = "tenant-123" - result = ExternalDatasetService.get_external_knowledge_api(mock_session, api_id, tenant_id) + result = ExternalDatasetService.get_external_knowledge_api(api_id, tenant_id, session=mock_db.session) # Assert assert result.id == api_id - def test_get_external_knowledge_api_not_found(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_get_external_knowledge_api_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test error when API is not found.""" # Arrange - mock_session = MagicMock() - mock_session.scalar.return_value = None + mock_db.session.scalar.return_value = None # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.get_external_knowledge_api(mock_session, "nonexistent-id", "tenant-123") + ExternalDatasetService.get_external_knowledge_api("nonexistent-id", "tenant-123", session=mock_db.session) class TestExternalDatasetServiceUpdateAPI: """Test update_external_knowledge_api operations.""" @patch("services.external_knowledge_service.naive_utc_now") + @patch("services.external_knowledge_service.db") def test_update_external_knowledge_api_success_all_fields( - self, mock_now, factory: ExternalDatasetServiceTestDataFactory + self, mock_db, mock_now, factory: ExternalDatasetServiceTestDataFactory ): """Test successful update with all fields.""" # Arrange - mock_session = MagicMock() api_id = "api-123" tenant_id = "tenant-123" user_id = "user-456" @@ -875,24 +889,26 @@ class TestExternalDatasetServiceUpdateAPI: "settings": {"endpoint": "https://new.example.com", "api_key": "new-key"}, } - mock_session.scalar.return_value = existing_api + mock_db.session.scalar.return_value = existing_api # Act - result = ExternalDatasetService.update_external_knowledge_api(mock_session, tenant_id, user_id, api_id, args) + result = ExternalDatasetService.update_external_knowledge_api( + tenant_id, user_id, api_id, args, session=mock_db.session + ) # Assert assert result.name == "Updated API" assert result.description == "Updated description" assert result.updated_by == user_id assert result.updated_at == current_time - mock_session.commit.assert_called_once() + mock_db.session.commit.assert_called_once() + @patch("services.external_knowledge_service.db") def test_update_external_knowledge_api_preserve_hidden_api_key( - self, factory: ExternalDatasetServiceTestDataFactory + self, mock_db, factory: ExternalDatasetServiceTestDataFactory ): """Test that hidden API key is preserved from existing settings.""" # Arrange - mock_session = MagicMock() api_id = "api-123" tenant_id = "tenant-123" @@ -907,47 +923,51 @@ class TestExternalDatasetServiceUpdateAPI: "settings": {"endpoint": "https://api.example.com", "api_key": HIDDEN_VALUE}, } - mock_session.scalar.return_value = existing_api + mock_db.session.scalar.return_value = existing_api # Act - result = ExternalDatasetService.update_external_knowledge_api(mock_session, tenant_id, "user-123", api_id, args) + result = ExternalDatasetService.update_external_knowledge_api( + tenant_id, "user-123", api_id, args, session=mock_db.session + ) # Assert settings = json.loads(result.settings) assert settings["api_key"] == "original-secret-key" - def test_update_external_knowledge_api_not_found(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_update_external_knowledge_api_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test error when API is not found.""" # Arrange - mock_session = MagicMock() - mock_session.scalar.return_value = None + mock_db.session.scalar.return_value = None args = {"name": "Updated API"} # Act & Assert with pytest.raises(ValueError, match="api template not found"): ExternalDatasetService.update_external_knowledge_api( - mock_session, "tenant-123", "user-123", "api-123", args + "tenant-123", "user-123", "api-123", args, session=mock_db.session ) - def test_update_external_knowledge_api_tenant_mismatch(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_update_external_knowledge_api_tenant_mismatch( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test error when tenant ID doesn't match.""" # Arrange - mock_session = MagicMock() - mock_session.scalar.return_value = None + mock_db.session.scalar.return_value = None args = {"name": "Updated API"} # Act & Assert with pytest.raises(ValueError, match="api template not found"): ExternalDatasetService.update_external_knowledge_api( - mock_session, "wrong-tenant", "user-123", "api-123", args + "wrong-tenant", "user-123", "api-123", args, session=mock_db.session ) - def test_update_external_knowledge_api_name_only(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_update_external_knowledge_api_name_only(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test updating only the name field.""" # Arrange - mock_session = MagicMock() existing_api = factory.create_external_knowledge_api_mock( description="Original description", settings={"endpoint": "https://api.example.com", "api_key": "key"}, @@ -955,11 +975,11 @@ class TestExternalDatasetServiceUpdateAPI: args = {"name": "New Name Only"} - mock_session.scalar.return_value = existing_api + mock_db.session.scalar.return_value = existing_api # Act result = ExternalDatasetService.update_external_knowledge_api( - mock_session, "tenant-123", "user-123", "api-123", args + "tenant-123", "user-123", "api-123", args, session=mock_db.session ) # Assert @@ -969,92 +989,104 @@ class TestExternalDatasetServiceUpdateAPI: class TestExternalDatasetServiceDeleteAPI: """Test delete_external_knowledge_api operations.""" - def test_delete_external_knowledge_api_success(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_delete_external_knowledge_api_success(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test successful deletion of external knowledge API.""" # Arrange - mock_session = MagicMock() api_id = "api-123" tenant_id = "tenant-123" existing_api = factory.create_external_knowledge_api_mock(api_id=api_id, tenant_id=tenant_id) - mock_session.scalar.return_value = existing_api + mock_db.session.scalar.return_value = existing_api # Act - ExternalDatasetService.delete_external_knowledge_api(mock_session, tenant_id, api_id) + ExternalDatasetService.delete_external_knowledge_api(tenant_id, api_id, session=mock_db.session) # Assert - mock_session.delete.assert_called_once_with(existing_api) - mock_session.commit.assert_called_once() + mock_db.session.delete.assert_called_once_with(existing_api) + mock_db.session.commit.assert_called_once() - def test_delete_external_knowledge_api_not_found(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_delete_external_knowledge_api_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test error when API is not found.""" # Arrange - mock_session = MagicMock() - mock_session.scalar.return_value = None + mock_db.session.scalar.return_value = None # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.delete_external_knowledge_api(mock_session, "tenant-123", "api-123") + ExternalDatasetService.delete_external_knowledge_api("tenant-123", "api-123", session=mock_db.session) - def test_delete_external_knowledge_api_tenant_mismatch(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_delete_external_knowledge_api_tenant_mismatch( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test error when tenant ID doesn't match.""" # Arrange - mock_session = MagicMock() - mock_session.scalar.return_value = None + mock_db.session.scalar.return_value = None # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.delete_external_knowledge_api(mock_session, "wrong-tenant", "api-123") + ExternalDatasetService.delete_external_knowledge_api("wrong-tenant", "api-123", session=mock_db.session) class TestExternalDatasetServiceAPIUseCheck: """Test external_knowledge_api_use_check operations.""" - def test_external_knowledge_api_use_check_in_use_single(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_external_knowledge_api_use_check_in_use_single( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test API use check when API has one binding.""" # Arrange - mock_session = MagicMock() api_id = "api-123" tenant_id = "tenant-123" - mock_session.scalar.return_value = 1 + mock_db.session.scalar.return_value = 1 # Act - in_use, count = ExternalDatasetService.external_knowledge_api_use_check(mock_session, api_id, tenant_id) + in_use, count = ExternalDatasetService.external_knowledge_api_use_check( + api_id, tenant_id, session=mock_db.session + ) # Assert assert in_use is True assert count == 1 - assert "tenant_id" in str(mock_session.scalar.call_args.args[0]) + assert "tenant_id" in str(mock_db.session.scalar.call_args.args[0]) - def test_external_knowledge_api_use_check_in_use_multiple(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_external_knowledge_api_use_check_in_use_multiple( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test API use check with multiple bindings.""" # Arrange - mock_session = MagicMock() api_id = "api-123" tenant_id = "tenant-123" - mock_session.scalar.return_value = 10 + mock_db.session.scalar.return_value = 10 # Act - in_use, count = ExternalDatasetService.external_knowledge_api_use_check(mock_session, api_id, tenant_id) + in_use, count = ExternalDatasetService.external_knowledge_api_use_check( + api_id, tenant_id, session=mock_db.session + ) # Assert assert in_use is True assert count == 10 - def test_external_knowledge_api_use_check_not_in_use(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_external_knowledge_api_use_check_not_in_use(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test API use check when API is not in use.""" # Arrange - mock_session = MagicMock() api_id = "api-123" tenant_id = "tenant-123" - mock_session.scalar.return_value = 0 + mock_db.session.scalar.return_value = 0 # Act - in_use, count = ExternalDatasetService.external_knowledge_api_use_check(mock_session, api_id, tenant_id) + in_use, count = ExternalDatasetService.external_knowledge_api_use_check( + api_id, tenant_id, session=mock_db.session + ) # Assert assert in_use is False @@ -1064,46 +1096,48 @@ class TestExternalDatasetServiceAPIUseCheck: class TestExternalDatasetServiceGetBinding: """Test get_external_knowledge_binding_with_dataset_id operations.""" - def test_get_external_knowledge_binding_success(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_get_external_knowledge_binding_success(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test successful retrieval of external knowledge binding.""" # Arrange - mock_session = MagicMock() tenant_id = "tenant-123" dataset_id = "dataset-123" expected_binding = factory.create_external_knowledge_binding_mock(tenant_id=tenant_id, dataset_id=dataset_id) - mock_session.scalar.return_value = expected_binding + mock_db.session.scalar.return_value = expected_binding # Act result = ExternalDatasetService.get_external_knowledge_binding_with_dataset_id( - mock_session, tenant_id, dataset_id + tenant_id, dataset_id, session=mock_db.session ) # Assert assert result.dataset_id == dataset_id assert result.tenant_id == tenant_id - def test_get_external_knowledge_binding_not_found(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_get_external_knowledge_binding_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test error when binding is not found.""" # Arrange - mock_session = MagicMock() - mock_session.scalar.return_value = None + mock_db.session.scalar.return_value = None # Act & Assert with pytest.raises(ValueError, match="external knowledge binding not found"): ExternalDatasetService.get_external_knowledge_binding_with_dataset_id( - mock_session, "tenant-123", "dataset-123" + "tenant-123", "dataset-123", session=mock_db.session ) class TestExternalDatasetServiceDocumentValidate: """Test document_create_args_validate operations.""" - def test_document_create_args_validate_success_all_params(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_document_create_args_validate_success_all_params( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test successful validation with all required parameters.""" # Arrange - mock_session = MagicMock() tenant_id = "tenant-123" api_id = "api-123" @@ -1117,17 +1151,21 @@ class TestExternalDatasetServiceDocumentValidate: api = factory.create_external_knowledge_api_mock(api_id=api_id, settings=[settings]) - mock_session.scalar.return_value = api + mock_db.session.scalar.return_value = api process_parameter = {"param1": "value1", "param2": "value2"} # Act & Assert - should not raise - ExternalDatasetService.document_create_args_validate(mock_session, tenant_id, api_id, process_parameter) + ExternalDatasetService.document_create_args_validate( + tenant_id, api_id, process_parameter, session=mock_db.session + ) - def test_document_create_args_validate_missing_required_param(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_document_create_args_validate_missing_required_param( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test validation fails when required parameter is missing.""" # Arrange - mock_session = MagicMock() tenant_id = "tenant-123" api_id = "api-123" @@ -1135,42 +1173,46 @@ class TestExternalDatasetServiceDocumentValidate: api = factory.create_external_knowledge_api_mock(api_id=api_id, settings=[settings]) - mock_session.scalar.return_value = api + mock_db.session.scalar.return_value = api process_parameter = {} # Act & Assert with pytest.raises(ValueError, match="required_param is required"): - ExternalDatasetService.document_create_args_validate(mock_session, tenant_id, api_id, process_parameter) + ExternalDatasetService.document_create_args_validate( + tenant_id, api_id, process_parameter, session=mock_db.session + ) - def test_document_create_args_validate_api_not_found(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_document_create_args_validate_api_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test validation fails when API is not found.""" # Arrange - mock_session = MagicMock() - mock_session.scalar.return_value = None + mock_db.session.scalar.return_value = None # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.document_create_args_validate(mock_session, "tenant-123", "api-123", {}) + ExternalDatasetService.document_create_args_validate("tenant-123", "api-123", {}, session=mock_db.session) - def test_document_create_args_validate_no_custom_parameters(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_document_create_args_validate_no_custom_parameters( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test validation succeeds when no custom parameters defined.""" # Arrange - mock_session = MagicMock() settings = {} api = factory.create_external_knowledge_api_mock(settings=[settings]) - mock_session.scalar.return_value = api + mock_db.session.scalar.return_value = api # Act & Assert - should not raise - ExternalDatasetService.document_create_args_validate(mock_session, "tenant-123", "api-123", {}) + ExternalDatasetService.document_create_args_validate("tenant-123", "api-123", {}, session=mock_db.session) + @patch("services.external_knowledge_service.db") def test_document_create_args_validate_optional_params_not_required( - self, factory: ExternalDatasetServiceTestDataFactory + self, mock_db, factory: ExternalDatasetServiceTestDataFactory ): """Test that optional parameters don't cause validation failure.""" # Arrange - mock_session = MagicMock() settings = { "document_process_setting": [ {"name": "required_param", "required": True}, @@ -1180,12 +1222,14 @@ class TestExternalDatasetServiceDocumentValidate: api = factory.create_external_knowledge_api_mock(settings=[settings]) - mock_session.scalar.return_value = api + mock_db.session.scalar.return_value = api process_parameter = {"required_param": "value"} # Act & Assert - should not raise - ExternalDatasetService.document_create_args_validate(mock_session, "tenant-123", "api-123", process_parameter) + ExternalDatasetService.document_create_args_validate( + "tenant-123", "api-123", process_parameter, session=mock_db.session + ) class TestExternalDatasetServiceProcessAPI: @@ -1475,10 +1519,10 @@ class TestExternalDatasetServiceGetSettings: class TestExternalDatasetServiceCreateDataset: """Test create_external_dataset operations.""" - def test_create_external_dataset_success_full(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_create_external_dataset_success_full(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test successful creation of external dataset with all fields.""" # Arrange - mock_session = MagicMock() tenant_id = "tenant-123" user_id = "user-123" args = { @@ -1491,84 +1535,90 @@ class TestExternalDatasetServiceCreateDataset: api = factory.create_external_knowledge_api_mock(api_id="api-123") - mock_session.scalar.side_effect = [None, api] + mock_db.session.scalar.side_effect = [None, api] # Act - result = ExternalDatasetService.create_external_dataset(tenant_id, user_id, args, mock_session) + result = ExternalDatasetService.create_external_dataset(tenant_id, user_id, args, session=mock_db.session) # Assert assert result.name == "Test External Dataset" assert result.description == "Comprehensive test description" assert result.provider == "external" assert result.created_by == user_id - mock_session.add.assert_called() - mock_session.commit.assert_called_once() + mock_db.session.add.assert_called() + mock_db.session.commit.assert_called_once() - def test_create_external_dataset_duplicate_name_error(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_create_external_dataset_duplicate_name_error( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test error when dataset name already exists.""" # Arrange - mock_session = MagicMock() existing_dataset = factory.create_dataset_mock(name="Duplicate Dataset") - mock_session.scalar.return_value = existing_dataset + mock_db.session.scalar.return_value = existing_dataset args = {"name": "Duplicate Dataset"} # Act & Assert with pytest.raises(DatasetNameDuplicateError): - ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, mock_session) + ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, session=mock_db.session) - def test_create_external_dataset_api_not_found_error(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_create_external_dataset_api_not_found_error(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test error when external knowledge API is not found.""" # Arrange - mock_session = MagicMock() - mock_session.scalar.side_effect = [None, None] + mock_db.session.scalar.side_effect = [None, None] args = {"name": "Test Dataset", "external_knowledge_api_id": "nonexistent-api"} # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, mock_session) + ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, session=mock_db.session) - def test_create_external_dataset_missing_knowledge_id_error(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_create_external_dataset_missing_knowledge_id_error( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test error when external_knowledge_id is missing.""" # Arrange - mock_session = MagicMock() api = factory.create_external_knowledge_api_mock() - mock_session.scalar.side_effect = [None, api] + mock_db.session.scalar.side_effect = [None, api] args = {"name": "Test Dataset", "external_knowledge_api_id": "api-123"} # Act & Assert with pytest.raises(ValueError, match="external_knowledge_id is required"): - ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, mock_session) + ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, session=mock_db.session) - def test_create_external_dataset_missing_api_id_error(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_create_external_dataset_missing_api_id_error( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test error when external_knowledge_api_id is missing.""" # Arrange - mock_session = MagicMock() api = factory.create_external_knowledge_api_mock() - mock_session.scalar.side_effect = [None, api] + mock_db.session.scalar.side_effect = [None, api] args = {"name": "Test Dataset", "external_knowledge_id": "knowledge-123"} # Act & Assert with pytest.raises(ValueError, match="external_knowledge_api_id is required"): - ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, mock_session) + ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, session=mock_db.session) class TestExternalDatasetServiceFetchRetrieval: """Test fetch_external_knowledge_retrieval operations.""" @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") + @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_success_with_results( - self, mock_process, factory: ExternalDatasetServiceTestDataFactory + self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory ): """Test successful external knowledge retrieval with results.""" # Arrange - mock_session = MagicMock() tenant_id = "tenant-123" dataset_id = "dataset-123" query = "test query for retrieval" @@ -1578,7 +1628,7 @@ class TestExternalDatasetServiceFetchRetrieval: ) api = factory.create_external_knowledge_api_mock(api_id="api-123") - mock_session.scalar.side_effect = [binding, api] + mock_db.session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 200 @@ -1594,7 +1644,11 @@ class TestExternalDatasetServiceFetchRetrieval: # Act result = ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, tenant_id, dataset_id, query, external_retrieval_parameters + tenant_id, + dataset_id, + query, + external_retrieval_parameters, + session=mock_db.session, ) # Assert @@ -1602,46 +1656,46 @@ class TestExternalDatasetServiceFetchRetrieval: assert result[0]["content"] == "result 1" assert result[1]["score"] == 0.8 + @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_binding_not_found_error( - self, factory: ExternalDatasetServiceTestDataFactory + self, mock_db, factory: ExternalDatasetServiceTestDataFactory ): """Test error when external knowledge binding is not found.""" # Arrange - mock_session = MagicMock() - mock_session.scalar.return_value = None + mock_db.session.scalar.return_value = None # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError, match="external knowledge binding not found"): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, "tenant-123", "dataset-123", "query", {} + "tenant-123", "dataset-123", "query", {}, session=mock_db.session ) + @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_cross_tenant_api_template_error( - self, factory: ExternalDatasetServiceTestDataFactory + self, mock_db, factory: ExternalDatasetServiceTestDataFactory ): """Test error when a binding points to an API template outside the dataset tenant.""" # Arrange - mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() - mock_session.scalar.side_effect = [binding, None] + mock_db.session.scalar.side_effect = [binding, None] # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError, match="external api template not found"): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, "tenant-123", "dataset-123", "query", {} + "tenant-123", "dataset-123", "query", {}, session=mock_db.session ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") + @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_empty_results( - self, mock_process, factory: ExternalDatasetServiceTestDataFactory + self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory ): """Test retrieval with empty results.""" # Arrange - mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_session.scalar.side_effect = [binding, api] + mock_db.session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 200 @@ -1650,23 +1704,27 @@ class TestExternalDatasetServiceFetchRetrieval: # Act result = ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, "tenant-123", "dataset-123", "query", {"top_k": 5} + "tenant-123", + "dataset-123", + "query", + {"top_k": 5}, + session=mock_db.session, ) # Assert assert len(result) == 0 @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") + @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_with_score_threshold( - self, mock_process, factory: ExternalDatasetServiceTestDataFactory + self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory ): """Test retrieval with score threshold enabled.""" # Arrange - mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_session.scalar.side_effect = [binding, api] + mock_db.session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 200 @@ -1681,7 +1739,11 @@ class TestExternalDatasetServiceFetchRetrieval: # Act result = ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, "tenant-123", "dataset-123", "query", external_retrieval_parameters + "tenant-123", + "dataset-123", + "query", + external_retrieval_parameters, + session=mock_db.session, ) # Assert @@ -1691,16 +1753,16 @@ class TestExternalDatasetServiceFetchRetrieval: assert call_args.params["retrieval_setting"]["score_threshold"] == 0.75 @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") + @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_non_200_status_raises_exception( - self, mock_process, factory: ExternalDatasetServiceTestDataFactory + self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory ): """Test that non-200 status code raises Exception with response text.""" # Arrange - mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_session.scalar.side_effect = [binding, api] + mock_db.session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 500 @@ -1710,7 +1772,11 @@ class TestExternalDatasetServiceFetchRetrieval: # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError, match="Internal Server Error: Database connection failed"): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, "tenant-123", "dataset-123", "query", {"top_k": 5} + "tenant-123", + "dataset-123", + "query", + {"top_k": 5}, + session=mock_db.session, ) @pytest.mark.parametrize( @@ -1727,12 +1793,12 @@ class TestExternalDatasetServiceFetchRetrieval: ], ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") + @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_various_error_status_codes( - self, mock_process, factory: ExternalDatasetServiceTestDataFactory, status_code, error_message + self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory, status_code, error_message ): """Test that various error status codes raise exceptions with response text.""" # Arrange - mock_session = MagicMock() tenant_id = "tenant-123" dataset_id = "dataset-123" @@ -1741,7 +1807,7 @@ class TestExternalDatasetServiceFetchRetrieval: ) api = factory.create_external_knowledge_api_mock(api_id="api-123") - mock_session.scalar.side_effect = [binding, api] + mock_db.session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = status_code @@ -1751,20 +1817,20 @@ class TestExternalDatasetServiceFetchRetrieval: # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError, match=re.escape(error_message)): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, tenant_id, dataset_id, "query", {"top_k": 5} + tenant_id, dataset_id, "query", {"top_k": 5}, session=mock_db.session ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") + @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_empty_response_text( - self, mock_process, factory: ExternalDatasetServiceTestDataFactory + self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory ): """Test exception with empty response text.""" # Arrange - mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_session.scalar.side_effect = [binding, api] + mock_db.session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 503 @@ -1774,17 +1840,21 @@ class TestExternalDatasetServiceFetchRetrieval: # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, "tenant-123", "dataset-123", "query", {"top_k": 5} + "tenant-123", + "dataset-123", + "query", + {"top_k": 5}, + session=mock_db.session, ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - def test_fetch_external_knowledge_retrieval_invalid_json_response(self, mock_process, factory): + @patch("services.external_knowledge_service.db") + def test_fetch_external_knowledge_retrieval_invalid_json_response(self, mock_db, mock_process, factory): """Test malformed JSON success responses are normalized to external retrieval errors.""" - mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_session.scalar.side_effect = [binding, api] + mock_db.session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 200 @@ -1793,17 +1863,21 @@ class TestExternalDatasetServiceFetchRetrieval: with pytest.raises(ExternalKnowledgeRetrievalError, match="invalid external knowledge response"): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, "tenant-123", "dataset-123", "query", {"top_k": 5} + "tenant-123", + "dataset-123", + "query", + {"top_k": 5}, + session=mock_db.session, ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - def test_fetch_external_knowledge_retrieval_invalid_success_payload_shape(self, mock_process, factory): + @patch("services.external_knowledge_service.db") + def test_fetch_external_knowledge_retrieval_invalid_success_payload_shape(self, mock_db, mock_process, factory): """Test malformed success payload shapes are normalized to external retrieval errors.""" - mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_session.scalar.side_effect = [binding, api] + mock_db.session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 200 @@ -1812,17 +1886,21 @@ class TestExternalDatasetServiceFetchRetrieval: with pytest.raises(ExternalKnowledgeRetrievalError, match="invalid external knowledge response"): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, "tenant-123", "dataset-123", "query", {"top_k": 5} + "tenant-123", + "dataset-123", + "query", + {"top_k": 5}, + session=mock_db.session, ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - def test_fetch_external_knowledge_retrieval_invalid_records_shape(self, mock_process, factory): + @patch("services.external_knowledge_service.db") + def test_fetch_external_knowledge_retrieval_invalid_records_shape(self, mock_db, mock_process, factory): """Test non-list records payloads are normalized to external retrieval errors.""" - mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_session.scalar.side_effect = [binding, api] + mock_db.session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 200 @@ -1831,20 +1909,28 @@ class TestExternalDatasetServiceFetchRetrieval: with pytest.raises(ExternalKnowledgeRetrievalError, match="invalid external knowledge response"): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, "tenant-123", "dataset-123", "query", {"top_k": 5} + "tenant-123", + "dataset-123", + "query", + {"top_k": 5}, + session=mock_db.session, ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - def test_fetch_external_knowledge_retrieval_wraps_transport_errors(self, mock_process, factory): + @patch("services.external_knowledge_service.db") + def test_fetch_external_knowledge_retrieval_wraps_transport_errors(self, mock_db, mock_process, factory): """Test transport/runtime failures are normalized to external retrieval errors.""" - mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_session.scalar.side_effect = [binding, api] + mock_db.session.scalar.side_effect = [binding, api] mock_process.side_effect = RuntimeError("connection reset by peer") with pytest.raises(ExternalKnowledgeRetrievalError, match="connection reset by peer"): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, "tenant-123", "dataset-123", "query", {"top_k": 5} + "tenant-123", + "dataset-123", + "query", + {"top_k": 5}, + session=mock_db.session, ) diff --git a/api/tests/unit_tests/services/test_file_service.py b/api/tests/unit_tests/services/test_file_service.py index b81fb823949..41b86fda0cb 100644 --- a/api/tests/unit_tests/services/test_file_service.py +++ b/api/tests/unit_tests/services/test_file_service.py @@ -377,7 +377,7 @@ class TestFileService: def test_get_upload_files_by_ids_empty(self): session = MagicMock() - result = FileService.get_upload_files_by_ids(session, "tenant_id", []) + result = FileService.get_upload_files_by_ids("tenant_id", [], session=session) assert result == {} def test_get_upload_files_by_ids(self): @@ -387,7 +387,9 @@ class TestFileService: session = MagicMock() session.scalars().all.return_value = [upload_file] - result = FileService.get_upload_files_by_ids(session, "tenant_id", ["550e8400-e29b-41d4-a716-446655440000"]) + result = FileService.get_upload_files_by_ids( + "tenant_id", ["550e8400-e29b-41d4-a716-446655440000"], session=session + ) assert result["550e8400-e29b-41d4-a716-446655440000"] == upload_file def test_sanitize_zip_entry_name(self): diff --git a/api/tests/unit_tests/services/test_message_service.py b/api/tests/unit_tests/services/test_message_service.py index 6588c8a8de6..13f340e9f4a 100644 --- a/api/tests/unit_tests/services/test_message_service.py +++ b/api/tests/unit_tests/services/test_message_service.py @@ -102,6 +102,7 @@ class TestMessageServicePaginationByFirstId: conversation_id="conv-001", first_id=None, limit=10, + session=MagicMock(), ) # Assert @@ -124,6 +125,7 @@ class TestMessageServicePaginationByFirstId: conversation_id="", first_id=None, limit=10, + session=MagicMock(), ) # Assert @@ -166,6 +168,7 @@ class TestMessageServicePaginationByFirstId: first_id=None, limit=10, order="desc", + session=mock_db.session, ) # Assert @@ -209,6 +212,7 @@ class TestMessageServicePaginationByFirstId: first_id=None, limit=10, order="asc", + session=mock_db.session, ) # Assert @@ -258,6 +262,7 @@ class TestMessageServicePaginationByFirstId: first_id="msg-005", limit=10, order="desc", + session=mock_db.session, ) # Assert @@ -288,6 +293,7 @@ class TestMessageServicePaginationByFirstId: conversation_id="conv-001", first_id="nonexistent-msg", limit=10, + session=mock_db.session, ) # Test 07: Has_more flag when results exceed limit @@ -323,6 +329,7 @@ class TestMessageServicePaginationByFirstId: conversation_id="conv-001", first_id=None, limit=10, + session=mock_db.session, ) # Assert @@ -353,6 +360,7 @@ class TestMessageServicePaginationByFirstId: conversation_id="conv-001", first_id=None, limit=10, + session=mock_db.session, ) # Assert @@ -389,6 +397,7 @@ class TestMessageServicePaginationByLastId: user=None, last_id=None, limit=10, + session=MagicMock(), ) # Assert @@ -421,6 +430,7 @@ class TestMessageServicePaginationByLastId: user=user, last_id=None, limit=10, + session=mock_db.session, ) # Assert @@ -459,6 +469,7 @@ class TestMessageServicePaginationByLastId: user=user, last_id="msg-005", limit=10, + session=mock_db.session, ) # Assert @@ -482,6 +493,7 @@ class TestMessageServicePaginationByLastId: user=user, last_id="nonexistent-msg", limit=10, + session=mock_db.session, ) # Test 13: Pagination with conversation_id filter @@ -516,6 +528,7 @@ class TestMessageServicePaginationByLastId: last_id=None, limit=10, conversation_id="conv-001", + session=mock_db.session, ) # Assert @@ -546,6 +559,7 @@ class TestMessageServicePaginationByLastId: last_id=None, limit=10, include_ids=["msg-001", "msg-003"], + session=mock_db.session, ) # Assert @@ -578,6 +592,7 @@ class TestMessageServicePaginationByLastId: user=user, last_id=None, limit=10, + session=mock_db.session, ) # Assert @@ -680,8 +695,8 @@ class TestMessageServiceGetMessage: mock_db.session.scalar.return_value = message - # Act - result = MessageService.get_message(app_model=app, user=user, message_id="msg-123") + # Act, + result = MessageService.get_message(app_model=app, user=user, message_id="msg-123", session=mock_db.session) # Assert assert result == message @@ -700,8 +715,8 @@ class TestMessageServiceGetMessage: mock_db.session.scalar.return_value = message - # Act - result = MessageService.get_message(app_model=app, user=user, message_id="msg-123") + # Act, + result = MessageService.get_message(app_model=app, user=user, message_id="msg-123", session=mock_db.session) # Assert assert result == message @@ -718,7 +733,7 @@ class TestMessageServiceGetMessage: # Act & Assert with pytest.raises(MessageNotExistsError): - MessageService.get_message(app_model=app, user=user, message_id="msg-123") + MessageService.get_message(app_model=app, user=user, message_id="msg-123", session=mock_db.session) class TestMessageServiceFeedback: @@ -748,6 +763,7 @@ class TestMessageServiceFeedback: user=user, rating=FeedbackRating.LIKE, content="Good answer", + session=mock_db.session, ) # Assert @@ -780,6 +796,7 @@ class TestMessageServiceFeedback: user=user, rating=FeedbackRating.DISLIKE, content="Bad answer", + session=mock_db.session, ) # Assert @@ -808,6 +825,7 @@ class TestMessageServiceFeedback: user=user, rating=None, content=None, + session=mock_db.session, ) # Assert @@ -826,8 +844,8 @@ class TestMessageServiceFeedback: mock_db.session.scalars.return_value.all.return_value = [feedback] - # Act - result = MessageService.get_all_messages_feedbacks(app_model=app, page=1, limit=10) + # Act, + result = MessageService.get_all_messages_feedbacks(app_model=app, page=1, limit=10, session=mock_db.session) # Assert assert result == [{"id": "fb-1"}] @@ -846,7 +864,11 @@ class TestMessageServiceSuggestedQuestions: app = factory.create_app_mock() with pytest.raises(ValueError, match="user cannot be None"): MessageService.get_suggested_questions_after_answer( - app_model=app, user=None, message_id="msg-123", invoke_from=MagicMock() + app_model=app, + user=None, + message_id="msg-123", + invoke_from=MagicMock(), + session=MagicMock(), ) # Test 28: get_suggested_questions_after_answer - Advanced Chat success @@ -890,7 +912,11 @@ class TestMessageServiceSuggestedQuestions: # Act result = MessageService.get_suggested_questions_after_answer( - app_model=app, user=user, message_id="msg-123", invoke_from=InvokeFrom.WEB_APP + app_model=app, + user=user, + message_id="msg-123", + invoke_from=InvokeFrom.WEB_APP, + session=MagicMock(), ) # Assert @@ -938,7 +964,11 @@ class TestMessageServiceSuggestedQuestions: # Act result = MessageService.get_suggested_questions_after_answer( - app_model=app, user=user, message_id="msg-123", invoke_from=MagicMock() + app_model=app, + user=user, + message_id="msg-123", + invoke_from=MagicMock(), + session=mock_db.session, ) # Assert @@ -996,6 +1026,7 @@ class TestMessageServiceSuggestedQuestions: user=user, message_id="msg-123", invoke_from=InvokeFrom.WEB_APP, + session=mock_db.session, ) assert result == ["Q1?"] @@ -1059,7 +1090,11 @@ class TestMessageServiceSuggestedQuestions: mock_llm_gen.generate_suggested_questions_after_answer.return_value = ["Q1?"] result = MessageService.get_suggested_questions_after_answer( - app_model=app, user=user, message_id="msg-123", invoke_from=MagicMock() + app_model=app, + user=user, + message_id="msg-123", + invoke_from=MagicMock(), + session=mock_db.session, ) assert result == ["Q1?"] @@ -1168,7 +1203,11 @@ class TestMessageServiceSuggestedQuestions: mock_llm_gen.generate_suggested_questions_after_answer.return_value = ["Q1?"] result = MessageService.get_suggested_questions_after_answer( - app_model=app, user=user, message_id="msg-123", invoke_from=MagicMock() + app_model=app, + user=user, + message_id="msg-123", + invoke_from=MagicMock(), + session=mock_db.session, ) assert result == ["Q1?"] @@ -1209,5 +1248,9 @@ class TestMessageServiceSuggestedQuestions: # Act & Assert with pytest.raises(SuggestedQuestionsAfterAnswerDisabledError): MessageService.get_suggested_questions_after_answer( - app_model=app, user=user, message_id="msg-123", invoke_from=MagicMock() + app_model=app, + user=user, + message_id="msg-123", + invoke_from=MagicMock(), + session=MagicMock(), ) diff --git a/api/tests/unit_tests/services/test_metadata_bug_complete.py b/api/tests/unit_tests/services/test_metadata_bug_complete.py index 6792243e9d0..00f16f75ac0 100644 --- a/api/tests/unit_tests/services/test_metadata_bug_complete.py +++ b/api/tests/unit_tests/services/test_metadata_bug_complete.py @@ -48,14 +48,14 @@ class TestMetadataBugCompleteValidation: account = _make_account() # Should crash with TypeError with pytest.raises(TypeError, match="object of type 'NoneType' has no len"): - MetadataService.create_metadata(Mock(), "dataset-123", mock_metadata_args, account, "tenant-123") + MetadataService.create_metadata("dataset-123", mock_metadata_args, account, "tenant-123", session=Mock()) # Test update method as well account = _make_account() none_name = cast(str, None) with pytest.raises(TypeError, match="object of type 'NoneType' has no len"): MetadataService.update_metadata_name( - Mock(), "dataset-123", "metadata-456", none_name, account, "tenant-123" + "dataset-123", "metadata-456", none_name, account, "tenant-123", session=Mock() ) def test_3_database_constraints_verification(self) -> None: @@ -99,7 +99,7 @@ class TestMetadataBugCompleteValidation: account = _make_account() with pytest.raises(TypeError, match="object of type 'NoneType' has no len"): - MetadataService.create_metadata(Mock(), "dataset-123", mock_metadata_args, account, "tenant-123") + MetadataService.create_metadata("dataset-123", mock_metadata_args, account, "tenant-123", session=Mock()) def test_7_end_to_end_validation_layers(self) -> None: """Test all validation layers work together correctly.""" diff --git a/api/tests/unit_tests/services/test_metadata_nullable_bug.py b/api/tests/unit_tests/services/test_metadata_nullable_bug.py index ae93fe5ef51..cfd3d034df2 100644 --- a/api/tests/unit_tests/services/test_metadata_nullable_bug.py +++ b/api/tests/unit_tests/services/test_metadata_nullable_bug.py @@ -37,7 +37,7 @@ class TestMetadataNullableBug: account = _make_account() # This should crash with TypeError when calling len(None) with pytest.raises(TypeError, match="object of type 'NoneType' has no len"): - MetadataService.create_metadata(Mock(), "dataset-123", mock_metadata_args, account, "tenant-123") + MetadataService.create_metadata("dataset-123", mock_metadata_args, account, "tenant-123", session=Mock()) def test_metadata_service_update_with_none_name_crashes(self) -> None: """Test that MetadataService.update_metadata_name crashes when name is None.""" @@ -46,7 +46,7 @@ class TestMetadataNullableBug: # This should crash with TypeError when calling len(None) with pytest.raises(TypeError, match="object of type 'NoneType' has no len"): MetadataService.update_metadata_name( - Mock(), "dataset-123", "metadata-456", none_name, account, "tenant-123" + "dataset-123", "metadata-456", none_name, account, "tenant-123", session=Mock() ) def test_api_layer_now_uses_pydantic_validation(self) -> None: diff --git a/api/tests/unit_tests/services/test_model_load_balancing_service.py b/api/tests/unit_tests/services/test_model_load_balancing_service.py index 827567f1afe..743e6e797a3 100644 --- a/api/tests/unit_tests/services/test_model_load_balancing_service.py +++ b/api/tests/unit_tests/services/test_model_load_balancing_service.py @@ -80,9 +80,9 @@ def service(mocker: MockerFixture) -> ModelLoadBalancingService: @pytest.fixture -def mock_db(mocker: MockerFixture) -> MagicMock: +def mock_db() -> MagicMock: # Arrange - mocked_db = mocker.patch("services.model_load_balancing_service.db") + mocked_db = MagicMock() mocked_db.session = MagicMock() return mocked_db @@ -159,7 +159,7 @@ def test_get_load_balancing_configs_should_raise_value_error_when_provider_missi # Act + Assert with pytest.raises(ValueError, match="Provider openai does not exist"): - service.get_load_balancing_configs("tenant-1", "openai", "gpt-4o-mini", ModelType.LLM) + service.get_load_balancing_configs("tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, session=MagicMock()) def test_get_load_balancing_configs_should_insert_inherit_config_when_missing_for_custom_provider( @@ -201,6 +201,7 @@ def test_get_load_balancing_configs_should_insert_inherit_config_when_missing_fo "openai", "gpt-4o-mini", ModelType.LLM, + session=mock_db.session, ) # Assert @@ -263,6 +264,7 @@ def test_get_load_balancing_configs_should_reorder_existing_inherit_and_tolerate "gpt-4o-mini", ModelType.LLM, config_from="predefined-model", + session=mock_db.session, ) # Assert @@ -282,7 +284,9 @@ def test_get_load_balancing_config_should_raise_value_error_when_provider_missin # Act + Assert with pytest.raises(ValueError, match="Provider openai does not exist"): - service.get_load_balancing_config("tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, "cfg-1") + service.get_load_balancing_config( + "tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, "cfg-1", session=MagicMock() + ) def test_get_load_balancing_config_should_return_none_when_config_not_found( @@ -295,7 +299,9 @@ def test_get_load_balancing_config_should_return_none_when_config_not_found( mock_db.session.scalar.return_value = None # Act - result = service.get_load_balancing_config("tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, "cfg-1") + result = service.get_load_balancing_config( + "tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, "cfg-1", session=mock_db.session + ) # Assert assert result is None @@ -315,7 +321,9 @@ def test_get_load_balancing_config_should_return_obfuscated_payload_when_config_ mock_db.session.scalar.return_value = config # Act - result = service.get_load_balancing_config("tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, "cfg-1") + result = service.get_load_balancing_config( + "tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, "cfg-1", session=mock_db.session + ) # Assert assert result == { @@ -334,7 +342,9 @@ def test_init_inherit_config_should_create_and_persist_inherit_configuration( model_type = ModelType.LLM # Act - inherit_config = service._init_inherit_config("tenant-1", "openai", "gpt-4o-mini", model_type) + inherit_config = service._init_inherit_config( + "tenant-1", "openai", "gpt-4o-mini", model_type, session=mock_db.session + ) # Assert assert inherit_config.tenant_id == "tenant-1" @@ -361,6 +371,7 @@ def test_update_load_balancing_configs_should_raise_value_error_when_provider_mi ModelType.LLM, [], "custom-model", + session=MagicMock(), ) @@ -380,6 +391,7 @@ def test_update_load_balancing_configs_should_raise_value_error_when_configs_is_ ModelType.LLM, cast(list[dict[str, object]], "invalid-configs"), "custom-model", + session=MagicMock(), ) @@ -401,6 +413,7 @@ def test_update_load_balancing_configs_should_raise_value_error_when_config_item ModelType.LLM, cast(list[dict[str, object]], ["bad-item"]), "custom-model", + session=mock_db.session, ) @@ -423,6 +436,7 @@ def test_update_load_balancing_configs_should_raise_value_error_when_credential_ ModelType.LLM, [{"credential_id": "cred-1", "enabled": True}], "predefined-model", + session=mock_db.session, ) @@ -444,6 +458,7 @@ def test_update_load_balancing_configs_should_raise_value_error_when_name_or_ena ModelType.LLM, [{"enabled": True}], "custom-model", + session=mock_db.session, ) with pytest.raises(ValueError, match="Invalid load balancing config enabled"): @@ -454,6 +469,7 @@ def test_update_load_balancing_configs_should_raise_value_error_when_name_or_ena ModelType.LLM, [{"name": "cfg-without-enabled"}], "custom-model", + session=mock_db.session, ) @@ -476,6 +492,7 @@ def test_update_load_balancing_configs_should_raise_value_error_when_existing_co ModelType.LLM, [{"id": "cfg-2", "name": "invalid", "enabled": True}], "custom-model", + session=mock_db.session, ) @@ -498,6 +515,7 @@ def test_update_load_balancing_configs_should_raise_value_error_when_credentials ModelType.LLM, [{"id": "cfg-1", "name": "new", "enabled": True, "credentials": "bad"}], "custom-model", + session=mock_db.session, ) with pytest.raises(ValueError, match="Invalid load balancing config credentials"): @@ -508,6 +526,7 @@ def test_update_load_balancing_configs_should_raise_value_error_when_credentials ModelType.LLM, [{"name": "new-config", "enabled": True, "credentials": "bad"}], "custom-model", + session=mock_db.session, ) @@ -548,6 +567,7 @@ def test_update_load_balancing_configs_should_update_existing_create_new_and_del {"name": "new-config", "enabled": True, "credentials": {"api_key": "plain"}}, ], "custom-model", + session=mock_db.session, ) # Assert @@ -579,6 +599,7 @@ def test_update_load_balancing_configs_should_raise_value_error_for_invalid_new_ ModelType.LLM, [{"name": "__inherit__", "enabled": True, "credentials": {"api_key": "x"}}], "custom-model", + session=mock_db.session, ) with pytest.raises(ValueError, match="Invalid load balancing config credentials"): @@ -589,6 +610,7 @@ def test_update_load_balancing_configs_should_raise_value_error_for_invalid_new_ ModelType.LLM, [{"name": "new", "enabled": True}], "custom-model", + session=mock_db.session, ) @@ -611,6 +633,7 @@ def test_update_load_balancing_configs_should_create_from_existing_provider_cred ModelType.LLM, [{"credential_id": "cred-1", "enabled": True}], "predefined-model", + session=mock_db.session, ) # Assert @@ -636,6 +659,7 @@ def test_validate_load_balancing_credentials_should_raise_value_error_when_provi "gpt-4o-mini", ModelType.LLM, {"api_key": "plain"}, + session=MagicMock(), ) @@ -657,6 +681,7 @@ def test_validate_load_balancing_credentials_should_raise_value_error_when_confi ModelType.LLM, {"api_key": "plain"}, config_id="cfg-1", + session=mock_db.session, ) @@ -680,6 +705,7 @@ def test_validate_load_balancing_credentials_should_delegate_to_custom_validate_ ModelType.LLM, {"api_key": "plain"}, config_id="cfg-1", + session=mock_db.session, ) service.validate_load_balancing_credentials( "tenant-1", @@ -687,6 +713,7 @@ def test_validate_load_balancing_credentials_should_delegate_to_custom_validate_ "gpt-4o-mini", ModelType.LLM, {"api_key": "plain"}, + session=mock_db.session, ) # Assert diff --git a/api/tests/unit_tests/services/test_oauth_device_flow.py b/api/tests/unit_tests/services/test_oauth_device_flow.py index fcb3f29a76f..00b2919240d 100644 --- a/api/tests/unit_tests/services/test_oauth_device_flow.py +++ b/api/tests/unit_tests/services/test_oauth_device_flow.py @@ -83,7 +83,7 @@ def test_revoke_oauth_token_invalidates_redis_cache_when_live_hash_seen(): redis = MagicMock() - revoke_oauth_token(session, redis, "token-id") + revoke_oauth_token(redis, "token-id", session=session) assert session.execute.called # UPDATE ... WHERE revoked_at IS NULL assert session.commit.called @@ -101,7 +101,7 @@ def test_revoke_oauth_token_is_idempotent_when_already_revoked(): redis = MagicMock() - revoke_oauth_token(session, redis, "token-id") + revoke_oauth_token(redis, "token-id", session=session) assert session.execute.called assert session.commit.called @@ -126,7 +126,7 @@ def test_list_active_sessions_returns_session_execute_rows(): fake_rows = [MagicMock(), MagicMock()] session.execute.return_value.scalars.return_value.all.return_value = fake_rows - out = list_active_sessions(session, _account_ctx(), datetime.now(UTC)) + out = list_active_sessions(_account_ctx(), datetime.now(UTC), session=session) assert out == fake_rows assert session.execute.called @@ -136,11 +136,11 @@ def test_token_belongs_to_subject_true_when_row_present(): session = MagicMock() session.execute.return_value.first.return_value = ("some-id",) - assert token_belongs_to_subject(session, "token-id", _account_ctx()) is True + assert token_belongs_to_subject("token-id", _account_ctx(), session=session) is True def test_token_belongs_to_subject_false_when_no_row(): session = MagicMock() session.execute.return_value.first.return_value = None - assert token_belongs_to_subject(session, "token-id", _account_ctx()) is False + assert token_belongs_to_subject("token-id", _account_ctx(), session=session) is False diff --git a/api/tests/unit_tests/services/test_summary_index_service.py b/api/tests/unit_tests/services/test_summary_index_service.py index 19418c43926..7ece6204ce3 100644 --- a/api/tests/unit_tests/services/test_summary_index_service.py +++ b/api/tests/unit_tests/services/test_summary_index_service.py @@ -118,7 +118,7 @@ def test_generate_summary_for_segment_raises_when_empty(monkeypatch: pytest.Monk SummaryIndexService.generate_summary_for_segment(_segment(), _dataset(), {"a": 1}) -def test_create_summary_record_updates_existing_and_reenables(monkeypatch: pytest.MonkeyPatch) -> None: +def test_create_summary_record_updates_existing_and_reenables() -> None: existing = _summary_record(summary_content="old", node_id="n1") existing.enabled = False existing.disabled_at = datetime(2024, 1, 1) @@ -127,13 +127,12 @@ def test_create_summary_record_updates_existing_and_reenables(monkeypatch: pytes session = MagicMock(name="session") session.scalar.return_value = existing - create_session_mock = MagicMock(return_value=_SessionContext(session)) - monkeypatch.setattr(summary_module, "session_factory", SimpleNamespace(create_session=create_session_mock)) - segment = _segment() dataset = _dataset() - result = SummaryIndexService.create_summary_record(segment, dataset, "new", status=SummaryStatus.GENERATING) + result = SummaryIndexService.create_summary_record( + segment, dataset, "new", status=SummaryStatus.GENERATING, session=session + ) assert result is existing assert existing.summary_content == "new" assert existing.status == SummaryStatus.GENERATING @@ -145,14 +144,13 @@ def test_create_summary_record_updates_existing_and_reenables(monkeypatch: pytes session.flush.assert_called_once() -def test_create_summary_record_creates_new(monkeypatch: pytest.MonkeyPatch) -> None: +def test_create_summary_record_creates_new() -> None: session = MagicMock(name="session") session.scalar.return_value = None - create_session_mock = MagicMock(return_value=_SessionContext(session)) - monkeypatch.setattr(summary_module, "session_factory", SimpleNamespace(create_session=create_session_mock)) - - record = SummaryIndexService.create_summary_record(_segment(), _dataset(), "new", status=SummaryStatus.GENERATING) + record = SummaryIndexService.create_summary_record( + _segment(), _dataset(), "new", status=SummaryStatus.GENERATING, session=session + ) assert record.dataset_id == "dataset-1" assert record.chunk_id == "seg-1" assert record.summary_content == "new" @@ -331,17 +329,12 @@ def test_generate_and_vectorize_summary_success(monkeypatch: pytest.MonkeyPatch) session = MagicMock() session.scalar.return_value = record - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) monkeypatch.setattr( SummaryIndexService, "generate_summary_for_segment", MagicMock(return_value=("sum", MagicMock(total_tokens=0))) ) monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(return_value=None)) - out = SummaryIndexService.generate_and_vectorize_summary(segment, dataset, {"enable": True}) + out = SummaryIndexService.generate_and_vectorize_summary(segment, dataset, {"enable": True}, session=session) assert out is record session.refresh.assert_called_once_with(record) session.commit.assert_called() @@ -355,18 +348,13 @@ def test_generate_and_vectorize_summary_vectorize_failure_sets_error(monkeypatch session = MagicMock() session.scalar.return_value = record - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) monkeypatch.setattr( SummaryIndexService, "generate_summary_for_segment", MagicMock(return_value=("sum", MagicMock(total_tokens=0))) ) monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(side_effect=RuntimeError("boom"))) with pytest.raises(RuntimeError, match="boom"): - SummaryIndexService.generate_and_vectorize_summary(segment, dataset, {"enable": True}) + SummaryIndexService.generate_and_vectorize_summary(segment, dataset, {"enable": True}, session=session) assert record.status == SummaryStatus.ERROR # Outer exception handler overwrites the error with the raw exception message. assert record.error == "boom" @@ -562,18 +550,12 @@ def test_generate_and_vectorize_summary_creates_missing_record_and_logs_usage( session = MagicMock() session.scalar.return_value = None - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - usage = MagicMock(total_tokens=4, prompt_tokens=1, completion_tokens=3) monkeypatch.setattr(SummaryIndexService, "generate_summary_for_segment", MagicMock(return_value=("sum", usage))) monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(return_value=None)) with caplog.at_level(logging.INFO, logger="services.summary_index_service"): - result = SummaryIndexService.generate_and_vectorize_summary(segment, dataset, {"enable": True}) + result = SummaryIndexService.generate_and_vectorize_summary(segment, dataset, {"enable": True}, session=session) assert result.status in {SummaryStatus.GENERATING, SummaryStatus.COMPLETED} assert any(r.levelno >= logging.INFO for r in caplog.records) @@ -833,11 +815,12 @@ def test_delete_summaries_for_segments_no_summaries_noop(monkeypatch: pytest.Mon def test_update_summary_for_segment_skip_conditions() -> None: + session = MagicMock() economy_dataset = _dataset(indexing_technique=IndexTechniqueType.ECONOMY) - assert SummaryIndexService.update_summary_for_segment(_segment(), economy_dataset, "x") is None + assert SummaryIndexService.update_summary_for_segment(_segment(), economy_dataset, "x", session=session) is None seg = _segment(has_document=True) seg.document.doc_form = IndexStructureType.QA_INDEX - assert SummaryIndexService.update_summary_for_segment(seg, _dataset(), "x") is None + assert SummaryIndexService.update_summary_for_segment(seg, _dataset(), "x", session=session) is None def test_update_summary_for_segment_empty_content_deletes_existing(monkeypatch: pytest.MonkeyPatch) -> None: @@ -850,13 +833,7 @@ def test_update_summary_for_segment_empty_content_deletes_existing(monkeypatch: vector_instance = MagicMock() monkeypatch.setattr(summary_module, "Vector", MagicMock(return_value=vector_instance)) - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - - assert SummaryIndexService.update_summary_for_segment(segment, dataset, " ") is None + assert SummaryIndexService.update_summary_for_segment(segment, dataset, " ", session=session) is None vector_instance.delete_by_ids.assert_called_once_with(["n1"]) session.delete.assert_called_once_with(record) session.commit.assert_called_once() @@ -872,18 +849,12 @@ def test_update_summary_for_segment_empty_content_delete_vector_warns( session = MagicMock() session.scalar.return_value = record - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - vector_instance = MagicMock() vector_instance.delete_by_ids.side_effect = RuntimeError("boom") monkeypatch.setattr(summary_module, "Vector", MagicMock(return_value=vector_instance)) with caplog.at_level(logging.WARNING, logger="services.summary_index_service"): - assert SummaryIndexService.update_summary_for_segment(segment, dataset, "") is None + assert SummaryIndexService.update_summary_for_segment(segment, dataset, "", session=session) is None assert any(r.levelno >= logging.WARNING for r in caplog.records) @@ -893,13 +864,7 @@ def test_update_summary_for_segment_empty_content_no_record_noop(monkeypatch: py session = MagicMock() session.scalar.return_value = None - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - - assert SummaryIndexService.update_summary_for_segment(segment, dataset, " ") is None + assert SummaryIndexService.update_summary_for_segment(segment, dataset, " ", session=session) is None def test_update_summary_for_segment_updates_existing_and_vectorizes(monkeypatch: pytest.MonkeyPatch) -> None: @@ -912,16 +877,10 @@ def test_update_summary_for_segment_updates_existing_and_vectorizes(monkeypatch: vector_instance = MagicMock() monkeypatch.setattr(summary_module, "Vector", MagicMock(return_value=vector_instance)) - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - vectorize_mock = MagicMock() monkeypatch.setattr(SummaryIndexService, "vectorize_summary", vectorize_mock) - out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new summary") + out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new summary", session=session) assert out is record vectorize_mock.assert_called_once() session.refresh.assert_called_once_with(record) @@ -938,19 +897,13 @@ def test_update_summary_for_segment_existing_vector_delete_warns( session = MagicMock() session.scalar.return_value = record - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - vector_instance = MagicMock() vector_instance.delete_by_ids.side_effect = RuntimeError("boom") monkeypatch.setattr(summary_module, "Vector", MagicMock(return_value=vector_instance)) monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(return_value=None)) with caplog.at_level(logging.WARNING, logger="services.summary_index_service"): - SummaryIndexService.update_summary_for_segment(segment, dataset, "new") + SummaryIndexService.update_summary_for_segment(segment, dataset, "new", session=session) assert any(r.levelno >= logging.WARNING for r in caplog.records) @@ -963,14 +916,9 @@ def test_update_summary_for_segment_existing_vectorize_failure_returns_error_rec session = MagicMock() session.scalar.return_value = record - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(side_effect=RuntimeError("boom"))) - out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new") + out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new", session=session) assert out is record assert out.status == SummaryStatus.ERROR assert "Vectorization failed" in (out.error or "") @@ -982,18 +930,11 @@ def test_update_summary_for_segment_new_record_success(monkeypatch: pytest.Monke session = MagicMock() session.scalar.return_value = None - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - created = _summary_record(summary_content="new", node_id=None) monkeypatch.setattr(SummaryIndexService, "create_summary_record", MagicMock(return_value=created)) - session.merge.return_value = created monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(return_value=None)) - out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new") + out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new", session=session) assert out is created session.refresh.assert_called() session.commit.assert_called() @@ -1007,81 +948,60 @@ def test_update_summary_for_segment_outer_exception_sets_error_and_reraises(monk session = MagicMock() session.scalar.return_value = record session.flush.side_effect = RuntimeError("flush boom") - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - with pytest.raises(RuntimeError, match="flush boom"): - SummaryIndexService.update_summary_for_segment(segment, dataset, "new") + SummaryIndexService.update_summary_for_segment(segment, dataset, "new", session=session) assert record.status == SummaryStatus.ERROR assert record.error == "flush boom" session.commit.assert_called() -def test_get_segment_summary_and_document_summaries(monkeypatch: pytest.MonkeyPatch) -> None: +def test_get_segment_summary_and_document_summaries() -> None: record = _summary_record(summary_content="sum", node_id="n1") session = MagicMock() session.scalar.return_value = record session.scalars.return_value.all.return_value = [record] - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - - assert SummaryIndexService.get_segment_summary("seg-1", "dataset-1") is record - assert SummaryIndexService.get_document_summaries("doc-1", "dataset-1", segment_ids=["seg-1"]) == [record] + assert SummaryIndexService.get_segment_summary("seg-1", "dataset-1", session=session) is record + assert SummaryIndexService.get_document_summaries("doc-1", "dataset-1", segment_ids=["seg-1"], session=session) == [ + record + ] -def test_get_segments_summaries_non_empty(monkeypatch: pytest.MonkeyPatch) -> None: +def test_get_segments_summaries_non_empty() -> None: record1 = _summary_record() record1.chunk_id = "seg-1" record2 = _summary_record() record2.chunk_id = "seg-2" session = MagicMock() session.scalars.return_value.all.return_value = [record1, record2] - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - out = SummaryIndexService.get_segments_summaries(["seg-1", "seg-2"], "dataset-1") + out = SummaryIndexService.get_segments_summaries(["seg-1", "seg-2"], "dataset-1", session=session) assert set(out.keys()) == {"seg-1", "seg-2"} -def test_get_document_summary_index_status_no_segments_returns_none(monkeypatch: pytest.MonkeyPatch) -> None: +def test_get_document_summary_index_status_no_segments_returns_none() -> None: session = MagicMock() session.scalars.return_value.all.return_value = [] - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), + assert ( + SummaryIndexService.get_document_summary_index_status("doc-1", "dataset-1", "tenant-1", session=session) is None ) - assert SummaryIndexService.get_document_summary_index_status("doc-1", "dataset-1", "tenant-1") is None -def test_get_documents_summary_index_status_empty_input(monkeypatch: pytest.MonkeyPatch) -> None: - assert SummaryIndexService.get_documents_summary_index_status([], "dataset-1", "tenant-1") == {} +def test_get_documents_summary_index_status_empty_input() -> None: + assert ( + SummaryIndexService.get_documents_summary_index_status([], "dataset-1", "tenant-1", session=MagicMock()) == {} + ) def test_get_documents_summary_index_status_no_pending_sets_none(monkeypatch: pytest.MonkeyPatch) -> None: session = MagicMock() session.execute.return_value.all.return_value = [SimpleNamespace(id="seg-1", document_id="doc-1")] - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) monkeypatch.setattr( SummaryIndexService, "get_segments_summaries", MagicMock(return_value={"seg-1": SimpleNamespace(status=SummaryStatus.COMPLETED)}), ) - result = SummaryIndexService.get_documents_summary_index_status(["doc-1"], "dataset-1", "tenant-1") + result = SummaryIndexService.get_documents_summary_index_status(["doc-1"], "dataset-1", "tenant-1", session=session) assert result["doc-1"] is None @@ -1094,26 +1014,19 @@ def test_update_summary_for_segment_creates_new_and_vectorize_fails_returns_erro session = MagicMock() session.scalar.return_value = None - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - created = _summary_record(summary_content="new", node_id=None) monkeypatch.setattr(SummaryIndexService, "create_summary_record", MagicMock(return_value=created)) - session.merge.return_value = created vectorize_mock = MagicMock(side_effect=RuntimeError("boom")) monkeypatch.setattr(SummaryIndexService, "vectorize_summary", vectorize_mock) - out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new") + out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new", session=session) assert out.status == SummaryStatus.ERROR assert "Vectorization failed" in (out.error or "") def test_get_segments_summaries_empty_list() -> None: - assert SummaryIndexService.get_segments_summaries([], "dataset-1") == {} + assert SummaryIndexService.get_segments_summaries([], "dataset-1", session=MagicMock()) == {} def test_get_document_summary_index_status_and_documents_status(monkeypatch: pytest.MonkeyPatch) -> None: @@ -1121,30 +1034,27 @@ def test_get_document_summary_index_status_and_documents_status(monkeypatch: pyt session = MagicMock() session.scalars.return_value.all.return_value = ["seg-1"] # get_document_summary_index_status returns IDs - create_session_mock = MagicMock(return_value=_SessionContext(session)) - monkeypatch.setattr(summary_module, "session_factory", SimpleNamespace(create_session=create_session_mock)) - monkeypatch.setattr( SummaryIndexService, "get_segments_summaries", MagicMock(return_value={"seg-1": SimpleNamespace(status=SummaryStatus.GENERATING)}), ) - assert SummaryIndexService.get_document_summary_index_status("doc-1", "dataset-1", "tenant-1") == "SUMMARIZING" + assert ( + SummaryIndexService.get_document_summary_index_status("doc-1", "dataset-1", "tenant-1", session=session) + == "SUMMARIZING" + ) # Multiple docs session2 = MagicMock() session2.execute.return_value.all.return_value = [seg_row] # get_documents_summary_index_status uses execute - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session2))), - ) monkeypatch.setattr( SummaryIndexService, "get_segments_summaries", MagicMock(return_value={"seg-1": SimpleNamespace(status=SummaryStatus.NOT_STARTED)}), ) - result = SummaryIndexService.get_documents_summary_index_status(["doc-1", "doc-2"], "dataset-1", "tenant-1") + result = SummaryIndexService.get_documents_summary_index_status( + ["doc-1", "doc-2"], "dataset-1", "tenant-1", session=session2 + ) assert result["doc-1"] == "SUMMARIZING" assert result["doc-2"] is None diff --git a/api/tests/unit_tests/services/test_trigger_provider_service.py b/api/tests/unit_tests/services/test_trigger_provider_service.py index 0a4452cf478..ff11bbb3035 100644 --- a/api/tests/unit_tests/services/test_trigger_provider_service.py +++ b/api/tests/unit_tests/services/test_trigger_provider_service.py @@ -444,7 +444,7 @@ def test_delete_trigger_provider_should_raise_error_when_subscription_missing( # Act + Assert with pytest.raises(ValueError, match="not found"): - TriggerProviderService.delete_trigger_provider(mock_session, "tenant-1", "sub-1") + TriggerProviderService.delete_trigger_provider("tenant-1", "sub-1", session=mock_session) def test_delete_trigger_provider_should_delete_and_clear_cache_even_if_unsubscribe_fails( @@ -476,7 +476,7 @@ def test_delete_trigger_provider_should_delete_and_clear_cache_even_if_unsubscri mock_delete_cache = mocker.patch("services.trigger.trigger_provider_service.delete_cache_for_subscription") # Act - TriggerProviderService.delete_trigger_provider(mock_session, "tenant-1", "sub-1") + TriggerProviderService.delete_trigger_provider("tenant-1", "sub-1", session=mock_session) # Assert mock_session.delete.assert_called_once_with(subscription) @@ -507,7 +507,7 @@ def test_delete_trigger_provider_should_skip_unsubscribe_for_unauthorized( ) # Act - TriggerProviderService.delete_trigger_provider(mock_session, "tenant-1", "sub-2") + TriggerProviderService.delete_trigger_provider("tenant-1", "sub-2", session=mock_session) # Assert mock_unsubscribe.assert_not_called() diff --git a/api/tests/unit_tests/services/test_vector_service.py b/api/tests/unit_tests/services/test_vector_service.py index e7ebada6bea..3659b85228b 100644 --- a/api/tests/unit_tests/services/test_vector_service.py +++ b/api/tests/unit_tests/services/test_vector_service.py @@ -98,7 +98,7 @@ def test_create_segments_vector_regular_indexing_loads_documents_and_keywords(mo factory_instance.init_index_processor.return_value = index_processor monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance)) - VectorService.create_segments_vector([["k1"]], [segment], dataset, IndexStructureType.PARAGRAPH_INDEX) + VectorService.create_segments_vector([["k1"]], [segment], dataset, IndexStructureType.PARAGRAPH_INDEX, MagicMock()) index_processor.load.assert_called_once() args, kwargs = index_processor.load.call_args @@ -123,7 +123,7 @@ def test_create_segments_vector_regular_indexing_loads_multimodal_documents(monk factory_instance.init_index_processor.return_value = index_processor monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance)) - VectorService.create_segments_vector([["k1"]], [segment], dataset, IndexStructureType.PARAGRAPH_INDEX) + VectorService.create_segments_vector([["k1"]], [segment], dataset, IndexStructureType.PARAGRAPH_INDEX, MagicMock()) assert index_processor.load.call_count == 2 first_args, first_kwargs = index_processor.load.call_args_list[0] @@ -145,7 +145,7 @@ def test_create_segments_vector_with_no_segments_does_not_load(monkeypatch: pyte factory_instance.init_index_processor.return_value = index_processor monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance)) - VectorService.create_segments_vector(None, [], dataset, IndexStructureType.PARAGRAPH_INDEX) + VectorService.create_segments_vector(None, [], dataset, IndexStructureType.PARAGRAPH_INDEX, MagicMock()) index_processor.load.assert_not_called() @@ -189,11 +189,7 @@ def test_create_segments_vector_parent_child_calls_generate_child_chunks_with_ex processing_rule = MagicMock(name="processing_rule") processing_rule.to_dict.return_value = {"rules": {}} - monkeypatch.setattr( - vector_service_module, - "db", - _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=processing_rule), - ) + db_mock = _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=processing_rule) embedding_model_instance = MagicMock(name="embedding_model_instance") model_manager_instance = MagicMock(name="model_manager_instance") @@ -211,12 +207,22 @@ def test_create_segments_vector_parent_child_calls_generate_child_chunks_with_ex monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance)) VectorService.create_segments_vector( - None, [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX + None, + [segment], + dataset, + vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, + db_mock.session, ) model_manager_instance.get_model_instance.assert_called_once() generate_child_chunks_mock.assert_called_once_with( - segment, dataset_document, dataset, embedding_model_instance, processing_rule, False + segment, + dataset_document, + dataset, + embedding_model_instance, + processing_rule, + db_mock.session, + False, ) index_processor.load.assert_not_called() @@ -239,11 +245,7 @@ def test_create_segments_vector_parent_child_uses_default_embedding_model_when_p processing_rule = MagicMock() processing_rule.to_dict.return_value = {"rules": {}} - monkeypatch.setattr( - vector_service_module, - "db", - _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=processing_rule), - ) + db_mock = _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=processing_rule) embedding_model_instance = MagicMock() model_manager_instance = MagicMock() @@ -261,7 +263,11 @@ def test_create_segments_vector_parent_child_uses_default_embedding_model_when_p monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance)) VectorService.create_segments_vector( - None, [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX + None, + [segment], + dataset, + vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, + db_mock.session, ) model_manager_instance.get_default_model_instance.assert_called_once() @@ -276,11 +282,7 @@ def test_create_segments_vector_parent_child_missing_document_logs_warning_and_c segment = _make_segment() processing_rule = MagicMock() - monkeypatch.setattr( - vector_service_module, - "db", - _mock_parent_child_queries(dataset_document=None, processing_rule=processing_rule), - ) + db_mock = _mock_parent_child_queries(dataset_document=None, processing_rule=processing_rule) index_processor = MagicMock() factory_instance = MagicMock() @@ -289,7 +291,11 @@ def test_create_segments_vector_parent_child_missing_document_logs_warning_and_c with caplog.at_level(logging.WARNING, logger="services.vector_service"): VectorService.create_segments_vector( - None, [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX + None, + [segment], + dataset, + vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, + db_mock.session, ) assert any(r.levelno >= logging.WARNING for r in caplog.records) index_processor.load.assert_not_called() @@ -301,15 +307,15 @@ def test_create_segments_vector_parent_child_missing_processing_rule_raises(monk dataset_document = MagicMock() dataset_document.dataset_process_rule_id = "rule-1" - monkeypatch.setattr( - vector_service_module, - "db", - _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=None), - ) + db_mock = _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=None) with pytest.raises(ValueError, match="No processing rule found"): VectorService.create_segments_vector( - None, [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX + None, + [segment], + dataset, + vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, + db_mock.session, ) @@ -322,15 +328,15 @@ def test_create_segments_vector_parent_child_non_high_quality_raises(monkeypatch dataset_document = MagicMock() dataset_document.dataset_process_rule_id = "rule-1" processing_rule = MagicMock() - monkeypatch.setattr( - vector_service_module, - "db", - _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=processing_rule), - ) + db_mock = _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=processing_rule) with pytest.raises(ValueError, match="not high quality"): VectorService.create_segments_vector( - None, [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX + None, + [segment], + dataset, + vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, + db_mock.session, ) @@ -404,10 +410,7 @@ def test_generate_child_chunks_regenerate_cleans_then_saves_children(monkeypatch child_chunk_ctor = MagicMock(side_effect=lambda **kwargs: kwargs) monkeypatch.setattr(vector_service_module, "ChildChunk", child_chunk_ctor) - db_mock = MagicMock() - db_mock.session.add = MagicMock() - db_mock.session.commit = MagicMock() - monkeypatch.setattr(vector_service_module, "db", db_mock) + session = MagicMock() VectorService.generate_child_chunks( segment=segment, @@ -415,6 +418,7 @@ def test_generate_child_chunks_regenerate_cleans_then_saves_children(monkeypatch dataset=dataset, embedding_model_instance=MagicMock(), processing_rule=processing_rule, + session=session, regenerate=True, ) @@ -422,8 +426,8 @@ def test_generate_child_chunks_regenerate_cleans_then_saves_children(monkeypatch _, transform_kwargs = index_processor.transform.call_args assert transform_kwargs["process_rule"]["rules"]["parent_mode"] == vector_service_module.ParentMode.FULL_DOC index_processor.load.assert_called_once() - assert db_mock.session.add.call_count == 2 - db_mock.session.commit.assert_called_once() + assert session.add.call_count == 2 + session.commit.assert_called_once() def test_generate_child_chunks_commits_even_when_no_children(monkeypatch: pytest.MonkeyPatch) -> None: @@ -442,8 +446,7 @@ def test_generate_child_chunks_commits_even_when_no_children(monkeypatch: pytest factory_instance.init_index_processor.return_value = index_processor monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance)) - db_mock = MagicMock() - monkeypatch.setattr(vector_service_module, "db", db_mock) + session = MagicMock() VectorService.generate_child_chunks( segment=segment, @@ -451,12 +454,13 @@ def test_generate_child_chunks_commits_even_when_no_children(monkeypatch: pytest dataset=dataset, embedding_model_instance=MagicMock(), processing_rule=processing_rule, + session=session, regenerate=False, ) index_processor.load.assert_not_called() - db_mock.session.add.assert_not_called() - db_mock.session.commit.assert_called_once() + session.add.assert_not_called() + session.commit.assert_called_once() def test_create_child_chunk_vector_high_quality_adds_texts(monkeypatch: pytest.MonkeyPatch) -> None: @@ -554,9 +558,10 @@ def test_update_multimodel_vector_returns_when_not_high_quality(monkeypatch: pyt vector_cls = MagicMock() db_mock = _mock_db_session_for_update_multimodel(upload_files=[]) monkeypatch.setattr(vector_service_module, "Vector", vector_cls) - monkeypatch.setattr(vector_service_module, "db", db_mock) - VectorService.update_multimodel_vector(segment=segment, attachment_ids=["a"], dataset=dataset) + VectorService.update_multimodel_vector( + segment=segment, attachment_ids=["a"], dataset=dataset, session=db_mock.session + ) vector_cls.assert_not_called() db_mock.session.query.assert_not_called() @@ -568,9 +573,10 @@ def test_update_multimodel_vector_returns_when_no_actual_change(monkeypatch: pyt vector_cls = MagicMock() db_mock = _mock_db_session_for_update_multimodel(upload_files=[]) monkeypatch.setattr(vector_service_module, "Vector", vector_cls) - monkeypatch.setattr(vector_service_module, "db", db_mock) - VectorService.update_multimodel_vector(segment=segment, attachment_ids=["b", "a"], dataset=dataset) + VectorService.update_multimodel_vector( + segment=segment, attachment_ids=["b", "a"], dataset=dataset, session=db_mock.session + ) vector_cls.assert_not_called() db_mock.session.query.assert_not_called() @@ -586,9 +592,8 @@ def test_update_multimodel_vector_deletes_bindings_and_commits_on_empty_new_ids( db_mock = _mock_db_session_for_update_multimodel(upload_files=[]) monkeypatch.setattr(vector_service_module, "Vector", vector_cls) - monkeypatch.setattr(vector_service_module, "db", db_mock) - VectorService.update_multimodel_vector(segment=segment, attachment_ids=[], dataset=dataset) + VectorService.update_multimodel_vector(segment=segment, attachment_ids=[], dataset=dataset, session=db_mock.session) vector_cls.assert_called_once_with(dataset=dataset) vector_instance.delete_by_ids.assert_called_once_with(["old-1", "old-2"]) @@ -605,9 +610,10 @@ def test_update_multimodel_vector_commits_when_no_upload_files_found(monkeypatch vector_instance = MagicMock() monkeypatch.setattr(vector_service_module, "Vector", MagicMock(return_value=vector_instance)) db_mock = _mock_db_session_for_update_multimodel(upload_files=[]) - monkeypatch.setattr(vector_service_module, "db", db_mock) - VectorService.update_multimodel_vector(segment=segment, attachment_ids=["new-1"], dataset=dataset) + VectorService.update_multimodel_vector( + segment=segment, attachment_ids=["new-1"], dataset=dataset, session=db_mock.session + ) db_mock.session.commit.assert_called_once() db_mock.session.add_all.assert_not_called() @@ -624,7 +630,6 @@ def test_update_multimodel_vector_adds_bindings_and_vectors_and_skips_missing_up vector_instance = MagicMock() monkeypatch.setattr(vector_service_module, "Vector", MagicMock(return_value=vector_instance)) db_mock = _mock_db_session_for_update_multimodel(upload_files=[_UploadFileStub(id="file-1", name="img.png")]) - monkeypatch.setattr(vector_service_module, "db", db_mock) binding_ctor = MagicMock(side_effect=lambda **kwargs: kwargs) monkeypatch.setattr(vector_service_module, "SegmentAttachmentBinding", binding_ctor) @@ -632,7 +637,12 @@ def test_update_multimodel_vector_adds_bindings_and_vectors_and_skips_missing_up monkeypatch.setattr(vector_service_module, "select", MagicMock()) with caplog.at_level(logging.WARNING, logger="services.vector_service"): - VectorService.update_multimodel_vector(segment=segment, attachment_ids=["file-1", "missing"], dataset=dataset) + VectorService.update_multimodel_vector( + segment=segment, + attachment_ids=["file-1", "missing"], + dataset=dataset, + session=db_mock.session, + ) assert any(r.levelno >= logging.WARNING for r in caplog.records) db_mock.session.add_all.assert_called_once() bindings = db_mock.session.add_all.call_args.args[0] @@ -656,14 +666,15 @@ def test_update_multimodel_vector_updates_bindings_without_multimodal_vector_ops vector_instance = MagicMock() monkeypatch.setattr(vector_service_module, "Vector", MagicMock(return_value=vector_instance)) db_mock = _mock_db_session_for_update_multimodel(upload_files=[_UploadFileStub(id="file-1", name="img.png")]) - monkeypatch.setattr(vector_service_module, "db", db_mock) monkeypatch.setattr( vector_service_module, "SegmentAttachmentBinding", MagicMock(side_effect=lambda **kwargs: kwargs) ) monkeypatch.setattr(vector_service_module, "delete", MagicMock()) monkeypatch.setattr(vector_service_module, "select", MagicMock()) - VectorService.update_multimodel_vector(segment=segment, attachment_ids=["file-1"], dataset=dataset) + VectorService.update_multimodel_vector( + segment=segment, attachment_ids=["file-1"], dataset=dataset, session=db_mock.session + ) vector_instance.delete_by_ids.assert_not_called() vector_instance.add_texts.assert_not_called() @@ -682,7 +693,6 @@ def test_update_multimodel_vector_rolls_back_and_reraises_on_error( monkeypatch.setattr(vector_service_module, "Vector", MagicMock(return_value=vector_instance)) db_mock = _mock_db_session_for_update_multimodel(upload_files=[_UploadFileStub(id="file-1", name="img.png")]) db_mock.session.commit.side_effect = RuntimeError("boom") - monkeypatch.setattr(vector_service_module, "db", db_mock) monkeypatch.setattr( vector_service_module, "SegmentAttachmentBinding", MagicMock(side_effect=lambda **kwargs: kwargs) ) @@ -691,7 +701,9 @@ def test_update_multimodel_vector_rolls_back_and_reraises_on_error( with caplog.at_level(logging.ERROR, logger="services.vector_service"): with pytest.raises(RuntimeError, match="boom"): - VectorService.update_multimodel_vector(segment=segment, attachment_ids=["file-1"], dataset=dataset) + VectorService.update_multimodel_vector( + segment=segment, attachment_ids=["file-1"], dataset=dataset, session=db_mock.session + ) assert any(r.levelno >= logging.ERROR for r in caplog.records) db_mock.session.rollback.assert_called_once() diff --git a/api/tests/unit_tests/services/test_workflow_collaboration_service.py b/api/tests/unit_tests/services/test_workflow_collaboration_service.py index a61e49c02fa..6b269443fa2 100644 --- a/api/tests/unit_tests/services/test_workflow_collaboration_service.py +++ b/api/tests/unit_tests/services/test_workflow_collaboration_service.py @@ -32,7 +32,7 @@ class TestWorkflowCollaborationService: patch.object(collaboration_service, "broadcast_online_users"), ): # Act - result = collaboration_service.authorize_and_join_workflow_room("wf-1", "sid-1") + result = collaboration_service.authorize_and_join_workflow_room("wf-1", "sid-1", session=Mock()) # Assert assert result == ("u-1", True) @@ -52,7 +52,7 @@ class TestWorkflowCollaborationService: socketio.get_session.return_value = {} # Act - result = collaboration_service.authorize_and_join_workflow_room("wf-1", "sid-1") + result = collaboration_service.authorize_and_join_workflow_room("wf-1", "sid-1", session=Mock()) # Assert assert result is None @@ -63,7 +63,7 @@ class TestWorkflowCollaborationService: collaboration_service, repository, socketio = service socketio.get_session.return_value = {"user_id": "u-1", "username": "Jane", "avatar": None} - result = collaboration_service.authorize_and_join_workflow_room("wf-1", "sid-1") + result = collaboration_service.authorize_and_join_workflow_room("wf-1", "sid-1", session=Mock()) assert result is None repository.set_session_info.assert_not_called() @@ -82,7 +82,7 @@ class TestWorkflowCollaborationService: } with patch.object(collaboration_service, "_can_access_workflow", return_value=False): - result = collaboration_service.authorize_and_join_workflow_room("wf-1", "sid-1") + result = collaboration_service.authorize_and_join_workflow_room("wf-1", "sid-1", session=Mock()) assert result is None repository.set_session_info.assert_not_called() @@ -106,21 +106,12 @@ class TestWorkflowCollaborationService: {"user_id": "u-1", "username": "Jane", "avatar": "avatar.png", "tenant_id": "t-1"}, ) - def test_can_access_workflow_uses_session_factory( - self, service: tuple[WorkflowCollaborationService, Mock, Mock] - ) -> None: + def test_can_access_workflow_uses_session(self, service: tuple[WorkflowCollaborationService, Mock, Mock]) -> None: collaboration_service, _repository, _socketio = service session = Mock() session.scalar.return_value = "wf-1" - session_context = Mock() - session_context.__enter__ = Mock(return_value=session) - session_context.__exit__ = Mock(return_value=False) - with patch( - "services.workflow_collaboration_service.session_factory.create_session", - return_value=session_context, - ): - result = collaboration_service._can_access_workflow("wf-1", "tenant-1") + result = collaboration_service._can_access_workflow("wf-1", "tenant-1", session=session) assert result is True session.scalar.assert_called_once() diff --git a/api/tests/unit_tests/services/test_workflow_service.py b/api/tests/unit_tests/services/test_workflow_service.py index 67b3e80da6b..0a75a0a8788 100644 --- a/api/tests/unit_tests/services/test_workflow_service.py +++ b/api/tests/unit_tests/services/test_workflow_service.py @@ -312,7 +312,7 @@ class TestWorkflowService: # Mock the database query to return True mock_db_session.session.execute.return_value.scalar_one.return_value = True - result = workflow_service.is_workflow_exist(app) + result = workflow_service.is_workflow_exist(app, session=mock_db_session.session) assert result is True @@ -323,7 +323,7 @@ class TestWorkflowService: # Mock the database query to return False mock_db_session.session.execute.return_value.scalar_one.return_value = False - result = workflow_service.is_workflow_exist(app) + result = workflow_service.is_workflow_exist(app, session=mock_db_session.session) assert result is False @@ -343,7 +343,7 @@ class TestWorkflowService: # Mock db.session.scalar() used by get_draft_workflow mock_db_session.session.scalar.return_value = mock_workflow - result = workflow_service.get_draft_workflow(app) + result = workflow_service.get_draft_workflow(app, session=mock_db_session.session) assert result == mock_workflow @@ -367,7 +367,7 @@ class TestWorkflowService: # Mock db.session.scalar() to return None mock_db_session.session.scalar.return_value = None - result = workflow_service.get_draft_workflow(app) + result = workflow_service.get_draft_workflow(app, session=mock_db_session.session) assert result is None @@ -380,7 +380,7 @@ class TestWorkflowService: # Mock db.session.scalar() used by get_published_workflow_by_id mock_db_session.session.scalar.return_value = mock_workflow - result = workflow_service.get_draft_workflow(app, workflow_id=workflow_id) + result = workflow_service.get_draft_workflow(app, workflow_id=workflow_id, session=mock_db_session.session) assert result == mock_workflow @@ -411,7 +411,7 @@ class TestWorkflowService: # Mock db.session.scalar() used by get_published_workflow_by_id mock_db_session.session.scalar.return_value = mock_workflow - result = workflow_service.get_published_workflow_by_id(app, workflow_id) + result = workflow_service.get_published_workflow_by_id(app, workflow_id, session=mock_db_session.session) assert result == mock_workflow @@ -432,7 +432,7 @@ class TestWorkflowService: mock_db_session.session.scalar.return_value = mock_workflow with pytest.raises(IsDraftWorkflowError): - workflow_service.get_published_workflow_by_id(app, workflow_id) + workflow_service.get_published_workflow_by_id(app, workflow_id, session=mock_db_session.session) def test_get_published_workflow_by_id_returns_none(self, workflow_service, mock_db_session): """Test get_published_workflow_by_id returns None when workflow not found.""" @@ -442,7 +442,7 @@ class TestWorkflowService: # Mock db.session.scalar() to return None mock_db_session.session.scalar.return_value = None - result = workflow_service.get_published_workflow_by_id(app, workflow_id) + result = workflow_service.get_published_workflow_by_id(app, workflow_id, session=mock_db_session.session) assert result is None @@ -455,7 +455,7 @@ class TestWorkflowService: # Mock db.session.scalar() used by get_published_workflow mock_db_session.session.scalar.return_value = mock_workflow - result = workflow_service.get_published_workflow(app) + result = workflow_service.get_published_workflow(app, session=mock_db_session.session) assert result == mock_workflow @@ -463,7 +463,7 @@ class TestWorkflowService: """Test get_published_workflow returns None when app has no workflow_id.""" app = TestWorkflowAssociatedDataFactory.create_app_mock(workflow_id=None) - result = workflow_service.get_published_workflow(app) + result = workflow_service.get_published_workflow(app, session=MagicMock()) assert result is None @@ -499,6 +499,7 @@ class TestWorkflowService: account=account, environment_variables=[], conversation_variables=[], + session=mock_db_session.session, ) # Verify workflow was added to session @@ -536,6 +537,7 @@ class TestWorkflowService: account=account, environment_variables=[], conversation_variables=[], + session=mock_db_session.session, ) # Verify workflow was updated @@ -571,6 +573,7 @@ class TestWorkflowService: account=account, environment_variables=[], conversation_variables=[], + session=mock_db_session.session, ) def test_restore_published_workflow_to_draft_keeps_source_features_unmodified( @@ -648,6 +651,7 @@ class TestWorkflowService: app_model=app, workflow_id=source_workflow.id, account=account, + session=mock_db_session.session, ) mock_validate_features.assert_called_once_with(app_model=app, features=normalized_features) @@ -761,6 +765,7 @@ class TestWorkflowService: app_model=app, environment_variables=variables, account=account, + session=mock_db_session.session, ) assert workflow.environment_variables == variables @@ -779,6 +784,7 @@ class TestWorkflowService: app_model=app, environment_variables=[], account=account, + session=MagicMock(), ) def test_update_draft_workflow_conversation_variables_updates_workflow(self, workflow_service, mock_db_session): @@ -796,6 +802,7 @@ class TestWorkflowService: app_model=app, conversation_variables=variables, account=account, + session=mock_db_session.session, ) assert workflow.conversation_variables == variables @@ -814,6 +821,7 @@ class TestWorkflowService: app_model=app, conversation_variables=[], account=account, + session=MagicMock(), ) # ==================== Publish Workflow Tests ==================== @@ -1429,7 +1437,7 @@ class TestWorkflowService: mock_new_app = TestWorkflowAssociatedDataFactory.create_app_mock(mode=AppMode.WORKFLOW) mock_converter.convert_to_workflow.return_value = mock_new_app - result = workflow_service.convert_to_workflow(app, account, args) + result = workflow_service.convert_to_workflow(app, account, args, session=MagicMock()) assert result == mock_new_app mock_converter.convert_to_workflow.assert_called_once() @@ -1451,7 +1459,7 @@ class TestWorkflowService: mock_new_app = TestWorkflowAssociatedDataFactory.create_app_mock(mode=AppMode.WORKFLOW) mock_converter.convert_to_workflow.return_value = mock_new_app - result = workflow_service.convert_to_workflow(app, account, args) + result = workflow_service.convert_to_workflow(app, account, args, session=MagicMock()) assert result == mock_new_app @@ -1467,7 +1475,7 @@ class TestWorkflowService: args = {} with pytest.raises(ValueError, match="not supported convert to workflow"): - workflow_service.convert_to_workflow(app, account, args) + workflow_service.convert_to_workflow(app, account, args, session=MagicMock()) # =========================================================================== @@ -1520,7 +1528,7 @@ class TestWorkflowServiceCredentialValidation: # Act + Assert with patch("core.helper.credential_utils.check_credential_policy_compliance") as mock_check: # Should not raise; mock allows the call - service._validate_workflow_credentials(workflow) + service._validate_workflow_credentials(workflow, session=MagicMock()) mock_check.assert_called_once() def test_validate_workflow_credentials_should_check_default_credential_when_no_credential_id( @@ -1541,10 +1549,11 @@ class TestWorkflowServiceCredentialValidation: # Act with patch.object(service, "_check_default_tool_credential") as mock_default: - service._validate_workflow_credentials(workflow) + session = MagicMock() + service._validate_workflow_credentials(workflow, session=session) # Assert - mock_default.assert_called_once_with("tenant-1", "my-provider") + mock_default.assert_called_once_with("tenant-1", "my-provider", session=session) def test_validate_workflow_credentials_should_skip_tool_node_without_provider( self, service: WorkflowService @@ -1556,7 +1565,7 @@ class TestWorkflowServiceCredentialValidation: # Act + Assert (no error raised) with patch.object(service, "_check_default_tool_credential") as mock_default: - service._validate_workflow_credentials(workflow) + service._validate_workflow_credentials(workflow, session=MagicMock()) mock_default.assert_not_called() def test_validate_workflow_credentials_should_validate_llm_node_with_model_config( @@ -1579,7 +1588,7 @@ class TestWorkflowServiceCredentialValidation: patch.object(service, "_validate_llm_model_config") as mock_llm, patch.object(service, "_validate_load_balancing_credentials"), ): - service._validate_workflow_credentials(workflow) + service._validate_workflow_credentials(workflow, session=MagicMock()) # Assert mock_llm.assert_called_once_with("tenant-1", "openai", "gpt-4") @@ -1599,7 +1608,7 @@ class TestWorkflowServiceCredentialValidation: # Act + Assert with pytest.raises(ValueError, match="Missing provider or model configuration"): - service._validate_workflow_credentials(workflow) + service._validate_workflow_credentials(workflow, session=MagicMock()) def test_validate_workflow_credentials_should_wrap_unexpected_exception_in_value_error( self, service: WorkflowService @@ -1620,7 +1629,7 @@ class TestWorkflowServiceCredentialValidation: # Act + Assert with patch.object(service, "_validate_llm_model_config", side_effect=RuntimeError("boom")): with pytest.raises(ValueError, match="boom"): - service._validate_workflow_credentials(workflow) + service._validate_workflow_credentials(workflow, session=MagicMock()) def test_validate_workflow_credentials_should_validate_agent_node_model(self, service: WorkflowService) -> None: # Arrange @@ -1643,7 +1652,7 @@ class TestWorkflowServiceCredentialValidation: patch.object(service, "_validate_llm_model_config") as mock_llm, patch.object(service, "_validate_load_balancing_credentials"), ): - service._validate_workflow_credentials(workflow) + service._validate_workflow_credentials(workflow, session=MagicMock()) # Assert mock_llm.assert_called_once_with("tenant-1", "openai", "gpt-4") @@ -1675,11 +1684,12 @@ class TestWorkflowServiceCredentialValidation: patch("core.helper.credential_utils.check_credential_policy_compliance") as mock_check, patch.object(service, "_check_default_tool_credential") as mock_default, ): - service._validate_workflow_credentials(workflow) + session = MagicMock() + service._validate_workflow_credentials(workflow, session=session) # Assert mock_check.assert_called_once() # provider-a has credential_id - mock_default.assert_called_once_with("tenant-1", "provider-b") + mock_default.assert_called_once_with("tenant-1", "provider-b", session=session) # --- _validate_llm_model_config --- @@ -1739,7 +1749,7 @@ class TestWorkflowServiceCredentialValidation: # Arrange with patch("services.workflow_service.db") as mock_db: # Act + Assert (should NOT raise) - service._check_default_tool_credential("tenant-1", "some-provider") + service._check_default_tool_credential("tenant-1", "some-provider", session=MagicMock()) def test_check_default_tool_credential_should_raise_when_compliance_fails(self, service: WorkflowService) -> None: # Arrange @@ -1751,7 +1761,7 @@ class TestWorkflowServiceCredentialValidation: ): # Act + Assert with pytest.raises(ValueError, match="Failed to validate default credential"): - service._check_default_tool_credential("tenant-1", "some-provider") + service._check_default_tool_credential("tenant-1", "some-provider", session=MagicMock()) # --- _is_load_balancing_enabled --- @@ -1811,7 +1821,7 @@ class TestWorkflowServiceCredentialValidation: side_effect=RuntimeError("fail"), ): # Act - result = service._get_load_balancing_configs("tenant-1", "openai", "gpt-4") + result = service._get_load_balancing_configs("tenant-1", "openai", "gpt-4", session=MagicMock()) # Assert assert result == [] @@ -1828,7 +1838,7 @@ class TestWorkflowServiceCredentialValidation: ], ): # Act - result = service._get_load_balancing_configs("tenant-1", "openai", "gpt-4") + result = service._get_load_balancing_configs("tenant-1", "openai", "gpt-4", session=MagicMock()) # Assert — only entries with a credential_id should be returned assert len(result) == 2 @@ -1845,7 +1855,7 @@ class TestWorkflowServiceCredentialValidation: node_data: dict[str, Any] = {} # no model key # Act + Assert (no error expected) - service._validate_load_balancing_credentials(workflow, node_data, "node-1") + service._validate_load_balancing_credentials(workflow, node_data, "node-1", session=MagicMock()) def test_validate_load_balancing_credentials_should_skip_when_lb_not_enabled( self, service: WorkflowService @@ -1856,7 +1866,7 @@ class TestWorkflowServiceCredentialValidation: # Act + Assert (no error expected) with patch.object(service, "_is_load_balancing_enabled", return_value=False): - service._validate_load_balancing_credentials(workflow, node_data, "node-1") + service._validate_load_balancing_credentials(workflow, node_data, "node-1", session=MagicMock()) def test_validate_load_balancing_credentials_should_raise_when_compliance_fails( self, service: WorkflowService @@ -1876,7 +1886,7 @@ class TestWorkflowServiceCredentialValidation: ), ): with pytest.raises(ValueError, match="Invalid load balancing credentials"): - service._validate_load_balancing_credentials(workflow, node_data, "node-1") + service._validate_load_balancing_credentials(workflow, node_data, "node-1", session=MagicMock()) # =========================================================================== @@ -2673,7 +2683,9 @@ class TestWorkflowServiceHumanInputOperations: def test_get_human_input_form_preview_should_raise_if_workflow_not_init(self, service: WorkflowService) -> None: service.get_draft_workflow = MagicMock(return_value=None) with pytest.raises(ValueError, match="Workflow not initialized"): - service.get_human_input_form_preview(app_model=MagicMock(), account=MagicMock(), node_id="node-1") + service.get_human_input_form_preview( + app_model=MagicMock(), account=MagicMock(), node_id="node-1", session=MagicMock() + ) def test_get_human_input_form_preview_should_raise_if_wrong_node_type(self, service: WorkflowService) -> None: draft = MagicMock() @@ -2681,7 +2693,9 @@ class TestWorkflowServiceHumanInputOperations: service.get_draft_workflow = MagicMock(return_value=draft) with patch("models.workflow.Workflow.get_node_type_from_node_config", return_value=BuiltinNodeTypes.LLM): with pytest.raises(ValueError, match="Node type must be human-input"): - service.get_human_input_form_preview(app_model=MagicMock(), account=MagicMock(), node_id="node-1") + service.get_human_input_form_preview( + app_model=MagicMock(), account=MagicMock(), node_id="node-1", session=MagicMock() + ) def test_get_human_input_form_preview_success(self, service: WorkflowService) -> None: app_model = MagicMock(spec=App) @@ -2716,7 +2730,9 @@ class TestWorkflowServiceHumanInputOperations: patch("services.workflow_service.HumanInputNode", return_value=mock_node), patch("services.workflow_service.HumanInputRequired") as mock_required_cls, ): - service.get_human_input_form_preview(app_model=app_model, account=account, node_id="node-1") + service.get_human_input_form_preview( + app_model=app_model, account=account, node_id="node-1", session=MagicMock() + ) mock_node.render_form_content_before_submission.assert_called_once() mock_required_cls.return_value.model_dump.assert_called_once() @@ -2760,7 +2776,12 @@ class TestWorkflowServiceHumanInputOperations: patch("services.workflow_service.DraftVariableSaver") as mock_saver_cls, ): result = service.submit_human_input_form_preview( - app_model=app_model, account=account, node_id="node-1", form_inputs={"field1": "val1"}, action="submit" + app_model=app_model, + account=account, + node_id="node-1", + form_inputs={"field1": "val1"}, + action="submit", + session=MagicMock(), ) assert result["__action_id"] == "submit" mock_validate.assert_called_once() @@ -2785,7 +2806,11 @@ class TestWorkflowServiceHumanInputOperations: ): mock_resolve.return_value = MagicMock() service.test_human_input_delivery( - app_model=MagicMock(), account=MagicMock(), node_id="node-1", delivery_method_id="method-1" + app_model=MagicMock(), + account=MagicMock(), + node_id="node-1", + delivery_method_id="method-1", + session=MagicMock(), ) mock_test_srv.return_value.send_test.assert_called_once() @@ -2801,7 +2826,11 @@ class TestWorkflowServiceHumanInputOperations: ): with pytest.raises(ValueError, match="Delivery method not found"): service.test_human_input_delivery( - app_model=MagicMock(), account=MagicMock(), node_id="node-1", delivery_method_id="none" + app_model=MagicMock(), + account=MagicMock(), + node_id="node-1", + delivery_method_id="none", + session=MagicMock(), ) def test_load_email_recipients_parsing_failure(self, service: WorkflowService) -> None: diff --git a/api/tests/unit_tests/services/tools/test_builtin_tools_manage_service.py b/api/tests/unit_tests/services/tools/test_builtin_tools_manage_service.py index c210db580e0..549f50cb370 100644 --- a/api/tests/unit_tests/services/tools/test_builtin_tools_manage_service.py +++ b/api/tests/unit_tests/services/tools/test_builtin_tools_manage_service.py @@ -354,7 +354,7 @@ class TestGetBuiltinToolProviderCredentialInfo: def test_returns_credential_info(self, mock_tm, mock_creds, mock_oauth): mock_tm.get_builtin_provider.return_value.get_supported_credential_types.return_value = ["api-key"] - result = BuiltinToolManageService.get_builtin_tool_provider_credential_info("t", "google") + result = BuiltinToolManageService.get_builtin_tool_provider_credential_info("t", "google", session=MagicMock()) assert result.credentials == [] assert result.supported_credential_types == ["api-key"] @@ -368,7 +368,7 @@ class TestGetBuiltinToolProviderCredentials: mock_db.session.no_autoflush.__exit__ = MagicMock(return_value=False) mock_db.session.scalars.return_value.all.return_value = [] - result = BuiltinToolManageService.get_builtin_tool_provider_credentials("t", "google") + result = BuiltinToolManageService.get_builtin_tool_provider_credentials("t", "google", session=mock_db.session) assert result == [] @@ -391,7 +391,7 @@ class TestGetBuiltinToolProviderCredentials: credential_entity = MagicMock() mock_transform.convert_builtin_provider_to_credential_entity.return_value = credential_entity - result = BuiltinToolManageService.get_builtin_tool_provider_credentials("t", "google") + result = BuiltinToolManageService.get_builtin_tool_provider_credentials("t", "google", session=mock_db.session) assert len(result) == 1 assert result[0] is credential_entity diff --git a/api/tests/unit_tests/services/workflow/test_node_output_inspector_service.py b/api/tests/unit_tests/services/workflow/test_node_output_inspector_service.py index 6f6c56fd67f..7f720575154 100644 --- a/api/tests/unit_tests/services/workflow/test_node_output_inspector_service.py +++ b/api/tests/unit_tests/services/workflow/test_node_output_inspector_service.py @@ -1,8 +1,8 @@ """Unit tests for NodeOutputInspectorService (Stage 4 §8). The service reads from postgres and resolves agent v2 bindings; this suite -mocks ``session_factory`` and the binding resolver so we exercise the -view-construction logic without DB / network access. +mocks the DB session and binding resolver so we exercise the view-construction +logic without DB / network access. """ from __future__ import annotations @@ -100,26 +100,17 @@ def _non_agent_node(*, node_id: str = "tool-node-1", node_type: str = "tool", ti } -def _patch_session( +def _mock_session( *, workflow_run: SimpleNamespace | None, executions: list[SimpleNamespace] | None = None, ): - """Patch ``session_factory.create_session`` to return the configured rows. - - Returns a context manager that the test uses with ``with``. - """ + """Build a mock DB session with the configured rows.""" executions = executions or [] - mock_session = MagicMock() - mock_session.scalar.return_value = workflow_run - mock_session.scalars.return_value.all.return_value = executions - cm = MagicMock() - cm.__enter__.return_value = mock_session - cm.__exit__.return_value = False - return patch( - "services.workflow.node_output_inspector_service.session_factory.create_session", - return_value=cm, - ) + session = MagicMock() + session.scalar.return_value = workflow_run + session.scalars.return_value.all.return_value = executions + return session def _stub_binding_resolver(*, declared_outputs: list[DeclaredOutputConfig]): @@ -149,9 +140,9 @@ def _make_service(declared_outputs: list[DeclaredOutputConfig] | None = None) -> def test_snapshot_404_when_workflow_run_missing(): service = _make_service() - with _patch_session(workflow_run=None): - with pytest.raises(NodeOutputInspectorError) as exc: - service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="missing") + session = _mock_session(workflow_run=None) + with pytest.raises(NodeOutputInspectorError) as exc: + service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="missing", session=session) assert exc.value.code == "workflow_run_not_found" @@ -162,8 +153,8 @@ def test_snapshot_accepts_published_run_d1_lifted(): nodes=[_agent_v2_node(node_id="agent-1")], triggered_from=WorkflowRunTriggeredFrom.APP_RUN, ) - with _patch_session(workflow_run=run, executions=[]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) assert snapshot.workflow_run_id == "run-1" assert [n.node_id for n in snapshot.node_outputs] == ["agent-1"] @@ -175,17 +166,17 @@ def test_snapshot_accepts_webhook_triggered_run(): nodes=[_agent_v2_node(node_id="agent-1")], triggered_from=WorkflowRunTriggeredFrom.WEBHOOK, ) - with _patch_session(workflow_run=run, executions=[]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) assert snapshot.workflow_run_id == "run-1" def test_node_detail_404_when_node_id_absent_from_graph(): service = _make_service() run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) - with _patch_session(workflow_run=run, executions=[]): - with pytest.raises(NodeOutputInspectorError) as exc: - service.node_detail(app_model=_app_model(), workflow_run_id="run-1", node_id="ghost") + session = _mock_session(workflow_run=run, executions=[]) + with pytest.raises(NodeOutputInspectorError) as exc: + service.node_detail(app_model=_app_model(), workflow_run_id="run-1", node_id="ghost", session=session) assert exc.value.code == "node_not_in_workflow_run" @@ -195,28 +186,30 @@ def test_output_preview_404_when_output_name_unknown(): ) run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) ex = _execution(node_id="agent-1", outputs={"text": "hello"}) - with _patch_session(workflow_run=run, executions=[ex]): - with pytest.raises(NodeOutputInspectorError) as exc: - service.output_preview( - app_model=_app_model(), - workflow_run_id="run-1", - node_id="agent-1", - output_name="missing", - ) + session = _mock_session(workflow_run=run, executions=[ex]) + with pytest.raises(NodeOutputInspectorError) as exc: + service.output_preview( + app_model=_app_model(), + workflow_run_id="run-1", + node_id="agent-1", + output_name="missing", + session=session, + ) assert exc.value.code == "node_output_not_declared" def test_output_preview_404_when_node_id_absent_from_graph(): service = _make_service() run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) - with _patch_session(workflow_run=run, executions=[]): - with pytest.raises(NodeOutputInspectorError) as exc: - service.output_preview( - app_model=_app_model(), - workflow_run_id="run-1", - node_id="ghost", - output_name="report", - ) + session = _mock_session(workflow_run=run, executions=[]) + with pytest.raises(NodeOutputInspectorError) as exc: + service.output_preview( + app_model=_app_model(), + workflow_run_id="run-1", + node_id="ghost", + output_name="report", + session=session, + ) assert exc.value.code == "node_not_in_workflow_run" @@ -230,8 +223,8 @@ def test_snapshot_status_pending_when_node_has_no_execution(): declared_outputs=[DeclaredOutputConfig(name="text", type=DeclaredOutputType.STRING)], ) run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) - with _patch_session(workflow_run=run, executions=[]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) assert len(snapshot.node_outputs) == 1 node = snapshot.node_outputs[0] @@ -245,8 +238,8 @@ def test_snapshot_status_running(): ) run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) ex = _execution(node_id="agent-1", status=WorkflowNodeExecutionStatus.RUNNING) - with _patch_session(workflow_run=run, executions=[ex]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[ex]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) assert snapshot.node_outputs[0].node_status == NodeStatus.RUNNING assert snapshot.node_outputs[0].outputs[0].status == NodeOutputStatus.RUNNING @@ -260,8 +253,8 @@ def test_snapshot_status_failed_node_marks_all_outputs_failed(): ) run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) ex = _execution(node_id="agent-1", status=WorkflowNodeExecutionStatus.FAILED) - with _patch_session(workflow_run=run, executions=[ex]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[ex]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) statuses = {o.name: o.status for o in snapshot.node_outputs[0].outputs} assert statuses == {"a": NodeOutputStatus.FAILED, "b": NodeOutputStatus.FAILED} @@ -272,8 +265,8 @@ def test_snapshot_status_ready_when_outputs_present_and_no_failure_metadata(): ) run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) ex = _execution(node_id="agent-1", outputs={"text": "hello"}) - with _patch_session(workflow_run=run, executions=[ex]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[ex]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) output = snapshot.node_outputs[0].outputs[0] assert output.status == NodeOutputStatus.READY assert output.value_preview == "hello" @@ -294,8 +287,8 @@ def test_snapshot_marks_type_check_failure(): } }, ) - with _patch_session(workflow_run=run, executions=[ex]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[ex]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) output = snapshot.node_outputs[0].outputs[0] assert output.status == NodeOutputStatus.TYPE_CHECK_FAILED assert output.type_check is not None @@ -324,14 +317,12 @@ def test_snapshot_marks_output_check_failure_when_type_check_passed(): }, }, ) - with ( - _patch_session(workflow_run=run, executions=[ex]), - patch( - "services.workflow.node_output_inspector_service.file_helpers.get_signed_file_url", - return_value="https://signed.example/x", - ), + session = _mock_session(workflow_run=run, executions=[ex]) + with patch( + "services.workflow.node_output_inspector_service.file_helpers.get_signed_file_url", + return_value="https://signed.example/x", ): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) output = snapshot.node_outputs[0].outputs[0] assert output.status == NodeOutputStatus.OUTPUT_CHECK_FAILED assert output.output_check is not None @@ -348,8 +339,8 @@ def test_snapshot_marks_not_produced_when_declared_output_missing_from_payload() ) run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) ex = _execution(node_id="agent-1", outputs={"text": "hi"}) # optional_meta missing - with _patch_session(workflow_run=run, executions=[ex]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[ex]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) statuses = {o.name: o.status for o in snapshot.node_outputs[0].outputs} assert statuses == {"text": NodeOutputStatus.READY, "optional_meta": NodeOutputStatus.NOT_PRODUCED} @@ -367,8 +358,8 @@ def test_non_agent_node_outputs_inferred_from_payload_keys(): node_type="tool", outputs={"message": "sent", "thread_ts": "1234"}, ) - with _patch_session(workflow_run=run, executions=[ex]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[ex]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) output_names = sorted(o.name for o in snapshot.node_outputs[0].outputs) assert output_names == ["message", "thread_ts"] # All inferred outputs should have ``type=None`` since we don't know the @@ -393,14 +384,12 @@ def test_file_output_preview_includes_signed_url(): "reference": build_file_reference(record_id="550e8400-e29b-41d4-a716-446655440000"), } ex = _execution(node_id="agent-1", outputs={"report": file_payload}) - with ( - _patch_session(workflow_run=run, executions=[ex]), - patch( - "services.workflow.node_output_inspector_service._resolve_preview_url", - return_value="https://signed.example/x.pdf", - ), + session = _mock_session(workflow_run=run, executions=[ex]) + with patch( + "services.workflow.node_output_inspector_service._resolve_preview_url", + return_value="https://signed.example/x.pdf", ): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) preview_value = snapshot.node_outputs[0].outputs[0].value_preview assert isinstance(preview_value, dict) assert preview_value["preview_url"] == "https://signed.example/x.pdf" @@ -419,18 +408,17 @@ def test_file_output_preview_endpoint_returns_full_value_with_signed_url(): "reference": build_file_reference(record_id="550e8400-e29b-41d4-a716-446655440000"), } ex = _execution(node_id="agent-1", outputs={"report": file_payload}) - with ( - _patch_session(workflow_run=run, executions=[ex]), - patch( - "services.workflow.node_output_inspector_service._resolve_preview_url", - return_value="https://signed.example/x.pdf", - ), + session = _mock_session(workflow_run=run, executions=[ex]) + with patch( + "services.workflow.node_output_inspector_service._resolve_preview_url", + return_value="https://signed.example/x.pdf", ): preview = service.output_preview( app_model=_app_model(), workflow_run_id="run-1", node_id="agent-1", output_name="report", + session=session, ) assert preview.output_name == "report" assert preview.status == NodeOutputStatus.READY @@ -484,26 +472,25 @@ def test_array_file_output_preview_includes_signed_urls_for_each_item(): }, ] ex = _execution(node_id="agent-1", outputs={"files": file_payloads}) - with ( - _patch_session(workflow_run=run, executions=[ex]), - patch( - "services.workflow.node_output_inspector_service._resolve_preview_url", - side_effect=[ - "https://signed.example/1.pdf", - "https://signed.example/2.pdf", - "https://signed.example/1-detail.pdf", - "https://signed.example/2-detail.pdf", - "https://signed.example/1-full.pdf", - "https://signed.example/2-full.pdf", - ], - ), + session = _mock_session(workflow_run=run, executions=[ex]) + with patch( + "services.workflow.node_output_inspector_service._resolve_preview_url", + side_effect=[ + "https://signed.example/1.pdf", + "https://signed.example/2.pdf", + "https://signed.example/1-detail.pdf", + "https://signed.example/2-detail.pdf", + "https://signed.example/1-full.pdf", + "https://signed.example/2-full.pdf", + ], ): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) preview = service.output_preview( app_model=_app_model(), workflow_run_id="run-1", node_id="agent-1", output_name="files", + session=session, ) snapshot_value = snapshot.node_outputs[0].outputs[0].value_preview @@ -531,14 +518,12 @@ def test_file_output_preview_uses_none_when_signed_url_resolution_fails(): "reference": build_file_reference(record_id="550e8400-e29b-41d4-a716-446655440000"), } ex = _execution(node_id="agent-1", outputs={"report": file_payload}) - with ( - _patch_session(workflow_run=run, executions=[ex]), - patch( - "services.workflow.node_output_inspector_service._resolve_preview_url", - side_effect=RuntimeError("boom"), - ), + session = _mock_session(workflow_run=run, executions=[ex]) + with patch( + "services.workflow.node_output_inspector_service._resolve_preview_url", + side_effect=RuntimeError("boom"), ): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) preview_value = snapshot.node_outputs[0].outputs[0].value_preview assert isinstance(preview_value, dict) @@ -557,19 +542,18 @@ def test_object_output_preview_does_not_augment_canonical_file_mapping_shape(): "reference": build_file_reference(record_id="550e8400-e29b-41d4-a716-446655440000"), } ex = _execution(node_id="agent-1", outputs={"meta": raw_value}) - with ( - _patch_session(workflow_run=run, executions=[ex]), - patch( - "services.workflow.node_output_inspector_service._resolve_preview_url", - return_value="https://signed.example/x.pdf", - ), + session = _mock_session(workflow_run=run, executions=[ex]) + with patch( + "services.workflow.node_output_inspector_service._resolve_preview_url", + return_value="https://signed.example/x.pdf", ): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) preview = service.output_preview( app_model=_app_model(), workflow_run_id="run-1", node_id="agent-1", output_name="meta", + session=session, ) assert snapshot.node_outputs[0].outputs[0].value_preview == raw_value @@ -591,8 +575,8 @@ def test_retried_count_pulled_from_attempt_metadata(): outputs={"text": "ok"}, execution_metadata={"attempt": 2}, ) - with _patch_session(workflow_run=run, executions=[ex]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[ex]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) assert snapshot.node_outputs[0].outputs[0].retried == 2 @@ -610,8 +594,8 @@ def test_keeps_latest_execution_per_node_by_index(): run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) older = _execution(node_id="agent-1", outputs={"text": "old"}, index=1) newer = _execution(node_id="agent-1", outputs={"text": "new"}, index=5) - with _patch_session(workflow_run=run, executions=[older, newer]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[older, newer]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) assert snapshot.node_outputs[0].outputs[0].value_preview == "new" @@ -632,8 +616,8 @@ def test_array_typed_output_with_array_item_renders_correctly(): ) run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) ex = _execution(node_id="agent-1", outputs={"files": []}) - with _patch_session(workflow_run=run, executions=[ex]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[ex]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) output = snapshot.node_outputs[0].outputs[0] assert output.type == DeclaredOutputType.ARRAY @@ -654,6 +638,6 @@ def test_unparseable_graph_blob_yields_empty_snapshot_not_500(): status=WorkflowExecutionStatus.RUNNING, graph="{not valid json", ) - with _patch_session(workflow_run=run, executions=[]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) assert snapshot.node_outputs == [] diff --git a/api/tests/unit_tests/services/workflow/test_workflow_converter_additional.py b/api/tests/unit_tests/services/workflow/test_workflow_converter_additional.py index 2aaf3bdf1d5..f471e4aeb56 100644 --- a/api/tests/unit_tests/services/workflow/test_workflow_converter_additional.py +++ b/api/tests/unit_tests/services/workflow/test_workflow_converter_additional.py @@ -118,6 +118,7 @@ def test__convert_to_http_request_node_for_chatbot(default_variables: list[Varia app_model=app_model, variables=default_variables, external_data_variables=external_data_variables, + session=MagicMock(), ) assert len(nodes) == 2 @@ -160,6 +161,7 @@ def test__convert_to_http_request_node_for_workflow_app(default_variables: list[ app_model=app_model, variables=default_variables, external_data_variables=external_data_variables, + session=MagicMock(), ) body = json.loads(nodes[0]["data"]["body"]["data"]) @@ -364,6 +366,7 @@ def test_convert_to_workflow_should_raise_when_app_model_config_is_missing(conve icon_type="emoji", icon="robot", icon_background="#fff", + session=MagicMock(), ) @@ -389,7 +392,6 @@ def test_convert_to_workflow_should_create_new_app_with_fallback_fields( monkeypatch.setattr(converter_module, "App", FakeApp) db_session = SimpleNamespace(add=MagicMock(), flush=MagicMock(), commit=MagicMock()) - monkeypatch.setattr(converter_module, "db", SimpleNamespace(session=db_session)) send_mock = MagicMock() monkeypatch.setattr(converter_module.app_was_created, "send", send_mock) @@ -417,6 +419,7 @@ def test_convert_to_workflow_should_create_new_app_with_fallback_fields( icon_type="", icon="", icon_background="", + session=db_session, ) assert new_app.name == "Source App(workflow)" @@ -501,12 +504,12 @@ def test_convert_app_model_config_to_workflow_should_build_advanced_chat_graph_a monkeypatch.setattr(converter_module, "Workflow", FakeWorkflow) db_session = SimpleNamespace(add=MagicMock(), commit=MagicMock()) - monkeypatch.setattr(converter_module, "db", SimpleNamespace(session=db_session)) workflow = converter.convert_app_model_config_to_workflow( app_model=app_model, app_model_config=_app_model_config(id="cfg"), account_id="account-1", + session=db_session, ) graph = json.loads(workflow.graph) @@ -568,12 +571,12 @@ def test_convert_app_model_config_to_workflow_should_build_workflow_mode_with_en monkeypatch.setattr(converter_module, "Workflow", FakeWorkflow) db_session = SimpleNamespace(add=MagicMock(), commit=MagicMock()) - monkeypatch.setattr(converter_module, "db", SimpleNamespace(session=db_session)) workflow = converter.convert_app_model_config_to_workflow( app_model=app_model, app_model_config=_app_model_config(id="cfg"), account_id="account-1", + session=db_session, ) graph = json.loads(workflow.graph) @@ -644,6 +647,7 @@ def test_convert_to_http_request_node_should_skip_non_api_and_missing_extension_ app_model=app_model, variables=[], external_data_variables=external_data_variables, + session=MagicMock(), ) assert nodes == [] @@ -810,10 +814,9 @@ def test_get_api_based_extension_should_raise_when_extension_not_found( monkeypatch: pytest.MonkeyPatch, ) -> None: db_session = SimpleNamespace(scalar=MagicMock(return_value=None)) - monkeypatch.setattr(converter_module, "db", SimpleNamespace(session=db_session)) with pytest.raises(ValueError, match="API Based Extension not found"): - converter._get_api_based_extension(tenant_id="tenant-1", api_based_extension_id="ext-1") + converter._get_api_based_extension(tenant_id="tenant-1", api_based_extension_id="ext-1", session=db_session) db_session.scalar.assert_called_once() @@ -823,9 +826,10 @@ def test_get_api_based_extension_should_return_entity_when_found( ) -> None: extension = SimpleNamespace(id="ext-1") db_session = SimpleNamespace(scalar=MagicMock(return_value=extension)) - monkeypatch.setattr(converter_module, "db", SimpleNamespace(session=db_session)) - result = converter._get_api_based_extension(tenant_id="tenant-1", api_based_extension_id="ext-1") + result = converter._get_api_based_extension( + tenant_id="tenant-1", api_based_extension_id="ext-1", session=db_session + ) assert result is extension db_session.scalar.assert_called_once() diff --git a/api/tests/unit_tests/services/workflow/test_workflow_human_input_delivery.py b/api/tests/unit_tests/services/workflow/test_workflow_human_input_delivery.py index 5bcb13c360c..cd97fa2e53b 100644 --- a/api/tests/unit_tests/services/workflow/test_workflow_human_input_delivery.py +++ b/api/tests/unit_tests/services/workflow/test_workflow_human_input_delivery.py @@ -64,6 +64,7 @@ def test_human_input_delivery_requires_draft_workflow(): account=account, node_id="node-1", delivery_method_id="delivery-1", + session=MagicMock(), ) @@ -98,6 +99,7 @@ def test_human_input_delivery_allows_disabled_method(monkeypatch: pytest.MonkeyP account=account, node_id="node-1", delivery_method_id=str(delivery_method.id), + session=MagicMock(), ) test_service_instance.send_test.assert_called_once() @@ -135,6 +137,7 @@ def test_human_input_delivery_dispatches_to_test_service(monkeypatch: pytest.Mon node_id="node-1", delivery_method_id=str(delivery_method.id), inputs={"#node-1.output#": "value"}, + session=MagicMock(), ) pool_args = service._build_human_input_variable_pool.call_args.kwargs @@ -173,6 +176,7 @@ def test_human_input_delivery_debug_mode_overrides_recipients(monkeypatch: pytes account=account, node_id="node-1", delivery_method_id=str(delivery_method.id), + session=MagicMock(), ) test_service_instance.send_test.assert_called_once() From 6c0aa3ed0d8188a658193f582eee135d985e60ae Mon Sep 17 00:00:00 2001 From: yyh <92089059+lyzno1@users.noreply.github.com> Date: Wed, 8 Jul 2026 11:24:22 +0800 Subject: [PATCH 37/70] test(dify-ui): remove brittle primitive assertions (#38529) --- .../src/avatar/__tests__/index.spec.tsx | 67 ++++++++++++++----- .../src/progress/__tests__/index.spec.tsx | 33 ++------- 2 files changed, 54 insertions(+), 46 deletions(-) diff --git a/packages/dify-ui/src/avatar/__tests__/index.spec.tsx b/packages/dify-ui/src/avatar/__tests__/index.spec.tsx index e427b09ed0d..2287358b32d 100644 --- a/packages/dify-ui/src/avatar/__tests__/index.spec.tsx +++ b/packages/dify-ui/src/avatar/__tests__/index.spec.tsx @@ -1,6 +1,30 @@ import { render } from 'vitest-browser-react' import { Avatar, AvatarFallback, AvatarImage, AvatarRoot } from '..' +function stubImageLoader() { + const originalImage = window.Image + const images: HTMLImageElement[] = [] + + function TestImage(_width?: number, _height?: number): HTMLImageElement { + const image = document.createElement('img') + images.push(image) + return image + } + + Object.defineProperty(window, 'Image', { + configurable: true, + value: TestImage, + writable: true, + }) + + return { + images, + restore: () => { + window.Image = originalImage + }, + } +} + describe('Avatar', () => { describe('Rendering', () => { it('should keep the fallback visible when avatar URL is provided before image load', async () => { @@ -69,7 +93,7 @@ describe('Avatar', () => { }) it('should handle empty string avatar as falsy value', async () => { - const screen = await render() + const screen = await render() expect(screen.container.querySelector('img')).not.toBeInTheDocument() await expect.element(screen.getByText('T')).toBeInTheDocument() @@ -77,25 +101,32 @@ describe('Avatar', () => { }) describe('onLoadingStatusChange', () => { - it('should render the fallback when avatar and onLoadingStatusChange are provided', async () => { - const screen = await render( - , - ) - - await expect.element(screen.getByText('J')).toBeInTheDocument() - }) - - it('should not render image when avatar is null even with onLoadingStatusChange', async () => { + it('should forward image loading status changes', async () => { + const { images, restore } = stubImageLoader() const onStatusChange = vi.fn() - const screen = await render( - , - ) - expect(screen.container.querySelector('img')).not.toBeInTheDocument() + try { + await render( + , + ) + + await vi.waitFor(() => { + expect(onStatusChange).toHaveBeenCalledWith('loading') + }) + + images[0]?.onload?.(new Event('load')) + + await vi.waitFor(() => { + expect(onStatusChange).toHaveBeenCalledWith('loaded') + }) + } + finally { + restore() + } }) }) }) diff --git a/packages/dify-ui/src/progress/__tests__/index.spec.tsx b/packages/dify-ui/src/progress/__tests__/index.spec.tsx index 987aef228d8..2e86633247b 100644 --- a/packages/dify-ui/src/progress/__tests__/index.spec.tsx +++ b/packages/dify-ui/src/progress/__tests__/index.spec.tsx @@ -24,34 +24,11 @@ describe('ProgressCircle', () => { }) it('renders indeterminate state when value is null', async () => { - const screen = await render() + const screen = await render() + const progress = screen.getByLabelText('Processing') - await expect.element(screen.getByTestId('progress')).toHaveAttribute('data-indeterminate') - await expect.element(screen.getByTestId('progress')).not.toHaveAttribute('aria-valuenow') - expect(screen.getByTestId('progress').element().querySelector('path')).toBeNull() - }) - - it('does not render a progress sector for zero progress', async () => { - const screen = await render() - - expect(screen.getByTestId('progress').element().querySelector('path')).toBeNull() - }) - - it('renders a deterministic progress sector', async () => { - const screen = await render() - - const path = screen.getByTestId('progress').element().querySelector('path')! - - expect(path.getAttribute('d')).toContain('A 6,6 0 1 1') - }) - - it('renders a closed circle sector for complete progress', async () => { - const screen = await render() - - const path = screen.getByTestId('progress').element().querySelector('path')! - const pathData = path.getAttribute('d')! - - expect(pathData).toContain('A 6,6 0 1 1 6,12') - expect(pathData).toContain('A 6,6 0 1 1 6,0') + await expect.element(progress).toHaveAttribute('role', 'progressbar') + await expect.element(progress).toHaveAttribute('data-indeterminate') + await expect.element(progress).not.toHaveAttribute('aria-valuenow') }) }) From 76a6cd3335b42c9270a5978fcc9c674619ac5488 Mon Sep 17 00:00:00 2001 From: Stephen Zhou Date: Wed, 8 Jul 2026 11:51:50 +0800 Subject: [PATCH 38/70] refactor(web): migrate app context consumers (#38530) --- .../apps/app-list-browsing-flow.test.tsx | 18 +++ web/__tests__/apps/create-app-flow.test.tsx | 18 +++ web/__tests__/utils/mock-app-context-state.ts | 153 ++++++++++++++++++ .../[appId]/__tests__/layout-main.spec.tsx | 39 +++-- .../(appDetailLayout)/[appId]/layout-main.tsx | 21 ++- .../overview/__tests__/card-view.spec.tsx | 31 +++- .../[appId]/overview/card-view.tsx | 7 +- .../[appId]/overview/chart-view.tsx | 7 +- .../[appId]/overview/tracing/panel.tsx | 7 +- .../[appId]/overview/view.tsx | 7 +- .../access-config/__tests__/index.spec.tsx | 12 ++ .../components/app/access-config/index.tsx | 10 +- .../add-member-or-group-pop.tsx | 5 +- .../dataset-config/__tests__/index.spec.tsx | 15 ++ .../configuration/dataset-config/index.tsx | 7 +- .../select-dataset/__tests__/index.spec.tsx | 14 ++ .../dataset-config/select-dataset/index.tsx | 5 +- .../debug-with-multiple-model/chat-item.tsx | 5 +- .../debug/debug-with-single-model/index.tsx | 5 +- .../__tests__/use-configuration.spec.tsx | 17 ++ .../configuration/hooks/use-configuration.ts | 17 +- .../app-list/__tests__/index.spec.tsx | 15 ++ .../app/create-app-dialog/app-list/index.tsx | 10 +- .../create-app-modal/__tests__/index.spec.tsx | 16 ++ .../components/app/create-app-modal/index.tsx | 10 +- .../__tests__/index.spec.tsx | 15 ++ .../app/create-from-dsl-modal/index.tsx | 14 +- web/app/components/app/log/empty-element.tsx | 7 +- web/app/components/app/overview/app-card.tsx | 7 +- .../app/overview/embedded/index.tsx | 5 +- .../components/app/overview/trigger-card.tsx | 7 +- .../apps/__tests__/app-card.spec.tsx | 12 ++ .../apps/__tests__/creators-filter.spec.tsx | 14 ++ .../components/apps/__tests__/index.spec.tsx | 14 ++ .../components/apps/__tests__/list.spec.tsx | 15 ++ web/app/components/apps/app-card.tsx | 11 +- web/app/components/apps/creators-filter.tsx | 8 +- web/app/components/apps/index.tsx | 5 +- web/app/components/apps/list.tsx | 5 +- web/app/components/apps/starred-app-card.tsx | 7 +- .../snippet-list/__tests__/index.spec.tsx | 17 ++ .../__tests__/access-surface-cards.spec.tsx | 25 +++ 42 files changed, 563 insertions(+), 96 deletions(-) create mode 100644 web/__tests__/utils/mock-app-context-state.ts diff --git a/web/__tests__/apps/app-list-browsing-flow.test.tsx b/web/__tests__/apps/app-list-browsing-flow.test.tsx index 59e53051e11..972f538abfa 100644 --- a/web/__tests__/apps/app-list-browsing-flow.test.tsx +++ b/web/__tests__/apps/app-list-browsing-flow.test.tsx @@ -75,6 +75,24 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateAtomMock(importOriginal, () => ({ + userProfile: { id: 'user-1' }, + currentWorkspace: { id: 'workspace-1' }, + isLoadingCurrentWorkspace: mockIsLoadingCurrentWorkspace, + isLoadingWorkspacePermissionKeys: mockIsLoadingCurrentWorkspace, + workspacePermissionKeys: mockWorkspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/provider-context', () => ({ useProviderContext: () => ({ onPlanInfoChanged: vi.fn(), diff --git a/web/__tests__/apps/create-app-flow.test.tsx b/web/__tests__/apps/create-app-flow.test.tsx index 4900593d505..b93de421d1b 100644 --- a/web/__tests__/apps/create-app-flow.test.tsx +++ b/web/__tests__/apps/create-app-flow.test.tsx @@ -61,6 +61,24 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateAtomMock(importOriginal, () => ({ + userProfile: { id: 'user-1' }, + currentWorkspace: { id: 'workspace-1' }, + isLoadingCurrentWorkspace: mockIsLoadingCurrentWorkspace, + isLoadingWorkspacePermissionKeys: mockIsLoadingCurrentWorkspace, + workspacePermissionKeys: mockWorkspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/provider-context', () => ({ useProviderContext: () => ({ onPlanInfoChanged: mockOnPlanInfoChanged, diff --git a/web/__tests__/utils/mock-app-context-state.ts b/web/__tests__/utils/mock-app-context-state.ts new file mode 100644 index 00000000000..99795f60766 --- /dev/null +++ b/web/__tests__/utils/mock-app-context-state.ts @@ -0,0 +1,153 @@ +import type { LangGeniusVersionResponse } from '@/models/common' + +const APP_CONTEXT_STATE_ATOM_KIND = Symbol('app-context-state-atom-kind') + +export type AppContextStateMockState = { + userProfile?: { + id?: string + name?: string + email?: string + avatar?: string + avatar_url?: string + is_password_set?: boolean + } | null + currentWorkspace?: { + id?: string + } | null + isLoadingCurrentWorkspace?: boolean + isLoadingWorkspacePermissionKeys?: boolean + workspacePermissionKeys?: string[] + langGeniusVersionInfo?: LangGeniusVersionResponse +} + +type AppContextStateAtomKind + = | 'userProfile' + | 'userProfileId' + | 'currentWorkspace' + | 'currentWorkspaceId' + | 'currentWorkspaceLoading' + | 'workspacePermissionKeys' + | 'workspacePermissionKeysLoading' + | 'langGeniusVersionInfo' + +type AppContextStateMockAtom = { + [APP_CONTEXT_STATE_ATOM_KIND]: AppContextStateAtomKind +} + +type AppContextStateMockRegistry = { + getState: () => AppContextStateMockState +} + +const defaultUserProfile = { + id: 'user-1', + name: 'User', + email: 'user@example.com', + avatar: '', + avatar_url: '', + is_password_set: true, +} + +const defaultCurrentWorkspace = { + id: 'workspace-1', +} + +const defaultLangGeniusVersionInfo = { + current_env: 'CLOUD', + current_version: '', + latest_version: '', + version: '', + release_date: '', + release_notes: '', + can_auto_update: false, +} satisfies LangGeniusVersionResponse + +let appContextStateMockRegistry: AppContextStateMockRegistry | undefined + +const createMockAtom = ( + kind: AppContextStateAtomKind, +): AppContextStateMockAtom => ({ + [APP_CONTEXT_STATE_ATOM_KIND]: kind, +}) + +const isAppContextStateMockAtom = (atom: unknown): atom is AppContextStateMockAtom => { + return typeof atom === 'object' && atom !== null && APP_CONTEXT_STATE_ATOM_KIND in atom +} + +const getUserProfile = (state: AppContextStateMockState) => ({ + ...defaultUserProfile, + ...state.userProfile, +}) + +const getCurrentWorkspace = (state: AppContextStateMockState) => ({ + ...defaultCurrentWorkspace, + ...state.currentWorkspace, +}) + +export const createAppContextStateAtomMock = async ( + importOriginal: () => Promise, + getState: () => AppContextStateMockState, +) => { + const actual = await importOriginal() + appContextStateMockRegistry = { + getState, + } + + return { + ...actual, + userProfileAtom: createMockAtom('userProfile'), + userProfileIdAtom: createMockAtom('userProfileId'), + currentWorkspaceAtom: createMockAtom('currentWorkspace'), + currentWorkspaceIdAtom: createMockAtom('currentWorkspaceId'), + currentWorkspaceLoadingAtom: createMockAtom('currentWorkspaceLoading'), + workspacePermissionKeysAtom: createMockAtom('workspacePermissionKeys'), + workspacePermissionKeysLoadingAtom: createMockAtom('workspacePermissionKeysLoading'), + langGeniusVersionInfoAtom: createMockAtom('langGeniusVersionInfo'), + } +} + +export const createAppContextStateJotaiMock = async ( + importOriginal: () => Promise, +) => { + const actual = await importOriginal() + + return { + ...actual, + useAtomValue: (atom: unknown) => { + if (!isAppContextStateMockAtom(atom)) + return actual.useAtomValue(atom as Parameters[0]) + + if (!appContextStateMockRegistry) + throw new Error('App context state atom mock is not initialized') + + const state = appContextStateMockRegistry.getState() + const userProfile = getUserProfile(state) + const currentWorkspace = getCurrentWorkspace(state) + + if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'userProfile') + return userProfile + + if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'userProfileId') + return userProfile.id + + if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'currentWorkspace') + return currentWorkspace + + if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'currentWorkspaceId') + return currentWorkspace.id + + if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'currentWorkspaceLoading') + return state.isLoadingCurrentWorkspace ?? false + + if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'workspacePermissionKeys') + return state.workspacePermissionKeys ?? [] + + if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'workspacePermissionKeysLoading') + return state.isLoadingWorkspacePermissionKeys ?? false + + if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'langGeniusVersionInfo') + return state.langGeniusVersionInfo ?? defaultLangGeniusVersionInfo + + throw new Error(`Unsupported app context state atom: ${atom[APP_CONTEXT_STATE_ATOM_KIND]}`) + }, + } +} diff --git a/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/__tests__/layout-main.spec.tsx b/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/__tests__/layout-main.spec.tsx index 8d4b0e2831c..ebf508b1345 100644 --- a/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/__tests__/layout-main.spec.tsx +++ b/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/__tests__/layout-main.spec.tsx @@ -10,8 +10,14 @@ import AppDetailLayout from '../layout-main' const mockReplace = vi.fn() let mockPathname = '/app/app-1/workflow' -let mockIsLoadingWorkspacePermissionKeys = false let mockIsRbacEnabled = true +const mockAppContextState = vi.hoisted(() => ({ + currentWorkspace: { id: 'workspace-1' }, + isLoadingCurrentWorkspace: false, + isLoadingWorkspacePermissionKeys: false, + userProfile: { id: 'user-1' }, + workspacePermissionKeys: [] as string[], +})) const render = (ui: Parameters[0]) => renderWithSystemFeatures(ui, { systemFeatures: { @@ -28,16 +34,17 @@ vi.mock('@/service/apps', () => ({ fetchAppDetailDirect: vi.fn(), })) -vi.mock('@/context/app-context', () => ({ - useAppContext: () => ({ - currentWorkspace: { id: 'workspace-1' }, - isCurrentWorkspaceDatasetOperator: false, - isLoadingCurrentWorkspace: false, - isLoadingWorkspacePermissionKeys: mockIsLoadingWorkspacePermissionKeys, - userProfile: { id: 'user-1' }, - workspacePermissionKeys: [], - }), -})) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateAtomMock(importOriginal, () => mockAppContextState) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateJotaiMock(importOriginal) +}) vi.mock('@/hooks/use-document-title', () => ({ default: vi.fn(), @@ -65,8 +72,12 @@ describe('AppDetailLayout', () => { beforeEach(() => { vi.clearAllMocks() mockPathname = '/app/app-1/workflow' - mockIsLoadingWorkspacePermissionKeys = false mockIsRbacEnabled = true + mockAppContextState.currentWorkspace = { id: 'workspace-1' } + mockAppContextState.isLoadingCurrentWorkspace = false + mockAppContextState.isLoadingWorkspacePermissionKeys = false + mockAppContextState.userProfile = { id: 'user-1' } + mockAppContextState.workspacePermissionKeys = [] mockUsePathname.mockImplementation(() => mockPathname) mockUseRouter.mockReturnValue({ back: vi.fn(), @@ -240,7 +251,7 @@ describe('AppDetailLayout', () => { }) it('should wait for workspace permission keys before redirecting restricted pages', async () => { - mockIsLoadingWorkspacePermissionKeys = true + mockAppContextState.isLoadingWorkspacePermissionKeys = true mockPathname = '/app/app-1/overview' mockFetchAppDetailDirect.mockResolvedValue(createAppDetail({ permission_keys: [AppACLPermission.ViewLayout] })) @@ -256,7 +267,7 @@ describe('AppDetailLayout', () => { expect(mockReplace).not.toHaveBeenCalled() expect(screen.queryByText('App page content')).not.toBeInTheDocument() - mockIsLoadingWorkspacePermissionKeys = false + mockAppContextState.isLoadingWorkspacePermissionKeys = false rerender(
    App page content
    diff --git a/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/layout-main.tsx b/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/layout-main.tsx index 7ef9c3e31ea..3bdf0113de2 100644 --- a/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/layout-main.tsx +++ b/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/layout-main.tsx @@ -3,13 +3,20 @@ import type { FC } from 'react' import type { App } from '@/types/app' import { cn } from '@langgenius/dify-ui/cn' import { useSuspenseQuery } from '@tanstack/react-query' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useEffect, useState } from 'react' import { useTranslation } from 'react-i18next' import { useShallow } from 'zustand/react/shallow' import { useStore } from '@/app/components/app/store' import Loading from '@/app/components/base/loading' -import { useAppContext } from '@/context/app-context' +import { + currentWorkspaceAtom, + currentWorkspaceLoadingAtom, + userProfileIdAtom, + workspacePermissionKeysAtom, + workspacePermissionKeysLoadingAtom, +} from '@/context/app-context-state' import { systemFeaturesQueryOptions } from '@/features/system-features/client' import useDocumentTitle from '@/hooks/use-document-title' import { usePathname, useRouter } from '@/next/navigation' @@ -39,7 +46,11 @@ const AppDetailLayout: FC = (props) => { const router = useRouter() const pathname = usePathname() const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) - const { isLoadingCurrentWorkspace, isLoadingWorkspacePermissionKeys, currentWorkspace, userProfile, workspacePermissionKeys } = useAppContext() + const isLoadingCurrentWorkspace = useAtomValue(currentWorkspaceLoadingAtom) + const isLoadingWorkspacePermissionKeys = useAtomValue(workspacePermissionKeysLoadingAtom) + const currentWorkspace = useAtomValue(currentWorkspaceAtom) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const isRbacEnabled = systemFeatures.rbac_enabled const { appDetail, setAppDetail } = useStore(useShallow(state => ({ appDetail: state.appDetail, @@ -96,7 +107,7 @@ const AppDetailLayout: FC = (props) => { return const appACLCapabilities = getAppACLCapabilities(routeAppDetail.permission_keys, { - currentUserId: userProfile?.id, + currentUserId, resourceMaintainer: routeAppDetail.maintainer, workspacePermissionKeys, isRbacEnabled, @@ -114,7 +125,7 @@ const AppDetailLayout: FC = (props) => { || (isAccessConfigPath && !appACLCapabilities.canAccessConfig) ) { router.replace(getRedirectionPath(routeAppDetail, { - currentUserId: userProfile?.id, + currentUserId, resourceMaintainer: routeAppDetail.maintainer, workspacePermissionKeys, isRbacEnabled, @@ -131,7 +142,7 @@ const AppDetailLayout: FC = (props) => { if (appDetailRes && appDetail?.id !== appDetailRes.id) setAppDetail({ ...appDetailRes, enable_sso: false }) - }, [appDetail?.id, appDetailRes, appId, currentWorkspace.id, isLoadingAppDetail, isLoadingCurrentWorkspace, isLoadingWorkspacePermissionKeys, isRbacEnabled, pathname, routeAppDetail, router, setAppDetail, userProfile?.id, workspacePermissionKeys]) + }, [appDetail?.id, appDetailRes, appId, currentUserId, currentWorkspace.id, isLoadingAppDetail, isLoadingCurrentWorkspace, isLoadingWorkspacePermissionKeys, isRbacEnabled, pathname, routeAppDetail, router, setAppDetail, workspacePermissionKeys]) const isWorkflowPage = pathname.endsWith('/workflow') const content = !appDetail diff --git a/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/__tests__/card-view.spec.tsx b/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/__tests__/card-view.spec.tsx index cf6360b4bff..1935201e0d3 100644 --- a/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/__tests__/card-view.spec.tsx +++ b/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/__tests__/card-view.spec.tsx @@ -33,11 +33,32 @@ vi.mock('@/service/apps', () => ({ updateAppSiteAccessToken: (...args: unknown[]) => mockUpdateAppSiteAccessToken(...args), })) -vi.mock('@tanstack/react-query', () => ({ - useQueryClient: () => ({ - setQueryData: mockSetQueryData, - }), -})) +vi.mock('@tanstack/react-query', async (importOriginal) => { + const actual = await importOriginal() + + return { + ...actual, + useQueryClient: () => ({ + setQueryData: mockSetQueryData, + }), + } +}) + +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateAtomMock(importOriginal, () => ({ + userProfile: { id: 'user-1' }, + currentWorkspace: { id: 'workspace-1' }, + workspacePermissionKeys: mockAppState.appDetail.permission_keys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateJotaiMock(importOriginal) +}) vi.mock('@/app/components/workflow/collaboration/core/collaboration-manager', () => ({ collaborationManager: { diff --git a/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/card-view.tsx b/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/card-view.tsx index 0a1859b71fd..44af5962d30 100644 --- a/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/card-view.tsx +++ b/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/card-view.tsx @@ -7,6 +7,7 @@ import type { App } from '@/types/app' import type { I18nKeysByPrefix } from '@/types/i18n' import { toast } from '@langgenius/dify-ui/toast' import { useQueryClient } from '@tanstack/react-query' +import { useAtomValue } from 'jotai' import { useCallback, useEffect, useMemo } from 'react' import { useTranslation } from 'react-i18next' import AppCard from '@/app/components/app/overview/app-card' @@ -18,7 +19,7 @@ import MCPServiceCard from '@/app/components/tools/mcp/mcp-service-card' import { collaborationManager } from '@/app/components/workflow/collaboration/core/collaboration-manager' import { webSocketClient } from '@/app/components/workflow/collaboration/core/websocket-manager' import { isTriggerNode } from '@/app/components/workflow/types' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { userProfileIdAtom, workspacePermissionKeysAtom } from '@/context/app-context-state' import { fetchAppDetail, updateAppSiteAccessToken, @@ -42,8 +43,8 @@ const CardView: FC = ({ appId, isInPanel, className }) => { const queryClient = useQueryClient() const appDetail = useAppStore(state => state.appDetail) const setAppDetail = useAppStore(state => state.setAppDetail) - const currentUserId = useAppContextWithSelector(state => state.userProfile?.id) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const canEditApp = useMemo(() => getAppACLCapabilities(appDetail?.permission_keys, { currentUserId, resourceMaintainer: appDetail?.maintainer, diff --git a/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/chart-view.tsx b/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/chart-view.tsx index a4ee7d06a92..6048493c081 100644 --- a/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/chart-view.tsx +++ b/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/chart-view.tsx @@ -3,6 +3,7 @@ import type { PeriodParams } from '@/app/components/app/overview/app-chart' import type { I18nKeysByPrefix } from '@/types/i18n' import dayjs from 'dayjs' import quarterOfYear from 'dayjs/plugin/quarterOfYear' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useState } from 'react' import { useTranslation } from 'react-i18next' @@ -10,7 +11,7 @@ import { TIME_PERIOD_MAPPING as LONG_TIME_PERIOD_MAPPING } from '@/app/component import { AvgResponseTime, AvgSessionInteractions, AvgUserInteractions, ConversationsChart, CostChart, EndUsersChart, MessagesChart, TokenPerSecond, UserSatisfactionRate, WorkflowCostChart, WorkflowDailyTerminalsChart, WorkflowMessagesChart } from '@/app/components/app/overview/app-chart' import { useStore as useAppStore } from '@/app/components/app/store' import { IS_CLOUD_EDITION } from '@/config' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { userProfileIdAtom, workspacePermissionKeysAtom } from '@/context/app-context-state' import { useDocLink } from '@/context/i18n' import { getAppACLCapabilities } from '@/utils/permission' import LongTimeRangePicker from './long-time-range-picker' @@ -39,8 +40,8 @@ export default function ChartView({ appId, headerRight }: IChartViewProps) { const { t } = useTranslation() const docLink = useDocLink() const appDetail = useAppStore(state => state.appDetail) - const currentUserId = useAppContextWithSelector(state => state.userProfile?.id) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const canMonitor = React.useMemo(() => getAppACLCapabilities(appDetail?.permission_keys, { currentUserId, resourceMaintainer: appDetail?.maintainer, diff --git a/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/tracing/panel.tsx b/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/tracing/panel.tsx index 6183110742b..ed6f7b78195 100644 --- a/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/tracing/panel.tsx +++ b/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/tracing/panel.tsx @@ -6,6 +6,7 @@ import { cn } from '@langgenius/dify-ui/cn' import { StatusDot } from '@langgenius/dify-ui/status-dot' import { toast } from '@langgenius/dify-ui/toast' import { useBoolean } from 'ahooks' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useEffect, useState } from 'react' import { useTranslation } from 'react-i18next' @@ -24,7 +25,7 @@ import { WeaveIcon, } from '@/app/components/base/icons/src/public/tracing' import Loading from '@/app/components/base/loading' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { userProfileIdAtom, workspacePermissionKeysAtom } from '@/context/app-context-state' import { usePathname } from '@/next/navigation' import { fetchTracingConfig as doFetchTracingConfig, fetchTracingStatus, updateTracingStatus } from '@/service/apps' import { getAppACLCapabilities } from '@/utils/permission' @@ -39,8 +40,8 @@ const Panel: FC = () => { const pathname = usePathname() const matched = /\/app\/([^/]+)/.exec(pathname) const appId = (matched?.length && matched[1]) ? matched[1] : '' - const currentUserId = useAppContextWithSelector(state => state.userProfile?.id) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const appDetail = useAppStore(s => s.appDetail) const appACLCapabilities = React.useMemo(() => getAppACLCapabilities(appDetail?.permission_keys, { currentUserId, diff --git a/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/view.tsx b/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/view.tsx index 9929dab3713..cb5ba455422 100644 --- a/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/view.tsx +++ b/web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/view.tsx @@ -1,9 +1,10 @@ 'use client' +import { useAtomValue } from 'jotai' import * as React from 'react' import ApikeyInfoPanel from '@/app/components/app/overview/apikey-info-panel' import { useStore as useAppStore } from '@/app/components/app/store' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { userProfileIdAtom, workspacePermissionKeysAtom } from '@/context/app-context-state' import { getAppACLCapabilities } from '@/utils/permission' import ChartView from './chart-view' import TracingPanel from './tracing/panel' @@ -14,8 +15,8 @@ type OverviewViewProps = { const OverviewView = ({ appId }: OverviewViewProps) => { const appDetail = useAppStore(state => state.appDetail) - const currentUserId = useAppContextWithSelector(state => state.userProfile?.id) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const appACLCapabilities = React.useMemo(() => getAppACLCapabilities(appDetail?.permission_keys, { currentUserId, resourceMaintainer: appDetail?.maintainer, diff --git a/web/app/components/app/access-config/__tests__/index.spec.tsx b/web/app/components/app/access-config/__tests__/index.spec.tsx index 63909979344..9530a49c2d2 100644 --- a/web/app/components/app/access-config/__tests__/index.spec.tsx +++ b/web/app/components/app/access-config/__tests__/index.spec.tsx @@ -47,6 +47,18 @@ vi.mock('@/context/app-context', () => ({ useAppContext: () => mockAppContext, })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateAtomMock(importOriginal, () => mockAppContext) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/service/access-control/use-app-access-config', () => ({ useAppAccessRules: vi.fn(() => ({ data: { items: mockAppAccessRules.items }, diff --git a/web/app/components/app/access-config/index.tsx b/web/app/components/app/access-config/index.tsx index 6efc718012c..e99ee5adb4c 100644 --- a/web/app/components/app/access-config/index.tsx +++ b/web/app/components/app/access-config/index.tsx @@ -3,11 +3,12 @@ import type { ResourceOpenScope } from '@/models/access-control' import { ScrollArea } from '@langgenius/dify-ui/scroll-area' import { useSuspenseQuery } from '@tanstack/react-query' +import { useAtomValue } from 'jotai' import { useCallback, useMemo, useState } from 'react' import { useTranslation } from 'react-i18next' import AccessRulesEditor from '@/app/components/access-rules-editor' import { useStore } from '@/app/components/app/store' -import { useAppContext } from '@/context/app-context' +import { userProfileIdAtom, workspacePermissionKeysAtom } from '@/context/app-context-state' import { useLocale } from '@/context/i18n' import { systemFeaturesQueryOptions } from '@/features/system-features/client' import { getAccessControlTemplateLanguage } from '@/i18n-config/language' @@ -106,15 +107,16 @@ const AppAccessConfigContent = ({ appId, maintainerId }: AppAccessConfigContentP const AppAccessConfigPage = ({ appId }: AppAccessConfigPageProps) => { const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) - const { userProfile, workspacePermissionKeys } = useAppContext() + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const isRbacEnabled = systemFeatures.rbac_enabled const appDetail = useStore(state => state.appDetail) const appACLCapabilities = useMemo(() => getAppACLCapabilities(appDetail?.permission_keys, { - currentUserId: userProfile?.id, + currentUserId, resourceMaintainer: appDetail?.maintainer, workspacePermissionKeys, isRbacEnabled, - }), [appDetail?.maintainer, appDetail?.permission_keys, isRbacEnabled, userProfile?.id, workspacePermissionKeys]) + }), [appDetail?.maintainer, appDetail?.permission_keys, currentUserId, isRbacEnabled, workspacePermissionKeys]) if (!appDetail || appDetail.id !== appId || !appACLCapabilities.canAccessConfig) return null diff --git a/web/app/components/app/app-access-control/add-member-or-group-pop.tsx b/web/app/components/app/app-access-control/add-member-or-group-pop.tsx index 13ff5a520b4..0c3cc93fc8e 100644 --- a/web/app/components/app/app-access-control/add-member-or-group-pop.tsx +++ b/web/app/components/app/app-access-control/add-member-or-group-pop.tsx @@ -18,9 +18,10 @@ import { } from '@langgenius/dify-ui/combobox' import { RiArrowRightSLine, RiOrganizationChart } from '@remixicon/react' import { useDebounce } from 'ahooks' +import { useAtomValue } from 'jotai' import { useEffect, useRef, useState } from 'react' import { useTranslation } from 'react-i18next' -import { useSelector } from '@/context/app-context' +import { userProfileAtom } from '@/context/app-context-state' import { SubjectType } from '@/models/access-control' import { useSearchForWhiteListCandidates } from '@/service/access-control' import useAccessControlStore from '../../../../context/access-control-store' @@ -308,7 +309,7 @@ type MemberItemProps = { subject: Subject } function MemberItem({ member, subject }: MemberItemProps) { - const currentUser = useSelector(s => s.userProfile) + const currentUser = useAtomValue(userProfileAtom) const { t } = useTranslation() const specificMembers = useAccessControlStore(s => s.specificMembers) const isChecked = specificMembers.some(m => m.id === member.id) diff --git a/web/app/components/app/configuration/dataset-config/__tests__/index.spec.tsx b/web/app/components/app/configuration/dataset-config/__tests__/index.spec.tsx index f1942acb029..9611da5c814 100644 --- a/web/app/components/app/configuration/dataset-config/__tests__/index.spec.tsx +++ b/web/app/components/app/configuration/dataset-config/__tests__/index.spec.tsx @@ -46,6 +46,21 @@ vi.mock('@/context/app-context', () => ({ })), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateAtomMock(importOriginal, () => ({ + userProfile: { id: 'user-123' }, + workspacePermissionKeys: [], + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/utils/permission', () => ({ DatasetACLPermission: { Readonly: 'dataset.acl.readonly', diff --git a/web/app/components/app/configuration/dataset-config/index.tsx b/web/app/components/app/configuration/dataset-config/index.tsx index aa5624099aa..0009ea796a0 100644 --- a/web/app/components/app/configuration/dataset-config/index.tsx +++ b/web/app/components/app/configuration/dataset-config/index.tsx @@ -11,6 +11,7 @@ import type { DataSet } from '@/models/datasets' import { cn } from '@langgenius/dify-ui/cn' import { intersectionBy } from 'es-toolkit/compat' import { produce } from 'immer' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useCallback, useMemo } from 'react' import { useTranslation } from 'react-i18next' @@ -28,7 +29,7 @@ import { getMultipleRetrievalConfig, getSelectedDatasetsMode, } from '@/app/components/workflow/nodes/knowledge-retrieval/utils' -import { useSelector as useAppContextSelector } from '@/context/app-context' +import { userProfileIdAtom, workspacePermissionKeysAtom } from '@/context/app-context-state' import ConfigContext from '@/context/debug-configuration' import { AppModeEnum } from '@/types/app' import { getDatasetACLCapabilities } from '@/utils/permission' @@ -45,8 +46,8 @@ type Props = Readonly<{ }> const DatasetConfig: FC = ({ readonly, hideMetadataFilter }) => { const { t } = useTranslation() - const currentUserId = useAppContextSelector(s => s.userProfile?.id) - const workspacePermissionKeys = useAppContextSelector(s => s.workspacePermissionKeys) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const { mode, dataSets: dataSet, diff --git a/web/app/components/app/configuration/dataset-config/select-dataset/__tests__/index.spec.tsx b/web/app/components/app/configuration/dataset-config/select-dataset/__tests__/index.spec.tsx index 46dac5351ce..d0331de79c6 100644 --- a/web/app/components/app/configuration/dataset-config/select-dataset/__tests__/index.spec.tsx +++ b/web/app/components/app/configuration/dataset-config/select-dataset/__tests__/index.spec.tsx @@ -40,6 +40,20 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockWorkspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/hooks/use-knowledge', () => ({ useKnowledge: () => ({ formatIndexingTechniqueAndMethod: (tech: string, method: string) => `${tech}:${method}`, diff --git a/web/app/components/app/configuration/dataset-config/select-dataset/index.tsx b/web/app/components/app/configuration/dataset-config/select-dataset/index.tsx index 60c4d79a1a6..0a8844bfee1 100644 --- a/web/app/components/app/configuration/dataset-config/select-dataset/index.tsx +++ b/web/app/components/app/configuration/dataset-config/select-dataset/index.tsx @@ -5,6 +5,7 @@ import { Button } from '@langgenius/dify-ui/button' import { cn } from '@langgenius/dify-ui/cn' import { Dialog, DialogCloseButton, DialogContent, DialogTitle } from '@langgenius/dify-ui/dialog' import { useInfiniteScroll } from 'ahooks' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useCallback, useMemo, useRef, useState } from 'react' import { useTranslation } from 'react-i18next' @@ -13,7 +14,7 @@ import Badge from '@/app/components/base/badge' import Loading from '@/app/components/base/loading' import { ModelFeatureEnum } from '@/app/components/header/account-setting/model-provider-page/declarations' import FeatureIcon from '@/app/components/header/account-setting/model-provider-page/model-selector/feature-icon' -import { useSelector as useAppContextSelector } from '@/context/app-context' +import { workspacePermissionKeysAtom } from '@/context/app-context-state' import { useKnowledge } from '@/hooks/use-knowledge' import Link from '@/next/link' import { useInfiniteDatasets } from '@/service/knowledge/use-dataset' @@ -38,7 +39,7 @@ const SelectDataSet: FC = ({ const [selectedIdsInModal, setSelectedIdsInModal] = useState(() => selectedIds) const canSelectMulti = true const { formatIndexingTechniqueAndMethod } = useKnowledge() - const workspacePermissionKeys = useAppContextSelector(state => state.workspacePermissionKeys) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const canCreateDataset = hasPermission(workspacePermissionKeys, 'dataset.create_and_management') const { data, isLoading, isFetchingNextPage, fetchNextPage, hasNextPage } = useInfiniteDatasets( { page: 1 }, diff --git a/web/app/components/app/configuration/debug/debug-with-multiple-model/chat-item.tsx b/web/app/components/app/configuration/debug/debug-with-multiple-model/chat-item.tsx index 39f29972ea4..dc397d6dc18 100644 --- a/web/app/components/app/configuration/debug/debug-with-multiple-model/chat-item.tsx +++ b/web/app/components/app/configuration/debug/debug-with-multiple-model/chat-item.tsx @@ -3,6 +3,7 @@ import type { ModelAndParameter } from '../types' import type { InputForm } from '@/app/components/base/chat/chat/type' import type { ChatConfig, OnSend } from '@/app/components/base/chat/types' import { Avatar } from '@langgenius/dify-ui/avatar' +import { useAtomValue } from 'jotai' import { memo, useCallback, @@ -13,7 +14,7 @@ import { useChat } from '@/app/components/base/chat/chat/hooks' import { getLastAnswer } from '@/app/components/base/chat/utils' import { useFeatures } from '@/app/components/base/features/hooks' import { ModelFeatureEnum } from '@/app/components/header/account-setting/model-provider-page/declarations' -import { useAppContext } from '@/context/app-context' +import { userProfileAtom } from '@/context/app-context-state' import { useDebugConfigurationContext } from '@/context/debug-configuration' import { useEventEmitterContextContext } from '@/context/event-emitter' import { useProviderContext } from '@/context/provider-context' @@ -38,7 +39,7 @@ type ChatItemProps = { const ChatItem: FC = ({ modelAndParameter, }) => { - const { userProfile } = useAppContext() + const userProfile = useAtomValue(userProfileAtom) const { modelConfig, appId, diff --git a/web/app/components/app/configuration/debug/debug-with-single-model/index.tsx b/web/app/components/app/configuration/debug/debug-with-single-model/index.tsx index 0670a097366..4346b76f949 100644 --- a/web/app/components/app/configuration/debug/debug-with-single-model/index.tsx +++ b/web/app/components/app/configuration/debug/debug-with-single-model/index.tsx @@ -2,6 +2,7 @@ import type { InputForm } from '@/app/components/base/chat/chat/type' import type { ChatConfig, ChatItem, OnSend } from '@/app/components/base/chat/types' import type { FileEntity } from '@/app/components/base/file-uploader/types' import { Avatar } from '@langgenius/dify-ui/avatar' +import { useAtomValue } from 'jotai' import { memo, useCallback, useImperativeHandle, useMemo } from 'react' import { useStore as useAppStore } from '@/app/components/app/store' import Chat from '@/app/components/base/chat/chat' @@ -9,7 +10,7 @@ import { useChat } from '@/app/components/base/chat/chat/hooks' import { getLastAnswer, isValidGeneratedAnswer } from '@/app/components/base/chat/utils' import { useFeatures } from '@/app/components/base/features/hooks' import { ModelFeatureEnum } from '@/app/components/header/account-setting/model-provider-page/declarations' -import { useAppContext } from '@/context/app-context' +import { userProfileAtom } from '@/context/app-context-state' import { useDebugConfigurationContext } from '@/context/debug-configuration' import { useProviderContext } from '@/context/provider-context' import { @@ -37,7 +38,7 @@ const DebugWithSingleModel = ( ref: React.RefObject }, ) => { - const { userProfile } = useAppContext() + const userProfile = useAtomValue(userProfileAtom) const { readonly, canTestAndRun = false, diff --git a/web/app/components/app/configuration/hooks/__tests__/use-configuration.spec.tsx b/web/app/components/app/configuration/hooks/__tests__/use-configuration.spec.tsx index 86869d513e8..af04a9572ee 100644 --- a/web/app/components/app/configuration/hooks/__tests__/use-configuration.spec.tsx +++ b/web/app/components/app/configuration/hooks/__tests__/use-configuration.spec.tsx @@ -63,6 +63,23 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateAtomMock(importOriginal, () => ({ + currentWorkspace: { id: 'workspace-1' }, + isLoadingCurrentWorkspace: false, + userProfile: { id: 'user-1' }, + workspacePermissionKeys: ['app.create_and_management'], + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/modal-context', () => ({ useModalContext: () => ({ setShowAccountSettingModal: mockSetShowAccountSettingModal, diff --git a/web/app/components/app/configuration/hooks/use-configuration.ts b/web/app/components/app/configuration/hooks/use-configuration.ts index f92b0b5c58d..abb589be9b4 100644 --- a/web/app/components/app/configuration/hooks/use-configuration.ts +++ b/web/app/components/app/configuration/hooks/use-configuration.ts @@ -25,6 +25,7 @@ import type { VisionSettings } from '@/types/app' import { useBoolean, useGetState } from 'ahooks' import { clone } from 'es-toolkit/object' import { produce } from 'immer' +import { useAtomValue } from 'jotai' import { useCallback, useEffect, useMemo, useRef, useState } from 'react' import { useTranslation } from 'react-i18next' import { useShallow } from 'zustand/react/shallow' @@ -43,7 +44,12 @@ import { } from '@/app/components/header/account-setting/model-provider-page/hooks' import { useIntegrationsSetting } from '@/app/components/header/account-setting/use-integrations-setting' import { ANNOTATION_DEFAULT, DATASET_DEFAULT, DEFAULT_AGENT_SETTING, DEFAULT_CHAT_PROMPT_CONFIG, DEFAULT_COMPLETION_PROMPT_CONFIG } from '@/config' -import { useAppContext } from '@/context/app-context' +import { + currentWorkspaceAtom, + currentWorkspaceLoadingAtom, + userProfileIdAtom, + workspacePermissionKeysAtom, +} from '@/context/app-context-state' import { useProviderContext } from '@/context/provider-context' import useBreakpoints, { MediaType } from '@/hooks/use-breakpoints' import { PromptMode } from '@/models/debug' @@ -110,7 +116,10 @@ export type ConfigurationViewModel = { export const useConfiguration = (): ConfigurationViewModel => { const { t } = useTranslation() - const { isLoadingCurrentWorkspace, currentWorkspace, userProfile, workspacePermissionKeys } = useAppContext() + const isLoadingCurrentWorkspace = useAtomValue(currentWorkspaceLoadingAtom) + const currentWorkspace = useAtomValue(currentWorkspaceAtom) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const openIntegrationsSetting = useIntegrationsSetting() const { appDetail, showAppConfigureFeaturesModal, setShowAppConfigureFeaturesModal } = useAppStore(useShallow(state => ({ @@ -123,10 +132,10 @@ export const useConfiguration = (): ConfigurationViewModel => { const { data: fileUploadConfigResponse } = useFileUploadConfig() const latestPublishedAt = useMemo(() => appDetail?.model_config?.updated_at, [appDetail]) const appACLCapabilities = useMemo(() => getAppACLCapabilities(appDetail?.permission_keys, { - currentUserId: userProfile?.id, + currentUserId, resourceMaintainer: appDetail?.maintainer, workspacePermissionKeys, - }), [appDetail?.maintainer, appDetail?.permission_keys, userProfile?.id, workspacePermissionKeys]) + }), [appDetail?.maintainer, appDetail?.permission_keys, currentUserId, workspacePermissionKeys]) const configurationReadonly = !appACLCapabilities.canEdit const [formattingChanged, setFormattingChanged] = useState(false) const [hasFetchedDetail, setHasFetchedDetail] = useState(false) diff --git a/web/app/components/app/create-app-dialog/app-list/__tests__/index.spec.tsx b/web/app/components/app/create-app-dialog/app-list/__tests__/index.spec.tsx index 14c4cc37312..711f6f3b37f 100644 --- a/web/app/components/app/create-app-dialog/app-list/__tests__/index.spec.tsx +++ b/web/app/components/app/create-app-dialog/app-list/__tests__/index.spec.tsx @@ -36,6 +36,21 @@ vi.mock('@/context/app-context', () => ({ workspacePermissionKeys: mockWorkspacePermissionKeys, }), })) + +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateAtomMock(importOriginal, () => ({ + userProfile: mockUserProfile, + workspacePermissionKeys: mockWorkspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateJotaiMock(importOriginal) +}) vi.mock('nuqs', () => ({ useQueryState: () => ['Recommended', vi.fn()], })) diff --git a/web/app/components/app/create-app-dialog/app-list/index.tsx b/web/app/components/app/create-app-dialog/app-list/index.tsx index 8efa3427f4a..5d7c664ab17 100644 --- a/web/app/components/app/create-app-dialog/app-list/index.tsx +++ b/web/app/components/app/create-app-dialog/app-list/index.tsx @@ -7,6 +7,7 @@ import { toast } from '@langgenius/dify-ui/toast' import { RiRobot2Line } from '@remixicon/react' import { useSuspenseQuery } from '@tanstack/react-query' import { useDebounceFn } from 'ahooks' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useMemo, useState } from 'react' import { useTranslation } from 'react-i18next' @@ -17,7 +18,7 @@ import Input from '@/app/components/base/input' import Loading from '@/app/components/base/loading' import CreateAppModal from '@/app/components/explore/create-app-modal' import { usePluginDependencies } from '@/app/components/workflow/plugin-dependency/hooks' -import { useAppContext } from '@/context/app-context' +import { userProfileIdAtom, workspacePermissionKeysAtom } from '@/context/app-context-state' import { systemFeaturesQueryOptions } from '@/features/system-features/client' import { DSLImportMode } from '@/models/app' import { useRouter } from '@/next/navigation' @@ -48,7 +49,8 @@ const Apps = ({ }: AppsProps) => { const { t } = useTranslation() const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) - const { userProfile, workspacePermissionKeys } = useAppContext() + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const isRbacEnabled = systemFeatures.rbac_enabled const canCreateAppFromTemplate = hasPermission(workspacePermissionKeys, 'app.create_and_management') const { push } = useRouter() @@ -165,8 +167,8 @@ const Apps = ({ invalidateAppList() if (app.app_id) { getRedirection({ id: app.app_id, mode: app.app_mode, permission_keys: app.permission_keys }, push, { - currentUserId: userProfile?.id, - resourceMaintainer: userProfile?.id, + currentUserId, + resourceMaintainer: currentUserId, workspacePermissionKeys, isRbacEnabled, }) diff --git a/web/app/components/app/create-app-modal/__tests__/index.spec.tsx b/web/app/components/app/create-app-modal/__tests__/index.spec.tsx index d583dadedda..3ab83834a5b 100644 --- a/web/app/components/app/create-app-modal/__tests__/index.spec.tsx +++ b/web/app/components/app/create-app-modal/__tests__/index.spec.tsx @@ -17,6 +17,10 @@ const ahooksMocks = vi.hoisted(() => ({ keyPressHandlers: [] as Array<() => void>, })) const mockInvalidateAppList = vi.hoisted(() => vi.fn()) +const mockAppContextState = vi.hoisted(() => ({ + userProfile: { id: 'user-1' }, + workspacePermissionKeys: ['app.create_and_management'] as string[], +})) vi.mock('ahooks', () => ({ useDebounceFn: unknown>(fn: T) => { @@ -73,6 +77,16 @@ vi.mock('@/context/provider-context', () => ({ vi.mock('@/context/app-context', () => ({ useAppContext: vi.fn(), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateAtomMock(importOriginal, () => mockAppContextState) +}) +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateJotaiMock(importOriginal) +}) vi.mock('@/context/i18n', () => ({ useDocLink: () => () => '/guides', })) @@ -135,6 +149,8 @@ describe('CreateAppModal', () => { userProfile: { id: 'user-1' }, workspacePermissionKeys: ['app.create_and_management'], } as unknown as ReturnType) + mockAppContextState.userProfile = { id: 'user-1' } + mockAppContextState.workspacePermissionKeys = ['app.create_and_management'] mockSetItem.mockClear() Object.defineProperty(window, 'localStorage', { value: { diff --git a/web/app/components/app/create-app-modal/index.tsx b/web/app/components/app/create-app-modal/index.tsx index a53f07d5b12..011e5292167 100644 --- a/web/app/components/app/create-app-modal/index.tsx +++ b/web/app/components/app/create-app-modal/index.tsx @@ -11,6 +11,7 @@ import { RiArrowRightLine, RiArrowRightSLine, RiExchange2Fill } from '@remixicon import { formatForDisplay, useHotkey } from '@tanstack/react-hotkeys' import { useSuspenseQuery } from '@tanstack/react-query' import { useDebounceFn } from 'ahooks' +import { useAtomValue } from 'jotai' import { useCallback, useRef, useState } from 'react' import { useTranslation } from 'react-i18next' import { useSetNeedRefreshAppList } from '@/app/components/apps/storage' @@ -19,7 +20,7 @@ import Divider from '@/app/components/base/divider' import { BubbleTextMod, ChatBot, ListSparkle, Logic } from '@/app/components/base/icons/src/vender/solid/communication' import Input from '@/app/components/base/input' import AppsFull from '@/app/components/billing/apps-full-in-dialog' -import { useAppContext } from '@/context/app-context' +import { userProfileIdAtom, workspacePermissionKeysAtom } from '@/context/app-context-state' import { useProviderContext } from '@/context/provider-context' import { systemFeaturesQueryOptions } from '@/features/system-features/client' import useTheme from '@/hooks/use-theme' @@ -59,7 +60,8 @@ function CreateApp({ onClose, onSuccess, onCreateFromTemplate, defaultAppMode }: const { plan, enableBilling } = useProviderContext() const isAppsFull = (enableBilling && plan.usage.buildApps >= plan.total.buildApps) const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) - const { userProfile, workspacePermissionKeys } = useAppContext() + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const isRbacEnabled = systemFeatures.rbac_enabled const canCreateApp = hasPermission(workspacePermissionKeys, 'app.create_and_management') const invalidateAppList = useInvalidateAppList() @@ -101,7 +103,7 @@ function CreateApp({ onClose, onSuccess, onCreateFromTemplate, defaultAppMode }: setNeedRefresh('1') invalidateAppList() getRedirection(app, push, { - currentUserId: userProfile?.id, + currentUserId, resourceMaintainer: app.maintainer, workspacePermissionKeys, isRbacEnabled, @@ -111,7 +113,7 @@ function CreateApp({ onClose, onSuccess, onCreateFromTemplate, defaultAppMode }: toast.error(error instanceof Error ? error.message : t('newApp.appCreateFailed', { ns: 'app' })) } isCreatingRef.current = false - }, [canCreateApp, name, t, appMode, appIcon, description, onSuccess, onClose, push, userProfile?.id, workspacePermissionKeys, isRbacEnabled, setNeedRefresh, invalidateAppList]) + }, [canCreateApp, currentUserId, name, t, appMode, appIcon, description, onSuccess, onClose, push, workspacePermissionKeys, isRbacEnabled, setNeedRefresh, invalidateAppList]) const { run: handleCreateApp } = useDebounceFn(onCreate, { wait: 300 }) useHotkey('Mod+Enter', () => { diff --git a/web/app/components/app/create-from-dsl-modal/__tests__/index.spec.tsx b/web/app/components/app/create-from-dsl-modal/__tests__/index.spec.tsx index 008225e9963..a6932f6daaf 100644 --- a/web/app/components/app/create-from-dsl-modal/__tests__/index.spec.tsx +++ b/web/app/components/app/create-from-dsl-modal/__tests__/index.spec.tsx @@ -93,6 +93,21 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateAtomMock(importOriginal, () => ({ + userProfile: mockUserProfile, + workspacePermissionKeys: mockWorkspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/provider-context', () => ({ useProviderContext: () => ({ plan: { diff --git a/web/app/components/app/create-from-dsl-modal/index.tsx b/web/app/components/app/create-from-dsl-modal/index.tsx index 191e1727120..d76b1511302 100644 --- a/web/app/components/app/create-from-dsl-modal/index.tsx +++ b/web/app/components/app/create-from-dsl-modal/index.tsx @@ -9,13 +9,14 @@ import { toast } from '@langgenius/dify-ui/toast' import { formatForDisplay, useHotkey } from '@tanstack/react-hotkeys' import { useSuspenseQuery } from '@tanstack/react-query' import { useDebounceFn } from 'ahooks' +import { useAtomValue } from 'jotai' import { useCallback, useEffect, useMemo, useRef, useState } from 'react' import { useTranslation } from 'react-i18next' import { useSetNeedRefreshAppList } from '@/app/components/apps/storage' import Input from '@/app/components/base/input' import AppsFull from '@/app/components/billing/apps-full-in-dialog' import { usePluginDependencies } from '@/app/components/workflow/plugin-dependency/hooks' -import { useAppContext } from '@/context/app-context' +import { userProfileIdAtom, workspacePermissionKeysAtom } from '@/context/app-context-state' import { useProviderContext } from '@/context/provider-context' import { systemFeaturesQueryOptions } from '@/features/system-features/client' import { @@ -60,7 +61,8 @@ const CreateFromDSLModal = ({ show, onSuccess, onClose, activeTab = CreateFromDS const setNeedRefresh = useSetNeedRefreshAppList() const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) const isRbacEnabled = systemFeatures.rbac_enabled - const { userProfile, workspacePermissionKeys } = useAppContext() + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const readFile = useCallback((file: File) => { const reader = new FileReader() @@ -136,8 +138,8 @@ const CreateFromDSLModal = ({ show, onSuccess, onClose, activeTab = CreateFromDS if (app_id) { await handleCheckPluginDependencies(app_id) getRedirection({ id: app_id, mode: app_mode, permission_keys }, push, { - currentUserId: userProfile?.id, - resourceMaintainer: userProfile?.id, + currentUserId, + resourceMaintainer: currentUserId, workspacePermissionKeys, isRbacEnabled, }) @@ -196,8 +198,8 @@ const CreateFromDSLModal = ({ show, onSuccess, onClose, activeTab = CreateFromDS invalidateAppList() if (app_id) { getRedirection({ id: app_id, mode: app_mode, permission_keys }, push, { - currentUserId: userProfile?.id, - resourceMaintainer: userProfile?.id, + currentUserId, + resourceMaintainer: currentUserId, workspacePermissionKeys, isRbacEnabled, }) diff --git a/web/app/components/app/log/empty-element.tsx b/web/app/components/app/log/empty-element.tsx index 70813290319..a810f53e8e3 100644 --- a/web/app/components/app/log/empty-element.tsx +++ b/web/app/components/app/log/empty-element.tsx @@ -2,9 +2,10 @@ import type { FC, SVGProps } from 'react' import type { App } from '@/types/app' import { useSuspenseQuery } from '@tanstack/react-query' +import { useAtomValue } from 'jotai' import * as React from 'react' import { Trans, useTranslation } from 'react-i18next' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { userProfileIdAtom, workspacePermissionKeysAtom } from '@/context/app-context-state' import { systemFeaturesQueryOptions } from '@/features/system-features/client' import Link from '@/next/link' import { AppModeEnum } from '@/types/app' @@ -21,8 +22,8 @@ const ThreeDotsIcon = ({ className }: SVGProps) => { const EmptyElement: FC<{ appDetail: App }> = ({ appDetail }) => { const { t } = useTranslation() - const currentUserId = useAppContextWithSelector(state => state.userProfile?.id) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) const isRbacEnabled = systemFeatures.rbac_enabled diff --git a/web/app/components/app/overview/app-card.tsx b/web/app/components/app/overview/app-card.tsx index 2994689af77..246b0c7a07b 100644 --- a/web/app/components/app/overview/app-card.tsx +++ b/web/app/components/app/overview/app-card.tsx @@ -7,13 +7,14 @@ import { Popover, PopoverContent, PopoverTrigger } from '@langgenius/dify-ui/pop import { StatusDot } from '@langgenius/dify-ui/status-dot' import { Switch } from '@langgenius/dify-ui/switch' import { useQueryClient, useSuspenseQuery } from '@tanstack/react-query' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useCallback, useMemo, useState } from 'react' import { useTranslation } from 'react-i18next' import AppBasic from '@/app/components/app-sidebar/basic' import { useStore as useAppStore } from '@/app/components/app/store' import SecretKeyButton from '@/app/components/develop/secret-key/secret-key-button' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { userProfileIdAtom, workspacePermissionKeysAtom } from '@/context/app-context-state' import { useDocLink } from '@/context/i18n' import { systemFeaturesQueryOptions } from '@/features/system-features/client' import { AccessMode } from '@/models/access-control' @@ -71,8 +72,8 @@ function AppCard({ const router = useRouter() const pathname = usePathname() const queryClient = useQueryClient() - const currentUserId = useAppContextWithSelector(state => state.userProfile?.id) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const appACLCapabilities = useMemo(() => getAppACLCapabilities(appInfo.permission_keys, { currentUserId, resourceMaintainer: appInfo.maintainer, diff --git a/web/app/components/app/overview/embedded/index.tsx b/web/app/components/app/overview/embedded/index.tsx index fd36938af79..f93901304b4 100644 --- a/web/app/components/app/overview/embedded/index.tsx +++ b/web/app/components/app/overview/embedded/index.tsx @@ -5,12 +5,13 @@ import { cn } from '@langgenius/dify-ui/cn' import { Dialog, DialogCloseButton, DialogContent, DialogTitle } from '@langgenius/dify-ui/dialog' import { Tooltip, TooltipContent, TooltipTrigger } from '@langgenius/dify-ui/tooltip' import copy from 'copy-to-clipboard' +import { useAtomValue } from 'jotai' import { Suspense, use, useEffect, useMemo, useRef, useState } from 'react' import { useTranslation } from 'react-i18next' import ActionButton from '@/app/components/base/action-button' import { useThemeContext } from '@/app/components/base/chat/embedded-chatbot/theme/theme-context' import { InputVarType } from '@/app/components/workflow/types' -import { useAppContext } from '@/context/app-context' +import { langGeniusVersionInfoAtom } from '@/context/app-context-state' import { basePath } from '@/utils/var' import { compressAndEncodeBase64, @@ -129,7 +130,7 @@ const EmbeddedContent = ({ ) const latestResolvedIframeUrlRef = useRef('') - const { langGeniusVersionInfo } = useAppContext() + const langGeniusVersionInfo = useAtomValue(langGeniusVersionInfoAtom) const themeBuilder = useThemeContext() const isTestEnv = langGeniusVersionInfo.current_env === 'TESTING' || langGeniusVersionInfo.current_env === 'DEVELOPMENT' diff --git a/web/app/components/app/overview/trigger-card.tsx b/web/app/components/app/overview/trigger-card.tsx index 5928f9140e7..a7112dfc3dd 100644 --- a/web/app/components/app/overview/trigger-card.tsx +++ b/web/app/components/app/overview/trigger-card.tsx @@ -6,12 +6,13 @@ import type { AppSSO } from '@/types/app' import type { I18nKeysByPrefix } from '@/types/i18n' import { StatusDot } from '@langgenius/dify-ui/status-dot' import { Switch } from '@langgenius/dify-ui/switch' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useTranslation } from 'react-i18next' import BlockIcon from '@/app/components/workflow/block-icon' import { useTriggerStatusStore } from '@/app/components/workflow/store/trigger-status' import { BlockEnum } from '@/app/components/workflow/types' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { userProfileIdAtom, workspacePermissionKeysAtom } from '@/context/app-context-state' import { useDocLink } from '@/context/i18n' import Link from '@/next/link' import { @@ -79,8 +80,8 @@ function TriggerCard({ appInfo, onToggleResult }: ITriggerCardProps) { const { t } = useTranslation() const docLink = useDocLink() const appId = appInfo.id - const currentUserId = useAppContextWithSelector(state => state.userProfile?.id) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const canEditApp = React.useMemo(() => getAppACLCapabilities(appInfo.permission_keys, { currentUserId, resourceMaintainer: appInfo.maintainer, diff --git a/web/app/components/apps/__tests__/app-card.spec.tsx b/web/app/components/apps/__tests__/app-card.spec.tsx index 3be6fa998d9..0c333204702 100644 --- a/web/app/components/apps/__tests__/app-card.spec.tsx +++ b/web/app/components/apps/__tests__/app-card.spec.tsx @@ -80,6 +80,18 @@ vi.mock('@/context/app-context', () => ({ useSelector: (selector: (state: typeof mockAppContext) => unknown) => selector(mockAppContext), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateAtomMock(importOriginal, () => mockAppContext) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateJotaiMock(importOriginal) +}) + // Mock provider context const mockOnPlanInfoChanged = vi.fn() vi.mock('@/context/provider-context', () => ({ diff --git a/web/app/components/apps/__tests__/creators-filter.spec.tsx b/web/app/components/apps/__tests__/creators-filter.spec.tsx index bf891e63dbf..3a95174f0a3 100644 --- a/web/app/components/apps/__tests__/creators-filter.spec.tsx +++ b/web/app/components/apps/__tests__/creators-filter.spec.tsx @@ -9,6 +9,20 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateAtomMock(importOriginal, () => ({ + userProfile: { id: 'member-2' }, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/service/use-common', () => ({ useMembers: () => ({ data: { diff --git a/web/app/components/apps/__tests__/index.spec.tsx b/web/app/components/apps/__tests__/index.spec.tsx index 7f6120afef0..8ea477bd7cb 100644 --- a/web/app/components/apps/__tests__/index.spec.tsx +++ b/web/app/components/apps/__tests__/index.spec.tsx @@ -78,6 +78,20 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockWorkspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/hooks/use-import-dsl', () => ({ useImportDSL: () => ({ handleImportDSL: mockHandleImportDSL, diff --git a/web/app/components/apps/__tests__/list.spec.tsx b/web/app/components/apps/__tests__/list.spec.tsx index 6fc9e47258a..47d40e1d9a8 100644 --- a/web/app/components/apps/__tests__/list.spec.tsx +++ b/web/app/components/apps/__tests__/list.spec.tsx @@ -87,6 +87,21 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateAtomMock(importOriginal, () => ({ + userProfile: { id: 'creator-1' }, + workspacePermissionKeys: mockWorkspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateJotaiMock(importOriginal) +}) + const mockOnPlanInfoChanged = vi.fn() vi.mock('@/context/provider-context', () => ({ useProviderContext: () => ({ diff --git a/web/app/components/apps/app-card.tsx b/web/app/components/apps/app-card.tsx index 39564412da4..46d908abc32 100644 --- a/web/app/components/apps/app-card.tsx +++ b/web/app/components/apps/app-card.tsx @@ -31,6 +31,7 @@ import { TooltipTrigger, } from '@langgenius/dify-ui/tooltip' import { useSuspenseQuery } from '@tanstack/react-query' +import { useAtomValue } from 'jotai' import { useCallback, useId, useMemo, useState } from 'react' import { Trans, useTranslation } from 'react-i18next' import { AppTypeIcon } from '@/app/components/app/type-selector' @@ -39,7 +40,7 @@ import AppIcon from '@/app/components/base/app-icon' import StarIcon from '@/app/components/base/icons/src/vender/Star' import { UserAvatarList } from '@/app/components/base/user-avatar-list' import { buildInstalledAppPath } from '@/app/components/explore/installed-app/routes' -import { useSelector as useAppContextSelector } from '@/context/app-context' +import { userProfileIdAtom, workspacePermissionKeysAtom } from '@/context/app-context-state' import { useProviderContext } from '@/context/provider-context' import { systemFeaturesQueryOptions } from '@/features/system-features/client' import { AppCardTags } from '@/features/tag-management/components/app-card-tags' @@ -311,8 +312,8 @@ type AppCardActionBarProps = { export function AppCardActionBar({ app, onRefresh }: AppCardActionBarProps) { const { t } = useTranslation() const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) - const currentUserId = useAppContextSelector(state => state.userProfile?.id) - const workspacePermissionKeys = useAppContextSelector(state => state.workspacePermissionKeys) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const isRbacEnabled = systemFeatures.rbac_enabled const { onPlanInfoChanged } = useProviderContext() const { push } = useRouter() @@ -744,8 +745,8 @@ export function AppCardActionBar({ app, onRefresh }: AppCardActionBarProps) { export function AppCard({ app, onlineUsers = [], onRefresh, onOpenTagManagement = () => {} }: AppCardProps) { const { t } = useTranslation() const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) - const currentUserId = useAppContextSelector(state => state.userProfile?.id) - const workspacePermissionKeys = useAppContextSelector(state => state.workspacePermissionKeys) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const isRbacEnabled = systemFeatures.rbac_enabled const { onPlanInfoChanged } = useProviderContext() const { push } = useRouter() diff --git a/web/app/components/apps/creators-filter.tsx b/web/app/components/apps/creators-filter.tsx index 6a20b9348eb..f78d335d631 100644 --- a/web/app/components/apps/creators-filter.tsx +++ b/web/app/components/apps/creators-filter.tsx @@ -9,9 +9,10 @@ import { DropdownMenuTrigger, } from '@langgenius/dify-ui/dropdown-menu' import { Input } from '@langgenius/dify-ui/input' +import { useAtomValue } from 'jotai' import { useCallback, useMemo, useState } from 'react' import { useTranslation } from 'react-i18next' -import { useAppContext } from '@/context/app-context' +import { userProfileIdAtom } from '@/context/app-context-state' import { useMembers } from '@/service/use-common' type CreatorsFilterProps = { @@ -33,12 +34,11 @@ const CreatorsFilter = ({ onChange, }: CreatorsFilterProps) => { const { t } = useTranslation() - const { userProfile } = useAppContext() + const currentUserId = useAtomValue(userProfileIdAtom) const { data: membersData } = useMembers() const [keywords, setKeywords] = useState('') const creatorOptions = useMemo(() => { - const currentUserId = userProfile?.id const members = membersData?.accounts ?? [] return [...members] @@ -56,7 +56,7 @@ const CreatorsFilter = ({ avatarUrl: member.avatar_url, isYou: member.id === currentUserId, })) - }, [membersData?.accounts, userProfile?.id]) + }, [currentUserId, membersData?.accounts]) const filteredCreators = useMemo(() => { const normalizedKeywords = keywords.trim().toLowerCase() diff --git a/web/app/components/apps/index.tsx b/web/app/components/apps/index.tsx index 66a9cd6b67f..400f11bf9a7 100644 --- a/web/app/components/apps/index.tsx +++ b/web/app/components/apps/index.tsx @@ -2,10 +2,11 @@ import type { CreateAppModalProps } from '../explore/create-app-modal' import type { TryAppSelection } from '@/types/try-app' import type { TrackCreateAppParams } from '@/utils/create-app-tracking' +import { useAtomValue } from 'jotai' import { useCallback, useRef, useState } from 'react' import { useTranslation } from 'react-i18next' import { useEducationInit } from '@/app/education-apply/hooks' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { workspacePermissionKeysAtom } from '@/context/app-context-state' import AppListContext from '@/context/app-list-context' import useDocumentTitle from '@/hooks/use-document-title' import { useImportDSL } from '@/hooks/use-import-dsl' @@ -26,7 +27,7 @@ const Apps = () => { const { t } = useTranslation() const searchParams = useSearchParams() const { replace } = useRouter() - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const canCreateApp = hasPermission(workspacePermissionKeys, 'app.create_and_management') const templateId = searchParams.get('template-id') const templateDismissedRef = useRef(false) diff --git a/web/app/components/apps/list.tsx b/web/app/components/apps/list.tsx index a1cf1a39c15..7b56e49611d 100644 --- a/web/app/components/apps/list.tsx +++ b/web/app/components/apps/list.tsx @@ -4,10 +4,11 @@ import type { GetAppsData } from '@dify/contracts/api/console/apps/types.gen' import { cn } from '@langgenius/dify-ui/cn' import { keepPreviousData, useInfiniteQuery, useQuery, useSuspenseQuery } from '@tanstack/react-query' import { useDebounce } from 'ahooks' +import { useAtomValue } from 'jotai' import { useCallback, useEffect, useMemo, useRef, useState } from 'react' import { useTranslation } from 'react-i18next' import { useNeedRefreshAppList } from '@/app/components/apps/storage' -import { useAppContext } from '@/context/app-context' +import { workspacePermissionKeysAtom } from '@/context/app-context-state' import { useProviderContext } from '@/context/provider-context' import { systemFeaturesQueryOptions } from '@/features/system-features/client' import { CheckModal } from '@/hooks/use-pay' @@ -43,7 +44,7 @@ function List({ }: Props) { const { t } = useTranslation() const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) - const { workspacePermissionKeys } = useAppContext() + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const { onPlanInfoChanged } = useProviderContext() // eslint-disable-next-line react/use-state -- custom URL query hook, not React.useState diff --git a/web/app/components/apps/starred-app-card.tsx b/web/app/components/apps/starred-app-card.tsx index 7865e78e13a..31266a8e57b 100644 --- a/web/app/components/apps/starred-app-card.tsx +++ b/web/app/components/apps/starred-app-card.tsx @@ -5,11 +5,12 @@ import type { App } from '@/types/app' import { cn } from '@langgenius/dify-ui/cn' import { toast } from '@langgenius/dify-ui/toast' import { useSuspenseQuery } from '@tanstack/react-query' +import { useAtomValue } from 'jotai' import { useCallback, useMemo } from 'react' import { useTranslation } from 'react-i18next' import { AppTypeIcon } from '@/app/components/app/type-selector' import AppIcon from '@/app/components/base/app-icon' -import { useSelector as useAppContextSelector } from '@/context/app-context' +import { userProfileIdAtom, workspacePermissionKeysAtom } from '@/context/app-context-state' import { systemFeaturesQueryOptions } from '@/features/system-features/client' import Link from '@/next/link' import { getRedirectionPath } from '@/utils/app-redirection' @@ -24,8 +25,8 @@ type StarredAppCardProps = { export function StarredAppCard({ app, onRefresh }: StarredAppCardProps) { const { t } = useTranslation() - const currentUserId = useAppContextSelector(state => state.userProfile?.id) - const workspacePermissionKeys = useAppContextSelector(state => state.workspacePermissionKeys) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) const isRbacEnabled = systemFeatures.rbac_enabled const isPreviewOnly = hasOnlyAppPreviewPermission(app.permission_keys) diff --git a/web/app/components/snippet-list/__tests__/index.spec.tsx b/web/app/components/snippet-list/__tests__/index.spec.tsx index d5bdbedd06b..a73a8f3965f 100644 --- a/web/app/components/snippet-list/__tests__/index.spec.tsx +++ b/web/app/components/snippet-list/__tests__/index.spec.tsx @@ -98,6 +98,23 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateAtomMock(importOriginal, () => ({ + userProfile: { id: 'creator-1' }, + currentWorkspace: { id: 'workspace-1' }, + isLoadingCurrentWorkspace: false, + workspacePermissionKeys: mockWorkspacePermissionKeys(), + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/service/use-common', () => ({ useMembers: () => ({ data: { diff --git a/web/features/agent-v2/agent-detail/access/components/__tests__/access-surface-cards.spec.tsx b/web/features/agent-v2/agent-detail/access/components/__tests__/access-surface-cards.spec.tsx index 0401954579d..fabf9a736ea 100644 --- a/web/features/agent-v2/agent-detail/access/components/__tests__/access-surface-cards.spec.tsx +++ b/web/features/agent-v2/agent-detail/access/components/__tests__/access-surface-cards.spec.tsx @@ -51,6 +51,31 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateAtomMock(importOriginal, () => ({ + userProfile: { id: 'user-1' }, + currentWorkspace: { id: 'workspace-1' }, + workspacePermissionKeys: ['app.acl.edit'], + langGeniusVersionInfo: { + current_env: 'PRODUCTION', + current_version: '', + latest_version: '', + version: '', + release_date: '', + release_notes: '', + can_auto_update: false, + }, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/service/client', () => ({ consoleQuery: { apps: { From 6ac7ff658636b3ed5f429294a8733b3da679208e Mon Sep 17 00:00:00 2001 From: yyh <92089059+lyzno1@users.noreply.github.com> Date: Wed, 8 Jul 2026 12:33:10 +0800 Subject: [PATCH 39/70] chore(deps): upgrade vite-plus toolchain (#38534) --- pnpm-lock.yaml | 1275 +++++++++++++++++++------------------------ pnpm-workspace.yaml | 6 +- 2 files changed, 565 insertions(+), 716 deletions(-) diff --git a/pnpm-lock.yaml b/pnpm-lock.yaml index d9ac763b741..d87be83e16a 100644 --- a/pnpm-lock.yaml +++ b/pnpm-lock.yaml @@ -598,8 +598,8 @@ catalogs: specifier: 12.0.0-beta.3 version: 12.0.0-beta.3 vite-plus: - specifier: 0.2.1 - version: 0.2.1 + specifier: 0.2.3 + version: 0.2.3 vitest-browser-react: specifier: 2.2.0 version: 2.2.0 @@ -635,7 +635,7 @@ overrides: solid-js: 1.9.13 string-width: ~8.2.1 tar@<=7.5.15: ^7.5.16 - vite: npm:@voidzero-dev/vite-plus-core@0.2.1 + vite: npm:@voidzero-dev/vite-plus-core@0.2.3 vitest: 4.1.9 ws@>=8.0.0 <8.20.1: ^8.21.0 yaml@>=2.0.0 <2.8.3: 2.9.0 @@ -647,7 +647,7 @@ importers: devDependencies: '@antfu/eslint-config': specifier: 'catalog:' - version: 9.1.0(@eslint-react/eslint-plugin@5.9.5(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2))(typescript@6.0.3))(@next/eslint-plugin-next@16.2.9)(@typescript-eslint/typescript-estree@8.62.0(typescript@6.0.3))(@typescript-eslint/utils@8.62.0(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2))(typescript@6.0.3))(eslint-plugin-jsx-a11y@6.10.2(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2)))(eslint-plugin-react-refresh@0.5.3(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2)))(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2))(oxlint@1.70.0(oxlint-tsgolint@0.23.0)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(supports-color@10.2.2)(typescript@6.0.3)(vitest@4.1.9) + version: 9.1.0(@eslint-react/eslint-plugin@5.9.5(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2))(typescript@6.0.3))(@next/eslint-plugin-next@16.2.9)(@typescript-eslint/typescript-estree@8.62.0(typescript@6.0.3))(@typescript-eslint/utils@8.62.0(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2))(typescript@6.0.3))(eslint-plugin-jsx-a11y@6.10.2(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2)))(eslint-plugin-react-refresh@0.5.3(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2)))(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2))(oxlint@1.72.0(oxlint-tsgolint@0.24.0)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(supports-color@10.2.2)(typescript@6.0.3)(vitest@4.1.9) concurrently: specifier: 'catalog:' version: 10.0.3 @@ -667,11 +667,11 @@ importers: specifier: runtime:^22.22.1 version: runtime:22.23.1 vite: - specifier: npm:@voidzero-dev/vite-plus-core@0.2.1 - version: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + specifier: npm:@voidzero-dev/vite-plus-core@0.2.3 + version: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' vite-plus: specifier: 'catalog:' - version: 0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) + version: 0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) cli: dependencies: @@ -752,14 +752,14 @@ importers: specifier: 'catalog:' version: 6.0.3 vite: - specifier: npm:@voidzero-dev/vite-plus-core@0.2.1 - version: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + specifier: npm:@voidzero-dev/vite-plus-core@0.2.3 + version: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' vite-plus: specifier: 'catalog:' - version: 0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) + version: 0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) vitest: specifier: 4.1.9 - version: 4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) + version: 4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) e2e: devDependencies: @@ -791,11 +791,11 @@ importers: specifier: 'catalog:' version: 6.0.3 vite: - specifier: npm:@voidzero-dev/vite-plus-core@0.2.1 - version: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + specifier: npm:@voidzero-dev/vite-plus-core@0.2.3 + version: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' vite-plus: specifier: 'catalog:' - version: 0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) + version: 0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) zod: specifier: 'catalog:' version: 4.4.3 @@ -835,7 +835,7 @@ importers: version: 6.0.3 vite-plus: specifier: 'catalog:' - version: 0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) + version: 0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) packages/dev-proxy: dependencies: @@ -862,14 +862,14 @@ importers: specifier: 'catalog:' version: 7.0.0-dev.20260627.2 vite: - specifier: npm:@voidzero-dev/vite-plus-core@0.2.1 - version: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + specifier: npm:@voidzero-dev/vite-plus-core@0.2.3 + version: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' vite-plus: specifier: 'catalog:' - version: 0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) + version: 0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) vitest: specifier: 4.1.9 - version: 4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) + version: 4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) packages/dify-ui: dependencies: @@ -885,7 +885,7 @@ importers: version: 1.6.0(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7) '@chromatic-com/storybook': specifier: 'catalog:' - version: 5.2.1(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + version: 5.2.1(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) '@dify/tsconfig': specifier: workspace:* version: link:../tsconfig @@ -897,25 +897,25 @@ importers: version: 1.2.10 '@storybook/addon-a11y': specifier: 'catalog:' - version: 10.4.6(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + version: 10.4.6(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) '@storybook/addon-docs': specifier: 'catalog:' - version: 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + version: 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) '@storybook/addon-links': specifier: 'catalog:' - version: 10.4.6(@types/react@19.2.17)(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + version: 10.4.6(@types/react@19.2.17)(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) '@storybook/addon-themes': specifier: 'catalog:' - version: 10.4.6(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + version: 10.4.6(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) '@storybook/addon-vitest': specifier: 'catalog:' - version: 10.4.6(@vitest/browser-playwright@4.1.9)(@vitest/browser@4.1.9)(@vitest/runner@4.1.9)(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(vitest@4.1.9) + version: 10.4.6(@vitest/browser-playwright@4.1.9)(@vitest/browser@4.1.9)(@vitest/runner@4.1.9)(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(vitest@4.1.9) '@storybook/react-vite': specifier: 'catalog:' - version: 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3) + version: 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3) '@tailwindcss/vite': specifier: 'catalog:' - version: 4.3.1(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + version: 4.3.1(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) '@tanstack/react-hotkeys': specifier: 'catalog:' version: 0.10.0(react-dom@19.2.7(react@19.2.7))(react@19.2.7) @@ -933,13 +933,13 @@ importers: version: 7.0.0-dev.20260627.2 '@vitejs/plugin-react': specifier: 'catalog:' - version: 6.0.3(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + version: 6.0.3(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) '@vitest/browser': specifier: 'catalog:' - version: 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) + version: 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) '@vitest/browser-playwright': specifier: 'catalog:' - version: 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(playwright@1.61.1)(vitest@4.1.9) + version: 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(playwright@1.61.1)(vitest@4.1.9) '@vitest/coverage-v8': specifier: 'catalog:' version: 4.1.9(@vitest/browser@4.1.9)(vitest@4.1.9) @@ -957,7 +957,7 @@ importers: version: 19.2.7(react@19.2.7) storybook: specifier: 'catalog:' - version: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + version: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) tailwindcss: specifier: 'catalog:' version: 4.3.1 @@ -965,14 +965,14 @@ importers: specifier: 'catalog:' version: 6.0.3 vite: - specifier: npm:@voidzero-dev/vite-plus-core@0.2.1 - version: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + specifier: npm:@voidzero-dev/vite-plus-core@0.2.3 + version: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' vite-plus: specifier: 'catalog:' - version: 0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) + version: 0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) vitest: specifier: 4.1.9 - version: 4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) + version: 4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) vitest-browser-react: specifier: 'catalog:' version: 2.2.0(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(vitest@4.1.9) @@ -1004,14 +1004,14 @@ importers: specifier: 'catalog:' version: 6.0.3 vite: - specifier: npm:@voidzero-dev/vite-plus-core@0.2.1 - version: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + specifier: npm:@voidzero-dev/vite-plus-core@0.2.3 + version: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' vite-plus: specifier: 'catalog:' - version: 0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) + version: 0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) vitest: specifier: 4.1.9 - version: 4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) + version: 4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) packages/migrate-no-unchecked-indexed-access: dependencies: @@ -1029,11 +1029,11 @@ importers: specifier: 'catalog:' version: 25.9.4 vite: - specifier: npm:@voidzero-dev/vite-plus-core@0.2.1 - version: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + specifier: npm:@voidzero-dev/vite-plus-core@0.2.3 + version: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' vite-plus: specifier: 'catalog:' - version: 0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) + version: 0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) packages/tsconfig: {} @@ -1067,14 +1067,14 @@ importers: specifier: 'catalog:' version: 6.0.3 vite: - specifier: npm:@voidzero-dev/vite-plus-core@0.2.1 - version: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + specifier: npm:@voidzero-dev/vite-plus-core@0.2.3 + version: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' vite-plus: specifier: 'catalog:' - version: 0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) + version: 0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) vitest: specifier: 4.1.9 - version: 4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) + version: 4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) web: dependencies: @@ -1408,10 +1408,10 @@ importers: devDependencies: '@antfu/eslint-config': specifier: 'catalog:' - version: 9.1.0(@eslint-react/eslint-plugin@5.9.5(eslint@10.6.0(jiti@2.7.0))(typescript@6.0.3))(@next/eslint-plugin-next@16.2.9)(@typescript-eslint/typescript-estree@8.62.0(typescript@6.0.3))(@typescript-eslint/utils@8.62.0(eslint@10.6.0(jiti@2.7.0))(typescript@6.0.3))(eslint-plugin-jsx-a11y@6.10.2(eslint@10.6.0(jiti@2.7.0)))(eslint-plugin-react-refresh@0.5.3(eslint@10.6.0(jiti@2.7.0)))(eslint@10.6.0(jiti@2.7.0))(oxlint@1.70.0(oxlint-tsgolint@0.23.0)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3)(vitest@4.1.9) + version: 9.1.0(@eslint-react/eslint-plugin@5.9.5(eslint@10.6.0(jiti@2.7.0))(typescript@6.0.3))(@next/eslint-plugin-next@16.2.9)(@typescript-eslint/typescript-estree@8.62.0(typescript@6.0.3))(@typescript-eslint/utils@8.62.0(eslint@10.6.0(jiti@2.7.0))(typescript@6.0.3))(eslint-plugin-jsx-a11y@6.10.2(eslint@10.6.0(jiti@2.7.0)))(eslint-plugin-react-refresh@0.5.3(eslint@10.6.0(jiti@2.7.0)))(eslint@10.6.0(jiti@2.7.0))(oxlint@1.72.0(oxlint-tsgolint@0.24.0)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3)(vitest@4.1.9) '@chromatic-com/storybook': specifier: 'catalog:' - version: 5.2.1(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + version: 5.2.1(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) '@dify/contracts': specifier: workspace:* version: link:../packages/contracts @@ -1459,28 +1459,28 @@ importers: version: 4.2.1 '@storybook/addon-docs': specifier: 'catalog:' - version: 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + version: 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) '@storybook/addon-links': specifier: 'catalog:' - version: 10.4.6(@types/react@19.2.17)(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + version: 10.4.6(@types/react@19.2.17)(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) '@storybook/addon-onboarding': specifier: 'catalog:' - version: 10.4.6(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + version: 10.4.6(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) '@storybook/addon-themes': specifier: 'catalog:' - version: 10.4.6(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + version: 10.4.6(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) '@storybook/nextjs-vite': specifier: 'catalog:' - version: 10.4.6(@babel/core@7.29.7)(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(next@16.2.9(@babel/core@7.29.7)(@playwright/test@1.61.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(supports-color@10.2.2)(typescript@6.0.3) + version: 10.4.6(@babel/core@7.29.7)(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(next@16.2.9(@babel/core@7.29.7)(@playwright/test@1.61.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(supports-color@10.2.2)(typescript@6.0.3) '@storybook/react': specifier: 'catalog:' - version: 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3) + version: 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3) '@tailwindcss/postcss': specifier: 'catalog:' version: 4.3.1 '@tailwindcss/vite': specifier: 'catalog:' - version: 4.3.1(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + version: 4.3.1(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) '@tanstack/eslint-plugin-query': specifier: 'catalog:' version: 5.101.2(eslint@10.6.0(jiti@2.7.0))(typescript@6.0.3) @@ -1537,10 +1537,10 @@ importers: version: 7.0.0-dev.20260627.2 '@vitejs/plugin-react': specifier: 'catalog:' - version: 6.0.3(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + version: 6.0.3(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) '@vitejs/plugin-rsc': specifier: 'catalog:' - version: 0.5.27(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(react-dom@19.2.7(react@19.2.7))(react-server-dom-webpack@19.2.7(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react@19.2.7) + version: 0.5.27(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(react-dom@19.2.7(react@19.2.7))(react-server-dom-webpack@19.2.7(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react@19.2.7) '@vitest/coverage-v8': specifier: 'catalog:' version: 4.1.9(@vitest/browser@4.1.9)(vitest@4.1.9) @@ -1558,7 +1558,7 @@ importers: version: 0.11.0(eslint@10.6.0(jiti@2.7.0)) eslint-plugin-better-tailwindcss: specifier: 'catalog:' - version: 4.6.0(eslint@10.6.0(jiti@2.7.0))(oxlint@1.70.0(oxlint-tsgolint@0.23.0)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(tailwindcss@4.3.1)(typescript@6.0.3) + version: 4.6.0(eslint@10.6.0(jiti@2.7.0))(oxlint@1.72.0(oxlint-tsgolint@0.24.0)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(tailwindcss@4.3.1)(typescript@6.0.3) eslint-plugin-hyoban: specifier: 'catalog:' version: 0.14.1(eslint@10.6.0(jiti@2.7.0)) @@ -1579,7 +1579,7 @@ importers: version: 4.1.0(eslint@10.6.0(jiti@2.7.0)) eslint-plugin-storybook: specifier: 'catalog:' - version: 10.4.6(eslint@10.6.0(jiti@2.7.0))(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3) + version: 10.4.6(eslint@10.6.0(jiti@2.7.0))(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3) happy-dom: specifier: 'catalog:' version: 20.10.6 @@ -1594,7 +1594,7 @@ importers: version: 19.2.7(react-dom@19.2.7(react@19.2.7))(react@19.2.7) storybook: specifier: 'catalog:' - version: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + version: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) tailwindcss: specifier: 'catalog:' version: 4.3.1 @@ -1609,19 +1609,19 @@ importers: version: 3.19.3 vinext: specifier: 'catalog:' - version: 0.1.8(@mdx-js/rollup@3.1.1)(@vitejs/plugin-react@6.0.3(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(@vitejs/plugin-rsc@0.5.27(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(react-dom@19.2.7(react@19.2.7))(react-server-dom-webpack@19.2.7(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react@19.2.7))(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(next@16.2.9(@babel/core@7.29.7)(@playwright/test@1.61.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react-dom@19.2.7(react@19.2.7))(react-server-dom-webpack@19.2.7(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react@19.2.7) + version: 0.1.8(@mdx-js/rollup@3.1.1)(@vitejs/plugin-react@6.0.3(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(@vitejs/plugin-rsc@0.5.27(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(react-dom@19.2.7(react@19.2.7))(react-server-dom-webpack@19.2.7(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react@19.2.7))(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(next@16.2.9(@babel/core@7.29.7)(@playwright/test@1.61.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react-dom@19.2.7(react@19.2.7))(react-server-dom-webpack@19.2.7(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react@19.2.7) vite: - specifier: npm:@voidzero-dev/vite-plus-core@0.2.1 - version: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + specifier: npm:@voidzero-dev/vite-plus-core@0.2.3 + version: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' vite-plugin-inspect: specifier: 'catalog:' - version: 12.0.0-beta.3(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(typescript@6.0.3) + version: 12.0.0-beta.3(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(typescript@6.0.3) vite-plus: specifier: 'catalog:' - version: 0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) + version: 0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) vitest: specifier: 4.1.9 - version: 4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) + version: 4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) vitest-canvas-mock: specifier: 'catalog:' version: 1.1.4(vitest@4.1.9) @@ -3611,8 +3611,8 @@ packages: cpu: [x64] os: [win32] - '@oxc-project/runtime@0.136.0': - resolution: {integrity: sha512-u0EutjK5y6NHJkl5jNJCs8zbup1z6A/UEWgajrYzqcEU3UX05HjqybhMQOLhSM0eKGISyM6WfSMMuklYSmH2wA==} + '@oxc-project/runtime@0.138.0': + resolution: {integrity: sha512-yHhoXsN8tYxgdJCdD91PbySNjEEaBX/tH2OQRDXJpsQv5b184oC4/qVbU7qlblvfil/JP15Lh2HW7+HN5DS90Q==} engines: {node: ^20.19.0 || >=22.12.0} '@oxc-project/types@0.127.0': @@ -3621,12 +3621,12 @@ packages: '@oxc-project/types@0.132.0': resolution: {integrity: sha512-FESMOxil5Se014ui/Eq8fT5uHJo6nIRwH0PfJrZJXs6Gek3ZVFOrpUv3YIZT20m+extU98Hg1Ym72U58rlsxUQ==} - '@oxc-project/types@0.136.0': - resolution: {integrity: sha512-39Al/B3v9esnHCX7S8l9Se2+s2tb9b2jcMd+bZ2L659VG73kNyGPpPrL5Zi/p0ty7p4pTTU2/Dd+g27hv94XCg==} - '@oxc-project/types@0.137.0': resolution: {integrity: sha512-WT+Gb24i8hmvo85AIv2oEYouEXkRlKAlT9WaCa3TfLgNCN+GhrJOGZuIlMouAh38Qe4QOx26eUOVsq70qXrywA==} + '@oxc-project/types@0.138.0': + resolution: {integrity: sha512-1a7ZKmrRTCoN1XMZ4L0PyyqrMnrNlLyPuOkdSX2MZg7IiIGRUyurNhAm73ptDOraoBcIordsIGKNPKUzy3ZmfA==} + '@oxc-resolver/binding-android-arm-eabi@11.21.3': resolution: {integrity: sha512-eNU11A2WNizh04v3uyaJCootrHIaS0B9aHYXvAvVnPNk4xYSjMUjHnhQ6dewPN2MRYDskV85d1N0Aw0WNWhcyg==} cpu: [arm] @@ -3730,276 +3730,276 @@ packages: cpu: [x64] os: [win32] - '@oxfmt/binding-android-arm-eabi@0.55.0': - resolution: {integrity: sha512-+rFDOqQe5LOWgxrAJaZgLRudr6GQm0wGI6gtu7vVkrdLGjNMUSGbAlaCr8j7F2H2Er97vYQCU8WDb30onqMM1g==} + '@oxfmt/binding-android-arm-eabi@0.57.0': + resolution: {integrity: sha512-qVBsEO+KugOsCmUHcO8iqNnqc65p7PCKpCs8M66mPZ+Ri+CWbcpoQOEJBg2OTu03+0qu++NK1jj6IzvQVs0Sig==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [arm] os: [android] - '@oxfmt/binding-android-arm64@0.55.0': - resolution: {integrity: sha512-ctulLq8s3x8Zmvw6+iccB09TIKERAklRSmbJ10gk8mlAn05qZxoyo52dj3Hi9IJcmDSwF54fQaTVh2CbL6PInw==} + '@oxfmt/binding-android-arm64@0.57.0': + resolution: {integrity: sha512-mp6PibWbao3aizijcheOeHQaYEhcUAt8pwLniYbtLfHxL/psFF0BykAwCj+s3c6qIpa8yN8keZICWrqtZ70w8g==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [arm64] os: [android] - '@oxfmt/binding-darwin-arm64@0.55.0': - resolution: {integrity: sha512-xDQczLH9pw/RBk1h/GH0qcGMm8hQtmtVHBNLSH3lk1gEIR09hZ4L+mJQl4VqiVAvPK9VG9PYrWWuSQLt7xTbiA==} + '@oxfmt/binding-darwin-arm64@0.57.0': + resolution: {integrity: sha512-T+0stuCBqmUVY+aMIvrgXhzGhHO3sD5tNiiEcYqgSdPsnukskQqn2u5qOVD0sv1l7RLdFS5Z/f5Wi9Ktyjr3Eg==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [arm64] os: [darwin] - '@oxfmt/binding-darwin-x64@0.55.0': - resolution: {integrity: sha512-JaNoFCkF2CJdGgpPSMbuO9HVyXyoNGIhMHPvp6NYAjeVKw9XEYc0HcUWJLPQa3Q69WV5wMa9m5jPMJPtbLtcRg==} + '@oxfmt/binding-darwin-x64@0.57.0': + resolution: {integrity: sha512-O+3JbqWs/mCI2oi4xfhRO2IVPFJNDDEBV8Odo+ZpmsUOeKJfjXoNH7nDmBEQcDgK7NfjDIyE7kRgYSZcTLDO0A==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [x64] os: [darwin] - '@oxfmt/binding-freebsd-x64@0.55.0': - resolution: {integrity: sha512-DNbszhpg6S2MIzax5azdHFTTBIVkR5xr8yyRZuA4yoDAwOkzIp3tmldgKZM2+VlT+hJIG0xUksA+elISzMEAfA==} + '@oxfmt/binding-freebsd-x64@0.57.0': + resolution: {integrity: sha512-pxwhxVC+JkLX9twOQ/8C/vbuOQcMZyKIDmiRDZfO7yITuVcIdZCiLRqqf4QOxb2+8FWrRXzQpm+1DBKcMpHSSQ==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [x64] os: [freebsd] - '@oxfmt/binding-linux-arm-gnueabihf@0.55.0': - resolution: {integrity: sha512-2snoaoRfFFyGnbOcKUK36rREBYxe/Xgz3uHbiA5zbCB/s6R4DQj4mHqYAaWWhgizCUSDxV8cE9zAZ0XleNpKGw==} + '@oxfmt/binding-linux-arm-gnueabihf@0.57.0': + resolution: {integrity: sha512-pxBU4zH2imB/MDBfth2rOMeVxXUMjRQLCazagwLARIFH3hVlxZJBlM4nSnHXaIHJK4/qezoFCIORN6AY8Mra4A==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [arm] os: [linux] - '@oxfmt/binding-linux-arm-musleabihf@0.55.0': - resolution: {integrity: sha512-q1aktHF/WRpSK81BX1dE/9vWrS2jGw1Nax2kb4DBLGAewubCLcoNyp4Zl/NSMgbv3vUS46Z33wIQkBVYOP3PYg==} + '@oxfmt/binding-linux-arm-musleabihf@0.57.0': + resolution: {integrity: sha512-JAprOzt8tycYou36ZgEw14DlRHTiN8qdtKANdV3VZIRIvTI/lh/cX13c9pJ/EnDk2GT3FASH7KvCgQ2AufAifQ==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [arm] os: [linux] - '@oxfmt/binding-linux-arm64-gnu@0.55.0': - resolution: {integrity: sha512-VD0y36aENezl/3tsclA/4G53Cc7iV+7Uoh7gz4yvcOTaEYBtJpQsE6PKDGTtUtOvGS4kv51ybfXY/nWZejO5IA==} + '@oxfmt/binding-linux-arm64-gnu@0.57.0': + resolution: {integrity: sha512-ajtjaxSaj9xl4BW7REt+Cef/ttzbAq00Bq4z7JUDZEfgFXdwSjH8K9bF+IcIJzZB9lKqMfQ4eHuSFOvvlvtqOg==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [arm64] os: [linux] libc: [glibc] - '@oxfmt/binding-linux-arm64-musl@0.55.0': - resolution: {integrity: sha512-r8xlKJFcsRmn0H5jZrdORae6RX9jDBrZVvOoxF+bCQtampQJClv80aZEHsv+NsLsp2KCE5ql79O7DpPVzYWpXA==} + '@oxfmt/binding-linux-arm64-musl@0.57.0': + resolution: {integrity: sha512-p4Y/+RYk9Bk5WO+zHSUXAClRmZ2fbJCejMuCAsU2HhyME4jqf6Ftt/mJYEwIah1wGCBDYOB7wEGV1x5bCEZ6hA==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [arm64] os: [linux] libc: [musl] - '@oxfmt/binding-linux-ppc64-gnu@0.55.0': - resolution: {integrity: sha512-GRKv/HXHcwIVld/WU61rF0g0R16hl5EJ+ScKdpjevT57lnLnagj/U2YUbXf2mT+2Pg1uCzWC+mvGicPV3CDdLQ==} + '@oxfmt/binding-linux-ppc64-gnu@0.57.0': + resolution: {integrity: sha512-By6tRALAZsno0F4zedmtG+wdMvJiJmJoXM4d3+A9zHE4HRXLqXITwRH8mgrlcXc5yJM2g2W3riRPwTYdgemZLQ==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [ppc64] os: [linux] libc: [glibc] - '@oxfmt/binding-linux-riscv64-gnu@0.55.0': - resolution: {integrity: sha512-rdv57enTiPtpSYRMKfAiEbQb0Puw5t9N7isVinDoo5qeLDScro2gznmZqSgSWbVZRzLisTeCTW8Qwgw0bOHv3A==} + '@oxfmt/binding-linux-riscv64-gnu@0.57.0': + resolution: {integrity: sha512-skYeG+RgvyzspqVEBsEprL90OYYZfoVNqB3HcCNR6QDJyXKOzfDRT3zncnHmUaFluIlBHuY23mU1b5WGgR98hA==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [riscv64] os: [linux] libc: [glibc] - '@oxfmt/binding-linux-riscv64-musl@0.55.0': - resolution: {integrity: sha512-7v1nNrlD43VY6+sYQ6efYyb3lE6QY182304PD/768ZxTjOmFd/3dQa3u/nGBUAXYdGSWOQc5N3PnS0QzUXyEIA==} + '@oxfmt/binding-linux-riscv64-musl@0.57.0': + resolution: {integrity: sha512-FFgACrZOXAXUh5KQh2mt1CDOVOZmn+QzHP71wM9QobNwyQvoFfyAeefVUltW83g3sm7LTiH3yfFqLLVUpA5ZFQ==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [riscv64] os: [linux] libc: [musl] - '@oxfmt/binding-linux-s390x-gnu@0.55.0': - resolution: {integrity: sha512-f4lJLUSPOgScjFl9LiflKCTocyNRwE25JmTMbN4XQdDjoZzEHjqf3wA3VESF1/csg7i8m7+EQLbrZyYDqe10UQ==} + '@oxfmt/binding-linux-s390x-gnu@0.57.0': + resolution: {integrity: sha512-Nm/BAOfQeFiiKd502mZn/GAVKJwtd0RdCg17G3Wz/WSOIQmDi3+7/SZH4BHn1Ye5KvTVH3ua8WvfwLLycNIuvA==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [s390x] os: [linux] libc: [glibc] - '@oxfmt/binding-linux-x64-gnu@0.55.0': - resolution: {integrity: sha512-MihqiPziJNoWy4MqNSV+jVA1g+07iQDjZiR0vaCaDoPgFEiJpCMsxamktzLV07cEeQsSJ04vQaU4CzCQwIvtDA==} + '@oxfmt/binding-linux-x64-gnu@0.57.0': + resolution: {integrity: sha512-BiSy5Ku3mQqyxS6YIqAJgd403wEUWvI7kerfzPxc2l/txZVmZM0pSj7oDM+4bGBExowxOi7o73jEam1W0EDTZg==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [x64] os: [linux] libc: [glibc] - '@oxfmt/binding-linux-x64-musl@0.55.0': - resolution: {integrity: sha512-Yqghym7KYAVjP9MmSrNZiDeerMuoejNjo0r3ox5H3GDKk8eAfl8VyJm9i+pWCLDCTnAbcTUMMN2ZKjUYXH1v3g==} + '@oxfmt/binding-linux-x64-musl@0.57.0': + resolution: {integrity: sha512-BCRkJiotz5s9afLYD2LuMvzAoDYx9H17E/YbDyu4xK7l4zHDPeny9ErSXL//i/nJyaOwRk08x4b8cgJC00+JDg==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [x64] os: [linux] libc: [musl] - '@oxfmt/binding-openharmony-arm64@0.55.0': - resolution: {integrity: sha512-s5SDvVVSbyQl1V5UU3Yl12M+XLUQ3rl5SglNqgAA2K4PXUtQhyNSS00wivONPEnNo5W01rCou8WkDNyvI/RGHg==} + '@oxfmt/binding-openharmony-arm64@0.57.0': + resolution: {integrity: sha512-4Oaxe1qrGgXfpCJ1C/ERJ2iCtV2rN1R79ga9fsfyVHfSQRu/hVW780u2KDqZWFZ/iGTHODJji0JemxqFZ63eIQ==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [arm64] os: [openharmony] - '@oxfmt/binding-win32-arm64-msvc@0.55.0': - resolution: {integrity: sha512-7p9FB5R32tw2KyyNX3wpQrR2WHwEHvMEiBlGXxeTCaRMCVNx3UtFMAUbaQ/pRNWIrEUZmYhJ6tcUH52uPTRYjQ==} + '@oxfmt/binding-win32-arm64-msvc@0.57.0': + resolution: {integrity: sha512-MYLAsDnhdNsSGheLYhWgbk0vfIrlS84iQYun/y21fX6u0jj8iBtYtbpZMdiqYeuf8U12eVPUjVY2xE2NrCfJ0g==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [arm64] os: [win32] - '@oxfmt/binding-win32-ia32-msvc@0.55.0': - resolution: {integrity: sha512-ZYqj3fDnOT1IaVGMP5kpmkQl4F3tQIm2ZyAxvqkJYmI0xgWWak4ss4XYwv3VDfM+TWXeC9K4uQ/wW5jm/5XABA==} + '@oxfmt/binding-win32-ia32-msvc@0.57.0': + resolution: {integrity: sha512-PBwdzZALJY/jcCx2E6is0yu+cuVXeySTDmwuseD+9j0mHqlRNxwlKgsyRTBed/woPeqfVfuXfWjoq4Cx2Zt3Eg==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [ia32] os: [win32] - '@oxfmt/binding-win32-x64-msvc@0.55.0': - resolution: {integrity: sha512-eEYT5tivGnGbPHuOHuQpi6CGLObhh0re/5jcNQHihD2GRYkTM85dyi5a19zjP8Q00t1uqAx+/QGLUGdHeqzWyg==} + '@oxfmt/binding-win32-x64-msvc@0.57.0': + resolution: {integrity: sha512-bQJdH9i4RRfw55jm7+8/xS7GzHLLTbHx4huhrrDxQJaJtbSDbsyOnODvP1ftT7EG0KFKAYO2S+q6AcioXODx8w==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [x64] os: [win32] - '@oxlint-tsgolint/darwin-arm64@0.23.0': - resolution: {integrity: sha512-gOs9PVr2wEg4ox9z0aJo+RKhhImW86YL5N6yav8BK/rgPsIrwN/igSZ+pbRr723NFvUNKde9fgMhRA6JrXAOZw==} + '@oxlint-tsgolint/darwin-arm64@0.24.0': + resolution: {integrity: sha512-C2uMmwK5Bc4ri4ysZ6sA8Rcu+A5zBQTp6ml2u0CLLbRZp4kMFPV3yWk8B5DK9Aw7y9bbjogIm75tUwGLFzlsYQ==} cpu: [arm64] os: [darwin] - '@oxlint-tsgolint/darwin-x64@0.23.0': - resolution: {integrity: sha512-kjJ8B+7n4tB9VJdxS5A9GdJt6/bYpzbu4lXp2uO1S3sRmCB5gDEABlGoiePNApRWaW+xqL4b4xgiE727jSLhuA==} + '@oxlint-tsgolint/darwin-x64@0.24.0': + resolution: {integrity: sha512-Wgvt/1lRbDxmoNqWQKKcL+UIiqLmdJ+EWLpQa1qzoNVAfNB0PJpa82/8dH1twT/3rSs4zrP5TXPWl4juB71WuQ==} cpu: [x64] os: [darwin] - '@oxlint-tsgolint/linux-arm64@0.23.0': - resolution: {integrity: sha512-6dCZuKNu135seMXilkRk9SpCx6i1XgmiipYGalLij5WVRX6ZYS8c4xI7preN/zv9fCXhsQclTIMDu2Y/cytTjw==} + '@oxlint-tsgolint/linux-arm64@0.24.0': + resolution: {integrity: sha512-PB1rxII7KV83+ASY4sSkXtqvpij6ME66+QCRL49uksi/ofs2Rf/UVboYr095n0Rkbl2wgvlsHGl6DHC361jQUQ==} cpu: [arm64] os: [linux] - '@oxlint-tsgolint/linux-x64@0.23.0': - resolution: {integrity: sha512-3bdilnyA7kmSTjK27rvjIjSxL5SIg3wt7vwNiRkouWB83ytssyKnuGvxSYJxgMEmFpSutzaBzcCUM2jDtPGcgA==} + '@oxlint-tsgolint/linux-x64@0.24.0': + resolution: {integrity: sha512-xcz3CxKmjTQLREtE/UShh+ruWmm9nAb7UM9zKcD65BStiuYgOakAKkPHl4YS5DztpVcDrE0+HqbOolTlRKYWmw==} cpu: [x64] os: [linux] - '@oxlint-tsgolint/win32-arm64@0.23.0': - resolution: {integrity: sha512-j+OEp44SVYiQ+ZD+uttsX7u6L9SvmbbQ77SO1pSFCcJlsVMeCk8qZsjhKfGKuT/jIA+ipOJMVs/+pqUfObBWNw==} + '@oxlint-tsgolint/win32-arm64@0.24.0': + resolution: {integrity: sha512-A2i6ZGBec3i20S7RaxkgHc6r3HYtD5Mn7j/mb22NkTz14u0JuudvTu6JggAnbGMcv8+dBKQI//EasxSPJLD8pw==} cpu: [arm64] os: [win32] - '@oxlint-tsgolint/win32-x64@0.23.0': - resolution: {integrity: sha512-5MyjFuqf+g8OUPJBSGWHJtmoWnzFJYyOg4To9WMQshZYEWig/vtu7JtJ03VWnzHv9LJkAUeApY0gVCOywFR/iQ==} + '@oxlint-tsgolint/win32-x64@0.24.0': + resolution: {integrity: sha512-0ZbGd9qRB6zs82moekaKdEvncRANq49EAwfNX62JpTS46feXUhKAuoyVDvZMj6Rywejylrmmu79Wo6faYCo4Ew==} cpu: [x64] os: [win32] - '@oxlint/binding-android-arm-eabi@1.70.0': - resolution: {integrity: sha512-zFh0P4cswmRvw6nkyb89dr18rRanuaCPAsEXsFDoQY8WdaquI8Pt4NWFjaMJg6L23cy5NeN8J9cBnREbWzZhaw==} + '@oxlint/binding-android-arm-eabi@1.72.0': + resolution: {integrity: sha512-zhCmvn+1Mj3UchAc/90i99S0t7jJUsHmFVSPg4UWrjO8b8eaSGwscgO6QAUtvHBstkjQwBttQNswEnAF1mIQdA==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [arm] os: [android] - '@oxlint/binding-android-arm64@1.70.0': - resolution: {integrity: sha512-qI8o4HZjeGiBrWv+pJv4lH0Yi2Gl/JSp/EumBUApezJprIKa5PS4nU0lQsQngtky8k+SplQIOjv6hwu0SSxeyg==} + '@oxlint/binding-android-arm64@1.72.0': + resolution: {integrity: sha512-mtH+aY/ozv1eZoCUC2owjFAtyNBKHpJHygKeEu9zXXnQGW1Q2/qOpvx+I+Lf23+TvTz66F4iiXUbl2cGvoLPCQ==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [arm64] os: [android] - '@oxlint/binding-darwin-arm64@1.70.0': - resolution: {integrity: sha512-8KjgVVHI5F9nVwHCRwwA78Ty7zNKP4Wd9OeN5PSv3iu/F/u1RVXoOCgLhWqust6HmwQG6xc8c+RCyaWENy24+w==} + '@oxlint/binding-darwin-arm64@1.72.0': + resolution: {integrity: sha512-EvnajNPDtfknB3ZieeOOyDTwJn9QXDiwfnF4ZDQqART6RG6hjY4WigQcZdGoK2dkB3e1vrmEzN9aYbQCUkh/gQ==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [arm64] os: [darwin] - '@oxlint/binding-darwin-x64@1.70.0': - resolution: {integrity: sha512-WVydssv5PSUBXFJTdNBWlmGkbNmvPGaFt/2SUT/EZRB6bq6bEOHmMlbnupZD5jmlEvi9+mZJHi8TCw15lyfSfQ==} + '@oxlint/binding-darwin-x64@1.72.0': + resolution: {integrity: sha512-ZkCdEa/G80A7vEHfeCDz/+L3m33DE73v32mDKhgOIgz8Uwf0DFcK7+uu6qC+7LEhmz5fpOe1osWKyjSNMydFIQ==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [x64] os: [darwin] - '@oxlint/binding-freebsd-x64@1.70.0': - resolution: {integrity: sha512-hJucmUf8OlinHNb1R7fI4Fw6WsAstOz7i8nmkWQfiHoZXtbufNm+MxiDTIMk1ggh2Ro4vLzgQ+bKvRY54MZoRA==} + '@oxlint/binding-freebsd-x64@1.72.0': + resolution: {integrity: sha512-NroXv2vh+sxVY1uya/rM5pjhx1hm8BzlYpx9q67QP0Xhw5MH2bf5GJylpvLEC+781p1Xli/317EoV9AlGwViag==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [x64] os: [freebsd] - '@oxlint/binding-linux-arm-gnueabihf@1.70.0': - resolution: {integrity: sha512-1BnS7wbCYDSXwWzJJ+mc3NURoha6m6m6RT5c6vgAY3oz7C3OVXP+S0awo2mRq97arrJkVvO3qRQfyAHL+76xtQ==} + '@oxlint/binding-linux-arm-gnueabihf@1.72.0': + resolution: {integrity: sha512-0NDywYgfj279Ou/BcQuCYSj7NJwBfmWn5qc5uGO/Ny7fUWmXyIpvawqX/8acQlWG6IXelJsJhj+JAy6sjsKj0A==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [arm] os: [linux] - '@oxlint/binding-linux-arm-musleabihf@1.70.0': - resolution: {integrity: sha512-yKy/UdbR55+M2yEcuiV5DCNC/gdQAjr/GioUy50QwBzSrKm8ueWADqyRLS9Xk+qjNeCYGg6A8FvUBds56ttfqg==} + '@oxlint/binding-linux-arm-musleabihf@1.72.0': + resolution: {integrity: sha512-4vpXB06h65Ezsy4hRyrGjGrfa1SkVPii09yaajiYhmVpgsFiLD+KNxIx/BNAY+XiO+i1yqp9HHdwqM8VTqa5XQ==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [arm] os: [linux] - '@oxlint/binding-linux-arm64-gnu@1.70.0': - resolution: {integrity: sha512-0A5XJ4alvmqFUFP/4oYSyaO+qLto/HrKEWTSaegiVl+HOufFngK2BjYw9x4RbwBt/du5QG6l5q1zeWiJYYG5yg==} + '@oxlint/binding-linux-arm64-gnu@1.72.0': + resolution: {integrity: sha512-immaN4g2ZGFiOkKrvRX9LvzZdd2GkQM5wR+UyzYyUuyhUTXGQ4HKUJH18xp4G8OfhCVaVAJfKZxwE1r8+4hhaQ==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [arm64] os: [linux] libc: [glibc] - '@oxlint/binding-linux-arm64-musl@1.70.0': - resolution: {integrity: sha512-JiylyurlB0CLSedNtx1gzv3FvfWPF1h/2Y3BJszPLNt5XQFlBsH5ke0Jle3iJb3uqu5m2e7A/DwzpuCAHdiU+A==} + '@oxlint/binding-linux-arm64-musl@1.72.0': + resolution: {integrity: sha512-JGHS9Mnr7iWyyLDxgCv1MhzVpAckgptg00F2gnxt/GD7lQ2SW1BRcxHqhSTaSdDpjWRrBkBxMMh4+Hn3aVtExg==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [arm64] os: [linux] libc: [musl] - '@oxlint/binding-linux-ppc64-gnu@1.70.0': - resolution: {integrity: sha512-J8VPG7I3/HmgaU4u8pNU2kFx2+0U+vPLS1dXFxXOaR/2TQ0f8AC7DRz0SRGRI1bfphnX2hVYTTtLuhL4nYKL+Q==} + '@oxlint/binding-linux-ppc64-gnu@1.72.0': + resolution: {integrity: sha512-AOYgBZqxNshrg83P9v0RYv+m8s10Cqkj4/PxXFDhcS3k7FqsIG5+CxErshZCIN7G8iy4Y+VGfAsuEdar8AcbBg==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [ppc64] os: [linux] libc: [glibc] - '@oxlint/binding-linux-riscv64-gnu@1.70.0': - resolution: {integrity: sha512-N2+4lV2KLN+oXTIIIwmWDhwkrnvqf5oX7Hw0zPjk+RuIVgiBQSOlJWF7uQoFx2siEYX0ZQ5cfSbEAHm+J3t7Wg==} + '@oxlint/binding-linux-riscv64-gnu@1.72.0': + resolution: {integrity: sha512-QMybPS5ij3/vrKG67mqzHwW++91sYxK/PPUVi6SBtNCEzW4niS52fVBdXbQ6nou0wWbUPEpx8Sl/ZjtgE3clXA==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [riscv64] os: [linux] libc: [glibc] - '@oxlint/binding-linux-riscv64-musl@1.70.0': - resolution: {integrity: sha512-1e2L7cFCvx9QDzq6NPP+0tABKb5z6nWHyddWTNKprEsjO9xNrAtPowuCGpjNXxkTdsMiZ4jc8YQ5SstZd4XK6g==} + '@oxlint/binding-linux-riscv64-musl@1.72.0': + resolution: {integrity: sha512-gOc3W7JV0PXRpIL7stUlLe3Wa9Gp0Kdlup87IT3gHDvPKck2xNgMIl/Gs2lldYY2lyXZDC4rWi3hmoLUobkgbQ==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [riscv64] os: [linux] libc: [musl] - '@oxlint/binding-linux-s390x-gnu@1.70.0': - resolution: {integrity: sha512-Kwu/l/8GcYibCWA9m9N5pRXMIKVSsL/YbgpLzYkqDhWTiqdRfnNJ/+nqIKRKQiFbHWsdlHEhzMwruJK+qcEruA==} + '@oxlint/binding-linux-s390x-gnu@1.72.0': + resolution: {integrity: sha512-rpGxph+FjjHcYI5q6uxB3Az+tnfmEnDbSA8+PK9ZE/VzyUAkvBOMeuY7ZQMhu5mpZH7YQDsTdW6Cx4kV/msc6w==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [s390x] os: [linux] libc: [glibc] - '@oxlint/binding-linux-x64-gnu@1.70.0': - resolution: {integrity: sha512-tap04CsHYOl0nSAQJfPNIuBxqEPB2HnhQqwaOXLg1jnp2XfRo8Fa814dA4QC4zpvTWXCjAAaCY1W5LOORkEQuQ==} + '@oxlint/binding-linux-x64-gnu@1.72.0': + resolution: {integrity: sha512-WND+uhf/Ko13SLqQMWQUgsZuLvYYEvL0ZKgg0tgGYfLqxG7l8Ju123fHDMJyYSDl5E3bUbpFUuii/OvMreFQzw==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [x64] os: [linux] libc: [glibc] - '@oxlint/binding-linux-x64-musl@1.70.0': - resolution: {integrity: sha512-hzJa/WgvtJpbBD9rgfy0qe+MjbxOXNUT0bfR1S6EQQzfTtBFA9xg5q8KSwRrQ2QfSS+TaP4j+4mVPQrfNc6UNg==} + '@oxlint/binding-linux-x64-musl@1.72.0': + resolution: {integrity: sha512-SrpbrUL70nG9vh6zP4/oKHWgLuHquwsr7MW9XOn0olBVgh10Uqr8qscKhQoBGEn6olK/IUpn5GSKcdQ5AjUhGA==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [x64] os: [linux] libc: [musl] - '@oxlint/binding-openharmony-arm64@1.70.0': - resolution: {integrity: sha512-xbsaNSNzVSnaJACCUYr1HQMyY/Q/Q1LkePmHG3UvZPvGCYGNxrsZp9OmtA6ick8xH47ltRRbRrPCM1YXYcyC+A==} + '@oxlint/binding-openharmony-arm64@1.72.0': + resolution: {integrity: sha512-qkrsEn6NmgFKr7U/QnezQMb+q/vzAy0Dd9Y95gQGQTyjzDLN+HRZMuM5u70iyH4nBLCfKBzhjMsYCehKay2jyg==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [arm64] os: [openharmony] - '@oxlint/binding-win32-arm64-msvc@1.70.0': - resolution: {integrity: sha512-icAEsUI7JbW1TMRdEXV83mVAInhRVQYuuAlPpxdGwJ95chNdnCzjloRW8GglT0WvzOEZSio6fnYSk2DJ2Hv7LQ==} + '@oxlint/binding-win32-arm64-msvc@1.72.0': + resolution: {integrity: sha512-LWR6ZlFZph+KPjXv8opgZsXRDCdrdQe8VL8Cg9zxCoBS73h6znzZpydVgmdnwj8mB9AuSM5jxEgDJDpQkjboeg==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [arm64] os: [win32] - '@oxlint/binding-win32-ia32-msvc@1.70.0': - resolution: {integrity: sha512-FHMSWbVsPVs/f+Jcl04ws4JJ2wUnauyTzlpxWRG/lSO/8GpX08Fo2gQZqdA6CrRFI+zvkxl+N/KwJGWfUwYVZA==} + '@oxlint/binding-win32-ia32-msvc@1.72.0': + resolution: {integrity: sha512-yt6HEh7IsHvtjRWtmeZRX134eaXKHq5Gnqlf1xBJdJl1JtdoRUEJw3nAxpZoUDS860cX/foKbztO441anVBtVQ==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [ia32] os: [win32] - '@oxlint/binding-win32-x64-msvc@1.70.0': - resolution: {integrity: sha512-ptOlKwCz7n4AKs5VweMqG6DAg677FmKOK+vBkkL9DMNgFATIQ+upqUYBTOEwRQyRAx1ncGlPlXleV2hIcm3z4g==} + '@oxlint/binding-win32-x64-msvc@1.72.0': + resolution: {integrity: sha512-b2eKFD2hX7tIwmo/cyH6TDq8vzWRZ2qNHrzoGntUTmq0h3zQh/uX3eTSHCwI8OB/ADQfJCRelLItK8BsxuucDA==} engines: {node: ^20.19.0 || >=22.12.0} cpu: [x64] os: [win32] @@ -5235,15 +5235,13 @@ packages: '@vitest/utils@4.1.9': resolution: {integrity: sha512-A51o8ymO5PpqlWNnBP9ZHPXDIpuMtTLlGSjN7la4US+LJzoUMyhwjA5QXlm39JexgwHKW4Xjs8Z2d3dLCXOeuA==} - '@voidzero-dev/vite-plus-core@0.2.1': - resolution: {integrity: sha512-iWdtOlLezgYcDqIzxZx1yOUhY93vUB+ob+mRYBNr7/3Hf80uRyTQbqVD1WtsYaANbzeUi81SQ1ZoUraXHO+u8A==} + '@voidzero-dev/vite-plus-core@0.2.3': + resolution: {integrity: sha512-ObavWwyDQOy7/P0vwCPo6HNg0Wfi9e/Y+D9q6iqrEoxBU7WfMrwuiiExVwJpJsTALRXR3U31etlHbOQXt+otSg==} engines: {node: ^20.19.0 || ^22.18.0 || >=24.11.0} peerDependencies: '@arethetypeswrong/core': ^0.18.1 - '@tsdown/css': 0.22.3 - '@tsdown/exe': 0.22.3 '@types/node': ^20.19.0 || >=22.12.0 - '@vitejs/devtools': ^0.1.18 + '@vitejs/devtools': ^0.3.0 esbuild: ^0.28.1 jiti: '>=1.21.0' less: ^4.0.0 @@ -5261,10 +5259,6 @@ packages: peerDependenciesMeta: '@arethetypeswrong/core': optional: true - '@tsdown/css': - optional: true - '@tsdown/exe': - optional: true '@types/node': optional: true '@vitejs/devtools': @@ -5298,55 +5292,55 @@ packages: yaml: optional: true - '@voidzero-dev/vite-plus-darwin-arm64@0.2.1': - resolution: {integrity: sha512-9AfN/5LKRks8gbTaHPiQHT0L4yboy2xB6x6vvCRWxQMWxPS6/ZJLf5kUIZeE7I1z33AEyLKKkDscsZZVMgMLgg==} - engines: {node: ^20.19.0 || ^22.18.0 || >=24.11.0} + '@voidzero-dev/vite-plus-darwin-arm64@0.2.3': + resolution: {integrity: sha512-k8RBnsutZltnIgD60xpOX/pZC6dReJltZOOAJlcjytT1xJT45WUVfbQ9Yqp2ig9Npjn2YIgyJ/E4T8vbCYcZGw==} + engines: {node: '>=20.0.0'} cpu: [arm64] os: [darwin] - '@voidzero-dev/vite-plus-darwin-x64@0.2.1': - resolution: {integrity: sha512-Q1vyimRbf4M82qIQSWRyr7NJaH9ag5G7vVEfGVVJlQHNprI+Q8zj2Phcs/PGf6QcyjcL8UclLznQTHU9NgnKZw==} - engines: {node: ^20.19.0 || ^22.18.0 || >=24.11.0} + '@voidzero-dev/vite-plus-darwin-x64@0.2.3': + resolution: {integrity: sha512-9uIf4p6FzV7bTCJ9q10sgXHSZ0/0vWKG/utM79xfSPwFa+wpKtGErPYOEcl0Sq20Kr/T6PV0D1TOlxwi34jEzQ==} + engines: {node: '>=20.0.0'} cpu: [x64] os: [darwin] - '@voidzero-dev/vite-plus-linux-arm64-gnu@0.2.1': - resolution: {integrity: sha512-WHW3DziqedRfhJ2upq6kC4y/pmdQWYt322DVB7+4Xb4oOa/CT9GtnSrWIiXVJ4PSO42v54+YsSTKPH2HC5RbtA==} - engines: {node: ^20.19.0 || ^22.18.0 || >=24.11.0} + '@voidzero-dev/vite-plus-linux-arm64-gnu@0.2.3': + resolution: {integrity: sha512-eFurt6lOf4iHpNdZZhjJl0uLMAeRhGAdl1M7CCqX+ZGErb4E0QWWb9ylUn83e6GmJ86l9D511szE8P6EcEepjQ==} + engines: {node: '>=20.0.0'} cpu: [arm64] os: [linux] libc: [glibc] - '@voidzero-dev/vite-plus-linux-arm64-musl@0.2.1': - resolution: {integrity: sha512-vUY7hYycZW0qEevpl7ImzZJFnOEKRYCaCOX4TBW0vk6MJZ+zj/xW7e0LOggzJcz2wbYAgLDqp5h+b8wV9dguDA==} - engines: {node: ^20.19.0 || ^22.18.0 || >=24.11.0} + '@voidzero-dev/vite-plus-linux-arm64-musl@0.2.3': + resolution: {integrity: sha512-LhKllbXEUuhxfnPD/KP0MBbbIHjjCy591FEr2wlyTodfUHkku+lpfWEe9OByB1yfMbHpTHXdKeZovPI7mQhbOw==} + engines: {node: '>=20.0.0'} cpu: [arm64] os: [linux] libc: [musl] - '@voidzero-dev/vite-plus-linux-x64-gnu@0.2.1': - resolution: {integrity: sha512-tFxpToEaykBGxMQHp8M/qmr1yruRRED+c9gA1h9kmplqot04OxuqzRCWu/IiIvMJ0v3JFdOP3gqkyjXLLJhxIA==} - engines: {node: ^20.19.0 || ^22.18.0 || >=24.11.0} + '@voidzero-dev/vite-plus-linux-x64-gnu@0.2.3': + resolution: {integrity: sha512-XRbQngrET+P+kE1pHf6/g/oUC7mT5hFxZh7ESH7SlcMMnQgPyck4PifBP4M56Bksc08K3iezu0QZU046PHM/6g==} + engines: {node: '>=20.0.0'} cpu: [x64] os: [linux] libc: [glibc] - '@voidzero-dev/vite-plus-linux-x64-musl@0.2.1': - resolution: {integrity: sha512-2scSS7wEbLO2758fqr1/bAULg7nLCFa5V8LO2b5w3g1CrTYdMTDt2WX1ghPesIi+70pYGydRbXo6iaaN43zfMg==} - engines: {node: ^20.19.0 || ^22.18.0 || >=24.11.0} + '@voidzero-dev/vite-plus-linux-x64-musl@0.2.3': + resolution: {integrity: sha512-0DLtmg5DB1ePYceotaTnhXvWuTYTwmwJ6BvRji033I7yugNkaGsChuu/bB/R+7NFhzejMhA8I4gniG1Qjm2MtA==} + engines: {node: '>=20.0.0'} cpu: [x64] os: [linux] libc: [musl] - '@voidzero-dev/vite-plus-win32-arm64-msvc@0.2.1': - resolution: {integrity: sha512-3+5FJYhi9SqBszjngI2LBmvoiqEwxJWyQ5UsOUtNz6/d+yDrDw+tOgHLl4OKIh5aVNZeIGXzxvP6h24kcEqIyg==} - engines: {node: ^20.19.0 || ^22.18.0 || >=24.11.0} + '@voidzero-dev/vite-plus-win32-arm64-msvc@0.2.3': + resolution: {integrity: sha512-69DG5DtHuaWlHmu3f8+Aqop4QF/5UF7PvoS+bMkUBSTsLJqW6MaVegRg9SrSeqlnC/5mS9CkUPU/7btRO/BJeQ==} + engines: {node: '>=20.0.0'} cpu: [arm64] os: [win32] - '@voidzero-dev/vite-plus-win32-x64-msvc@0.2.1': - resolution: {integrity: sha512-5sOEwEoU5PW7ObmJ5VCakU09Oh14rYCoLQJkFqvOph6PK30lN5iqWGk0KigEyfcd7Zv+fZg9EmcERDol/3Xl9w==} - engines: {node: ^20.19.0 || ^22.18.0 || >=24.11.0} + '@voidzero-dev/vite-plus-win32-x64-msvc@0.2.3': + resolution: {integrity: sha512-+rqNdy3Iimm5/cwPiIpPmucQd7+A3Y7lVRkH5LrOT1TDdpZMx/HpSYiFgb1PXj9jy9V4RaNxfiG0uv6t3o8+pA==} + engines: {node: '>=20.0.0'} cpu: [x64] os: [win32] @@ -8315,8 +8309,8 @@ packages: oxc-resolver@11.21.3: resolution: {integrity: sha512-2Mx3fKQz7+xgrBONjsxOgCGtMHOn38/HxMzW1I5efwXB5a4lRN0Vp40gYUJFBWJslcrvwoofTrqoTnLbwTd3pA==} - oxfmt@0.55.0: - resolution: {integrity: sha512-jSj2wCTakwgPMxkfiVZX0jf+nX+Nz6xlyAZjqNE0qXTFdCBPYlP6JAN+ODjmealw7DXBjOzYbdsqwBMAZnPZ6A==} + oxfmt@0.57.0: + resolution: {integrity: sha512-ZB7Bi+rGDSqmVIo9jwcLyFgjxXvQhDdU+jx+ZrVy6VRiVXK2+CHc4hO3J4dUQjHe7V0ymHB+MDuv5z+NhK07HA==} engines: {node: ^20.19.0 || >=22.12.0} hasBin: true peerDependencies: @@ -8328,12 +8322,12 @@ packages: vite-plus: optional: true - oxlint-tsgolint@0.23.0: - resolution: {integrity: sha512-3mBv3CoPbh8dFbzfDGIWa2ytZjn2v+3EX4aKRXjIhsoGFzG8GCjfRirz3rwZf1wYbZzsNLTSgpw8VjQuWdp/jA==} + oxlint-tsgolint@0.24.0: + resolution: {integrity: sha512-giCk5sEvG02d5tzPmFMX3hem8ndzEEu1xvGYS5OwNfO2WGl6ZVxt5LjE0yiMDoz94INI7XkXwgFAQiydPvVHDw==} hasBin: true - oxlint@1.70.0: - resolution: {integrity: sha512-D6JgHtzkhRwvEC+A0Nw5AEc5bk8x5i1pHzvZIEf/a0C4hOzmAACNGtkDGPyFaxxX3ZVGxCPeig3P3rMM8XU3/g==} + oxlint@1.72.0: + resolution: {integrity: sha512-1rhdZIP/EvoI91ABIwNU5Q8+bWf8mjrS5UzIOZld4d4bXxJvtlUhlQvaoTogIGin/qdErMOrwaIJvCSIAKTLhA==} engines: {node: ^20.19.0 || >=22.12.0} hasBin: true peerDependencies: @@ -9672,8 +9666,8 @@ packages: storybook: ^0.0.0-0 || ^9.0.0 || ^10.0.0 || ^10.0.0-0 || ^10.1.0-0 || ^10.2.0-0 || ^10.3.0-0 || ^10.4.0-0 || ^10.5.0-0 || ^10.6.0-0 vite: ^5.0.0 || ^6.0.0 || ^7.0.0 || ^8.0.0 - vite-plus@0.2.1: - resolution: {integrity: sha512-q5q/Y38UkWFsNg1JO+RyRdPUqoewaSqIlMyK2p83GKNUvf4D38Ntb3PToRTDZbTRh7mWt+B+d0DQBv4nCDpMcQ==} + vite-plus@0.2.3: + resolution: {integrity: sha512-Rv2jyRNpvZR8MlxJJTmQXrt6VdRSe8Yw9QI1ZzMt5rKMPFAGhnMMktRMjfuZpgnRA+d5eY7g72Ypp2AjV74lyQ==} engines: {node: ^20.19.0 || ^22.18.0 || >=24.11.0} hasBin: true peerDependencies: @@ -10137,11 +10131,11 @@ snapshots: idb: 8.0.0 tslib: 2.8.1 - '@antfu/eslint-config@9.1.0(@eslint-react/eslint-plugin@5.9.5(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2))(typescript@6.0.3))(@next/eslint-plugin-next@16.2.9)(@typescript-eslint/typescript-estree@8.62.0(typescript@6.0.3))(@typescript-eslint/utils@8.62.0(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2))(typescript@6.0.3))(eslint-plugin-jsx-a11y@6.10.2(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2)))(eslint-plugin-react-refresh@0.5.3(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2)))(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2))(oxlint@1.70.0(oxlint-tsgolint@0.23.0)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(supports-color@10.2.2)(typescript@6.0.3)(vitest@4.1.9)': + '@antfu/eslint-config@9.1.0(@eslint-react/eslint-plugin@5.9.5(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2))(typescript@6.0.3))(@next/eslint-plugin-next@16.2.9)(@typescript-eslint/typescript-estree@8.62.0(typescript@6.0.3))(@typescript-eslint/utils@8.62.0(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2))(typescript@6.0.3))(eslint-plugin-jsx-a11y@6.10.2(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2)))(eslint-plugin-react-refresh@0.5.3(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2)))(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2))(oxlint@1.72.0(oxlint-tsgolint@0.24.0)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(supports-color@10.2.2)(typescript@6.0.3)(vitest@4.1.9)': dependencies: '@antfu/install-pkg': 1.1.0 '@clack/prompts': 1.6.0 - '@e18e/eslint-plugin': 0.5.1(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2))(oxlint@1.70.0(oxlint-tsgolint@0.23.0)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + '@e18e/eslint-plugin': 0.5.1(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2))(oxlint@1.72.0(oxlint-tsgolint@0.24.0)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) '@eslint-community/eslint-plugin-eslint-comments': 4.7.2(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2)) '@eslint/markdown': 8.0.2 '@stylistic/eslint-plugin': 5.10.0(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2)) @@ -10193,11 +10187,11 @@ snapshots: - typescript - vitest - '@antfu/eslint-config@9.1.0(@eslint-react/eslint-plugin@5.9.5(eslint@10.6.0(jiti@2.7.0))(typescript@6.0.3))(@next/eslint-plugin-next@16.2.9)(@typescript-eslint/typescript-estree@8.62.0(typescript@6.0.3))(@typescript-eslint/utils@8.62.0(eslint@10.6.0(jiti@2.7.0))(typescript@6.0.3))(eslint-plugin-jsx-a11y@6.10.2(eslint@10.6.0(jiti@2.7.0)))(eslint-plugin-react-refresh@0.5.3(eslint@10.6.0(jiti@2.7.0)))(eslint@10.6.0(jiti@2.7.0))(oxlint@1.70.0(oxlint-tsgolint@0.23.0)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3)(vitest@4.1.9)': + '@antfu/eslint-config@9.1.0(@eslint-react/eslint-plugin@5.9.5(eslint@10.6.0(jiti@2.7.0))(typescript@6.0.3))(@next/eslint-plugin-next@16.2.9)(@typescript-eslint/typescript-estree@8.62.0(typescript@6.0.3))(@typescript-eslint/utils@8.62.0(eslint@10.6.0(jiti@2.7.0))(typescript@6.0.3))(eslint-plugin-jsx-a11y@6.10.2(eslint@10.6.0(jiti@2.7.0)))(eslint-plugin-react-refresh@0.5.3(eslint@10.6.0(jiti@2.7.0)))(eslint@10.6.0(jiti@2.7.0))(oxlint@1.72.0(oxlint-tsgolint@0.24.0)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3)(vitest@4.1.9)': dependencies: '@antfu/install-pkg': 1.1.0 '@clack/prompts': 1.6.0 - '@e18e/eslint-plugin': 0.5.1(eslint@10.6.0(jiti@2.7.0))(oxlint@1.70.0(oxlint-tsgolint@0.23.0)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + '@e18e/eslint-plugin': 0.5.1(eslint@10.6.0(jiti@2.7.0))(oxlint@1.72.0(oxlint-tsgolint@0.24.0)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) '@eslint-community/eslint-plugin-eslint-comments': 4.7.2(eslint@10.6.0(jiti@2.7.0)) '@eslint/markdown': 8.0.2 '@stylistic/eslint-plugin': 5.10.0(eslint@10.6.0(jiti@2.7.0)) @@ -10389,24 +10383,24 @@ snapshots: '@chevrotain/types@11.1.2': {} - '@chromatic-com/storybook@5.2.1(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': + '@chromatic-com/storybook@5.2.1(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': dependencies: '@neoconfetti/react': 1.0.0 chromatic: 16.10.0 jsonfile: 6.2.1 - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) strip-ansi: 7.2.0 transitivePeerDependencies: - '@chromatic-com/cypress' - '@chromatic-com/playwright' - '@chromatic-com/vitest' - '@chromatic-com/storybook@5.2.1(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': + '@chromatic-com/storybook@5.2.1(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': dependencies: '@neoconfetti/react': 1.0.0 chromatic: 16.10.0 jsonfile: 6.2.1 - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) strip-ansi: 7.2.0 transitivePeerDependencies: - '@chromatic-com/cypress' @@ -10573,11 +10567,11 @@ snapshots: '@cucumber/tag-expressions@9.1.0': {} - '@devframes/hub@0.5.4(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(devframe@0.5.4(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(typescript@6.0.3))(esbuild@0.28.1)': + '@devframes/hub@0.5.4(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(devframe@0.5.4(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(typescript@6.0.3))(esbuild@0.28.1)': dependencies: birpc: 4.0.0 - devframe: 0.5.4(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(typescript@6.0.3) - nostics: 0.2.0(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1) + devframe: 0.5.4(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(typescript@6.0.3) + nostics: 0.2.0(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1) pathe: 2.0.3 perfect-debounce: 2.1.0 tinyexec: 1.2.4 @@ -10592,23 +10586,23 @@ snapshots: - vite - webpack - '@e18e/eslint-plugin@0.5.1(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2))(oxlint@1.70.0(oxlint-tsgolint@0.23.0)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': + '@e18e/eslint-plugin@0.5.1(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2))(oxlint@1.72.0(oxlint-tsgolint@0.24.0)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': dependencies: empathic: 2.0.1 module-replacements: 3.0.0-beta.8 semver: 7.8.5 optionalDependencies: eslint: 10.6.0(jiti@2.7.0)(supports-color@10.2.2) - oxlint: 1.70.0(oxlint-tsgolint@0.23.0)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + oxlint: 1.72.0(oxlint-tsgolint@0.24.0)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) - '@e18e/eslint-plugin@0.5.1(eslint@10.6.0(jiti@2.7.0))(oxlint@1.70.0(oxlint-tsgolint@0.23.0)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': + '@e18e/eslint-plugin@0.5.1(eslint@10.6.0(jiti@2.7.0))(oxlint@1.72.0(oxlint-tsgolint@0.24.0)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': dependencies: empathic: 2.0.1 module-replacements: 3.0.0-beta.8 semver: 7.8.5 optionalDependencies: eslint: 10.6.0(jiti@2.7.0) - oxlint: 1.70.0(oxlint-tsgolint@0.23.0)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + oxlint: 1.72.0(oxlint-tsgolint@0.24.0)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) '@egoist/tailwindcss-icons@1.9.2(tailwindcss@4.3.1)': dependencies: @@ -11411,19 +11405,19 @@ snapshots: dependencies: minipass: 7.1.3 - '@joshwooding/vite-plugin-react-docgen-typescript@0.7.0(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(typescript@6.0.3)': + '@joshwooding/vite-plugin-react-docgen-typescript@0.7.0(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(typescript@6.0.3)': dependencies: glob: 13.0.6 react-docgen-typescript: 2.4.0(typescript@6.0.3) - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' optionalDependencies: typescript: 6.0.3 - '@joshwooding/vite-plugin-react-docgen-typescript@0.7.0(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(typescript@6.0.3)': + '@joshwooding/vite-plugin-react-docgen-typescript@0.7.0(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(typescript@6.0.3)': dependencies: glob: 13.0.6 react-docgen-typescript: 2.4.0(typescript@6.0.3) - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' optionalDependencies: typescript: 6.0.3 @@ -12127,16 +12121,16 @@ snapshots: '@oxc-parser/binding-win32-x64-msvc@0.137.0': optional: true - '@oxc-project/runtime@0.136.0': {} + '@oxc-project/runtime@0.138.0': {} '@oxc-project/types@0.127.0': {} '@oxc-project/types@0.132.0': {} - '@oxc-project/types@0.136.0': {} - '@oxc-project/types@0.137.0': {} + '@oxc-project/types@0.138.0': {} + '@oxc-resolver/binding-android-arm-eabi@11.21.3': optional: true @@ -12198,136 +12192,136 @@ snapshots: '@oxc-resolver/binding-win32-x64-msvc@11.21.3': optional: true - '@oxfmt/binding-android-arm-eabi@0.55.0': + '@oxfmt/binding-android-arm-eabi@0.57.0': optional: true - '@oxfmt/binding-android-arm64@0.55.0': + '@oxfmt/binding-android-arm64@0.57.0': optional: true - '@oxfmt/binding-darwin-arm64@0.55.0': + '@oxfmt/binding-darwin-arm64@0.57.0': optional: true - '@oxfmt/binding-darwin-x64@0.55.0': + '@oxfmt/binding-darwin-x64@0.57.0': optional: true - '@oxfmt/binding-freebsd-x64@0.55.0': + '@oxfmt/binding-freebsd-x64@0.57.0': optional: true - '@oxfmt/binding-linux-arm-gnueabihf@0.55.0': + '@oxfmt/binding-linux-arm-gnueabihf@0.57.0': optional: true - '@oxfmt/binding-linux-arm-musleabihf@0.55.0': + '@oxfmt/binding-linux-arm-musleabihf@0.57.0': optional: true - '@oxfmt/binding-linux-arm64-gnu@0.55.0': + '@oxfmt/binding-linux-arm64-gnu@0.57.0': optional: true - '@oxfmt/binding-linux-arm64-musl@0.55.0': + '@oxfmt/binding-linux-arm64-musl@0.57.0': optional: true - '@oxfmt/binding-linux-ppc64-gnu@0.55.0': + '@oxfmt/binding-linux-ppc64-gnu@0.57.0': optional: true - '@oxfmt/binding-linux-riscv64-gnu@0.55.0': + '@oxfmt/binding-linux-riscv64-gnu@0.57.0': optional: true - '@oxfmt/binding-linux-riscv64-musl@0.55.0': + '@oxfmt/binding-linux-riscv64-musl@0.57.0': optional: true - '@oxfmt/binding-linux-s390x-gnu@0.55.0': + '@oxfmt/binding-linux-s390x-gnu@0.57.0': optional: true - '@oxfmt/binding-linux-x64-gnu@0.55.0': + '@oxfmt/binding-linux-x64-gnu@0.57.0': optional: true - '@oxfmt/binding-linux-x64-musl@0.55.0': + '@oxfmt/binding-linux-x64-musl@0.57.0': optional: true - '@oxfmt/binding-openharmony-arm64@0.55.0': + '@oxfmt/binding-openharmony-arm64@0.57.0': optional: true - '@oxfmt/binding-win32-arm64-msvc@0.55.0': + '@oxfmt/binding-win32-arm64-msvc@0.57.0': optional: true - '@oxfmt/binding-win32-ia32-msvc@0.55.0': + '@oxfmt/binding-win32-ia32-msvc@0.57.0': optional: true - '@oxfmt/binding-win32-x64-msvc@0.55.0': + '@oxfmt/binding-win32-x64-msvc@0.57.0': optional: true - '@oxlint-tsgolint/darwin-arm64@0.23.0': + '@oxlint-tsgolint/darwin-arm64@0.24.0': optional: true - '@oxlint-tsgolint/darwin-x64@0.23.0': + '@oxlint-tsgolint/darwin-x64@0.24.0': optional: true - '@oxlint-tsgolint/linux-arm64@0.23.0': + '@oxlint-tsgolint/linux-arm64@0.24.0': optional: true - '@oxlint-tsgolint/linux-x64@0.23.0': + '@oxlint-tsgolint/linux-x64@0.24.0': optional: true - '@oxlint-tsgolint/win32-arm64@0.23.0': + '@oxlint-tsgolint/win32-arm64@0.24.0': optional: true - '@oxlint-tsgolint/win32-x64@0.23.0': + '@oxlint-tsgolint/win32-x64@0.24.0': optional: true - '@oxlint/binding-android-arm-eabi@1.70.0': + '@oxlint/binding-android-arm-eabi@1.72.0': optional: true - '@oxlint/binding-android-arm64@1.70.0': + '@oxlint/binding-android-arm64@1.72.0': optional: true - '@oxlint/binding-darwin-arm64@1.70.0': + '@oxlint/binding-darwin-arm64@1.72.0': optional: true - '@oxlint/binding-darwin-x64@1.70.0': + '@oxlint/binding-darwin-x64@1.72.0': optional: true - '@oxlint/binding-freebsd-x64@1.70.0': + '@oxlint/binding-freebsd-x64@1.72.0': optional: true - '@oxlint/binding-linux-arm-gnueabihf@1.70.0': + '@oxlint/binding-linux-arm-gnueabihf@1.72.0': optional: true - '@oxlint/binding-linux-arm-musleabihf@1.70.0': + '@oxlint/binding-linux-arm-musleabihf@1.72.0': optional: true - '@oxlint/binding-linux-arm64-gnu@1.70.0': + '@oxlint/binding-linux-arm64-gnu@1.72.0': optional: true - '@oxlint/binding-linux-arm64-musl@1.70.0': + '@oxlint/binding-linux-arm64-musl@1.72.0': optional: true - '@oxlint/binding-linux-ppc64-gnu@1.70.0': + '@oxlint/binding-linux-ppc64-gnu@1.72.0': optional: true - '@oxlint/binding-linux-riscv64-gnu@1.70.0': + '@oxlint/binding-linux-riscv64-gnu@1.72.0': optional: true - '@oxlint/binding-linux-riscv64-musl@1.70.0': + '@oxlint/binding-linux-riscv64-musl@1.72.0': optional: true - '@oxlint/binding-linux-s390x-gnu@1.70.0': + '@oxlint/binding-linux-s390x-gnu@1.72.0': optional: true - '@oxlint/binding-linux-x64-gnu@1.70.0': + '@oxlint/binding-linux-x64-gnu@1.72.0': optional: true - '@oxlint/binding-linux-x64-musl@1.70.0': + '@oxlint/binding-linux-x64-musl@1.72.0': optional: true - '@oxlint/binding-openharmony-arm64@1.70.0': + '@oxlint/binding-openharmony-arm64@1.72.0': optional: true - '@oxlint/binding-win32-arm64-msvc@1.70.0': + '@oxlint/binding-win32-arm64-msvc@1.72.0': optional: true - '@oxlint/binding-win32-ia32-msvc@1.70.0': + '@oxlint/binding-win32-ia32-msvc@1.72.0': optional: true - '@oxlint/binding-win32-x64-msvc@1.70.0': + '@oxlint/binding-win32-x64-msvc@1.72.0': optional: true '@oxlint/plugins@1.68.0': {} @@ -12670,21 +12664,21 @@ snapshots: '@standard-schema/spec@1.1.0': {} - '@storybook/addon-a11y@10.4.6(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': + '@storybook/addon-a11y@10.4.6(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': dependencies: '@storybook/global': 5.0.0 axe-core: 4.12.1 - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) - '@storybook/addon-docs@10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': + '@storybook/addon-docs@10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': dependencies: '@mdx-js/react': 3.1.1(@types/react@19.2.17)(react@19.2.7) - '@storybook/csf-plugin': 10.4.6(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + '@storybook/csf-plugin': 10.4.6(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) '@storybook/icons': 2.1.0(react@19.2.7) - '@storybook/react-dom-shim': 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + '@storybook/react-dom-shim': 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) react: 19.2.7 react-dom: 19.2.7(react@19.2.7) - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) ts-dedent: 2.3.0 optionalDependencies: '@types/react': 19.2.17 @@ -12695,15 +12689,15 @@ snapshots: - vite - webpack - '@storybook/addon-docs@10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': + '@storybook/addon-docs@10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': dependencies: '@mdx-js/react': 3.1.1(@types/react@19.2.17)(react@19.2.7) - '@storybook/csf-plugin': 10.4.6(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + '@storybook/csf-plugin': 10.4.6(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) '@storybook/icons': 2.1.0(react@19.2.7) - '@storybook/react-dom-shim': 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + '@storybook/react-dom-shim': 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) react: 19.2.7 react-dom: 19.2.7(react@19.2.7) - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) ts-dedent: 2.3.0 optionalDependencies: '@types/react': 19.2.17 @@ -12714,86 +12708,86 @@ snapshots: - vite - webpack - '@storybook/addon-links@10.4.6(@types/react@19.2.17)(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': + '@storybook/addon-links@10.4.6(@types/react@19.2.17)(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': dependencies: '@storybook/global': 5.0.0 - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) optionalDependencies: '@types/react': 19.2.17 react: 19.2.7 - '@storybook/addon-links@10.4.6(@types/react@19.2.17)(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': + '@storybook/addon-links@10.4.6(@types/react@19.2.17)(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': dependencies: '@storybook/global': 5.0.0 - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) optionalDependencies: '@types/react': 19.2.17 react: 19.2.7 - '@storybook/addon-onboarding@10.4.6(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': + '@storybook/addon-onboarding@10.4.6(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': dependencies: - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) - '@storybook/addon-themes@10.4.6(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': + '@storybook/addon-themes@10.4.6(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': dependencies: - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) ts-dedent: 2.3.0 - '@storybook/addon-themes@10.4.6(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': + '@storybook/addon-themes@10.4.6(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': dependencies: - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) ts-dedent: 2.3.0 - '@storybook/addon-vitest@10.4.6(@vitest/browser-playwright@4.1.9)(@vitest/browser@4.1.9)(@vitest/runner@4.1.9)(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(vitest@4.1.9)': + '@storybook/addon-vitest@10.4.6(@vitest/browser-playwright@4.1.9)(@vitest/browser@4.1.9)(@vitest/runner@4.1.9)(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(vitest@4.1.9)': dependencies: '@storybook/global': 5.0.0 '@storybook/icons': 2.1.0(react@19.2.7) - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) optionalDependencies: - '@vitest/browser': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) - '@vitest/browser-playwright': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(playwright@1.61.1)(vitest@4.1.9) + '@vitest/browser': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) + '@vitest/browser-playwright': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(playwright@1.61.1)(vitest@4.1.9) '@vitest/runner': 4.1.9 - vitest: 4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) + vitest: 4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) transitivePeerDependencies: - react - '@storybook/builder-vite@10.4.6(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': + '@storybook/builder-vite@10.4.6(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': dependencies: - '@storybook/csf-plugin': 10.4.6(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + '@storybook/csf-plugin': 10.4.6(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) ts-dedent: 2.3.0 - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' transitivePeerDependencies: - esbuild - rollup - webpack - '@storybook/builder-vite@10.4.6(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': + '@storybook/builder-vite@10.4.6(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': dependencies: - '@storybook/csf-plugin': 10.4.6(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + '@storybook/csf-plugin': 10.4.6(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) ts-dedent: 2.3.0 - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' transitivePeerDependencies: - esbuild - rollup - webpack - '@storybook/csf-plugin@10.4.6(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': + '@storybook/csf-plugin@10.4.6(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': dependencies: - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) unplugin: 2.3.11 optionalDependencies: esbuild: 0.28.1 - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' - '@storybook/csf-plugin@10.4.6(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': + '@storybook/csf-plugin@10.4.6(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': dependencies: - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) unplugin: 2.3.11 optionalDependencies: esbuild: 0.28.1 - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' '@storybook/global@5.0.0': {} @@ -12801,18 +12795,18 @@ snapshots: dependencies: react: 19.2.7 - '@storybook/nextjs-vite@10.4.6(@babel/core@7.29.7)(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(next@16.2.9(@babel/core@7.29.7)(@playwright/test@1.61.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(supports-color@10.2.2)(typescript@6.0.3)': + '@storybook/nextjs-vite@10.4.6(@babel/core@7.29.7)(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(next@16.2.9(@babel/core@7.29.7)(@playwright/test@1.61.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(supports-color@10.2.2)(typescript@6.0.3)': dependencies: - '@storybook/builder-vite': 10.4.6(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) - '@storybook/react': 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3) - '@storybook/react-vite': 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3) + '@storybook/builder-vite': 10.4.6(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + '@storybook/react': 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3) + '@storybook/react-vite': 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3) next: 16.2.9(@babel/core@7.29.7)(@playwright/test@1.61.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7) react: 19.2.7 react-dom: 19.2.7(react@19.2.7) - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) styled-jsx: 5.1.6(@babel/core@7.29.7)(react@19.2.7) - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' - vite-plugin-storybook-nextjs: 3.3.0(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(next@16.2.9(@babel/core@7.29.7)(@playwright/test@1.61.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(supports-color@10.2.2)(typescript@6.0.3) + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite-plugin-storybook-nextjs: 3.3.0(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(next@16.2.9(@babel/core@7.29.7)(@playwright/test@1.61.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(supports-color@10.2.2)(typescript@6.0.3) optionalDependencies: '@types/react': 19.2.17 '@types/react-dom': 19.2.3(@types/react@19.2.17) @@ -12825,39 +12819,39 @@ snapshots: - supports-color - webpack - '@storybook/react-dom-shim@10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': + '@storybook/react-dom-shim@10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': dependencies: react: 19.2.7 react-dom: 19.2.7(react@19.2.7) - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) optionalDependencies: '@types/react': 19.2.17 '@types/react-dom': 19.2.3(@types/react@19.2.17) - '@storybook/react-dom-shim@10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': + '@storybook/react-dom-shim@10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))': dependencies: react: 19.2.7 react-dom: 19.2.7(react@19.2.7) - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) optionalDependencies: '@types/react': 19.2.17 '@types/react-dom': 19.2.3(@types/react@19.2.17) - '@storybook/react-vite@10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3)': + '@storybook/react-vite@10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3)': dependencies: - '@joshwooding/vite-plugin-react-docgen-typescript': 0.7.0(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(typescript@6.0.3) + '@joshwooding/vite-plugin-react-docgen-typescript': 0.7.0(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(typescript@6.0.3) '@rollup/pluginutils': 5.4.0 - '@storybook/builder-vite': 10.4.6(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) - '@storybook/react': 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3) + '@storybook/builder-vite': 10.4.6(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + '@storybook/react': 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3) empathic: 2.0.1 magic-string: 0.30.21 react: 19.2.7 react-docgen: 8.0.3 react-dom: 19.2.7(react@19.2.7) resolve: 1.22.12 - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) tsconfig-paths: 4.2.0 - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' transitivePeerDependencies: - '@types/react' - '@types/react-dom' @@ -12867,21 +12861,21 @@ snapshots: - typescript - webpack - '@storybook/react-vite@10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3)': + '@storybook/react-vite@10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3)': dependencies: - '@joshwooding/vite-plugin-react-docgen-typescript': 0.7.0(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(typescript@6.0.3) + '@joshwooding/vite-plugin-react-docgen-typescript': 0.7.0(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(typescript@6.0.3) '@rollup/pluginutils': 5.4.0 - '@storybook/builder-vite': 10.4.6(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) - '@storybook/react': 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3) + '@storybook/builder-vite': 10.4.6(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + '@storybook/react': 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3) empathic: 2.0.1 magic-string: 0.30.21 react: 19.2.7 react-docgen: 8.0.3 react-dom: 19.2.7(react@19.2.7) resolve: 1.22.12 - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) tsconfig-paths: 4.2.0 - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' transitivePeerDependencies: - '@types/react' - '@types/react-dom' @@ -12891,15 +12885,15 @@ snapshots: - typescript - webpack - '@storybook/react@10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3)': + '@storybook/react@10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3)': dependencies: '@storybook/global': 5.0.0 - '@storybook/react-dom-shim': 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + '@storybook/react-dom-shim': 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) react: 19.2.7 react-docgen: 8.0.3 react-docgen-typescript: 2.4.0(typescript@6.0.3) react-dom: 19.2.7(react@19.2.7) - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) optionalDependencies: '@types/react': 19.2.17 '@types/react-dom': 19.2.3(@types/react@19.2.17) @@ -12907,15 +12901,15 @@ snapshots: transitivePeerDependencies: - supports-color - '@storybook/react@10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3)': + '@storybook/react@10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3)': dependencies: '@storybook/global': 5.0.0 - '@storybook/react-dom-shim': 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) + '@storybook/react-dom-shim': 10.4.6(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))) react: 19.2.7 react-docgen: 8.0.3 react-docgen-typescript: 2.4.0(typescript@6.0.3) react-dom: 19.2.7(react@19.2.7) - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) optionalDependencies: '@types/react': 19.2.17 '@types/react-dom': 19.2.3(@types/react@19.2.17) @@ -13046,19 +13040,19 @@ snapshots: postcss-selector-parser: 6.1.4 tailwindcss: 4.3.1 - '@tailwindcss/vite@4.3.1(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))': + '@tailwindcss/vite@4.3.1(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))': dependencies: '@tailwindcss/node': 4.3.1 '@tailwindcss/oxide': 4.3.1 tailwindcss: 4.3.1 - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' - '@tailwindcss/vite@4.3.1(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))': + '@tailwindcss/vite@4.3.1(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))': dependencies: '@tailwindcss/node': 4.3.1 '@tailwindcss/oxide': 4.3.1 tailwindcss: 4.3.1 - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' '@tanstack/devtools-event-client@0.4.4': {} @@ -13634,17 +13628,17 @@ snapshots: '@resvg/resvg-wasm': 2.4.0 satori: 0.16.0 - '@vitejs/devtools-kit@0.3.3(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(typescript@6.0.3)': + '@vitejs/devtools-kit@0.3.3(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(typescript@6.0.3)': dependencies: - '@devframes/hub': 0.5.4(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(devframe@0.5.4(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(typescript@6.0.3))(esbuild@0.28.1) + '@devframes/hub': 0.5.4(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(devframe@0.5.4(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(typescript@6.0.3))(esbuild@0.28.1) birpc: 4.0.0 - devframe: 0.5.4(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(typescript@6.0.3) + devframe: 0.5.4(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(typescript@6.0.3) mlly: 1.8.2 - nostics: 0.3.0(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1) + nostics: 0.3.0(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1) pathe: 2.0.3 perfect-debounce: 2.1.0 tinyexec: 1.2.4 - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' transitivePeerDependencies: - '@farmfe/core' - '@modelcontextprotocol/sdk' @@ -13660,17 +13654,17 @@ snapshots: - utf-8-validate - webpack - '@vitejs/plugin-react@6.0.3(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))': + '@vitejs/plugin-react@6.0.3(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))': dependencies: '@rolldown/pluginutils': 1.0.1 - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' - '@vitejs/plugin-react@6.0.3(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))': + '@vitejs/plugin-react@6.0.3(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))': dependencies: '@rolldown/pluginutils': 1.0.1 - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' - '@vitejs/plugin-rsc@0.5.27(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(react-dom@19.2.7(react@19.2.7))(react-server-dom-webpack@19.2.7(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react@19.2.7)': + '@vitejs/plugin-rsc@0.5.27(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(react-dom@19.2.7(react@19.2.7))(react-server-dom-webpack@19.2.7(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react@19.2.7)': dependencies: '@rolldown/pluginutils': 1.0.1 es-module-lexer: 2.1.0 @@ -13681,18 +13675,18 @@ snapshots: srvx: 0.11.17 strip-literal: 3.1.0 turbo-stream: 3.2.0 - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' - vitefu: 1.1.3(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vitefu: 1.1.3(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) optionalDependencies: react-server-dom-webpack: 19.2.7(react-dom@19.2.7(react@19.2.7))(react@19.2.7) - '@vitest/browser-playwright@4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(playwright@1.61.1)(vitest@4.1.9)': + '@vitest/browser-playwright@4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(playwright@1.61.1)(vitest@4.1.9)': dependencies: - '@vitest/browser': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) - '@vitest/mocker': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + '@vitest/browser': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) + '@vitest/mocker': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) playwright: 1.61.1 tinyrainbow: 3.1.0 - vitest: 4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) + vitest: 4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) transitivePeerDependencies: - bufferutil - msw @@ -13700,53 +13694,53 @@ snapshots: - vite optional: true - '@vitest/browser-playwright@4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(playwright@1.61.1)(vitest@4.1.9)': + '@vitest/browser-playwright@4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(playwright@1.61.1)(vitest@4.1.9)': dependencies: - '@vitest/browser': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) - '@vitest/mocker': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + '@vitest/browser': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) + '@vitest/mocker': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) playwright: 1.61.1 tinyrainbow: 3.1.0 - vitest: 4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) + vitest: 4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) transitivePeerDependencies: - bufferutil - msw - utf-8-validate - vite - '@vitest/browser-preview@4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9)': + '@vitest/browser-preview@4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9)': dependencies: '@testing-library/dom': 10.4.1 '@testing-library/user-event': 14.6.1(@testing-library/dom@10.4.1) - '@vitest/browser': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) - vitest: 4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) + '@vitest/browser': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) + vitest: 4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) transitivePeerDependencies: - bufferutil - msw - utf-8-validate - vite - '@vitest/browser-preview@4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9)': + '@vitest/browser-preview@4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9)': dependencies: '@testing-library/dom': 10.4.1 '@testing-library/user-event': 14.6.1(@testing-library/dom@10.4.1) - '@vitest/browser': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) - vitest: 4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) + '@vitest/browser': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) + vitest: 4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) transitivePeerDependencies: - bufferutil - msw - utf-8-validate - vite - '@vitest/browser@4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9)': + '@vitest/browser@4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9)': dependencies: '@blazediff/core': 1.9.1 - '@vitest/mocker': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + '@vitest/mocker': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) '@vitest/utils': 4.1.9 magic-string: 0.30.21 pngjs: 7.0.0 sirv: 3.0.2 tinyrainbow: 3.1.0 - vitest: 4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) + vitest: 4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) ws: 8.21.0 transitivePeerDependencies: - bufferutil @@ -13754,16 +13748,16 @@ snapshots: - utf-8-validate - vite - '@vitest/browser@4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9)': + '@vitest/browser@4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9)': dependencies: '@blazediff/core': 1.9.1 - '@vitest/mocker': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + '@vitest/mocker': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) '@vitest/utils': 4.1.9 magic-string: 0.30.21 pngjs: 7.0.0 sirv: 3.0.2 tinyrainbow: 3.1.0 - vitest: 4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) + vitest: 4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) ws: 8.21.0 transitivePeerDependencies: - bufferutil @@ -13783,9 +13777,9 @@ snapshots: obug: 2.1.3 std-env: 4.1.0 tinyrainbow: 3.1.0 - vitest: 4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) + vitest: 4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) optionalDependencies: - '@vitest/browser': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) + '@vitest/browser': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) '@vitest/eslint-plugin@1.6.20(@typescript-eslint/eslint-plugin@8.62.0(@typescript-eslint/parser@8.62.0(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2))(typescript@6.0.3))(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2))(typescript@6.0.3))(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2))(typescript@6.0.3)(vitest@4.1.9)': dependencies: @@ -13795,7 +13789,7 @@ snapshots: optionalDependencies: '@typescript-eslint/eslint-plugin': 8.62.0(@typescript-eslint/parser@8.62.0(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2))(typescript@6.0.3))(eslint@10.6.0(jiti@2.7.0)(supports-color@10.2.2))(typescript@6.0.3) typescript: 6.0.3 - vitest: 4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) + vitest: 4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) transitivePeerDependencies: - supports-color @@ -13807,7 +13801,7 @@ snapshots: optionalDependencies: '@typescript-eslint/eslint-plugin': 8.62.0(@typescript-eslint/parser@8.62.0(eslint@10.6.0(jiti@2.7.0))(supports-color@10.2.2)(typescript@6.0.3))(eslint@10.6.0(jiti@2.7.0))(supports-color@10.2.2)(typescript@6.0.3) typescript: 6.0.3 - vitest: 4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) + vitest: 4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) transitivePeerDependencies: - supports-color @@ -13828,21 +13822,21 @@ snapshots: chai: 6.2.2 tinyrainbow: 3.1.0 - '@vitest/mocker@4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))': + '@vitest/mocker@4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))': dependencies: '@vitest/spy': 4.1.9 estree-walker: 3.0.3 magic-string: 0.30.21 optionalDependencies: - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' - '@vitest/mocker@4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))': + '@vitest/mocker@4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))': dependencies: '@vitest/spy': 4.1.9 estree-walker: 3.0.3 magic-string: 0.30.21 optionalDependencies: - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' '@vitest/pretty-format@3.2.4': dependencies: @@ -13882,10 +13876,10 @@ snapshots: convert-source-map: 2.0.0 tinyrainbow: 3.1.0 - '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)': + '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)': dependencies: - '@oxc-project/runtime': 0.136.0 - '@oxc-project/types': 0.136.0 + '@oxc-project/runtime': 0.138.0 + '@oxc-project/types': 0.138.0 lightningcss: 1.32.0 postcss: 8.5.16 optionalDependencies: @@ -13897,10 +13891,10 @@ snapshots: typescript: 6.0.3 yaml: 2.9.0 - '@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)': + '@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)': dependencies: - '@oxc-project/runtime': 0.136.0 - '@oxc-project/types': 0.136.0 + '@oxc-project/runtime': 0.138.0 + '@oxc-project/types': 0.138.0 lightningcss: 1.32.0 postcss: 8.5.16 optionalDependencies: @@ -13912,28 +13906,28 @@ snapshots: typescript: 6.0.3 yaml: 2.9.0 - '@voidzero-dev/vite-plus-darwin-arm64@0.2.1': + '@voidzero-dev/vite-plus-darwin-arm64@0.2.3': optional: true - '@voidzero-dev/vite-plus-darwin-x64@0.2.1': + '@voidzero-dev/vite-plus-darwin-x64@0.2.3': optional: true - '@voidzero-dev/vite-plus-linux-arm64-gnu@0.2.1': + '@voidzero-dev/vite-plus-linux-arm64-gnu@0.2.3': optional: true - '@voidzero-dev/vite-plus-linux-arm64-musl@0.2.1': + '@voidzero-dev/vite-plus-linux-arm64-musl@0.2.3': optional: true - '@voidzero-dev/vite-plus-linux-x64-gnu@0.2.1': + '@voidzero-dev/vite-plus-linux-x64-gnu@0.2.3': optional: true - '@voidzero-dev/vite-plus-linux-x64-musl@0.2.1': + '@voidzero-dev/vite-plus-linux-x64-musl@0.2.3': optional: true - '@voidzero-dev/vite-plus-win32-arm64-msvc@0.2.1': + '@voidzero-dev/vite-plus-win32-arm64-msvc@0.2.3': optional: true - '@voidzero-dev/vite-plus-win32-x64-msvc@0.2.1': + '@voidzero-dev/vite-plus-win32-x64-msvc@0.2.3': optional: true '@volar/language-core@2.4.28': @@ -14758,14 +14752,14 @@ snapshots: detect-node-es@1.1.0: {} - devframe@0.5.4(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(typescript@6.0.3): + devframe@0.5.4(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(typescript@6.0.3): dependencies: '@valibot/to-json-schema': 1.7.1(valibot@1.4.2(typescript@6.0.3)) birpc: 4.0.0 cac: 7.0.0 h3: 2.0.1-rc.22 mrmime: 2.0.1 - nostics: 0.2.0(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1) + nostics: 0.2.0(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1) pathe: 2.0.3 valibot: 1.4.2(typescript@6.0.3) ws: 8.21.0 @@ -15145,7 +15139,7 @@ snapshots: dependencies: eslint: 10.6.0(jiti@2.7.0) - eslint-plugin-better-tailwindcss@4.6.0(eslint@10.6.0(jiti@2.7.0))(oxlint@1.70.0(oxlint-tsgolint@0.23.0)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(tailwindcss@4.3.1)(typescript@6.0.3): + eslint-plugin-better-tailwindcss@4.6.0(eslint@10.6.0(jiti@2.7.0))(oxlint@1.72.0(oxlint-tsgolint@0.24.0)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(tailwindcss@4.3.1)(typescript@6.0.3): dependencies: '@eslint/css-tree': 4.0.4 '@valibot/to-json-schema': 1.7.1(valibot@1.4.2(typescript@6.0.3)) @@ -15158,7 +15152,7 @@ snapshots: valibot: 1.4.2(typescript@6.0.3) optionalDependencies: eslint: 10.6.0(jiti@2.7.0) - oxlint: 1.70.0(oxlint-tsgolint@0.23.0)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + oxlint: 1.72.0(oxlint-tsgolint@0.24.0)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) transitivePeerDependencies: - '@eslint/css' - typescript @@ -15680,11 +15674,11 @@ snapshots: typescript: 6.0.3 yaml: 2.9.0 - eslint-plugin-storybook@10.4.6(eslint@10.6.0(jiti@2.7.0))(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3): + eslint-plugin-storybook@10.4.6(eslint@10.6.0(jiti@2.7.0))(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(typescript@6.0.3): dependencies: '@typescript-eslint/utils': 8.62.0(eslint@10.6.0(jiti@2.7.0))(typescript@6.0.3) eslint: 10.6.0(jiti@2.7.0) - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) transitivePeerDependencies: - supports-color - typescript @@ -17605,11 +17599,11 @@ snapshots: normalize-wheel@1.0.1: {} - nostics@0.2.0(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1): + nostics@0.2.0(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1): dependencies: magic-string: 0.30.21 oxc-parser: 0.132.0 - unplugin: 3.2.0(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1) + unplugin: 3.2.0(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1) transitivePeerDependencies: - '@farmfe/core' - '@rspack/core' @@ -17621,11 +17615,11 @@ snapshots: - vite - webpack - nostics@0.3.0(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1): + nostics@0.3.0(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1): dependencies: magic-string: 0.30.21 oxc-parser: 0.132.0 - unplugin: 3.2.0(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1) + unplugin: 3.2.0(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1) transitivePeerDependencies: - '@farmfe/core' - '@rspack/core' @@ -17840,161 +17834,112 @@ snapshots: '@oxc-resolver/binding-win32-arm64-msvc': 11.21.3 '@oxc-resolver/binding-win32-x64-msvc': 11.21.3 - oxfmt@0.55.0(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)): + oxfmt@0.57.0(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)): dependencies: tinypool: 2.1.0 optionalDependencies: - '@oxfmt/binding-android-arm-eabi': 0.55.0 - '@oxfmt/binding-android-arm64': 0.55.0 - '@oxfmt/binding-darwin-arm64': 0.55.0 - '@oxfmt/binding-darwin-x64': 0.55.0 - '@oxfmt/binding-freebsd-x64': 0.55.0 - '@oxfmt/binding-linux-arm-gnueabihf': 0.55.0 - '@oxfmt/binding-linux-arm-musleabihf': 0.55.0 - '@oxfmt/binding-linux-arm64-gnu': 0.55.0 - '@oxfmt/binding-linux-arm64-musl': 0.55.0 - '@oxfmt/binding-linux-ppc64-gnu': 0.55.0 - '@oxfmt/binding-linux-riscv64-gnu': 0.55.0 - '@oxfmt/binding-linux-riscv64-musl': 0.55.0 - '@oxfmt/binding-linux-s390x-gnu': 0.55.0 - '@oxfmt/binding-linux-x64-gnu': 0.55.0 - '@oxfmt/binding-linux-x64-musl': 0.55.0 - '@oxfmt/binding-openharmony-arm64': 0.55.0 - '@oxfmt/binding-win32-arm64-msvc': 0.55.0 - '@oxfmt/binding-win32-ia32-msvc': 0.55.0 - '@oxfmt/binding-win32-x64-msvc': 0.55.0 - vite-plus: 0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) + '@oxfmt/binding-android-arm-eabi': 0.57.0 + '@oxfmt/binding-android-arm64': 0.57.0 + '@oxfmt/binding-darwin-arm64': 0.57.0 + '@oxfmt/binding-darwin-x64': 0.57.0 + '@oxfmt/binding-freebsd-x64': 0.57.0 + '@oxfmt/binding-linux-arm-gnueabihf': 0.57.0 + '@oxfmt/binding-linux-arm-musleabihf': 0.57.0 + '@oxfmt/binding-linux-arm64-gnu': 0.57.0 + '@oxfmt/binding-linux-arm64-musl': 0.57.0 + '@oxfmt/binding-linux-ppc64-gnu': 0.57.0 + '@oxfmt/binding-linux-riscv64-gnu': 0.57.0 + '@oxfmt/binding-linux-riscv64-musl': 0.57.0 + '@oxfmt/binding-linux-s390x-gnu': 0.57.0 + '@oxfmt/binding-linux-x64-gnu': 0.57.0 + '@oxfmt/binding-linux-x64-musl': 0.57.0 + '@oxfmt/binding-openharmony-arm64': 0.57.0 + '@oxfmt/binding-win32-arm64-msvc': 0.57.0 + '@oxfmt/binding-win32-ia32-msvc': 0.57.0 + '@oxfmt/binding-win32-x64-msvc': 0.57.0 + vite-plus: 0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) - oxfmt@0.55.0(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)): + oxfmt@0.57.0(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)): dependencies: tinypool: 2.1.0 optionalDependencies: - '@oxfmt/binding-android-arm-eabi': 0.55.0 - '@oxfmt/binding-android-arm64': 0.55.0 - '@oxfmt/binding-darwin-arm64': 0.55.0 - '@oxfmt/binding-darwin-x64': 0.55.0 - '@oxfmt/binding-freebsd-x64': 0.55.0 - '@oxfmt/binding-linux-arm-gnueabihf': 0.55.0 - '@oxfmt/binding-linux-arm-musleabihf': 0.55.0 - '@oxfmt/binding-linux-arm64-gnu': 0.55.0 - '@oxfmt/binding-linux-arm64-musl': 0.55.0 - '@oxfmt/binding-linux-ppc64-gnu': 0.55.0 - '@oxfmt/binding-linux-riscv64-gnu': 0.55.0 - '@oxfmt/binding-linux-riscv64-musl': 0.55.0 - '@oxfmt/binding-linux-s390x-gnu': 0.55.0 - '@oxfmt/binding-linux-x64-gnu': 0.55.0 - '@oxfmt/binding-linux-x64-musl': 0.55.0 - '@oxfmt/binding-openharmony-arm64': 0.55.0 - '@oxfmt/binding-win32-arm64-msvc': 0.55.0 - '@oxfmt/binding-win32-ia32-msvc': 0.55.0 - '@oxfmt/binding-win32-x64-msvc': 0.55.0 - vite-plus: 0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) + '@oxfmt/binding-android-arm-eabi': 0.57.0 + '@oxfmt/binding-android-arm64': 0.57.0 + '@oxfmt/binding-darwin-arm64': 0.57.0 + '@oxfmt/binding-darwin-x64': 0.57.0 + '@oxfmt/binding-freebsd-x64': 0.57.0 + '@oxfmt/binding-linux-arm-gnueabihf': 0.57.0 + '@oxfmt/binding-linux-arm-musleabihf': 0.57.0 + '@oxfmt/binding-linux-arm64-gnu': 0.57.0 + '@oxfmt/binding-linux-arm64-musl': 0.57.0 + '@oxfmt/binding-linux-ppc64-gnu': 0.57.0 + '@oxfmt/binding-linux-riscv64-gnu': 0.57.0 + '@oxfmt/binding-linux-riscv64-musl': 0.57.0 + '@oxfmt/binding-linux-s390x-gnu': 0.57.0 + '@oxfmt/binding-linux-x64-gnu': 0.57.0 + '@oxfmt/binding-linux-x64-musl': 0.57.0 + '@oxfmt/binding-openharmony-arm64': 0.57.0 + '@oxfmt/binding-win32-arm64-msvc': 0.57.0 + '@oxfmt/binding-win32-ia32-msvc': 0.57.0 + '@oxfmt/binding-win32-x64-msvc': 0.57.0 + vite-plus: 0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) - oxfmt@0.55.0(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)): - dependencies: - tinypool: 2.1.0 + oxlint-tsgolint@0.24.0: optionalDependencies: - '@oxfmt/binding-android-arm-eabi': 0.55.0 - '@oxfmt/binding-android-arm64': 0.55.0 - '@oxfmt/binding-darwin-arm64': 0.55.0 - '@oxfmt/binding-darwin-x64': 0.55.0 - '@oxfmt/binding-freebsd-x64': 0.55.0 - '@oxfmt/binding-linux-arm-gnueabihf': 0.55.0 - '@oxfmt/binding-linux-arm-musleabihf': 0.55.0 - '@oxfmt/binding-linux-arm64-gnu': 0.55.0 - '@oxfmt/binding-linux-arm64-musl': 0.55.0 - '@oxfmt/binding-linux-ppc64-gnu': 0.55.0 - '@oxfmt/binding-linux-riscv64-gnu': 0.55.0 - '@oxfmt/binding-linux-riscv64-musl': 0.55.0 - '@oxfmt/binding-linux-s390x-gnu': 0.55.0 - '@oxfmt/binding-linux-x64-gnu': 0.55.0 - '@oxfmt/binding-linux-x64-musl': 0.55.0 - '@oxfmt/binding-openharmony-arm64': 0.55.0 - '@oxfmt/binding-win32-arm64-msvc': 0.55.0 - '@oxfmt/binding-win32-ia32-msvc': 0.55.0 - '@oxfmt/binding-win32-x64-msvc': 0.55.0 - vite-plus: 0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) + '@oxlint-tsgolint/darwin-arm64': 0.24.0 + '@oxlint-tsgolint/darwin-x64': 0.24.0 + '@oxlint-tsgolint/linux-arm64': 0.24.0 + '@oxlint-tsgolint/linux-x64': 0.24.0 + '@oxlint-tsgolint/win32-arm64': 0.24.0 + '@oxlint-tsgolint/win32-x64': 0.24.0 - oxlint-tsgolint@0.23.0: + oxlint@1.72.0(oxlint-tsgolint@0.24.0)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)): optionalDependencies: - '@oxlint-tsgolint/darwin-arm64': 0.23.0 - '@oxlint-tsgolint/darwin-x64': 0.23.0 - '@oxlint-tsgolint/linux-arm64': 0.23.0 - '@oxlint-tsgolint/linux-x64': 0.23.0 - '@oxlint-tsgolint/win32-arm64': 0.23.0 - '@oxlint-tsgolint/win32-x64': 0.23.0 + '@oxlint/binding-android-arm-eabi': 1.72.0 + '@oxlint/binding-android-arm64': 1.72.0 + '@oxlint/binding-darwin-arm64': 1.72.0 + '@oxlint/binding-darwin-x64': 1.72.0 + '@oxlint/binding-freebsd-x64': 1.72.0 + '@oxlint/binding-linux-arm-gnueabihf': 1.72.0 + '@oxlint/binding-linux-arm-musleabihf': 1.72.0 + '@oxlint/binding-linux-arm64-gnu': 1.72.0 + '@oxlint/binding-linux-arm64-musl': 1.72.0 + '@oxlint/binding-linux-ppc64-gnu': 1.72.0 + '@oxlint/binding-linux-riscv64-gnu': 1.72.0 + '@oxlint/binding-linux-riscv64-musl': 1.72.0 + '@oxlint/binding-linux-s390x-gnu': 1.72.0 + '@oxlint/binding-linux-x64-gnu': 1.72.0 + '@oxlint/binding-linux-x64-musl': 1.72.0 + '@oxlint/binding-openharmony-arm64': 1.72.0 + '@oxlint/binding-win32-arm64-msvc': 1.72.0 + '@oxlint/binding-win32-ia32-msvc': 1.72.0 + '@oxlint/binding-win32-x64-msvc': 1.72.0 + oxlint-tsgolint: 0.24.0 + vite-plus: 0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) - oxlint@1.70.0(oxlint-tsgolint@0.23.0)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)): + oxlint@1.72.0(oxlint-tsgolint@0.24.0)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)): optionalDependencies: - '@oxlint/binding-android-arm-eabi': 1.70.0 - '@oxlint/binding-android-arm64': 1.70.0 - '@oxlint/binding-darwin-arm64': 1.70.0 - '@oxlint/binding-darwin-x64': 1.70.0 - '@oxlint/binding-freebsd-x64': 1.70.0 - '@oxlint/binding-linux-arm-gnueabihf': 1.70.0 - '@oxlint/binding-linux-arm-musleabihf': 1.70.0 - '@oxlint/binding-linux-arm64-gnu': 1.70.0 - '@oxlint/binding-linux-arm64-musl': 1.70.0 - '@oxlint/binding-linux-ppc64-gnu': 1.70.0 - '@oxlint/binding-linux-riscv64-gnu': 1.70.0 - '@oxlint/binding-linux-riscv64-musl': 1.70.0 - '@oxlint/binding-linux-s390x-gnu': 1.70.0 - '@oxlint/binding-linux-x64-gnu': 1.70.0 - '@oxlint/binding-linux-x64-musl': 1.70.0 - '@oxlint/binding-openharmony-arm64': 1.70.0 - '@oxlint/binding-win32-arm64-msvc': 1.70.0 - '@oxlint/binding-win32-ia32-msvc': 1.70.0 - '@oxlint/binding-win32-x64-msvc': 1.70.0 - oxlint-tsgolint: 0.23.0 - vite-plus: 0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) - - oxlint@1.70.0(oxlint-tsgolint@0.23.0)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)): - optionalDependencies: - '@oxlint/binding-android-arm-eabi': 1.70.0 - '@oxlint/binding-android-arm64': 1.70.0 - '@oxlint/binding-darwin-arm64': 1.70.0 - '@oxlint/binding-darwin-x64': 1.70.0 - '@oxlint/binding-freebsd-x64': 1.70.0 - '@oxlint/binding-linux-arm-gnueabihf': 1.70.0 - '@oxlint/binding-linux-arm-musleabihf': 1.70.0 - '@oxlint/binding-linux-arm64-gnu': 1.70.0 - '@oxlint/binding-linux-arm64-musl': 1.70.0 - '@oxlint/binding-linux-ppc64-gnu': 1.70.0 - '@oxlint/binding-linux-riscv64-gnu': 1.70.0 - '@oxlint/binding-linux-riscv64-musl': 1.70.0 - '@oxlint/binding-linux-s390x-gnu': 1.70.0 - '@oxlint/binding-linux-x64-gnu': 1.70.0 - '@oxlint/binding-linux-x64-musl': 1.70.0 - '@oxlint/binding-openharmony-arm64': 1.70.0 - '@oxlint/binding-win32-arm64-msvc': 1.70.0 - '@oxlint/binding-win32-ia32-msvc': 1.70.0 - '@oxlint/binding-win32-x64-msvc': 1.70.0 - oxlint-tsgolint: 0.23.0 - vite-plus: 0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) - - oxlint@1.70.0(oxlint-tsgolint@0.23.0)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)): - optionalDependencies: - '@oxlint/binding-android-arm-eabi': 1.70.0 - '@oxlint/binding-android-arm64': 1.70.0 - '@oxlint/binding-darwin-arm64': 1.70.0 - '@oxlint/binding-darwin-x64': 1.70.0 - '@oxlint/binding-freebsd-x64': 1.70.0 - '@oxlint/binding-linux-arm-gnueabihf': 1.70.0 - '@oxlint/binding-linux-arm-musleabihf': 1.70.0 - '@oxlint/binding-linux-arm64-gnu': 1.70.0 - '@oxlint/binding-linux-arm64-musl': 1.70.0 - '@oxlint/binding-linux-ppc64-gnu': 1.70.0 - '@oxlint/binding-linux-riscv64-gnu': 1.70.0 - '@oxlint/binding-linux-riscv64-musl': 1.70.0 - '@oxlint/binding-linux-s390x-gnu': 1.70.0 - '@oxlint/binding-linux-x64-gnu': 1.70.0 - '@oxlint/binding-linux-x64-musl': 1.70.0 - '@oxlint/binding-openharmony-arm64': 1.70.0 - '@oxlint/binding-win32-arm64-msvc': 1.70.0 - '@oxlint/binding-win32-ia32-msvc': 1.70.0 - '@oxlint/binding-win32-x64-msvc': 1.70.0 - oxlint-tsgolint: 0.23.0 - vite-plus: 0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) + '@oxlint/binding-android-arm-eabi': 1.72.0 + '@oxlint/binding-android-arm64': 1.72.0 + '@oxlint/binding-darwin-arm64': 1.72.0 + '@oxlint/binding-darwin-x64': 1.72.0 + '@oxlint/binding-freebsd-x64': 1.72.0 + '@oxlint/binding-linux-arm-gnueabihf': 1.72.0 + '@oxlint/binding-linux-arm-musleabihf': 1.72.0 + '@oxlint/binding-linux-arm64-gnu': 1.72.0 + '@oxlint/binding-linux-arm64-musl': 1.72.0 + '@oxlint/binding-linux-ppc64-gnu': 1.72.0 + '@oxlint/binding-linux-riscv64-gnu': 1.72.0 + '@oxlint/binding-linux-riscv64-musl': 1.72.0 + '@oxlint/binding-linux-s390x-gnu': 1.72.0 + '@oxlint/binding-linux-x64-gnu': 1.72.0 + '@oxlint/binding-linux-x64-musl': 1.72.0 + '@oxlint/binding-openharmony-arm64': 1.72.0 + '@oxlint/binding-win32-arm64-msvc': 1.72.0 + '@oxlint/binding-win32-ia32-msvc': 1.72.0 + '@oxlint/binding-win32-x64-msvc': 1.72.0 + oxlint-tsgolint: 0.24.0 + vite-plus: 0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) p-limit@3.1.0: dependencies: @@ -18921,7 +18866,7 @@ snapshots: es-errors: 1.3.0 internal-slot: 1.1.0 - storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)): + storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)): dependencies: '@storybook/global': 5.0.0 '@storybook/icons': 2.1.0(react@19.2.7) @@ -18940,14 +18885,14 @@ snapshots: ws: 8.21.0 optionalDependencies: '@types/react': 19.2.17 - vite-plus: 0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) + vite-plus: 0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) transitivePeerDependencies: - '@testing-library/dom' - bufferutil - react - utf-8-validate - storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)): + storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)): dependencies: '@storybook/global': 5.0.0 '@storybook/icons': 2.1.0(react@19.2.7) @@ -18966,7 +18911,7 @@ snapshots: ws: 8.21.0 optionalDependencies: '@types/react': 19.2.17 - vite-plus: 0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) + vite-plus: 0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) transitivePeerDependencies: - '@testing-library/dom' - bufferutil @@ -19380,14 +19325,14 @@ snapshots: picomatch: 4.0.4 webpack-virtual-modules: 0.6.2 - unplugin@3.2.0(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1): + unplugin@3.2.0(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1): dependencies: '@jridgewell/remapping': 2.3.5 picomatch: 4.0.4 webpack-virtual-modules: 0.6.2 optionalDependencies: esbuild: 0.28.1 - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' update-browserslist-db@1.2.3(browserslist@4.28.4): dependencies: @@ -19474,22 +19419,22 @@ snapshots: '@types/unist': 3.0.3 vfile-message: 4.0.3 - vinext@0.1.8(@mdx-js/rollup@3.1.1)(@vitejs/plugin-react@6.0.3(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(@vitejs/plugin-rsc@0.5.27(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(react-dom@19.2.7(react@19.2.7))(react-server-dom-webpack@19.2.7(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react@19.2.7))(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(next@16.2.9(@babel/core@7.29.7)(@playwright/test@1.61.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react-dom@19.2.7(react@19.2.7))(react-server-dom-webpack@19.2.7(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react@19.2.7): + vinext@0.1.8(@mdx-js/rollup@3.1.1)(@vitejs/plugin-react@6.0.3(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(@vitejs/plugin-rsc@0.5.27(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(react-dom@19.2.7(react@19.2.7))(react-server-dom-webpack@19.2.7(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react@19.2.7))(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(next@16.2.9(@babel/core@7.29.7)(@playwright/test@1.61.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react-dom@19.2.7(react@19.2.7))(react-server-dom-webpack@19.2.7(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react@19.2.7): dependencies: '@unpic/react': 1.0.2(next@16.2.9(@babel/core@7.29.7)(@playwright/test@1.61.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react-dom@19.2.7(react@19.2.7))(react@19.2.7) '@vercel/og': 0.8.6 - '@vitejs/plugin-react': 6.0.3(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + '@vitejs/plugin-react': 6.0.3(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) image-size: 2.0.2 ipaddr.js: 2.4.0 magic-string: 0.30.21 react: 19.2.7 react-dom: 19.2.7(react@19.2.7) - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' vite-plugin-commonjs: 0.10.4 web-vitals: 4.2.4 optionalDependencies: '@mdx-js/rollup': 3.1.1 - '@vitejs/plugin-rsc': 0.5.27(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(react-dom@19.2.7(react@19.2.7))(react-server-dom-webpack@19.2.7(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react@19.2.7) + '@vitejs/plugin-rsc': 0.5.27(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(react-dom@19.2.7(react@19.2.7))(react-server-dom-webpack@19.2.7(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(react@19.2.7) react-server-dom-webpack: 19.2.7(react-dom@19.2.7(react@19.2.7))(react@19.2.7) transitivePeerDependencies: - next @@ -19507,9 +19452,9 @@ snapshots: fast-glob: 3.3.3 magic-string: 0.30.21 - vite-plugin-inspect@12.0.0-beta.3(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(typescript@6.0.3): + vite-plugin-inspect@12.0.0-beta.3(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(typescript@6.0.3): dependencies: - '@vitejs/devtools-kit': 0.3.3(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(typescript@6.0.3) + '@vitejs/devtools-kit': 0.3.3(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(typescript@6.0.3) ansis: 4.3.1 error-stack-parser-es: 1.0.5 obug: 2.1.3 @@ -19518,7 +19463,7 @@ snapshots: perfect-debounce: 2.1.0 sirv: 3.0.2 unplugin-utils: 0.3.1 - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' transitivePeerDependencies: - '@farmfe/core' - '@modelcontextprotocol/sdk' @@ -19534,55 +19479,53 @@ snapshots: - utf-8-validate - webpack - vite-plugin-storybook-nextjs@3.3.0(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(next@16.2.9(@babel/core@7.29.7)(@playwright/test@1.61.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(supports-color@10.2.2)(typescript@6.0.3): + vite-plugin-storybook-nextjs@3.3.0(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(next@16.2.9(@babel/core@7.29.7)(@playwright/test@1.61.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7))(storybook@10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)))(supports-color@10.2.2)(typescript@6.0.3): dependencies: '@next/env': 16.0.0 image-size: 2.0.2 magic-string: 0.30.21 module-alias: 2.3.4 next: 16.2.9(@babel/core@7.29.7)(@playwright/test@1.61.1)(react-dom@19.2.7(react@19.2.7))(react@19.2.7) - storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + storybook: 10.4.6(@testing-library/dom@10.4.1)(@types/react@19.2.17)(react@19.2.7)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) ts-dedent: 2.3.0 - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' - vite-tsconfig-paths: 5.1.4(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(supports-color@10.2.2)(typescript@6.0.3) + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite-tsconfig-paths: 5.1.4(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(supports-color@10.2.2)(typescript@6.0.3) transitivePeerDependencies: - supports-color - typescript - vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0): + vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0): dependencies: - '@oxc-project/types': 0.136.0 + '@oxc-project/types': 0.138.0 '@oxlint/plugins': 1.68.0 - '@vitest/browser': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) - '@vitest/browser-preview': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) + '@vitest/browser': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) + '@vitest/browser-preview': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) '@vitest/expect': 4.1.9 - '@vitest/mocker': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + '@vitest/mocker': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) '@vitest/pretty-format': 4.1.9 '@vitest/runner': 4.1.9 '@vitest/snapshot': 4.1.9 '@vitest/spy': 4.1.9 '@vitest/utils': 4.1.9 - '@voidzero-dev/vite-plus-core': 0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) - oxfmt: 0.55.0(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) - oxlint: 1.70.0(oxlint-tsgolint@0.23.0)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) - oxlint-tsgolint: 0.23.0 - vitest: 4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) + '@voidzero-dev/vite-plus-core': 0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) + oxfmt: 0.57.0(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + oxlint: 1.72.0(oxlint-tsgolint@0.24.0)(vite-plus@0.2.3(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + oxlint-tsgolint: 0.24.0 + vitest: 4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) optionalDependencies: - '@vitest/browser-playwright': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(playwright@1.61.1)(vitest@4.1.9) - '@voidzero-dev/vite-plus-darwin-arm64': 0.2.1 - '@voidzero-dev/vite-plus-darwin-x64': 0.2.1 - '@voidzero-dev/vite-plus-linux-arm64-gnu': 0.2.1 - '@voidzero-dev/vite-plus-linux-arm64-musl': 0.2.1 - '@voidzero-dev/vite-plus-linux-x64-gnu': 0.2.1 - '@voidzero-dev/vite-plus-linux-x64-musl': 0.2.1 - '@voidzero-dev/vite-plus-win32-arm64-msvc': 0.2.1 - '@voidzero-dev/vite-plus-win32-x64-msvc': 0.2.1 + '@vitest/browser-playwright': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(playwright@1.61.1)(vitest@4.1.9) + '@voidzero-dev/vite-plus-darwin-arm64': 0.2.3 + '@voidzero-dev/vite-plus-darwin-x64': 0.2.3 + '@voidzero-dev/vite-plus-linux-arm64-gnu': 0.2.3 + '@voidzero-dev/vite-plus-linux-arm64-musl': 0.2.3 + '@voidzero-dev/vite-plus-linux-x64-gnu': 0.2.3 + '@voidzero-dev/vite-plus-linux-x64-musl': 0.2.3 + '@voidzero-dev/vite-plus-win32-arm64-msvc': 0.2.3 + '@voidzero-dev/vite-plus-win32-x64-msvc': 0.2.3 transitivePeerDependencies: - '@arethetypeswrong/core' - '@edge-runtime/vm' - '@opentelemetry/api' - - '@tsdown/css' - - '@tsdown/exe' - '@types/node' - '@vitejs/devtools' - '@vitest/coverage-istanbul' @@ -19610,40 +19553,38 @@ snapshots: - vite - yaml - vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0): + vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0): dependencies: - '@oxc-project/types': 0.136.0 + '@oxc-project/types': 0.138.0 '@oxlint/plugins': 1.68.0 - '@vitest/browser': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) - '@vitest/browser-preview': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) + '@vitest/browser': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) + '@vitest/browser-preview': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) '@vitest/expect': 4.1.9 - '@vitest/mocker': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + '@vitest/mocker': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) '@vitest/pretty-format': 4.1.9 '@vitest/runner': 4.1.9 '@vitest/snapshot': 4.1.9 '@vitest/spy': 4.1.9 '@vitest/utils': 4.1.9 - '@voidzero-dev/vite-plus-core': 0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) - oxfmt: 0.55.0(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) - oxlint: 1.70.0(oxlint-tsgolint@0.23.0)(vite-plus@0.2.1(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) - oxlint-tsgolint: 0.23.0 - vitest: 4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) + '@voidzero-dev/vite-plus-core': 0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) + oxfmt: 0.57.0(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + oxlint: 1.72.0(oxlint-tsgolint@0.24.0)(vite-plus@0.2.3(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + oxlint-tsgolint: 0.24.0 + vitest: 4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) optionalDependencies: - '@vitest/browser-playwright': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(playwright@1.61.1)(vitest@4.1.9) - '@voidzero-dev/vite-plus-darwin-arm64': 0.2.1 - '@voidzero-dev/vite-plus-darwin-x64': 0.2.1 - '@voidzero-dev/vite-plus-linux-arm64-gnu': 0.2.1 - '@voidzero-dev/vite-plus-linux-arm64-musl': 0.2.1 - '@voidzero-dev/vite-plus-linux-x64-gnu': 0.2.1 - '@voidzero-dev/vite-plus-linux-x64-musl': 0.2.1 - '@voidzero-dev/vite-plus-win32-arm64-msvc': 0.2.1 - '@voidzero-dev/vite-plus-win32-x64-msvc': 0.2.1 + '@vitest/browser-playwright': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(playwright@1.61.1)(vitest@4.1.9) + '@voidzero-dev/vite-plus-darwin-arm64': 0.2.3 + '@voidzero-dev/vite-plus-darwin-x64': 0.2.3 + '@voidzero-dev/vite-plus-linux-arm64-gnu': 0.2.3 + '@voidzero-dev/vite-plus-linux-arm64-musl': 0.2.3 + '@voidzero-dev/vite-plus-linux-x64-gnu': 0.2.3 + '@voidzero-dev/vite-plus-linux-x64-musl': 0.2.3 + '@voidzero-dev/vite-plus-win32-arm64-msvc': 0.2.3 + '@voidzero-dev/vite-plus-win32-x64-msvc': 0.2.3 transitivePeerDependencies: - '@arethetypeswrong/core' - '@edge-runtime/vm' - '@opentelemetry/api' - - '@tsdown/css' - - '@tsdown/exe' - '@types/node' - '@vitejs/devtools' - '@vitest/coverage-istanbul' @@ -19671,87 +19612,26 @@ snapshots: - vite - yaml - vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0): - dependencies: - '@oxc-project/types': 0.136.0 - '@oxlint/plugins': 1.68.0 - '@vitest/browser': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) - '@vitest/browser-preview': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) - '@vitest/expect': 4.1.9 - '@vitest/mocker': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) - '@vitest/pretty-format': 4.1.9 - '@vitest/runner': 4.1.9 - '@vitest/snapshot': 4.1.9 - '@vitest/spy': 4.1.9 - '@vitest/utils': 4.1.9 - '@voidzero-dev/vite-plus-core': 0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0) - oxfmt: 0.55.0(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) - oxlint: 1.70.0(oxlint-tsgolint@0.23.0)(vite-plus@0.2.1(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(esbuild@0.28.1)(happy-dom@20.10.6)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) - oxlint-tsgolint: 0.23.0 - vitest: 4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) - optionalDependencies: - '@vitest/browser-playwright': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(playwright@1.61.1)(vitest@4.1.9) - '@voidzero-dev/vite-plus-darwin-arm64': 0.2.1 - '@voidzero-dev/vite-plus-darwin-x64': 0.2.1 - '@voidzero-dev/vite-plus-linux-arm64-gnu': 0.2.1 - '@voidzero-dev/vite-plus-linux-arm64-musl': 0.2.1 - '@voidzero-dev/vite-plus-linux-x64-gnu': 0.2.1 - '@voidzero-dev/vite-plus-linux-x64-musl': 0.2.1 - '@voidzero-dev/vite-plus-win32-arm64-msvc': 0.2.1 - '@voidzero-dev/vite-plus-win32-x64-msvc': 0.2.1 - transitivePeerDependencies: - - '@arethetypeswrong/core' - - '@edge-runtime/vm' - - '@opentelemetry/api' - - '@tsdown/css' - - '@tsdown/exe' - - '@types/node' - - '@vitejs/devtools' - - '@vitest/coverage-istanbul' - - '@vitest/coverage-v8' - - '@vitest/ui' - - bufferutil - - esbuild - - happy-dom - - jiti - - jsdom - - less - - msw - - publint - - sass - - sass-embedded - - stylus - - sugarss - - svelte - - terser - - tsx - - typescript - - unplugin-unused - - unrun - - utf-8-validate - - vite - - yaml - - vite-tsconfig-paths@5.1.4(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(supports-color@10.2.2)(typescript@6.0.3): + vite-tsconfig-paths@5.1.4(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(supports-color@10.2.2)(typescript@6.0.3): dependencies: debug: 4.4.3(supports-color@10.2.2) globrex: 0.1.2 tsconfck: 3.1.6(typescript@6.0.3) optionalDependencies: - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' transitivePeerDependencies: - supports-color - typescript - vitefu@1.1.3(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)): + vitefu@1.1.3(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)): optionalDependencies: - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' vitest-browser-react@2.2.0(@types/react-dom@19.2.3(@types/react@19.2.17))(@types/react@19.2.17)(react-dom@19.2.7(react@19.2.7))(react@19.2.7)(vitest@4.1.9): dependencies: react: 19.2.7 react-dom: 19.2.7(react@19.2.7) - vitest: 4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) + vitest: 4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) optionalDependencies: '@types/react': 19.2.17 '@types/react-dom': 19.2.3(@types/react@19.2.17) @@ -19760,12 +19640,12 @@ snapshots: dependencies: cssfontparser: 1.2.1 moo-color: 1.0.3 - vitest: 4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) + vitest: 4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6) - vitest@4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6): + vitest@4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6): dependencies: '@vitest/expect': 4.1.9 - '@vitest/mocker': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + '@vitest/mocker': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) '@vitest/pretty-format': 4.1.9 '@vitest/runner': 4.1.9 '@vitest/snapshot': 4.1.9 @@ -19782,21 +19662,21 @@ snapshots: tinyexec: 1.2.4 tinyglobby: 0.2.17 tinyrainbow: 3.1.0 - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' why-is-node-running: 2.3.0 optionalDependencies: '@types/node': 25.9.4 - '@vitest/browser-playwright': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(playwright@1.61.1)(vitest@4.1.9) - '@vitest/browser-preview': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) + '@vitest/browser-playwright': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(playwright@1.61.1)(vitest@4.1.9) + '@vitest/browser-preview': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@25.9.4)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) '@vitest/coverage-v8': 4.1.9(@vitest/browser@4.1.9)(vitest@4.1.9) happy-dom: 20.10.6 transitivePeerDependencies: - msw - vitest@4.1.9(@types/node@25.9.4)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6): + vitest@4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6): dependencies: '@vitest/expect': 4.1.9 - '@vitest/mocker': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) + '@vitest/mocker': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) '@vitest/pretty-format': 4.1.9 '@vitest/runner': 4.1.9 '@vitest/snapshot': 4.1.9 @@ -19813,43 +19693,12 @@ snapshots: tinyexec: 1.2.4 tinyglobby: 0.2.17 tinyrainbow: 3.1.0 - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' - why-is-node-running: 2.3.0 - optionalDependencies: - '@types/node': 25.9.4 - '@vitest/browser-playwright': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(playwright@1.61.1)(vitest@4.1.9) - '@vitest/browser-preview': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) - '@vitest/coverage-v8': 4.1.9(@vitest/browser@4.1.9)(vitest@4.1.9) - happy-dom: 20.10.6 - transitivePeerDependencies: - - msw - - vitest@4.1.9(@types/node@26.0.1)(@vitest/browser-playwright@4.1.9)(@vitest/browser-preview@4.1.9)(@vitest/coverage-v8@4.1.9)(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(happy-dom@20.10.6): - dependencies: - '@vitest/expect': 4.1.9 - '@vitest/mocker': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)) - '@vitest/pretty-format': 4.1.9 - '@vitest/runner': 4.1.9 - '@vitest/snapshot': 4.1.9 - '@vitest/spy': 4.1.9 - '@vitest/utils': 4.1.9 - es-module-lexer: 2.1.0 - expect-type: 1.4.0 - magic-string: 0.30.21 - obug: 2.1.3 - pathe: 2.0.3 - picomatch: 4.0.4 - std-env: 4.1.0 - tinybench: 2.9.0 - tinyexec: 1.2.4 - tinyglobby: 0.2.17 - tinyrainbow: 3.1.0 - vite: '@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' + vite: '@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0)' why-is-node-running: 2.3.0 optionalDependencies: '@types/node': 26.0.1 - '@vitest/browser-playwright': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(playwright@1.61.1)(vitest@4.1.9) - '@vitest/browser-preview': 4.1.9(@voidzero-dev/vite-plus-core@0.2.1(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) + '@vitest/browser-playwright': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(playwright@1.61.1)(vitest@4.1.9) + '@vitest/browser-preview': 4.1.9(@voidzero-dev/vite-plus-core@0.2.3(@types/node@26.0.1)(esbuild@0.28.1)(jiti@2.7.0)(tsx@4.22.4)(typescript@6.0.3)(yaml@2.9.0))(vitest@4.1.9) '@vitest/coverage-v8': 4.1.9(@vitest/browser@4.1.9)(vitest@4.1.9) happy-dom: 20.10.6 transitivePeerDependencies: @@ -20137,7 +19986,7 @@ time: '@vitest/browser-playwright@4.1.9': '2026-06-15T07:21:50.860Z' '@vitest/browser@4.1.9': '2026-06-15T07:21:58.373Z' '@vitest/coverage-v8@4.1.9': '2026-06-15T07:21:14.145Z' - '@voidzero-dev/vite-plus-core@0.2.1': '2026-06-18T05:33:26.237Z' + '@voidzero-dev/vite-plus-core@0.2.3': '2026-07-07T12:11:35.431Z' abcjs@6.6.3: '2026-04-24T17:38:01.079Z' agentation@3.0.2: '2026-03-25T16:24:19.682Z' ahooks@3.9.7: '2026-03-23T15:49:13.605Z' @@ -20251,7 +20100,7 @@ time: uuid@14.0.1: '2026-06-20T11:56:02.499Z' vinext@0.1.8: '2026-06-23T12:54:50.783Z' vite-plugin-inspect@12.0.0-beta.3: '2026-05-29T06:16:55.694Z' - vite-plus@0.2.1: '2026-06-18T05:33:32.399Z' + vite-plus@0.2.3: '2026-07-07T12:11:40.691Z' vitest-browser-react@2.2.0: '2026-04-05T06:56:34.635Z' vitest-canvas-mock@1.1.4: '2026-03-24T14:42:39.285Z' vitest@4.1.9: '2026-06-15T07:23:00.326Z' diff --git a/pnpm-workspace.yaml b/pnpm-workspace.yaml index 62c36b70183..df50f9c4228 100644 --- a/pnpm-workspace.yaml +++ b/pnpm-workspace.yaml @@ -45,7 +45,7 @@ overrides: solid-js: 1.9.13 string-width: ~8.2.1 tar@<=7.5.15: ^7.5.16 - vite: npm:@voidzero-dev/vite-plus-core@0.2.1 + vite: npm:@voidzero-dev/vite-plus-core@0.2.3 vitest: 4.1.9 ws@>=8.0.0 <8.20.1: ^8.21.0 yaml@>=2.0.0 <2.8.3: 2.9.0 @@ -248,9 +248,9 @@ catalog: use-context-selector: 2.0.0 uuid: 14.0.1 vinext: 0.1.8 - vite: npm:@voidzero-dev/vite-plus-core@0.2.1 + vite: npm:@voidzero-dev/vite-plus-core@0.2.3 vite-plugin-inspect: 12.0.0-beta.3 - vite-plus: 0.2.1 + vite-plus: 0.2.3 vitest: 4.1.9 vitest-browser-react: 2.2.0 vitest-canvas-mock: 1.1.4 From 41fb47995726611bead6abb0dd31db01effd6aa1 Mon Sep 17 00:00:00 2001 From: yyh <92089059+lyzno1@users.noreply.github.com> Date: Wed, 8 Jul 2026 14:16:59 +0800 Subject: [PATCH 40/70] fix(web): cache server console context (#38535) --- web/service/server.ts | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/web/service/server.ts b/web/service/server.ts index 5e031b2f1e8..16d412e28fe 100644 --- a/web/service/server.ts +++ b/web/service/server.ts @@ -5,6 +5,7 @@ import type { JsonifiedClient } from '@orpc/openapi-client' import { createORPCClient, onError } from '@orpc/client' import { OpenAPILink } from '@orpc/openapi-client/fetch' import { createTanstackQueryUtils } from '@orpc/tanstack-query' +import { cache } from 'react' import { API_PREFIX, CSRF_COOKIE_NAME, @@ -93,7 +94,7 @@ function createServerConsoleOpenAPILink(contract: AnyContractRouter): ServerCons }) } -export const getServerConsoleClientContext = async (): Promise => { +export const getServerConsoleClientContext = cache(async (): Promise => { const { cookies, headers } = await import('@/next/headers') const requestHeaders = await headers() const cookieStore = await cookies() @@ -102,7 +103,7 @@ export const getServerConsoleClientContext = async (): Promise createServerConsoleRequestHeaders(await getServerConsoleClientContext()) From 10cf9be3c760c374fdfe497848db986791c6fb0b Mon Sep 17 00:00:00 2001 From: yyh <92089059+lyzno1@users.noreply.github.com> Date: Wed, 8 Jul 2026 14:17:14 +0800 Subject: [PATCH 41/70] test(web): align auth e2e with console home (#38538) --- e2e/features/auth/session-refresh.feature | 9 ++++ e2e/features/auth/sign-in.feature | 4 +- .../smoke/authenticated-entry.feature | 8 +-- .../smoke/unauthenticated-entry.feature | 6 +-- .../auth/session-refresh.steps.ts | 43 +++++++++++++++ .../step-definitions/auth/sign-in.steps.ts | 9 +--- .../common/navigation.steps.ts | 9 ++++ e2e/fixtures/auth.ts | 8 +-- e2e/support/apps.ts | 9 +++- e2e/support/home.ts | 12 +++++ web/app/auth/refresh/__tests__/route.spec.ts | 53 ++++++++++++++++++- web/app/signin/__tests__/normal-form.spec.tsx | 29 ++++++++-- 12 files changed, 171 insertions(+), 28 deletions(-) create mode 100644 e2e/features/auth/session-refresh.feature create mode 100644 e2e/features/step-definitions/auth/session-refresh.steps.ts create mode 100644 e2e/support/home.ts diff --git a/e2e/features/auth/session-refresh.feature b/e2e/features/auth/session-refresh.feature new file mode 100644 index 00000000000..f565a6aff69 --- /dev/null +++ b/e2e/features/auth/session-refresh.feature @@ -0,0 +1,9 @@ +@auth @core @authenticated +Feature: Console session refresh + + Scenario: Refresh the console session during server-side navigation + Given I am signed in as the default E2E admin + And my console session requires token refresh + When I open the default console entry after the access token expires + Then I should be on the console home + And I should not see the "Sign in" button diff --git a/e2e/features/auth/sign-in.feature b/e2e/features/auth/sign-in.feature index a9a1e13626d..ffa8a02ad71 100644 --- a/e2e/features/auth/sign-in.feature +++ b/e2e/features/auth/sign-in.feature @@ -1,8 +1,8 @@ @auth @smoke @core @unauthenticated Feature: Sign in - Scenario: Sign in with valid credentials and reach the apps console + Scenario: Sign in with valid credentials and reach the console home Given I am not signed in When I open the sign-in page And I sign in as the default E2E admin - Then I should be on the apps console + Then I should be on the console home diff --git a/e2e/features/smoke/authenticated-entry.feature b/e2e/features/smoke/authenticated-entry.feature index 53d72bd667c..955279bd247 100644 --- a/e2e/features/smoke/authenticated-entry.feature +++ b/e2e/features/smoke/authenticated-entry.feature @@ -1,7 +1,7 @@ @smoke @authenticated -Feature: Authenticated app console - Scenario: Open the apps console with the shared authenticated state +Feature: Authenticated console home + Scenario: Open the default console entry with the shared authenticated state Given I am signed in as the default E2E admin - When I open the apps console - Then I should stay on the apps console + When I open the default console entry + Then I should be on the console home And I should not see the "Sign in" button diff --git a/e2e/features/smoke/unauthenticated-entry.feature b/e2e/features/smoke/unauthenticated-entry.feature index a2783c1cba2..604ea4970dd 100644 --- a/e2e/features/smoke/unauthenticated-entry.feature +++ b/e2e/features/smoke/unauthenticated-entry.feature @@ -1,7 +1,7 @@ @smoke @unauthenticated -Feature: Unauthenticated app console entry - Scenario: Redirect to the sign-in page when opening the apps console without logging in +Feature: Unauthenticated console home entry + Scenario: Redirect to the sign-in page when opening the default console entry without logging in Given I am not signed in - When I open the apps console + When I open the default console entry Then I should be redirected to the signin page And I should see the "Sign in" button diff --git a/e2e/features/step-definitions/auth/session-refresh.steps.ts b/e2e/features/step-definitions/auth/session-refresh.steps.ts new file mode 100644 index 00000000000..f6468bfaf3d --- /dev/null +++ b/e2e/features/step-definitions/auth/session-refresh.steps.ts @@ -0,0 +1,43 @@ +import type { DifyWorld } from '../../support/world' +import { Given, When } from '@cucumber/cucumber' +import { expect } from '@playwright/test' + +const consoleAccessTokenCookieName = /^(?:__Host-)?access_token$/ +const consoleRefreshTokenCookieName = /^(?:__Host-)?refresh_token$/ + +Given('my console session requires token refresh', async function (this: DifyWorld) { + if (!this.context) + throw new Error('Playwright browser context has not been initialized for this scenario.') + + const cookies = await this.context.cookies() + const hasAccessToken = cookies.some(cookie => consoleAccessTokenCookieName.test(cookie.name)) + const hasRefreshToken = cookies.some(cookie => consoleRefreshTokenCookieName.test(cookie.name)) + + expect(hasAccessToken, 'Expected the authenticated E2E session to include a console access token.').toBe(true) + expect(hasRefreshToken, 'Expected the authenticated E2E session to include a console refresh token.').toBe(true) + + await this.context.clearCookies({ name: consoleAccessTokenCookieName }) + + const remainingCookies = await this.context.cookies() + expect( + remainingCookies.some(cookie => consoleAccessTokenCookieName.test(cookie.name)), + 'Expected the console access token to be removed before opening the default console entry.', + ).toBe(false) + expect( + remainingCookies.some(cookie => consoleRefreshTokenCookieName.test(cookie.name)), + 'Expected the console refresh token to remain available for server-side refresh.', + ).toBe(true) +}) + +When('I open the default console entry after the access token expires', async function (this: DifyWorld) { + const page = this.getPage() + const refreshRequestPromise = page.waitForRequest((request) => { + const url = new URL(request.url()) + return url.pathname.endsWith('/auth/refresh') && url.searchParams.get('redirect_url') === '/' + }) + + await page.goto('/') + + const refreshRequest = await refreshRequestPromise + this.attach(`Session refresh request: ${refreshRequest.url()}`, 'text/plain') +}) diff --git a/e2e/features/step-definitions/auth/sign-in.steps.ts b/e2e/features/step-definitions/auth/sign-in.steps.ts index 469203d8bfd..fcf106048b9 100644 --- a/e2e/features/step-definitions/auth/sign-in.steps.ts +++ b/e2e/features/step-definitions/auth/sign-in.steps.ts @@ -1,10 +1,9 @@ import type { DifyWorld } from '../../support/world' -import { Then, When } from '@cucumber/cucumber' -import { expect } from '@playwright/test' +import { When } from '@cucumber/cucumber' import { adminCredentials } from '../../../fixtures/auth' When('I open the sign-in page', async function (this: DifyWorld) { - await this.getPage().goto('/signin?redirect_url=%2Fapps') + await this.getPage().goto('/signin') }) When('I sign in as the default E2E admin', async function (this: DifyWorld) { @@ -14,7 +13,3 @@ When('I sign in as the default E2E admin', async function (this: DifyWorld) { await page.getByLabel('Password', { exact: true }).fill(adminCredentials.password) await page.getByRole('button', { name: 'Sign in' }).click() }) - -Then('I should be on the apps console', async function (this: DifyWorld) { - await expect(this.getPage()).toHaveURL(/\/apps(?:\?.*)?$/, { timeout: 30_000 }) -}) diff --git a/e2e/features/step-definitions/common/navigation.steps.ts b/e2e/features/step-definitions/common/navigation.steps.ts index a558d96f937..10ee7b789f0 100644 --- a/e2e/features/step-definitions/common/navigation.steps.ts +++ b/e2e/features/step-definitions/common/navigation.steps.ts @@ -2,6 +2,11 @@ import type { DifyWorld } from '../../support/world' import { Then, When } from '@cucumber/cucumber' import { expect } from '@playwright/test' import { waitForAppsConsole } from '../../../support/apps' +import { waitForConsoleHome } from '../../../support/home' + +When('I open the default console entry', async function (this: DifyWorld) { + await this.getPage().goto('/') +}) When('I open the apps console', async function (this: DifyWorld) { await this.getPage().goto('/apps') @@ -15,6 +20,10 @@ Then('I should stay on the apps console', async function (this: DifyWorld) { await waitForAppsConsole(this.getPage()) }) +Then('I should be on the console home', async function (this: DifyWorld) { + await waitForConsoleHome(this.getPage()) +}) + Then('I should be redirected to the signin page', async function (this: DifyWorld) { await expect(this.getPage()).toHaveURL(/\/signin(?:\?.*)?$/) }) diff --git a/e2e/fixtures/auth.ts b/e2e/fixtures/auth.ts index 9039f97483d..62d2880a465 100644 --- a/e2e/fixtures/auth.ts +++ b/e2e/fixtures/auth.ts @@ -3,7 +3,7 @@ import { Buffer } from 'node:buffer' import { mkdir, readFile, writeFile } from 'node:fs/promises' import path from 'node:path' import { fileURLToPath } from 'node:url' -import { waitForAppsConsole } from '../support/apps' +import { waitForConsoleHome } from '../support/home' import { apiURL, defaultBaseURL, defaultLocale } from '../test-env' export type AuthSessionMetadata = { @@ -150,12 +150,12 @@ export const ensureAuthenticatedState = async (browser: Browser, configuredBaseU const { mode, usedInitPassword } = await ensureAdminAccount(context, deadline) await loginAdmin(context, deadline) - console.warn('[e2e] auth bootstrap: verifying apps console') - await page.goto(appURL(baseURL, '/apps'), { + console.warn('[e2e] auth bootstrap: verifying console home') + await page.goto(appURL(baseURL, '/'), { timeout: getRemainingTimeout(deadline), waitUntil: 'domcontentloaded', }) - await waitForAppsConsole(page, getRemainingTimeout(deadline)) + await waitForConsoleHome(page, getRemainingTimeout(deadline)) await context.storageState({ path: authStatePath }) diff --git a/e2e/support/apps.ts b/e2e/support/apps.ts index 3c3af547a35..07d0f15716a 100644 --- a/e2e/support/apps.ts +++ b/e2e/support/apps.ts @@ -1,10 +1,15 @@ import type { Page } from '@playwright/test' import { expect } from '@playwright/test' +const getExpectOptions = (timeout?: number) => + timeout === undefined ? undefined : { timeout } + export const waitForAppsConsole = async (page: Page, timeout?: number) => { - await expect(page).toHaveURL(/\/apps(?:\?.*)?$/, timeout === undefined ? undefined : { timeout }) + const options = getExpectOptions(timeout) + + await expect(page).toHaveURL(/\/apps(?:\?.*)?$/, options) await expect(page.getByRole('heading', { name: 'Studio' })).toBeVisible( - timeout === undefined ? undefined : { timeout }, + options, ) } diff --git a/e2e/support/home.ts b/e2e/support/home.ts new file mode 100644 index 00000000000..83fe9b475c0 --- /dev/null +++ b/e2e/support/home.ts @@ -0,0 +1,12 @@ +import type { Page } from '@playwright/test' +import { expect } from '@playwright/test' + +const getExpectOptions = (timeout?: number) => + timeout === undefined ? undefined : { timeout } + +export const waitForConsoleHome = async (page: Page, timeout?: number) => { + const options = getExpectOptions(timeout) + + await expect.poll(() => new URL(page.url()).pathname, options).toBe('/') + await expect(page.getByRole('link', { name: 'Home' })).toHaveAttribute('aria-current', 'page', options) +} diff --git a/web/app/auth/refresh/__tests__/route.spec.ts b/web/app/auth/refresh/__tests__/route.spec.ts index 82e8d402884..ac1b5e6747f 100644 --- a/web/app/auth/refresh/__tests__/route.spec.ts +++ b/web/app/auth/refresh/__tests__/route.spec.ts @@ -2,6 +2,10 @@ import { beforeEach, describe, expect, it, vi } from 'vitest' +const mocks = vi.hoisted(() => ({ + basePath: '', +})) + vi.mock('@/config', () => ({ API_PREFIX: 'http://localhost:5001/console/api', CSRF_COOKIE_NAME: () => 'csrf_token', @@ -15,7 +19,9 @@ vi.mock('@/config/server', () => ({ })) vi.mock('@/utils/var', () => ({ - basePath: '', + get basePath() { + return mocks.basePath + }, })) const getSetCookieHeaders = (headers: Headers) => { @@ -38,7 +44,9 @@ const createRequest = (url: string, cookie?: string) => ({ describe('auth refresh route', () => { beforeEach(() => { vi.clearAllMocks() + vi.resetModules() vi.unstubAllGlobals() + mocks.basePath = '' }) it('should refresh cookies and redirect back to the requested path', async () => { @@ -131,4 +139,47 @@ describe('auth refresh route', () => { expect(response.status).toBe(303) expect(response.headers.get('location')).toBe('/signin?redirect_url=%2F') }) + + it('should preserve base path when refreshing and redirecting back', async () => { + mocks.basePath = '/console' + const headers = new Headers() + Object.defineProperty(headers, 'getSetCookie', { + value: () => [ + 'access_token=new-access; Path=/console; HttpOnly', + 'refresh_token=new-refresh; Path=/console; HttpOnly', + ], + }) + const fetchMock = vi.fn().mockResolvedValue({ + ok: true, + headers, + } as Response) + vi.stubGlobal('fetch', fetchMock) + const { GET } = await import('../route') + + const response = await GET(createRequest( + 'http://localhost:3000/console/auth/refresh?redirect_url=%2Fconsole%2Fapps%3Fcategory%3Dworkflow', + 'refresh_token=old-refresh', + )) + + expect(response.status).toBe(303) + expect(response.headers.get('location')).toBe('/console/apps?category=workflow') + expect(getSetCookieHeaders(response.headers)).toEqual([ + 'access_token=new-access; Path=/console; HttpOnly', + 'refresh_token=new-refresh; Path=/console; HttpOnly', + ]) + }) + + it('should fall back to the base path home when base path refresh redirects to itself', async () => { + mocks.basePath = '/console' + vi.stubGlobal('fetch', vi.fn().mockResolvedValue(new Response(null, { status: 401 }))) + const { GET } = await import('../route') + + const response = await GET(createRequest( + 'http://localhost:3000/console/auth/refresh?redirect_url=%2Fconsole%2Fauth%2Frefresh', + 'refresh_token=expired', + )) + + expect(response.status).toBe(303) + expect(response.headers.get('location')).toBe('/console/signin?redirect_url=%2Fconsole%2F') + }) }) diff --git a/web/app/signin/__tests__/normal-form.spec.tsx b/web/app/signin/__tests__/normal-form.spec.tsx index 6dfdfbb372f..0e03b7c3270 100644 --- a/web/app/signin/__tests__/normal-form.spec.tsx +++ b/web/app/signin/__tests__/normal-form.spec.tsx @@ -41,10 +41,6 @@ vi.mock('@/service/common', async () => { } }) -vi.mock('./utils/post-login-redirect', () => ({ - resolvePostLoginRedirect: vi.fn(() => null), -})) - const mockReplace = vi.fn() const mockUseQuery = vi.mocked(useQuery) const mockUseSuspenseQuery = vi.mocked(useSuspenseQuery) @@ -74,11 +70,17 @@ const invitationQueryResult = { }, } +const nonInviteQueryResult = { + isPending: false, + isError: false, + data: undefined, +} + describe('NormalForm', () => { beforeEach(() => { vi.clearAllMocks() mockUseRouter.mockReturnValue({ replace: mockReplace }) - mockUseSearchParams.mockReturnValue(new URLSearchParams('invite_token=invite-token')) + mockUseSearchParams.mockReturnValue(new URLSearchParams()) mockUseSuspenseQuery.mockReturnValue({ data: { enable_social_oauth_login: false, @@ -97,8 +99,25 @@ describe('NormalForm', () => { } as unknown as ReturnType) }) + describe('Default Redirects', () => { + it('should send logged-in visitors without a redirect target to the console home', async () => { + const searchParams = new URLSearchParams() + mockUseSearchParams.mockReturnValue(searchParams) + mockUseQuery + .mockReturnValueOnce(loggedInQueryResult as unknown as ReturnType) + .mockReturnValueOnce(nonInviteQueryResult as unknown as ReturnType) + + render() + + await waitFor(() => { + expect(mockReplace).toHaveBeenCalledWith('/') + }) + }) + }) + describe('Invite Redirects', () => { it('should send logged-in invite visitors to the invite confirmation page', async () => { + mockUseSearchParams.mockReturnValue(new URLSearchParams('invite_token=invite-token')) mockUseQuery .mockReturnValueOnce(loggedInQueryResult as unknown as ReturnType) .mockReturnValueOnce(invitationQueryResult as unknown as ReturnType) From af4538e942390c23d8a9d6c044aeaadc66f71ba3 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?i=E6=99=9F?= <15772650@qq.com> Date: Wed, 8 Jul 2026 14:20:10 +0800 Subject: [PATCH 42/70] fix: raise clear error on unsupported language in execute_code (#38448) Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- api/core/helper/code_executor/code_executor.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/api/core/helper/code_executor/code_executor.py b/api/core/helper/code_executor/code_executor.py index ab2a959a1fc..c30afc0e745 100644 --- a/api/core/helper/code_executor/code_executor.py +++ b/api/core/helper/code_executor/code_executor.py @@ -74,12 +74,16 @@ class CodeExecutor: :param code: code :return: """ + running_language = cls.code_language_to_running_language.get(language) + if running_language is None: + raise CodeExecutionError(f"Unsupported language {language}") + url = code_execution_endpoint_url / "v1" / "sandbox" / "run" headers = {"X-Api-Key": dify_config.CODE_EXECUTION_API_KEY} data = { - "language": cls.code_language_to_running_language.get(language), + "language": running_language, "code": code, "preload": preload, "enable_network": True, From 98f67e9c827e0418eaee1a791ea6e0b8ef86e3e2 Mon Sep 17 00:00:00 2001 From: Evan <2869018789@qq.com> Date: Wed, 8 Jul 2026 14:22:59 +0800 Subject: [PATCH 43/70] fix(ci): make no-new-getattr guard stable in shallow PR checkouts (#38480) Co-authored-by: QuantumGhost Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- .github/workflows/main-ci.yml | 6 + .github/workflows/style.yml | 15 +- .../commands/test_check_no_new_getattr.py | 246 ++++++++++++------ scripts/check_no_new_getattr.py | 29 ++- 4 files changed, 196 insertions(+), 100 deletions(-) diff --git a/.github/workflows/main-ci.yml b/.github/workflows/main-ci.yml index 8016ff4db8d..0dad5915d00 100644 --- a/.github/workflows/main-ci.yml +++ b/.github/workflows/main-ci.yml @@ -53,6 +53,10 @@ jobs: filters: | api: - 'api/**' + - 'scripts/check_no_new_getattr.py' + - 'scripts/ast_grep_rules/no_new_getattr.yml' + - '.github/workflows/style.yml' + - '.github/workflows/main-ci.yml' - '.github/workflows/api-tests.yml' - 'docker/.env.example' - 'docker/envs/middleware.env.example' @@ -380,6 +384,8 @@ jobs: needs: pre_job if: needs.pre_job.outputs.should_skip != 'true' uses: ./.github/workflows/style.yml + with: + base-rev: ${{ github.event.pull_request.base.sha || github.event.merge_group.base_sha }} vdb-tests-run: name: Run VDB Tests diff --git a/.github/workflows/style.yml b/.github/workflows/style.yml index ceef6a855b7..eecd7f83e5e 100644 --- a/.github/workflows/style.yml +++ b/.github/workflows/style.yml @@ -2,6 +2,10 @@ name: Style check on: workflow_call: + inputs: + base-rev: + required: true + type: string concurrency: group: style-${{ github.head_ref || github.run_id }} @@ -33,6 +37,7 @@ jobs: scripts/check_no_new_getattr.py scripts/ast_grep_rules/no_new_getattr.yml .github/workflows/style.yml + .github/workflows/main-ci.yml - name: Setup UV and Python if: steps.changed-files.outputs.any_changed == 'true' @@ -54,17 +59,9 @@ jobs: if: steps.changed-files.outputs.any_changed == 'true' run: uv run --project api --dev python api/dev/lint_response_contracts.py --fail-on-mismatch - - name: Fetch merge target ref for getattr guard - if: steps.changed-files.outputs.any_changed == 'true' - run: git fetch --no-tags --depth=1 origin +refs/heads/main:refs/remotes/origin/main - - - name: Bind merge target branch for getattr guard - if: steps.changed-files.outputs.any_changed == 'true' - run: git show-ref --verify --quiet refs/heads/main || git branch main origin/main - - name: Run No New Getattr Guard if: steps.changed-files.outputs.any_changed == 'true' - run: uv run --project api python scripts/check_no_new_getattr.py --mode ci --merge-target main + run: uv run --project api python scripts/check_no_new_getattr.py --base-rev "${{ inputs.base-rev }}" - name: Run Type Checks if: steps.changed-files.outputs.any_changed == 'true' diff --git a/api/tests/unit_tests/commands/test_check_no_new_getattr.py b/api/tests/unit_tests/commands/test_check_no_new_getattr.py index 2fbaea62d64..72e72a631aa 100644 --- a/api/tests/unit_tests/commands/test_check_no_new_getattr.py +++ b/api/tests/unit_tests/commands/test_check_no_new_getattr.py @@ -77,6 +77,10 @@ def assert_has_actionable_violation(stderr: str, path: str) -> None: assert "no-new-getattr" in stderr +def main_branch_rev(repo: Path) -> str: + return git(repo, "rev-parse", "main") + + def test_resolve_ast_grep_command_prefers_ast_grep(monkeypatch: pytest.MonkeyPatch) -> None: module = load_guard_module() monkeypatch.setattr( @@ -130,8 +134,46 @@ def test_resolve_ast_grep_command_raises_without_explicit_binary(monkeypatch: py module.resolve_ast_grep_command() +def test_cli_requires_explicit_diff_source(tmp_path: Path) -> None: + result = run_script(tmp_path) + + assert result.returncode == 2 + assert "one of the arguments --staged --base-rev is required" in result.stderr + + +def test_cli_rejects_mixed_diff_sources(tmp_path: Path) -> None: + result = run_script(tmp_path, "--staged", "--base-rev", "deadbeef") + + assert result.returncode == 2 + assert "not allowed with argument" in result.stderr + + +def test_cli_help_exposes_only_new_diff_source_flags(tmp_path: Path) -> None: + help_result = run_script(tmp_path, "--help") + + assert help_result.returncode == 0 + assert "--staged" in help_result.stdout + assert "--base-rev" in help_result.stdout + assert "--mode" not in help_result.stdout + assert "--merge-target" not in help_result.stdout + + result = run_script(tmp_path, "--staged", "--mode", "ci") + + assert result.returncode == 2 + assert "unrecognized arguments: --mode ci" in result.stderr + + result = run_script(tmp_path, "--base-rev", "deadbeef", "--merge-target", "main") + + assert result.returncode == 2 + assert "unrecognized arguments: --merge-target main" in result.stderr + + def test_style_workflow_wires_no_new_getattr_guard() -> None: workflow = (REPO_ROOT / ".github" / "workflows" / "style.yml").read_text(encoding="utf-8") + assert re.search( + r"(?ms)^on:\n workflow_call:\n inputs:\n base-rev:\n required: true\n type: string\n", + workflow, + ) python_style_job = re.search( r"(?ms)^ python-style:\n(?P.*?)(?=^ [a-z0-9-]+:\n|\Z)", workflow, @@ -157,8 +199,9 @@ def test_style_workflow_wires_no_new_getattr_guard() -> None: assert "scripts/check_no_new_getattr.py\n" in files_block assert "scripts/ast_grep_rules/no_new_getattr.yml\n" in files_block assert ".github/workflows/style.yml\n" in files_block + assert ".github/workflows/main-ci.yml\n" in files_block - guard_command = "scripts/check_no_new_getattr.py --mode ci --merge-target main" + guard_command = 'scripts/check_no_new_getattr.py --base-rev "${{ inputs.base-rev }}"' assert guard_command in job_text guard_step = re.search( @@ -168,55 +211,35 @@ def test_style_workflow_wires_no_new_getattr_guard() -> None: ) assert guard_step is not None - pre_guard_text = job_text[: guard_step.start()] - step_pattern = r"(?ms)^ - name: [^\n]*\n(?P.*?)(?=^ - name: |\Z)" - fetch_step_text = next( - ( - match.group("step") - for match in re.finditer(step_pattern, pre_guard_text) - if any( - re.search(pattern, line) - for line in match.group("step").splitlines() - for pattern in ( - r"git fetch .*refs/heads/main:refs/remotes/origin/main", - r"git fetch .*main:refs/remotes/origin/main", - r"git fetch .*refs/remotes/origin/main", - ) - ) - ), - "", - ) - assert fetch_step_text - assert "git fetch" in fetch_step_text - assert "origin" in fetch_step_text - assert any( - re.search(pattern, line) - for line in fetch_step_text.splitlines() - for pattern in ( - r"git fetch .*refs/heads/main:refs/remotes/origin/main", - r"git fetch .*main:refs/remotes/origin/main", - r"git fetch .*refs/remotes/origin/main", - ) - ) - - bind_step = re.search( - r"(?ms)^ - name: Bind merge target branch for getattr guard\n(?P.*?)(?=^ - name: |\Z)", - pre_guard_text, - ) - assert bind_step is not None - bind_step_text = bind_step.group("step") - assert any( - command in bind_step_text - for command in ( - "git branch main origin/main", - "git checkout -B main origin/main", - "git switch -C main origin/main", - "git update-ref refs/heads/main refs/remotes/origin/main", - ) - ) + assert "GITHUB_BASE_SHA" not in guard_step.group("step") -def test_ci_mode_passes_when_only_legacy_getattr_exists(tmp_path: Path) -> None: +def test_main_ci_passes_style_base_rev_input() -> None: + workflow = (REPO_ROOT / ".github" / "workflows" / "main-ci.yml").read_text(encoding="utf-8") + style_job = re.search( + r"(?ms)^ style-check:\n(?P.*?)(?=^ [a-z0-9-]+:\n|\Z)", + workflow, + ) + assert style_job is not None + assert "uses: ./.github/workflows/style.yml" in style_job.group("job") + assert ( + "base-rev: ${{ github.event.pull_request.base.sha || github.event.merge_group.base_sha }}" + in style_job.group("job") + ) + + api_filter = re.search( + r"(?ms)^ api:\n(?P(?:^ - '[^']+'\n)+)", + workflow, + ) + assert api_filter is not None + filter_text = api_filter.group("filter") + assert "scripts/check_no_new_getattr.py" in filter_text + assert "scripts/ast_grep_rules/no_new_getattr.yml" in filter_text + assert ".github/workflows/style.yml" in filter_text + assert ".github/workflows/main-ci.yml" in filter_text + + +def test_base_rev_mode_passes_when_only_legacy_getattr_exists(tmp_path: Path) -> None: init_repo(tmp_path) write_repo_file( tmp_path, @@ -239,12 +262,12 @@ def test_ci_mode_passes_when_only_legacy_getattr_exists(tmp_path: Path) -> None: ) commit_all(tmp_path, "unrelated change") - result = run_script(tmp_path, "--mode", "ci", "--merge-target", "main") + result = run_script(tmp_path, "--base-rev", main_branch_rev(tmp_path)) assert result.returncode == 0, result.stderr -def test_ci_mode_fails_for_new_file_with_getattr(tmp_path: Path) -> None: +def test_base_rev_mode_fails_for_new_file_with_getattr(tmp_path: Path) -> None: init_repo(tmp_path) write_repo_file( tmp_path, @@ -267,13 +290,13 @@ def test_ci_mode_fails_for_new_file_with_getattr(tmp_path: Path) -> None: ) commit_all(tmp_path, "add new getattr usage") - result = run_script(tmp_path, "--mode", "ci", "--merge-target", "main") + result = run_script(tmp_path, "--base-rev", main_branch_rev(tmp_path)) assert result.returncode == 1 assert_has_actionable_violation(result.stderr, "pkg/new_usage.py") -def test_ci_mode_fails_for_new_file_with_two_arg_getattr(tmp_path: Path) -> None: +def test_base_rev_mode_fails_for_new_file_with_two_arg_getattr(tmp_path: Path) -> None: init_repo(tmp_path) write_repo_file( tmp_path, @@ -296,13 +319,13 @@ def test_ci_mode_fails_for_new_file_with_two_arg_getattr(tmp_path: Path) -> None ) commit_all(tmp_path, "add new two-arg getattr usage") - result = run_script(tmp_path, "--mode", "ci", "--merge-target", "main") + result = run_script(tmp_path, "--base-rev", main_branch_rev(tmp_path)) assert result.returncode == 1 assert_has_actionable_violation(result.stderr, "pkg/new_usage.py") -def test_ci_mode_fails_for_new_file_with_builtins_getattr(tmp_path: Path) -> None: +def test_base_rev_mode_fails_for_new_file_with_builtins_getattr(tmp_path: Path) -> None: init_repo(tmp_path) write_repo_file( tmp_path, @@ -328,13 +351,13 @@ def test_ci_mode_fails_for_new_file_with_builtins_getattr(tmp_path: Path) -> Non ) commit_all(tmp_path, "add new builtins getattr usage") - result = run_script(tmp_path, "--mode", "ci", "--merge-target", "main") + result = run_script(tmp_path, "--base-rev", main_branch_rev(tmp_path)) assert result.returncode == 1 assert_has_actionable_violation(result.stderr, "pkg/new_usage.py") -def test_ci_mode_fails_for_new_file_with_two_arg_builtins_getattr(tmp_path: Path) -> None: +def test_base_rev_mode_fails_for_new_file_with_two_arg_builtins_getattr(tmp_path: Path) -> None: init_repo(tmp_path) write_repo_file( tmp_path, @@ -360,13 +383,13 @@ def test_ci_mode_fails_for_new_file_with_two_arg_builtins_getattr(tmp_path: Path ) commit_all(tmp_path, "add new two-arg builtins getattr usage") - result = run_script(tmp_path, "--mode", "ci", "--merge-target", "main") + result = run_script(tmp_path, "--base-rev", main_branch_rev(tmp_path)) assert result.returncode == 1 assert_has_actionable_violation(result.stderr, "pkg/new_usage.py") -def test_ci_mode_fails_for_new_file_with_dunder_builtins_getattr(tmp_path: Path) -> None: +def test_base_rev_mode_fails_for_new_file_with_dunder_builtins_getattr(tmp_path: Path) -> None: init_repo(tmp_path) write_repo_file( tmp_path, @@ -389,13 +412,13 @@ def test_ci_mode_fails_for_new_file_with_dunder_builtins_getattr(tmp_path: Path) ) commit_all(tmp_path, "add new dunder builtins getattr usage") - result = run_script(tmp_path, "--mode", "ci", "--merge-target", "main") + result = run_script(tmp_path, "--base-rev", main_branch_rev(tmp_path)) assert result.returncode == 1 assert_has_actionable_violation(result.stderr, "pkg/new_usage.py") -def test_ci_mode_fails_for_new_file_with_two_arg_dunder_builtins_getattr(tmp_path: Path) -> None: +def test_base_rev_mode_fails_for_new_file_with_two_arg_dunder_builtins_getattr(tmp_path: Path) -> None: init_repo(tmp_path) write_repo_file( tmp_path, @@ -418,13 +441,13 @@ def test_ci_mode_fails_for_new_file_with_two_arg_dunder_builtins_getattr(tmp_pat ) commit_all(tmp_path, "add new two-arg dunder builtins getattr usage") - result = run_script(tmp_path, "--mode", "ci", "--merge-target", "main") + result = run_script(tmp_path, "--base-rev", main_branch_rev(tmp_path)) assert result.returncode == 1 assert_has_actionable_violation(result.stderr, "pkg/new_usage.py") -def test_ci_mode_uses_merge_base_against_main_not_just_head_parent(tmp_path: Path) -> None: +def test_base_rev_mode_uses_provided_base_revision_not_head_parent(tmp_path: Path) -> None: init_repo(tmp_path) write_repo_file( tmp_path, @@ -435,6 +458,7 @@ def test_ci_mode_uses_merge_base_against_main_not_just_head_parent(tmp_path: Pat """, ) commit_all(tmp_path, "baseline") + base_rev = main_branch_rev(tmp_path) checkout_feature_branch(tmp_path) write_repo_file( @@ -457,13 +481,75 @@ def test_ci_mode_uses_merge_base_against_main_not_just_head_parent(tmp_path: Pat ) commit_all(tmp_path, "later feature commit does not touch violating file") - result = run_script(tmp_path, "--mode", "ci", "--merge-target", "main") + result = run_script(tmp_path, "--base-rev", base_rev) assert result.returncode == 1 assert_has_actionable_violation(result.stderr, "pkg/introduced_earlier.py") -def test_pre_commit_mode_reads_staged_content_only(tmp_path: Path) -> None: +def test_base_rev_mode_works_without_local_main_branch(tmp_path: Path) -> None: + init_repo(tmp_path) + write_repo_file( + tmp_path, + "pkg/existing.py", + """ + def stable() -> str: + return "ok" + """, + ) + commit_all(tmp_path, "baseline") + base_rev = main_branch_rev(tmp_path) + checkout_feature_branch(tmp_path) + + write_repo_file( + tmp_path, + "pkg/other.py", + """ + def meaning() -> int: + return 42 + """, + ) + commit_all(tmp_path, "feature change") + + git(tmp_path, "checkout", "--detach", "HEAD") + git(tmp_path, "branch", "-D", "main") + + result = run_script(tmp_path, "--base-rev", base_rev) + + assert result.returncode == 0, result.stderr + + +def test_base_rev_mode_ignores_github_base_sha_environment(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + init_repo(tmp_path) + write_repo_file( + tmp_path, + "pkg/existing.py", + """ + def stable() -> str: + return "ok" + """, + ) + commit_all(tmp_path, "baseline") + base_rev = main_branch_rev(tmp_path) + checkout_feature_branch(tmp_path) + + write_repo_file( + tmp_path, + "pkg/other.py", + """ + def meaning() -> int: + return 42 + """, + ) + commit_all(tmp_path, "feature change") + monkeypatch.setenv("GITHUB_BASE_SHA", "deadbeefdeadbeefdeadbeefdeadbeefdeadbeef") + + result = run_script(tmp_path, "--base-rev", base_rev) + + assert result.returncode == 0, result.stderr + + +def test_staged_mode_reads_staged_content_only(tmp_path: Path) -> None: init_repo(tmp_path) write_repo_file( tmp_path, @@ -494,12 +580,12 @@ def test_pre_commit_mode_reads_staged_content_only(tmp_path: Path) -> None: """, ) - result = run_script(tmp_path, "--mode", "pre-commit") + result = run_script(tmp_path, "--staged") assert result.returncode == 0, result.stderr -def test_pre_commit_mode_fails_for_staged_two_arg_getattr(tmp_path: Path) -> None: +def test_staged_mode_fails_for_staged_two_arg_getattr(tmp_path: Path) -> None: init_repo(tmp_path) write_repo_file( tmp_path, @@ -521,13 +607,13 @@ def test_pre_commit_mode_fails_for_staged_two_arg_getattr(tmp_path: Path) -> Non ) git(tmp_path, "add", "pkg/module.py") - result = run_script(tmp_path, "--mode", "pre-commit") + result = run_script(tmp_path, "--staged") assert result.returncode == 1 assert_has_actionable_violation(result.stderr, "pkg/module.py") -def test_pre_commit_mode_fails_for_staged_builtins_getattr(tmp_path: Path) -> None: +def test_staged_mode_fails_for_staged_builtins_getattr(tmp_path: Path) -> None: init_repo(tmp_path) write_repo_file( tmp_path, @@ -552,13 +638,13 @@ def test_pre_commit_mode_fails_for_staged_builtins_getattr(tmp_path: Path) -> No ) git(tmp_path, "add", "pkg/module.py") - result = run_script(tmp_path, "--mode", "pre-commit") + result = run_script(tmp_path, "--staged") assert result.returncode == 1 assert_has_actionable_violation(result.stderr, "pkg/module.py") -def test_pre_commit_mode_fails_for_staged_two_arg_builtins_getattr(tmp_path: Path) -> None: +def test_staged_mode_fails_for_staged_two_arg_builtins_getattr(tmp_path: Path) -> None: init_repo(tmp_path) write_repo_file( tmp_path, @@ -583,7 +669,7 @@ def test_pre_commit_mode_fails_for_staged_two_arg_builtins_getattr(tmp_path: Pat ) git(tmp_path, "add", "pkg/module.py") - result = run_script(tmp_path, "--mode", "pre-commit") + result = run_script(tmp_path, "--staged") assert result.returncode == 1 assert_has_actionable_violation(result.stderr, "pkg/module.py") @@ -603,6 +689,7 @@ def test_modified_hunk_with_same_getattr_count_is_allowed(tmp_path: Path) -> Non """, ) commit_all(tmp_path, "baseline") + base_rev = main_branch_rev(tmp_path) checkout_feature_branch(tmp_path) write_repo_file( @@ -618,7 +705,7 @@ def test_modified_hunk_with_same_getattr_count_is_allowed(tmp_path: Path) -> Non ) commit_all(tmp_path, "touch legacy getattr hunk") - result = run_script(tmp_path, "--mode", "ci", "--merge-target", "main") + result = run_script(tmp_path, "--base-rev", base_rev) assert result.returncode == 0, result.stderr @@ -635,6 +722,7 @@ def test_modified_hunk_with_decreased_getattr_count_is_allowed(tmp_path: Path) - """, ) commit_all(tmp_path, "baseline") + base_rev = main_branch_rev(tmp_path) checkout_feature_branch(tmp_path) write_repo_file( @@ -647,7 +735,7 @@ def test_modified_hunk_with_decreased_getattr_count_is_allowed(tmp_path: Path) - ) commit_all(tmp_path, "remove one legacy getattr") - result = run_script(tmp_path, "--mode", "ci", "--merge-target", "main") + result = run_script(tmp_path, "--base-rev", base_rev) assert result.returncode == 0, result.stderr @@ -663,6 +751,7 @@ def test_modified_hunk_with_increased_getattr_count_fails(tmp_path: Path) -> Non """, ) commit_all(tmp_path, "baseline") + base_rev = main_branch_rev(tmp_path) checkout_feature_branch(tmp_path) write_repo_file( @@ -676,7 +765,7 @@ def test_modified_hunk_with_increased_getattr_count_fails(tmp_path: Path) -> Non ) commit_all(tmp_path, "add one more getattr") - result = run_script(tmp_path, "--mode", "ci", "--merge-target", "main") + result = run_script(tmp_path, "--base-rev", base_rev) assert result.returncode == 1 assert_has_actionable_violation(result.stderr, "pkg/sample.py") @@ -694,6 +783,7 @@ def test_inline_noqa_suppression_with_explanatory_text_skips_added_getattr(tmp_p """, ) commit_all(tmp_path, "baseline") + base_rev = main_branch_rev(tmp_path) checkout_feature_branch(tmp_path) write_repo_file( @@ -706,7 +796,7 @@ def test_inline_noqa_suppression_with_explanatory_text_skips_added_getattr(tmp_p ) commit_all(tmp_path, "add suppressed getattr") - result = run_script(tmp_path, "--mode", "ci", "--merge-target", "main") + result = run_script(tmp_path, "--base-rev", base_rev) assert "no-new-getattr needed for plugin-defined attributes" in (tmp_path / "pkg/existing.py").read_text( encoding="utf-8" @@ -725,6 +815,7 @@ def test_inline_noqa_without_explanatory_text_is_not_sufficient(tmp_path: Path) """, ) commit_all(tmp_path, "baseline") + base_rev = main_branch_rev(tmp_path) checkout_feature_branch(tmp_path) write_repo_file( @@ -737,7 +828,7 @@ def test_inline_noqa_without_explanatory_text_is_not_sufficient(tmp_path: Path) ) commit_all(tmp_path, "add bare noqa getattr") - result = run_script(tmp_path, "--mode", "ci", "--merge-target", "main") + result = run_script(tmp_path, "--base-rev", base_rev) assert result.returncode == 1 assert_has_actionable_violation(result.stderr, "pkg/existing.py") @@ -753,6 +844,7 @@ def test_non_python_file_with_getattr_text_does_not_fail_guard(tmp_path: Path) - """, ) commit_all(tmp_path, "baseline") + base_rev = main_branch_rev(tmp_path) checkout_feature_branch(tmp_path) write_repo_file( @@ -764,6 +856,6 @@ def test_non_python_file_with_getattr_text_does_not_fail_guard(tmp_path: Path) - ) commit_all(tmp_path, "document getattr example") - result = run_script(tmp_path, "--mode", "ci", "--merge-target", "main") + result = run_script(tmp_path, "--base-rev", base_rev) assert result.returncode == 0, result.stderr diff --git a/scripts/check_no_new_getattr.py b/scripts/check_no_new_getattr.py index 098084fc1dd..8bf437c5367 100644 --- a/scripts/check_no_new_getattr.py +++ b/scripts/check_no_new_getattr.py @@ -1,5 +1,5 @@ #!/usr/bin/env python3 -"""Block net-new getattr() usage in changed Python hunks.""" +"""Block net-new getattr() usage in staged changes or against an explicit base revision.""" from __future__ import annotations @@ -58,8 +58,16 @@ class Violation: def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) - parser.add_argument("--mode", choices=("pre-commit", "ci"), required=True) - parser.add_argument("--merge-target", default="main") + diff_source_group = parser.add_mutually_exclusive_group(required=True) + diff_source_group.add_argument( + "--staged", + action="store_true", + help="Inspect staged changes against HEAD, for pre-commit style usage.", + ) + diff_source_group.add_argument( + "--base-rev", + help="Inspect changes between the provided git revision and HEAD.", + ) return parser.parse_args() @@ -86,21 +94,15 @@ def git_output(*args: str, allow_missing: bool = False) -> str: raise RuntimeError(completed.stderr.strip() or completed.stdout.strip() or "git command failed") -def resolve_ci_base(merge_target: str) -> str: - merge_base = git_output("merge-base", merge_target, "HEAD", allow_missing=True).strip() - return merge_base or merge_target - - def collect_diff_text(args: argparse.Namespace) -> str: - if args.mode == "pre-commit": + if args.staged: return git_output("diff", "--cached", "--unified=0", "--diff-filter=AM", "--no-ext-diff") - compare_base = resolve_ci_base(args.merge_target) return git_output( "diff", "--unified=0", "--diff-filter=AM", "--no-ext-diff", - f"{compare_base}..HEAD", + f"{args.base_rev}..HEAD", ) @@ -145,15 +147,14 @@ def is_python_source_path(path: str) -> bool: def load_file_versions(path: str, args: argparse.Namespace) -> tuple[str, str]: - if args.mode == "pre-commit": + if args.staged: return ( git_output("show", f"HEAD:{path}", allow_missing=True), git_output("show", f":{path}"), ) - compare_base = resolve_ci_base(args.merge_target) return ( - git_output("show", f"{compare_base}:{path}", allow_missing=True), + git_output("show", f"{args.base_rev}:{path}", allow_missing=True), git_output("show", f"HEAD:{path}"), ) From 503d80be1db2226df9cdddd5db712225cbcdf320 Mon Sep 17 00:00:00 2001 From: Stephen Zhou Date: Wed, 8 Jul 2026 14:26:56 +0800 Subject: [PATCH 44/70] refactor(web): migrate plugins and tools app context consumers (#38533) Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- eslint-suppressions.json | 8 -- .../plugins/plugin-page-shell-flow.test.tsx | 33 +++++ .../tools/tool-provider-detail-flow.test.tsx | 13 ++ web/__tests__/utils/mock-app-context-state.ts | 15 +++ .../app-publisher/__tests__/index.spec.tsx | 13 ++ .../account-setting/__tests__/index.spec.tsx | 22 +++- .../__tests__/card.spec.tsx | 12 ++ .../add-credential-in-load-balancing.spec.tsx | 12 ++ .../__tests__/add-custom-model.spec.tsx | 12 ++ ...itch-credential-in-load-balancing.spec.tsx | 19 ++- .../authorized/__tests__/index.spec.tsx | 12 ++ .../model-modal/__tests__/index.spec.tsx | 12 ++ .../__tests__/popup-item.spec.tsx | 12 ++ .../__tests__/index.spec.tsx | 13 ++ .../__tests__/api-key-section.spec.tsx | 12 ++ .../__tests__/dropdown-content.spec.tsx | 12 ++ .../__tests__/tool-provider-list.spec.tsx | 12 ++ .../plugins/card/__tests__/index.spec.tsx | 14 ++- web/app/components/plugins/card/index.tsx | 5 +- .../base/__tests__/use-get-icon.spec.ts | 18 ++- .../install-plugin/base/use-get-icon.ts | 9 +- ...orkspace-plugin-install-permission.spec.ts | 24 +++- ...use-workspace-plugin-install-permission.ts | 9 +- .../steps/__tests__/install.spec.tsx | 31 +++-- .../steps/install.tsx | 5 +- .../steps/__tests__/install.spec.tsx | 30 +++-- .../steps/install.tsx | 5 +- ...arketplace-install-permission-provider.tsx | 5 +- .../__tests__/plugin-auth-in-agent.spec.tsx | 14 +++ .../__tests__/plugin-auth.spec.tsx | 12 ++ .../authorize/__tests__/index.spec.tsx | 12 ++ .../authorized/__tests__/index.spec.tsx | 13 ++ .../authorized/__tests__/item.spec.tsx | 22 ++-- .../plugins/plugin-auth/authorized/item.tsx | 7 +- .../datasource-action-list.tsx | 2 - .../plugin-item/__tests__/index.spec.tsx | 39 ++++-- .../components/plugins/plugin-item/index.tsx | 5 +- .../plugin-page/__tests__/index.spec.tsx | 33 ++++- .../__tests__/use-reference-setting.spec.ts | 117 +++++++++++------- .../plugin-page/use-reference-setting.ts | 27 ++-- .../tools/hooks/use-tool-permissions.ts | 7 +- .../tools/mcp/__tests__/create-card.spec.tsx | 26 +++- .../tools/mcp/__tests__/index.spec.tsx | 24 +++- .../mcp/__tests__/provider-card.spec.tsx | 24 +++- .../mcp/detail/__tests__/content.spec.tsx | 31 +++-- .../tools/mcp/hooks/use-mcp-service-card.ts | 7 +- .../__tests__/custom-create-card.spec.tsx | 26 +++- .../tools/provider/__tests__/detail.spec.tsx | 14 +++ .../__tests__/tool-picker.spec.tsx | 12 ++ .../use-node-plugin-installation.spec.ts | 12 ++ web/hooks/use-credential-permissions.spec.ts | 20 +-- web/hooks/use-credential-permissions.ts | 5 +- web/service/__tests__/use-plugins.spec.tsx | 19 ++- web/service/use-plugins.ts | 5 +- 54 files changed, 745 insertions(+), 189 deletions(-) diff --git a/eslint-suppressions.json b/eslint-suppressions.json index 176341002bf..81c7678008c 100644 --- a/eslint-suppressions.json +++ b/eslint-suppressions.json @@ -3457,14 +3457,6 @@ "count": 2 } }, - "web/app/components/header/account-setting/model-provider-page/model-auth/__tests__/switch-credential-in-load-balancing.spec.tsx": { - "jsx-a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx-a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/header/account-setting/model-provider-page/model-auth/add-custom-model.tsx": { "jsx-a11y/click-events-have-key-events": { "count": 2 diff --git a/web/__tests__/plugins/plugin-page-shell-flow.test.tsx b/web/__tests__/plugins/plugin-page-shell-flow.test.tsx index 483df7c8ec7..14eebf2554a 100644 --- a/web/__tests__/plugins/plugin-page-shell-flow.test.tsx +++ b/web/__tests__/plugins/plugin-page-shell-flow.test.tsx @@ -44,7 +44,13 @@ vi.mock('@/context/app-context', () => ({ isCurrentWorkspaceManager: false, isCurrentWorkspaceOwner: false, langGeniusVersionInfo: { + current_env: 'CLOUD', current_version: '1.0.0', + latest_version: '1.0.0', + version: '1.0.0', + release_date: '', + release_notes: '', + can_auto_update: false, }, workspacePermissionKeys: [ 'plugin.install', @@ -54,6 +60,33 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + isCurrentWorkspaceManager: false, + isCurrentWorkspaceOwner: false, + langGeniusVersionInfo: { + current_env: 'CLOUD', + current_version: '1.0.0', + latest_version: '1.0.0', + version: '1.0.0', + release_date: '', + release_notes: '', + can_auto_update: false, + }, + workspacePermissionKeys: [ + 'plugin.install', + 'plugin.delete', + 'plugin.plugin_preferences', + ], + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/service/use-plugins', () => ({ hasPluginPermission: () => true, useReferenceSettings: () => ({ diff --git a/web/__tests__/tools/tool-provider-detail-flow.test.tsx b/web/__tests__/tools/tool-provider-detail-flow.test.tsx index 8cbfae8d693..fe7b55189bc 100644 --- a/web/__tests__/tools/tool-provider-detail-flow.test.tsx +++ b/web/__tests__/tools/tool-provider-detail-flow.test.tsx @@ -53,6 +53,19 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + isCurrentWorkspaceManager: true, + workspacePermissionKeys: ['tool.manage', 'credential.create', 'credential.manage', 'credential.use'], + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + const mockSetShowModelModal = vi.fn() vi.mock('@/context/modal-context', () => ({ useModalContext: () => ({ diff --git a/web/__tests__/utils/mock-app-context-state.ts b/web/__tests__/utils/mock-app-context-state.ts index 99795f60766..0c546de1471 100644 --- a/web/__tests__/utils/mock-app-context-state.ts +++ b/web/__tests__/utils/mock-app-context-state.ts @@ -14,6 +14,10 @@ export type AppContextStateMockState = { currentWorkspace?: { id?: string } | null + isCurrentWorkspaceManager?: boolean + isCurrentWorkspaceOwner?: boolean + isCurrentWorkspaceEditor?: boolean + isCurrentWorkspaceDatasetOperator?: boolean isLoadingCurrentWorkspace?: boolean isLoadingWorkspacePermissionKeys?: boolean workspacePermissionKeys?: string[] @@ -25,6 +29,7 @@ type AppContextStateAtomKind | 'userProfileId' | 'currentWorkspace' | 'currentWorkspaceId' + | 'workspaceRoleFlags' | 'currentWorkspaceLoading' | 'workspacePermissionKeys' | 'workspacePermissionKeysLoading' @@ -98,6 +103,7 @@ export const createAppContextStateAtomMock = async ( userProfileIdAtom: createMockAtom('userProfileId'), currentWorkspaceAtom: createMockAtom('currentWorkspace'), currentWorkspaceIdAtom: createMockAtom('currentWorkspaceId'), + workspaceRoleFlagsAtom: createMockAtom('workspaceRoleFlags'), currentWorkspaceLoadingAtom: createMockAtom('currentWorkspaceLoading'), workspacePermissionKeysAtom: createMockAtom('workspacePermissionKeys'), workspacePermissionKeysLoadingAtom: createMockAtom('workspacePermissionKeysLoading'), @@ -135,6 +141,15 @@ export const createAppContextStateJotaiMock = async ( if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'currentWorkspaceId') return currentWorkspace.id + if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'workspaceRoleFlags') { + return { + isCurrentWorkspaceManager: state.isCurrentWorkspaceManager ?? false, + isCurrentWorkspaceOwner: state.isCurrentWorkspaceOwner ?? false, + isCurrentWorkspaceEditor: state.isCurrentWorkspaceEditor ?? false, + isCurrentWorkspaceDatasetOperator: state.isCurrentWorkspaceDatasetOperator ?? false, + } + } + if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'currentWorkspaceLoading') return state.isLoadingCurrentWorkspace ?? false diff --git a/web/app/components/app/app-publisher/__tests__/index.spec.tsx b/web/app/components/app/app-publisher/__tests__/index.spec.tsx index 5bbe67b6890..d6fb10c8163 100644 --- a/web/app/components/app/app-publisher/__tests__/index.spec.tsx +++ b/web/app/components/app/app-publisher/__tests__/index.spec.tsx @@ -116,6 +116,19 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + isCurrentWorkspaceManager: true, + workspacePermissionKeys: mockWorkspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@langgenius/dify-ui/toast', () => ({ toast: { error: (...args: unknown[]) => mockToastError(...args), diff --git a/web/app/components/header/account-setting/__tests__/index.spec.tsx b/web/app/components/header/account-setting/__tests__/index.spec.tsx index 3d580673d5a..b4635473532 100644 --- a/web/app/components/header/account-setting/__tests__/index.spec.tsx +++ b/web/app/components/header/account-setting/__tests__/index.spec.tsx @@ -31,6 +31,16 @@ vi.mock('@/context/app-context', async (importOriginal) => { } }) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppContextState.current ?? {}) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/next/navigation', () => ({ useRouter: vi.fn(() => ({ push: vi.fn(), @@ -77,10 +87,14 @@ vi.mock('@/service/use-datasource', () => ({ useGetDataSourceListAuth: vi.fn(() => ({ data: { result: [] } })), })) -vi.mock('@/service/use-common', () => ({ - useMembers: vi.fn(() => ({ data: { accounts: [] }, refetch: vi.fn() })), - useProviderContext: vi.fn(), -})) +vi.mock('@/service/use-common', async (importOriginal) => { + const actual = await importOriginal() + return { + ...actual, + useMembers: vi.fn(() => ({ data: { accounts: [] }, refetch: vi.fn() })), + useProviderContext: vi.fn(), + } +}) vi.mock('@/service/client', async (importOriginal) => { const actual = await importOriginal() diff --git a/web/app/components/header/account-setting/data-source-page-new/__tests__/card.spec.tsx b/web/app/components/header/account-setting/data-source-page-new/__tests__/card.spec.tsx index 0318645ea9f..464e4d7f367 100644 --- a/web/app/components/header/account-setting/data-source-page-new/__tests__/card.spec.tsx +++ b/web/app/components/header/account-setting/data-source-page-new/__tests__/card.spec.tsx @@ -19,6 +19,18 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockWorkspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/app/components/plugins/plugin-auth', () => ({ ApiKeyModal: vi.fn(({ onClose, onUpdate, onRemove, disabled, editValues }: { onClose: () => void, onUpdate: () => void, onRemove: () => void, disabled: boolean, editValues: Record }) => (
    diff --git a/web/app/components/header/account-setting/model-provider-page/model-auth/__tests__/add-credential-in-load-balancing.spec.tsx b/web/app/components/header/account-setting/model-provider-page/model-auth/__tests__/add-credential-in-load-balancing.spec.tsx index c8ec3a1f2a2..499380c1e1a 100644 --- a/web/app/components/header/account-setting/model-provider-page/model-auth/__tests__/add-credential-in-load-balancing.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/model-auth/__tests__/add-credential-in-load-balancing.spec.tsx @@ -13,6 +13,18 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockWorkspacePermissionKeys.value, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/app/components/header/account-setting/model-provider-page/model-auth', () => ({ Authorized: ({ renderTrigger, diff --git a/web/app/components/header/account-setting/model-provider-page/model-auth/__tests__/add-custom-model.spec.tsx b/web/app/components/header/account-setting/model-provider-page/model-auth/__tests__/add-custom-model.spec.tsx index 7230be997f4..515081c22f8 100644 --- a/web/app/components/header/account-setting/model-provider-page/model-auth/__tests__/add-custom-model.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/model-auth/__tests__/add-custom-model.spec.tsx @@ -34,6 +34,18 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockWorkspacePermissionKeys.value, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + // Mock components vi.mock('../../model-icon', () => ({ default: () =>
    , diff --git a/web/app/components/header/account-setting/model-provider-page/model-auth/__tests__/switch-credential-in-load-balancing.spec.tsx b/web/app/components/header/account-setting/model-provider-page/model-auth/__tests__/switch-credential-in-load-balancing.spec.tsx index 3199e26692c..902a2000425 100644 --- a/web/app/components/header/account-setting/model-provider-page/model-auth/__tests__/switch-credential-in-load-balancing.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/model-auth/__tests__/switch-credential-in-load-balancing.spec.tsx @@ -14,6 +14,18 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockWorkspacePermissionKeys.value, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + // Mock components vi.mock('../authorized', () => ({ default: ({ @@ -37,7 +49,12 @@ vi.mock('../authorized', () => ({ data-hide-add-action={String(!!hideAddAction)} data-trigger-only-open-modal={String(!!triggerOnlyOpenModal)} > -
    onItemClick?.(items[0]!.credentials[0])}> +
    onItemClick?.(items[0]!.credentials[0])} + onKeyDown={() => undefined} + > {renderTrigger()}
    diff --git a/web/app/components/header/account-setting/model-provider-page/model-auth/authorized/__tests__/index.spec.tsx b/web/app/components/header/account-setting/model-provider-page/model-auth/authorized/__tests__/index.spec.tsx index ea1ff14d18a..bab7f605997 100644 --- a/web/app/components/header/account-setting/model-provider-page/model-auth/authorized/__tests__/index.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/model-auth/authorized/__tests__/index.spec.tsx @@ -19,6 +19,18 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockWorkspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('../../hooks', () => ({ useAuth: () => ({ openConfirmDelete: mockOpenConfirmDelete, diff --git a/web/app/components/header/account-setting/model-provider-page/model-modal/__tests__/index.spec.tsx b/web/app/components/header/account-setting/model-provider-page/model-modal/__tests__/index.spec.tsx index 0fb0feaef02..655220ee6ac 100644 --- a/web/app/components/header/account-setting/model-provider-page/model-modal/__tests__/index.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/model-modal/__tests__/index.spec.tsx @@ -80,6 +80,18 @@ vi.mock('@/context/app-context', () => ({ selector({ workspacePermissionKeys: mockState.workspacePermissionKeys }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockState.workspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/hooks/use-i18n', () => ({ useRenderI18nObject: () => (value: { en_US: string }) => value.en_US, })) diff --git a/web/app/components/header/account-setting/model-provider-page/model-selector/__tests__/popup-item.spec.tsx b/web/app/components/header/account-setting/model-provider-page/model-selector/__tests__/popup-item.spec.tsx index 6fd2876c9cb..29d8f728693 100644 --- a/web/app/components/header/account-setting/model-provider-page/model-selector/__tests__/popup-item.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/model-selector/__tests__/popup-item.spec.tsx @@ -94,6 +94,18 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockWorkspacePermissionKeys.value, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + const makeModelItem = (overrides: Partial = {}): ModelItem => ({ model: 'gpt-4', label: { en_US: 'GPT-4', zh_Hans: 'GPT-4' }, diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/index.spec.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/index.spec.tsx index 0756a172e59..9aee6fd9ead 100644 --- a/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/index.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/index.spec.tsx @@ -43,6 +43,19 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + isCurrentWorkspaceManager: mockIsCurrentWorkspaceManager, + workspacePermissionKeys: mockWorkspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + // Mock internal components to simplify testing of the index file vi.mock('../credential-panel', () => ({ default: () =>
    , diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-auth-dropdown/__tests__/api-key-section.spec.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-auth-dropdown/__tests__/api-key-section.spec.tsx index 0ad4a494285..82e4477b4cf 100644 --- a/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-auth-dropdown/__tests__/api-key-section.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-auth-dropdown/__tests__/api-key-section.spec.tsx @@ -9,6 +9,18 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: ['credential.use', 'credential.create', 'credential.manage'], + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + const createCredential = (overrides: Partial = {}): Credential => ({ credential_id: 'cred-1', credential_name: 'Test API Key', diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-auth-dropdown/__tests__/dropdown-content.spec.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-auth-dropdown/__tests__/dropdown-content.spec.tsx index f50d654f82e..fe3f55e68dd 100644 --- a/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-auth-dropdown/__tests__/dropdown-content.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-auth-dropdown/__tests__/dropdown-content.spec.tsx @@ -41,6 +41,18 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockWorkspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('../../../model-auth/authorized/credential-item', () => ({ default: ({ credential, disabled, disableEdit, disableDelete, onItemClick, onEdit, onDelete }: { credential: { credential_id: string, credential_name: string } diff --git a/web/app/components/integrations/__tests__/tool-provider-list.spec.tsx b/web/app/components/integrations/__tests__/tool-provider-list.spec.tsx index 91f0a437d24..c43d728916d 100644 --- a/web/app/components/integrations/__tests__/tool-provider-list.spec.tsx +++ b/web/app/components/integrations/__tests__/tool-provider-list.spec.tsx @@ -103,6 +103,18 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockAppContextState.workspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + let mockCheckedInstalledData: { plugins: { id: string, name: string }[] } | null = null const mockInvalidateInstalledPluginList = vi.fn() vi.mock('@/service/use-plugins', () => ({ diff --git a/web/app/components/plugins/card/__tests__/index.spec.tsx b/web/app/components/plugins/card/__tests__/index.spec.tsx index b6c97dfd026..9b7319ca580 100644 --- a/web/app/components/plugins/card/__tests__/index.spec.tsx +++ b/web/app/components/plugins/card/__tests__/index.spec.tsx @@ -41,11 +41,17 @@ vi.mock('@/utils/format', () => ({ formatNumber: (num: number) => num.toLocaleString(), })) -vi.mock('@/context/app-context', () => ({ - useSelector: (selector: (value: { currentWorkspace: { id: string } }) => string) => selector({ +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ currentWorkspace: { id: 'workspace-123' }, - }), -})) + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) vi.mock('@/utils/mcp', () => ({ shouldUseMcpIcon: (src: unknown) => typeof src === 'object' && src !== null && (src as { content?: string })?.content === '🔗', diff --git a/web/app/components/plugins/card/index.tsx b/web/app/components/plugins/card/index.tsx index 48a3f450641..d4707936eea 100644 --- a/web/app/components/plugins/card/index.tsx +++ b/web/app/components/plugins/card/index.tsx @@ -1,9 +1,10 @@ 'use client' import type { Plugin } from '../types' import { cn } from '@langgenius/dify-ui/cn' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useTranslation } from '#i18n' -import { useSelector } from '@/context/app-context' +import { currentWorkspaceIdAtom } from '@/context/app-context-state' import { useGetLanguage } from '@/context/i18n' import useTheme from '@/hooks/use-theme' import { @@ -61,7 +62,7 @@ const Card = ({ const locale = useGetLanguage() const { t } = useTranslation() const { categoriesMap } = useCategories(true) - const currentWorkspaceId = useSelector(s => s.currentWorkspace.id) + const currentWorkspaceId = useAtomValue(currentWorkspaceIdAtom) const { category, type, name, org, label, brief, icon, icon_dark, verified, from } = payload const badges = payload.badges ?? [] const { theme } = useTheme() diff --git a/web/app/components/plugins/install-plugin/base/__tests__/use-get-icon.spec.ts b/web/app/components/plugins/install-plugin/base/__tests__/use-get-icon.spec.ts index c5364ec47ff..ec7ba4c34db 100644 --- a/web/app/components/plugins/install-plugin/base/__tests__/use-get-icon.spec.ts +++ b/web/app/components/plugins/install-plugin/base/__tests__/use-get-icon.spec.ts @@ -2,15 +2,27 @@ import { renderHook } from '@testing-library/react' import { describe, expect, it, vi } from 'vitest' import useGetIcon from '../use-get-icon' +const mockCurrentWorkspaceIdAtom = vi.hoisted(() => Symbol('currentWorkspaceIdAtom')) + vi.mock('@/config', () => ({ API_PREFIX: 'https://api.example.com', })) -vi.mock('@/context/app-context', () => ({ - useSelector: (selector: (state: { currentWorkspace: { id: string } }) => string | { id: string }) => - selector({ currentWorkspace: { id: 'workspace-123' } }), +vi.mock('@/context/app-context-state', () => ({ + currentWorkspaceIdAtom: mockCurrentWorkspaceIdAtom, })) +vi.mock('jotai', () => { + return { + useAtomValue: (atom: unknown) => { + if (atom === mockCurrentWorkspaceIdAtom) + return 'workspace-123' + + throw new Error('Unexpected atom') + }, + } +}) + describe('useGetIcon', () => { it('builds icon url with current workspace id', () => { const { result } = renderHook(() => useGetIcon()) diff --git a/web/app/components/plugins/install-plugin/base/use-get-icon.ts b/web/app/components/plugins/install-plugin/base/use-get-icon.ts index 23d41aff07e..b390a7e951f 100644 --- a/web/app/components/plugins/install-plugin/base/use-get-icon.ts +++ b/web/app/components/plugins/install-plugin/base/use-get-icon.ts @@ -1,12 +1,13 @@ +import { useAtomValue } from 'jotai' import { useCallback } from 'react' import { API_PREFIX } from '@/config' -import { useSelector } from '@/context/app-context' +import { currentWorkspaceIdAtom } from '@/context/app-context-state' const useGetIcon = () => { - const currentWorkspace = useSelector(s => s.currentWorkspace) + const currentWorkspaceId = useAtomValue(currentWorkspaceIdAtom) const getIconUrl = useCallback((fileName: string) => { - return `${API_PREFIX}/workspaces/current/plugin/icon?tenant_id=${currentWorkspace.id}&filename=${fileName}` - }, [currentWorkspace.id]) + return `${API_PREFIX}/workspaces/current/plugin/icon?tenant_id=${currentWorkspaceId}&filename=${fileName}` + }, [currentWorkspaceId]) return { getIconUrl, diff --git a/web/app/components/plugins/install-plugin/hooks/__tests__/use-workspace-plugin-install-permission.spec.ts b/web/app/components/plugins/install-plugin/hooks/__tests__/use-workspace-plugin-install-permission.spec.ts index f3df9915058..f7ab89645e2 100644 --- a/web/app/components/plugins/install-plugin/hooks/__tests__/use-workspace-plugin-install-permission.spec.ts +++ b/web/app/components/plugins/install-plugin/hooks/__tests__/use-workspace-plugin-install-permission.spec.ts @@ -4,12 +4,26 @@ import useWorkspacePluginInstallPermission from '../use-workspace-plugin-install let mockWorkspacePermissionKeys: string[] = [] -vi.mock('@/context/app-context', () => ({ - useAppContext: () => ({ - langGeniusVersionInfo: { current_version: '1.0.0' }, +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + langGeniusVersionInfo: { + current_env: '', + current_version: '1.0.0', + latest_version: '', + release_date: '', + release_notes: '', + version: '', + can_auto_update: false, + }, workspacePermissionKeys: mockWorkspacePermissionKeys, - }), -})) + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) describe('useWorkspacePluginInstallPermission', () => { beforeEach(() => { diff --git a/web/app/components/plugins/install-plugin/hooks/use-workspace-plugin-install-permission.ts b/web/app/components/plugins/install-plugin/hooks/use-workspace-plugin-install-permission.ts index efaba878813..5e408886b23 100644 --- a/web/app/components/plugins/install-plugin/hooks/use-workspace-plugin-install-permission.ts +++ b/web/app/components/plugins/install-plugin/hooks/use-workspace-plugin-install-permission.ts @@ -1,12 +1,11 @@ +import { useAtomValue } from 'jotai' import { useMemo } from 'react' -import { useAppContext } from '@/context/app-context' +import { langGeniusVersionInfoAtom, workspacePermissionKeysAtom } from '@/context/app-context-state' import { hasPermission } from '@/utils/permission' const useWorkspacePluginInstallPermission = () => { - const { - langGeniusVersionInfo, - workspacePermissionKeys, - } = useAppContext() + const langGeniusVersionInfo = useAtomValue(langGeniusVersionInfoAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const canInstallPlugin = useMemo(() => { return hasPermission(workspacePermissionKeys, 'plugin.install') diff --git a/web/app/components/plugins/install-plugin/install-from-local-package/steps/__tests__/install.spec.tsx b/web/app/components/plugins/install-plugin/install-from-local-package/steps/__tests__/install.spec.tsx index 9c0e10d9059..3b2057b4c05 100644 --- a/web/app/components/plugins/install-plugin/install-from-local-package/steps/__tests__/install.spec.tsx +++ b/web/app/components/plugins/install-plugin/install-from-local-package/steps/__tests__/install.spec.tsx @@ -57,11 +57,22 @@ vi.mock('../../../base/check-task-status', () => ({ }), })) -const mockLangGeniusVersionInfo = { current_version: '1.0.0' } -vi.mock('@/context/app-context', () => ({ - useAppContext: () => ({ - langGeniusVersionInfo: mockLangGeniusVersionInfo, - }), +const mockAppContextState = vi.hoisted(() => ({ + langGeniusVersionInfoAtom: Symbol('langGeniusVersionInfoAtom'), + langGeniusVersionInfo: { current_version: '1.0.0' as string | undefined }, +})) + +vi.mock('@/context/app-context-state', () => ({ + langGeniusVersionInfoAtom: mockAppContextState.langGeniusVersionInfoAtom, +})) + +vi.mock('jotai', () => ({ + useAtomValue: (atom: unknown) => { + if (atom === mockAppContextState.langGeniusVersionInfoAtom) + return mockAppContextState.langGeniusVersionInfo + + throw new Error('Unexpected atom') + }, })) vi.mock('../../../../card', () => ({ @@ -466,7 +477,7 @@ describe('Install', () => { // ================================ describe('Dify Version Compatibility', () => { it('should not show warning when dify version is compatible', () => { - mockLangGeniusVersionInfo.current_version = '1.0.0' + mockAppContextState.langGeniusVersionInfo.current_version = '1.0.0' const payload = createMockManifest({ meta: { version: '1.0.0', minimum_dify_version: '0.8.0' } }) render() @@ -475,7 +486,7 @@ describe('Install', () => { }) it('should show warning when dify version is incompatible', () => { - mockLangGeniusVersionInfo.current_version = '1.0.0' + mockAppContextState.langGeniusVersionInfo.current_version = '1.0.0' const payload = createMockManifest({ meta: { version: '1.0.0', minimum_dify_version: '2.0.0' } }) render() @@ -484,7 +495,7 @@ describe('Install', () => { }) it('should be compatible when minimum_dify_version is undefined', () => { - mockLangGeniusVersionInfo.current_version = '1.0.0' + mockAppContextState.langGeniusVersionInfo.current_version = '1.0.0' const payload = createMockManifest({ meta: { version: '1.0.0' } }) render() @@ -493,7 +504,7 @@ describe('Install', () => { }) it('should be compatible when current_version is empty', () => { - mockLangGeniusVersionInfo.current_version = '' + mockAppContextState.langGeniusVersionInfo.current_version = '' const payload = createMockManifest({ meta: { version: '1.0.0', minimum_dify_version: '2.0.0' } }) render() @@ -503,7 +514,7 @@ describe('Install', () => { }) it('should be compatible when current_version is undefined', () => { - mockLangGeniusVersionInfo.current_version = undefined as unknown as string + mockAppContextState.langGeniusVersionInfo.current_version = undefined as unknown as string const payload = createMockManifest({ meta: { version: '1.0.0', minimum_dify_version: '2.0.0' } }) render() diff --git a/web/app/components/plugins/install-plugin/install-from-local-package/steps/install.tsx b/web/app/components/plugins/install-plugin/install-from-local-package/steps/install.tsx index 80a9eff7241..6fda49807a8 100644 --- a/web/app/components/plugins/install-plugin/install-from-local-package/steps/install.tsx +++ b/web/app/components/plugins/install-plugin/install-from-local-package/steps/install.tsx @@ -3,11 +3,12 @@ import type { FC } from 'react' import type { PluginDeclaration } from '../../../types' import { Button } from '@langgenius/dify-ui/button' import { RiLoader2Line } from '@remixicon/react' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useEffect, useMemo } from 'react' import { Trans, useTranslation } from 'react-i18next' import useCheckInstalled from '@/app/components/plugins/install-plugin/hooks/use-check-installed' -import { useAppContext } from '@/context/app-context' +import { langGeniusVersionInfoAtom } from '@/context/app-context-state' import { uninstallPlugin } from '@/service/plugins' import { useInstallPackageFromLocal, usePluginTaskList } from '@/service/use-plugins' import { isEqualOrLaterThanVersion } from '@/utils/semver' @@ -108,7 +109,7 @@ const Installed: FC = ({ } } - const { langGeniusVersionInfo } = useAppContext() + const langGeniusVersionInfo = useAtomValue(langGeniusVersionInfoAtom) const isDifyVersionCompatible = useMemo(() => { if (!langGeniusVersionInfo.current_version) return true diff --git a/web/app/components/plugins/install-plugin/install-from-marketplace/steps/__tests__/install.spec.tsx b/web/app/components/plugins/install-plugin/install-from-marketplace/steps/__tests__/install.spec.tsx index 222029ab013..7f878293e95 100644 --- a/web/app/components/plugins/install-plugin/install-from-marketplace/steps/__tests__/install.spec.tsx +++ b/web/app/components/plugins/install-plugin/install-from-marketplace/steps/__tests__/install.spec.tsx @@ -60,7 +60,10 @@ const mockStopTaskStatus = vi.fn() const mockHandleInstallTaskStart = vi.fn() let mockPluginDeclaration: { manifest: { meta: { minimum_dify_version: string } } } | undefined let mockCanInstall = true -let mockLangGeniusVersionInfo = { current_version: '1.0.0' } +const mockAppContextState = vi.hoisted(() => ({ + langGeniusVersionInfoAtom: Symbol('langGeniusVersionInfoAtom'), + langGeniusVersionInfo: { current_version: '1.0.0' as string | null }, +})) // Mock useCheckInstalled vi.mock('@/app/components/plugins/install-plugin/hooks/use-check-installed', () => ({ @@ -71,10 +74,17 @@ vi.mock('@/app/components/plugins/install-plugin/hooks/use-check-installed', () }), })) -vi.mock('@/context/app-context', () => ({ - useAppContext: () => ({ - langGeniusVersionInfo: mockLangGeniusVersionInfo, - }), +vi.mock('@/context/app-context-state', () => ({ + langGeniusVersionInfoAtom: mockAppContextState.langGeniusVersionInfoAtom, +})) + +vi.mock('jotai', () => ({ + useAtomValue: (atom: unknown) => { + if (atom === mockAppContextState.langGeniusVersionInfoAtom) + return mockAppContextState.langGeniusVersionInfo + + throw new Error('Unexpected atom') + }, })) // Mock service hooks @@ -104,7 +114,7 @@ vi.mock('../../../base/check-task-status', () => ({ vi.mock('@/app/components/plugins/install-plugin/hooks/use-plugin-install-permission', () => ({ default: () => ({ canInstallPlugin: true, - currentDifyVersion: mockLangGeniusVersionInfo.current_version, + currentDifyVersion: mockAppContextState.langGeniusVersionInfo.current_version, }), })) @@ -170,7 +180,7 @@ describe('Install Component (steps/install.tsx)', () => { mockIsLoading = false mockPluginDeclaration = undefined mockCanInstall = true - mockLangGeniusVersionInfo = { current_version: '1.0.0' } + mockAppContextState.langGeniusVersionInfo = { current_version: '1.0.0' } mockInstallPackageFromMarketPlace.mockResolvedValue({ all_installed: false, task_id: 'task-123', @@ -281,7 +291,7 @@ describe('Install Component (steps/install.tsx)', () => { }) it('should not show warning when dify version is compatible', () => { - mockLangGeniusVersionInfo = { current_version: '2.0.0' } + mockAppContextState.langGeniusVersionInfo = { current_version: '2.0.0' } mockPluginDeclaration = { manifest: { meta: { minimum_dify_version: '1.0.0' } }, } @@ -291,7 +301,7 @@ describe('Install Component (steps/install.tsx)', () => { }) it('should show warning when dify version is incompatible', () => { - mockLangGeniusVersionInfo = { current_version: '1.0.0' } + mockAppContextState.langGeniusVersionInfo = { current_version: '1.0.0' } mockPluginDeclaration = { manifest: { meta: { minimum_dify_version: '2.0.0' } }, } @@ -749,7 +759,7 @@ describe('Install Component (steps/install.tsx)', () => { }) it('should handle null current_version in langGeniusVersionInfo', () => { - mockLangGeniusVersionInfo = { current_version: null as unknown as string } + mockAppContextState.langGeniusVersionInfo = { current_version: null as unknown as string } mockPluginDeclaration = { manifest: { meta: { minimum_dify_version: '1.0.0' } }, } diff --git a/web/app/components/plugins/install-plugin/install-from-marketplace/steps/install.tsx b/web/app/components/plugins/install-plugin/install-from-marketplace/steps/install.tsx index b1c47be46bb..c8d934a39d4 100644 --- a/web/app/components/plugins/install-plugin/install-from-marketplace/steps/install.tsx +++ b/web/app/components/plugins/install-plugin/install-from-marketplace/steps/install.tsx @@ -3,11 +3,12 @@ import type { FC } from 'react' import type { InstallPackageResponse, Plugin, PluginManifestInMarket } from '../../../types' import { Button } from '@langgenius/dify-ui/button' import { RiLoader2Line } from '@remixicon/react' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useEffect, useMemo } from 'react' import { useTranslation } from 'react-i18next' import useCheckInstalled from '@/app/components/plugins/install-plugin/hooks/use-check-installed' -import { useAppContext } from '@/context/app-context' +import { langGeniusVersionInfoAtom } from '@/context/app-context-state' import { useInstallPackageFromMarketPlace, usePluginDeclarationFromMarketPlace, usePluginTaskList, useUpdatePackageFromMarketPlace } from '@/service/use-plugins' import { isEqualOrLaterThanVersion } from '@/utils/semver' import Card from '../../../card' @@ -133,7 +134,7 @@ const Installed: FC = ({ } } - const { langGeniusVersionInfo } = useAppContext() + const langGeniusVersionInfo = useAtomValue(langGeniusVersionInfoAtom) const { data: pluginDeclaration } = usePluginDeclarationFromMarketPlace(uniqueIdentifier) const isDifyVersionCompatible = useMemo(() => { if (!pluginDeclaration || !langGeniusVersionInfo.current_version) diff --git a/web/app/components/plugins/marketplace/marketplace-install-permission-provider.tsx b/web/app/components/plugins/marketplace/marketplace-install-permission-provider.tsx index ffa83c2a0ef..0e25810c6de 100644 --- a/web/app/components/plugins/marketplace/marketplace-install-permission-provider.tsx +++ b/web/app/components/plugins/marketplace/marketplace-install-permission-provider.tsx @@ -1,8 +1,9 @@ 'use client' import type { ReactNode } from 'react' +import { useAtomValue } from 'jotai' import { PluginInstallPermissionProvider } from '@/app/components/plugins/install-plugin/components/plugin-install-permission-provider' -import { useAppContext } from '@/context/app-context' +import { workspacePermissionKeysAtom } from '@/context/app-context-state' import { hasPermission } from '@/utils/permission' type MarketplaceInstallPermissionProviderProps = { @@ -12,7 +13,7 @@ type MarketplaceInstallPermissionProviderProps = { const MarketplaceInstallPermissionProvider = ({ children, }: MarketplaceInstallPermissionProviderProps) => { - const { workspacePermissionKeys } = useAppContext() + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const canInstallPlugin = hasPermission(workspacePermissionKeys, 'plugin.install') return ( diff --git a/web/app/components/plugins/plugin-auth/__tests__/plugin-auth-in-agent.spec.tsx b/web/app/components/plugins/plugin-auth/__tests__/plugin-auth-in-agent.spec.tsx index c7b66119c79..43b9ddf62dc 100644 --- a/web/app/components/plugins/plugin-auth/__tests__/plugin-auth-in-agent.spec.tsx +++ b/web/app/components/plugins/plugin-auth/__tests__/plugin-auth-in-agent.spec.tsx @@ -50,6 +50,20 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + userProfile: mockUserProfile, + isCurrentWorkspaceManager: mockIsCurrentWorkspaceManager(), + workspacePermissionKeys: ['credential.use', 'credential.create', 'credential.manage'], + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/hooks/use-oauth', () => ({ openOAuthPopup: vi.fn(), })) diff --git a/web/app/components/plugins/plugin-auth/__tests__/plugin-auth.spec.tsx b/web/app/components/plugins/plugin-auth/__tests__/plugin-auth.spec.tsx index b7bd28875c5..a4775a37e41 100644 --- a/web/app/components/plugins/plugin-auth/__tests__/plugin-auth.spec.tsx +++ b/web/app/components/plugins/plugin-auth/__tests__/plugin-auth.spec.tsx @@ -30,6 +30,18 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockAppContext.workspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/modal-context', () => ({ useModalContext: () => ({ setShowAccountSettingModal: mockSetShowAccountSettingModal, diff --git a/web/app/components/plugins/plugin-auth/authorize/__tests__/index.spec.tsx b/web/app/components/plugins/plugin-auth/authorize/__tests__/index.spec.tsx index 64813aac2a4..1c1f56e0b74 100644 --- a/web/app/components/plugins/plugin-auth/authorize/__tests__/index.spec.tsx +++ b/web/app/components/plugins/plugin-auth/authorize/__tests__/index.spec.tsx @@ -71,6 +71,18 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockAppContext.workspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + // Mock service/use-triggers - API service vi.mock('@/service/use-triggers', () => ({ useTriggerPluginDynamicOptions: () => ({ diff --git a/web/app/components/plugins/plugin-auth/authorized/__tests__/index.spec.tsx b/web/app/components/plugins/plugin-auth/authorized/__tests__/index.spec.tsx index 2ab81fceb94..00a50f9b06d 100644 --- a/web/app/components/plugins/plugin-auth/authorized/__tests__/index.spec.tsx +++ b/web/app/components/plugins/plugin-auth/authorized/__tests__/index.spec.tsx @@ -86,6 +86,19 @@ vi.mock('@/context/app-context', () => ({ }) => unknown) => selector(mockAppContext), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + userProfile: mockAppContext.userProfile, + workspacePermissionKeys: mockAppContext.workspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + // Mock service/use-triggers vi.mock('@/service/use-triggers', () => ({ useTriggerPluginDynamicOptions: () => ({ diff --git a/web/app/components/plugins/plugin-auth/authorized/__tests__/item.spec.tsx b/web/app/components/plugins/plugin-auth/authorized/__tests__/item.spec.tsx index 58fbe6dfeac..d2b1c694105 100644 --- a/web/app/components/plugins/plugin-auth/authorized/__tests__/item.spec.tsx +++ b/web/app/components/plugins/plugin-auth/authorized/__tests__/item.spec.tsx @@ -5,16 +5,18 @@ import { beforeEach, describe, expect, it, vi } from 'vitest' import { CredentialTypeEnum } from '../../types' import Item from '../item' -// Item uses useAppContextWithSelector(state => state.userProfile) for the -// borrowed-row heuristic; provide a minimal mock so the selector resolves. -const mockUserProfile = { id: 'test-user', name: 'Test User', email: 'test@example.com', avatar_url: '' } -vi.mock('@/context/app-context', () => ({ - useSelector: (selector: (state: { userProfile: typeof mockUserProfile, workspacePermissionKeys: string[] }) => unknown) => - selector({ - userProfile: mockUserProfile, - workspacePermissionKeys: ['credential.use', 'credential.create', 'credential.manage'], - }), -})) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + userProfile: { id: 'test-user' }, + workspacePermissionKeys: ['credential.use', 'credential.create', 'credential.manage'], + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) // ==================== Test Utilities ==================== diff --git a/web/app/components/plugins/plugin-auth/authorized/item.tsx b/web/app/components/plugins/plugin-auth/authorized/item.tsx index 03f8695449b..8db0a72f15e 100644 --- a/web/app/components/plugins/plugin-auth/authorized/item.tsx +++ b/web/app/components/plugins/plugin-auth/authorized/item.tsx @@ -6,6 +6,7 @@ import { Tooltip, TooltipContent, TooltipTrigger } from '@langgenius/dify-ui/too import { RiInformationLine, } from '@remixicon/react' +import { useAtomValue } from 'jotai' import { memo, useMemo, @@ -15,7 +16,7 @@ import { useTranslation } from 'react-i18next' import ActionButton from '@/app/components/base/action-button' import Badge from '@/app/components/base/badge' import Input from '@/app/components/base/input' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { userProfileIdAtom } from '@/context/app-context-state' import { useCredentialPermissions } from '@/hooks/use-credential-permissions' import { CredentialTypeEnum } from '../types' @@ -58,14 +59,14 @@ const Item = ({ const { canUseCredential, canManageCredential } = useCredentialPermissions() const isOAuth = credential.credential_type === CredentialTypeEnum.OAUTH2 const isPersonal = credential.visibility === 'only_me' - const userProfile = useAppContextWithSelector(state => state.userProfile) + const currentUserId = useAtomValue(userProfileIdAtom) // Borrowed-from-teammate: the backend explicitly flagged this row as another member's // only_me credential, returned only because the current node still references it. // Fallback heuristic (created_by mismatch on a selected row) is kept for backends // that don't yet emit the flag. const isSelected = showSelectedIcon && selectedCredentialId === credential.id const isConfiguredByOther - = !!credential.created_by && !!userProfile?.id && credential.created_by !== userProfile.id + = !!credential.created_by && !!currentUserId && credential.created_by !== currentUserId const isBorrowed = !!credential.from_other_member || (isSelected && isConfiguredByOther && isPersonal) const showSwitchAwayHint = isBorrowed diff --git a/web/app/components/plugins/plugin-detail-panel/datasource-action-list.tsx b/web/app/components/plugins/plugin-detail-panel/datasource-action-list.tsx index e52ee795da0..87850a3ec5b 100644 --- a/web/app/components/plugins/plugin-detail-panel/datasource-action-list.tsx +++ b/web/app/components/plugins/plugin-detail-panel/datasource-action-list.tsx @@ -1,4 +1,3 @@ -// import { useAppContext } from '@/context/app-context' // import { Button } from '@langgenius/dify-ui/button' // import { StatusDot } from '@langgenius/dify-ui/status-dot' // import ToolItem from '@/app/components/tools/provider/tool-item' @@ -18,7 +17,6 @@ const ActionList = ({ detail, }: Props) => { const { t } = useTranslation() - // const { isCurrentWorkspaceManager } = useAppContext() // const providerBriefInfo = detail.declaration.datasource?.identity // const providerKey = `${detail.plugin_id}/${providerBriefInfo?.name}` const { data: dataSourceList } = useDataSourceList(true) diff --git a/web/app/components/plugins/plugin-item/__tests__/index.spec.tsx b/web/app/components/plugins/plugin-item/__tests__/index.spec.tsx index 9ff2592c7ed..50c24df5a14 100644 --- a/web/app/components/plugins/plugin-item/__tests__/index.spec.tsx +++ b/web/app/components/plugins/plugin-item/__tests__/index.spec.tsx @@ -55,13 +55,36 @@ vi.mock('@/app/components/plugins/install-plugin/hooks/use-refresh-plugin-list', })) const mockLangGeniusVersionInfo = vi.fn(() => ({ + current_env: '', current_version: '1.0.0', + latest_version: '', + release_date: '', + release_notes: '', + version: '', + can_auto_update: false, })) -vi.mock('@/context/app-context', () => ({ - useAppContext: () => ({ + +const createLangGeniusVersionInfo = (currentVersion: string) => ({ + current_env: '', + current_version: currentVersion, + latest_version: '', + release_date: '', + release_notes: '', + version: '', + can_auto_update: false, +}) + +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ langGeniusVersionInfo: mockLangGeniusVersionInfo(), - }), -})) + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) vi.mock('../action', () => ({ default: ({ onDelete, pluginName }: { onDelete: () => void, pluginName: string }) => ( @@ -162,7 +185,7 @@ describe('PluginItem', () => { mockTheme.mockReturnValue('light') mockCurrentPluginID.mockReturnValue(undefined) mockEnableMarketplace.mockReturnValue(true) - mockLangGeniusVersionInfo.mockReturnValue({ current_version: '1.0.0' }) + mockLangGeniusVersionInfo.mockReturnValue(createLangGeniusVersionInfo('1.0.0')) mockGetValueFromI18nObject.mockImplementation((obj: Record) => obj?.en_US || '') }) @@ -359,7 +382,7 @@ describe('PluginItem', () => { describe('Version Compatibility', () => { it('should show warning icon when Dify version is not compatible', () => { // Arrange - mockLangGeniusVersionInfo.mockReturnValue({ current_version: '0.3.0' }) + mockLangGeniusVersionInfo.mockReturnValue(createLangGeniusVersionInfo('0.3.0')) const plugin = createPluginDetail({ declaration: createPluginDeclaration({ meta: { version: '1.0.0', minimum_dify_version: '0.5.0' }, @@ -376,7 +399,7 @@ describe('PluginItem', () => { it('should not show warning when Dify version is compatible', () => { // Arrange - mockLangGeniusVersionInfo.mockReturnValue({ current_version: '1.0.0' }) + mockLangGeniusVersionInfo.mockReturnValue(createLangGeniusVersionInfo('1.0.0')) const plugin = createPluginDetail({ declaration: createPluginDeclaration({ meta: { version: '1.0.0', minimum_dify_version: '0.5.0' }, @@ -393,7 +416,7 @@ describe('PluginItem', () => { it('should handle missing current_version gracefully', () => { // Arrange - mockLangGeniusVersionInfo.mockReturnValue({ current_version: '' }) + mockLangGeniusVersionInfo.mockReturnValue(createLangGeniusVersionInfo('')) const plugin = createPluginDetail() // Act diff --git a/web/app/components/plugins/plugin-item/index.tsx b/web/app/components/plugins/plugin-item/index.tsx index 91a0a5a5ef2..ac3087503ed 100644 --- a/web/app/components/plugins/plugin-item/index.tsx +++ b/web/app/components/plugins/plugin-item/index.tsx @@ -11,12 +11,13 @@ import { RiLoginCircleLine, } from '@remixicon/react' import { useSuspenseQuery } from '@tanstack/react-query' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useCallback, useMemo } from 'react' import { useTranslation } from 'react-i18next' import useRefreshPluginList from '@/app/components/plugins/install-plugin/hooks/use-refresh-plugin-list' import { API_PREFIX } from '@/config' -import { useAppContext } from '@/context/app-context' +import { langGeniusVersionInfoAtom } from '@/context/app-context-state' import { systemFeaturesQueryOptions } from '@/features/system-features/client' import { useRenderI18nObject } from '@/hooks/use-i18n' import useTheme from '@/hooks/use-theme' @@ -69,7 +70,7 @@ const PluginItem: FC = ({ return [PluginSource.github, PluginSource.marketplace].includes(source) ? author : '' }, [source, author]) - const { langGeniusVersionInfo } = useAppContext() + const langGeniusVersionInfo = useAtomValue(langGeniusVersionInfoAtom) const isDifyVersionCompatible = useMemo(() => { if (!langGeniusVersionInfo.current_version) diff --git a/web/app/components/plugins/plugin-page/__tests__/index.spec.tsx b/web/app/components/plugins/plugin-page/__tests__/index.spec.tsx index 75c8bbf61c5..80ea6e03ab4 100644 --- a/web/app/components/plugins/plugin-page/__tests__/index.spec.tsx +++ b/web/app/components/plugins/plugin-page/__tests__/index.spec.tsx @@ -51,11 +51,42 @@ vi.mock('@/context/app-context', () => ({ useAppContext: () => ({ isCurrentWorkspaceManager: true, isCurrentWorkspaceOwner: false, - langGeniusVersionInfo: { current_version: '1.0.0' }, + langGeniusVersionInfo: { + current_env: 'CLOUD', + current_version: '1.0.0', + latest_version: '1.0.0', + version: '1.0.0', + release_date: '', + release_notes: '', + can_auto_update: false, + }, workspacePermissionKeys: ['plugin.install', 'plugin.delete', 'plugin.debug', 'plugin.plugin_preferences'], }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + isCurrentWorkspaceManager: true, + isCurrentWorkspaceOwner: false, + langGeniusVersionInfo: { + current_env: 'CLOUD', + current_version: '1.0.0', + latest_version: '1.0.0', + version: '1.0.0', + release_date: '', + release_notes: '', + can_auto_update: false, + }, + workspacePermissionKeys: ['plugin.install', 'plugin.delete', 'plugin.debug', 'plugin.plugin_preferences'], + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/service/use-plugins', () => ({ hasPluginPermission: (permission: string | undefined, isAdmin: boolean) => { if (!permission) diff --git a/web/app/components/plugins/plugin-page/__tests__/use-reference-setting.spec.ts b/web/app/components/plugins/plugin-page/__tests__/use-reference-setting.spec.ts index 01e635d3a62..699738e5127 100644 --- a/web/app/components/plugins/plugin-page/__tests__/use-reference-setting.spec.ts +++ b/web/app/components/plugins/plugin-page/__tests__/use-reference-setting.spec.ts @@ -1,20 +1,49 @@ // Import mocks for assertions +import type { AppContextStateMockState } from '@/__tests__/utils/mock-app-context-state' +import type { LangGeniusVersionResponse } from '@/models/common' import { toast } from '@langgenius/dify-ui/toast' import { waitFor } from '@testing-library/react' import { beforeEach, describe, expect, it, vi } from 'vitest' import { renderHookWithSystemFeatures as renderHook } from '@/__tests__/utils/mock-system-features' -import { useAppContext } from '@/context/app-context' import { useInvalidateReferenceSettings, useMutationPluginPermissionSettings, useMutationReferenceSettings, usePluginAutoUpgradeSettings, usePluginPermissionSettings } from '@/service/use-plugins' import { PermissionType, PluginCategoryEnum } from '../../types' import useReferenceSetting, { useCanInstallPluginFromMarketplace } from '../use-reference-setting' -vi.mock('@/context/app-context', async () => { - const actual = await vi.importActual('@/context/app-context') - return { - ...actual, - useAppContext: vi.fn(), +const defaultLangGeniusVersionInfo: LangGeniusVersionResponse = { + current_env: '', + current_version: '1.0.0', + latest_version: '', + release_date: '', + release_notes: '', + version: '', + can_auto_update: false, +} + +type MockAppContextState = Omit & { + langGeniusVersionInfo?: Partial +} + +let mockAppContextState: AppContextStateMockState = {} + +const setAppContextState = (state: MockAppContextState) => { + mockAppContextState = { + ...state, + langGeniusVersionInfo: { + ...defaultLangGeniusVersionInfo, + ...state.langGeniusVersionInfo, + }, } +} + +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppContextState) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) }) vi.mock('@/service/use-plugins', () => ({ @@ -33,12 +62,12 @@ describe('useReferenceSetting Hook', () => { toastSuccessSpy.mockClear() // Default mocks - vi.mocked(useAppContext).mockReturnValue({ + setAppContextState({ isCurrentWorkspaceManager: false, isCurrentWorkspaceOwner: false, langGeniusVersionInfo: { current_version: '1.0.0', latest_version: '', version: '' }, workspacePermissionKeys: [] as string[], - } as ReturnType) + }) vi.mocked(usePluginAutoUpgradeSettings).mockReturnValue({ data: { @@ -109,12 +138,12 @@ describe('useReferenceSetting Hook', () => { }) it('should allow install and debug when plugin permission keys are present', () => { - vi.mocked(useAppContext).mockReturnValue({ + setAppContextState({ isCurrentWorkspaceManager: false, isCurrentWorkspaceOwner: false, langGeniusVersionInfo: { current_version: '1.0.0', latest_version: '', version: '' }, workspacePermissionKeys: ['plugin.install', 'plugin.debug'], - } as ReturnType) + }) vi.mocked(usePluginPermissionSettings).mockReturnValue({ data: { install_permission: PermissionType.everyone, @@ -129,12 +158,12 @@ describe('useReferenceSetting Hook', () => { }) it('should allow debug for managers with legacy admin permission when RBAC is disabled', () => { - vi.mocked(useAppContext).mockReturnValue({ + setAppContextState({ isCurrentWorkspaceManager: true, isCurrentWorkspaceOwner: false, langGeniusVersionInfo: { current_version: '1.0.0', latest_version: '', version: '' }, workspacePermissionKeys: [] as string[], - } as ReturnType) + }) vi.mocked(usePluginPermissionSettings).mockReturnValue({ data: { @@ -150,12 +179,12 @@ describe('useReferenceSetting Hook', () => { }) it('should allow debug for owners with legacy admin permission when RBAC is disabled', () => { - vi.mocked(useAppContext).mockReturnValue({ + setAppContextState({ isCurrentWorkspaceManager: false, isCurrentWorkspaceOwner: true, langGeniusVersionInfo: { current_version: '1.0.0', latest_version: '', version: '' }, workspacePermissionKeys: [] as string[], - } as ReturnType) + }) vi.mocked(usePluginPermissionSettings).mockReturnValue({ data: { @@ -171,12 +200,12 @@ describe('useReferenceSetting Hook', () => { }) it('should allow debug for normal users when legacy debug permission is everyone and RBAC is disabled', () => { - vi.mocked(useAppContext).mockReturnValue({ + setAppContextState({ isCurrentWorkspaceManager: false, isCurrentWorkspaceOwner: false, langGeniusVersionInfo: { current_version: '1.0.0', latest_version: '', version: '' }, workspacePermissionKeys: ['plugin.install'], - } as ReturnType) + }) vi.mocked(usePluginPermissionSettings).mockReturnValue({ data: { @@ -194,12 +223,12 @@ describe('useReferenceSetting Hook', () => { }) it('should use plugin keys even when legacy admin permission is configured and RBAC is enabled', () => { - vi.mocked(useAppContext).mockReturnValue({ + setAppContextState({ isCurrentWorkspaceManager: false, isCurrentWorkspaceOwner: false, langGeniusVersionInfo: { current_version: '1.0.0', latest_version: '', version: '' }, workspacePermissionKeys: ['plugin.install', 'plugin.debug'], - } as ReturnType) + }) vi.mocked(usePluginPermissionSettings).mockReturnValue({ data: { @@ -217,12 +246,12 @@ describe('useReferenceSetting Hook', () => { }) it('should apply legacy noOne plugin permissions when RBAC is disabled', () => { - vi.mocked(useAppContext).mockReturnValue({ + setAppContextState({ isCurrentWorkspaceManager: true, isCurrentWorkspaceOwner: false, langGeniusVersionInfo: { current_version: '1.0.0', latest_version: '', version: '' }, workspacePermissionKeys: ['plugin.install', 'plugin.delete', 'plugin.debug'], - } as ReturnType) + }) vi.mocked(usePluginPermissionSettings).mockReturnValue({ data: { install_permission: PermissionType.noOne, @@ -245,12 +274,12 @@ describe('useReferenceSetting Hook', () => { describe('canSetPermissions', () => { it('should be true with plugin preferences permission when RBAC is disabled', () => { - vi.mocked(useAppContext).mockReturnValue({ + setAppContextState({ isCurrentWorkspaceManager: false, isCurrentWorkspaceOwner: false, langGeniusVersionInfo: { current_version: '1.0.0', latest_version: '', version: '' }, workspacePermissionKeys: ['plugin.plugin_preferences'], - } as ReturnType) + }) const { result } = renderHook(() => useReferenceSetting(PluginCategoryEnum.tool)) @@ -258,12 +287,12 @@ describe('useReferenceSetting Hook', () => { }) it('should be false when RBAC is enabled even with plugin preferences permission', () => { - vi.mocked(useAppContext).mockReturnValue({ + setAppContextState({ isCurrentWorkspaceManager: false, isCurrentWorkspaceOwner: true, langGeniusVersionInfo: { current_version: '1.0.0', latest_version: '', version: '' }, workspacePermissionKeys: ['plugin.plugin_preferences'], - } as ReturnType) + }) const { result } = renderHook(() => useReferenceSetting(PluginCategoryEnum.tool), { systemFeatures: { rbac_enabled: true }, @@ -274,12 +303,12 @@ describe('useReferenceSetting Hook', () => { }) it('should be false without plugin preferences permission', () => { - vi.mocked(useAppContext).mockReturnValue({ + setAppContextState({ isCurrentWorkspaceManager: true, isCurrentWorkspaceOwner: false, langGeniusVersionInfo: { current_version: '1.0.0', latest_version: '', version: '' }, workspacePermissionKeys: [] as string[], - } as ReturnType) + }) const { result } = renderHook(() => useReferenceSetting(PluginCategoryEnum.tool)) @@ -353,12 +382,12 @@ describe('useReferenceSetting Hook', () => { }) it('should keep permission key access available when reference setting data is still loading', () => { - vi.mocked(useAppContext).mockReturnValue({ + setAppContextState({ isCurrentWorkspaceManager: false, isCurrentWorkspaceOwner: false, langGeniusVersionInfo: { current_version: '1.0.0', latest_version: '', version: '' }, workspacePermissionKeys: ['plugin.install', 'plugin.debug'], - } as ReturnType) + }) vi.mocked(usePluginAutoUpgradeSettings).mockReturnValue({ data: undefined, } as unknown as ReturnType) @@ -371,13 +400,13 @@ describe('useReferenceSetting Hook', () => { }) it('should keep permission state loading while workspace permission keys are loading', () => { - vi.mocked(useAppContext).mockReturnValue({ + setAppContextState({ isCurrentWorkspaceManager: false, isCurrentWorkspaceOwner: false, isLoadingWorkspacePermissionKeys: true, langGeniusVersionInfo: { current_version: '1.0.0', latest_version: '', version: '' }, workspacePermissionKeys: [] as string[], - } as ReturnType) + }) const { result } = renderHook(() => useReferenceSetting(PluginCategoryEnum.tool)) @@ -386,13 +415,13 @@ describe('useReferenceSetting Hook', () => { }) it('should keep permission state loading while current workspace is loading', () => { - vi.mocked(useAppContext).mockReturnValue({ + setAppContextState({ isCurrentWorkspaceManager: false, isCurrentWorkspaceOwner: false, isLoadingCurrentWorkspace: true, langGeniusVersionInfo: { current_version: '1.0.0', latest_version: '', version: '' }, workspacePermissionKeys: ['plugin.install'], - } as ReturnType) + }) const { result } = renderHook(() => useReferenceSetting(PluginCategoryEnum.tool)) @@ -402,7 +431,7 @@ describe('useReferenceSetting Hook', () => { describe('RBAC permissions', () => { it('should use workspace permission keys when RBAC is enabled', () => { - vi.mocked(useAppContext).mockReturnValue({ + setAppContextState({ isCurrentWorkspaceManager: false, isCurrentWorkspaceOwner: false, langGeniusVersionInfo: { current_version: '1.0.0', latest_version: '', version: '' }, @@ -412,7 +441,7 @@ describe('useReferenceSetting Hook', () => { 'plugin.debug', 'plugin.plugin_preferences', ], - } as ReturnType) + }) vi.mocked(usePluginPermissionSettings).mockReturnValue({ data: { install_permission: PermissionType.noOne, @@ -435,12 +464,12 @@ describe('useReferenceSetting Hook', () => { }) it('should ignore legacy plugin permission settings when RBAC is enabled', () => { - vi.mocked(useAppContext).mockReturnValue({ + setAppContextState({ isCurrentWorkspaceManager: true, isCurrentWorkspaceOwner: false, langGeniusVersionInfo: { current_version: '1.0.0', latest_version: '', version: '' }, workspacePermissionKeys: [] as string[], - } as ReturnType) + }) const { result } = renderHook(() => useReferenceSetting(PluginCategoryEnum.tool), { systemFeatures: { rbac_enabled: true }, @@ -462,12 +491,12 @@ describe('useCanInstallPluginFromMarketplace Hook', () => { beforeEach(() => { vi.clearAllMocks() - vi.mocked(useAppContext).mockReturnValue({ + setAppContextState({ isCurrentWorkspaceManager: true, isCurrentWorkspaceOwner: false, langGeniusVersionInfo: { current_version: '1.0.0', latest_version: '', version: '' }, workspacePermissionKeys: ['plugin.install'], - } as ReturnType) + }) vi.mocked(usePluginPermissionSettings).mockReturnValue({ data: { @@ -501,12 +530,12 @@ describe('useCanInstallPluginFromMarketplace Hook', () => { }) it('should return false without plugin.install', () => { - vi.mocked(useAppContext).mockReturnValue({ + setAppContextState({ isCurrentWorkspaceManager: true, isCurrentWorkspaceOwner: false, langGeniusVersionInfo: { current_version: '1.0.0', latest_version: '', version: '' }, workspacePermissionKeys: [] as string[], - } as ReturnType) + }) const { result } = renderHook(() => useCanInstallPluginFromMarketplace(), { systemFeatures: { enable_marketplace: true }, @@ -516,12 +545,12 @@ describe('useCanInstallPluginFromMarketplace Hook', () => { }) it('should return false when both marketplace is disabled and plugin.install is missing', () => { - vi.mocked(useAppContext).mockReturnValue({ + setAppContextState({ isCurrentWorkspaceManager: true, isCurrentWorkspaceOwner: false, langGeniusVersionInfo: { current_version: '1.0.0', latest_version: '', version: '' }, workspacePermissionKeys: [] as string[], - } as ReturnType) + }) const { result } = renderHook(() => useCanInstallPluginFromMarketplace(), { systemFeatures: { enable_marketplace: false }, @@ -558,12 +587,12 @@ describe('useCanInstallPluginFromMarketplace Hook', () => { }) it('should use plugin.install when marketplace and RBAC are enabled', () => { - vi.mocked(useAppContext).mockReturnValue({ + setAppContextState({ isCurrentWorkspaceManager: false, isCurrentWorkspaceOwner: false, langGeniusVersionInfo: { current_version: '1.0.0', latest_version: '', version: '' }, workspacePermissionKeys: ['plugin.install'], - } as ReturnType) + }) vi.mocked(usePluginPermissionSettings).mockReturnValue({ data: { install_permission: PermissionType.noOne, diff --git a/web/app/components/plugins/plugin-page/use-reference-setting.ts b/web/app/components/plugins/plugin-page/use-reference-setting.ts index 681a791ffd4..545e7609b0e 100644 --- a/web/app/components/plugins/plugin-page/use-reference-setting.ts +++ b/web/app/components/plugins/plugin-page/use-reference-setting.ts @@ -1,16 +1,23 @@ import type { PluginCategoryEnum } from '../types' import { toast } from '@langgenius/dify-ui/toast' import { useSuspenseQuery } from '@tanstack/react-query' +import { useAtomValue } from 'jotai' import { useMemo } from 'react' import { useTranslation } from 'react-i18next' -import { useAppContext } from '@/context/app-context' +import { + currentWorkspaceLoadingAtom, + langGeniusVersionInfoAtom, + workspacePermissionKeysAtom, + workspacePermissionKeysLoadingAtom, + workspaceRoleFlagsAtom, +} from '@/context/app-context-state' import { systemFeaturesQueryOptions } from '@/features/system-features/client' import { useInvalidateReferenceSettings, useMutationPluginPermissionSettings, useMutationReferenceSettings, usePluginAutoUpgradeSettings, usePluginPermissionSettings } from '@/service/use-plugins' import { hasPermission } from '@/utils/permission' import { hasLegacyPluginPermissionAccess } from '../plugin-permissions' const useCanSetPluginSettings = () => { - const { workspacePermissionKeys } = useAppContext() + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const { data: rbacEnabled } = useSuspenseQuery({ ...systemFeaturesQueryOptions(), select: s => s.rbac_enabled, @@ -28,11 +35,11 @@ export const usePluginSettingsAccess = () => { const { isCurrentWorkspaceManager, isCurrentWorkspaceOwner, - isLoadingCurrentWorkspace, - isLoadingWorkspacePermissionKeys, - workspacePermissionKeys, - langGeniusVersionInfo, - } = useAppContext() + } = useAtomValue(workspaceRoleFlagsAtom) + const isLoadingCurrentWorkspace = useAtomValue(currentWorkspaceLoadingAtom) + const isLoadingWorkspacePermissionKeys = useAtomValue(workspacePermissionKeysLoadingAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) + const langGeniusVersionInfo = useAtomValue(langGeniusVersionInfoAtom) const { data: rbacEnabled } = useSuspenseQuery({ ...systemFeaturesQueryOptions(), select: s => s.rbac_enabled, @@ -125,7 +132,11 @@ export const useCanInstallPluginFromMarketplace = () => { const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) const marketplaceAccess = systemFeatures.enable_marketplace const rbacEnabled = systemFeatures.rbac_enabled - const { isCurrentWorkspaceManager, isCurrentWorkspaceOwner, workspacePermissionKeys } = useAppContext() + const { + isCurrentWorkspaceManager, + isCurrentWorkspaceOwner, + } = useAtomValue(workspaceRoleFlagsAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const permissionQuery = usePluginPermissionSettings() const { data: permissions } = permissionQuery const legacyCanInstallPlugin = hasLegacyPluginPermissionAccess({ diff --git a/web/app/components/tools/hooks/use-tool-permissions.ts b/web/app/components/tools/hooks/use-tool-permissions.ts index 2d393007a29..b5d39bf7aae 100644 --- a/web/app/components/tools/hooks/use-tool-permissions.ts +++ b/web/app/components/tools/hooks/use-tool-permissions.ts @@ -1,16 +1,17 @@ 'use client' -import { useSelector as useAppContextSelector } from '@/context/app-context' +import { useAtomValue } from 'jotai' +import { workspacePermissionKeysAtom } from '@/context/app-context-state' import { hasPermission } from '@/utils/permission' export const useCanManageTools = () => { - const workspacePermissionKeys = useAppContextSelector(state => state.workspacePermissionKeys) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) return hasPermission(workspacePermissionKeys, 'tool.manage') } export const useCanManageMCP = () => { - const workspacePermissionKeys = useAppContextSelector(state => state.workspacePermissionKeys) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) return hasPermission(workspacePermissionKeys, 'mcp.manage') } diff --git a/web/app/components/tools/mcp/__tests__/create-card.spec.tsx b/web/app/components/tools/mcp/__tests__/create-card.spec.tsx index 0d8c0b9bc2f..daa2a273000 100644 --- a/web/app/components/tools/mcp/__tests__/create-card.spec.tsx +++ b/web/app/components/tools/mcp/__tests__/create-card.spec.tsx @@ -40,14 +40,30 @@ vi.mock('../modal', () => ({ }, })) -let mockWorkspacePermissionKeys: string[] = ['mcp.manage'] +const mockAppContextState = vi.hoisted(() => ({ + workspacePermissionKeys: ['mcp.manage'] as string[], + workspacePermissionKeysAtom: Symbol('workspacePermissionKeysAtom'), +})) vi.mock('@/context/app-context', () => ({ useSelector: (selector: (state: { workspacePermissionKeys: string[] }) => unknown) => selector({ - workspacePermissionKeys: mockWorkspacePermissionKeys, + workspacePermissionKeys: mockAppContextState.workspacePermissionKeys, }), })) +vi.mock('@/context/app-context-state', () => ({ + workspacePermissionKeysAtom: mockAppContextState.workspacePermissionKeysAtom, +})) + +vi.mock('jotai', () => ({ + useAtomValue: (atom: unknown) => { + if (atom === mockAppContextState.workspacePermissionKeysAtom) + return mockAppContextState.workspacePermissionKeys + + throw new Error('Unexpected atom') + }, +})) + // Mock the plugins service vi.mock('@/service/use-plugins', () => ({ useInstalledPluginList: () => ({ @@ -84,7 +100,7 @@ describe('NewMCPCard', () => { beforeEach(() => { mockCreateMCP.mockClear() - mockWorkspacePermissionKeys = ['mcp.manage'] + mockAppContextState.workspacePermissionKeys = ['mcp.manage'] }) describe('Rendering', () => { @@ -152,7 +168,7 @@ describe('NewMCPCard', () => { describe('mcp.manage Permission', () => { it('should not render card when user lacks mcp.manage', () => { - mockWorkspacePermissionKeys = [] + mockAppContextState.workspacePermissionKeys = [] render(, { wrapper: createWrapper() }) @@ -160,7 +176,7 @@ describe('NewMCPCard', () => { }) it('should not render toolbar button when user lacks mcp.manage', () => { - mockWorkspacePermissionKeys = [] + mockAppContextState.workspacePermissionKeys = [] render(, { wrapper: createWrapper() }) diff --git a/web/app/components/tools/mcp/__tests__/index.spec.tsx b/web/app/components/tools/mcp/__tests__/index.spec.tsx index 6352bb146d5..0b55d4d4310 100644 --- a/web/app/components/tools/mcp/__tests__/index.spec.tsx +++ b/web/app/components/tools/mcp/__tests__/index.spec.tsx @@ -15,7 +15,10 @@ const mockRefetch = vi.fn() const mockUseAllToolProviders = vi.fn() let mockProviders: MockProvider[] = [] let mockIsLoadingToolProviders = false -let mockWorkspacePermissionKeys = ['mcp.manage'] +const mockAppContextState = vi.hoisted(() => ({ + workspacePermissionKeys: ['mcp.manage'] as string[], + workspacePermissionKeysAtom: Symbol('workspacePermissionKeysAtom'), +})) vi.mock('@/service/use-tools', () => ({ useAllToolProviders: (enabled?: boolean) => { @@ -30,10 +33,23 @@ vi.mock('@/service/use-tools', () => ({ vi.mock('@/context/app-context', () => ({ useSelector: (selector: (state: { workspacePermissionKeys: string[] }) => unknown) => selector({ - workspacePermissionKeys: mockWorkspacePermissionKeys, + workspacePermissionKeys: mockAppContextState.workspacePermissionKeys, }), })) +vi.mock('@/context/app-context-state', () => ({ + workspacePermissionKeysAtom: mockAppContextState.workspacePermissionKeysAtom, +})) + +vi.mock('jotai', () => ({ + useAtomValue: (atom: unknown) => { + if (atom === mockAppContextState.workspacePermissionKeysAtom) + return mockAppContextState.workspacePermissionKeys + + throw new Error('Unexpected atom') + }, +})) + vi.mock('@/app/components/tools/provider/tool-card-skeleton', () => ({ default: ({ variant }: { variant?: string }) => ( <> @@ -89,7 +105,7 @@ describe('MCPList', () => { vi.useFakeTimers() mockProviders = [] mockIsLoadingToolProviders = false - mockWorkspacePermissionKeys = ['mcp.manage'] + mockAppContextState.workspacePermissionKeys = ['mcp.manage'] mockRefetch.mockResolvedValue(undefined) }) @@ -111,7 +127,7 @@ describe('MCPList', () => { }) it('should render providers read-only when user lacks mcp.manage', () => { - mockWorkspacePermissionKeys = [] + mockAppContextState.workspacePermissionKeys = [] mockProviders = [ { id: '1', name: 'Provider 1', type: 'mcp' }, ] diff --git a/web/app/components/tools/mcp/__tests__/provider-card.spec.tsx b/web/app/components/tools/mcp/__tests__/provider-card.spec.tsx index 6ff54b6ec1c..a399057b08a 100644 --- a/web/app/components/tools/mcp/__tests__/provider-card.spec.tsx +++ b/web/app/components/tools/mcp/__tests__/provider-card.spec.tsx @@ -81,14 +81,30 @@ vi.mock('../detail/operation-dropdown', () => ({ ), })) -let mockWorkspacePermissionKeys: string[] = ['mcp.manage'] +const mockAppContextState = vi.hoisted(() => ({ + workspacePermissionKeys: ['mcp.manage'] as string[], + workspacePermissionKeysAtom: Symbol('workspacePermissionKeysAtom'), +})) vi.mock('@/context/app-context', () => ({ useSelector: (selector: (state: { workspacePermissionKeys: string[] }) => unknown) => selector({ - workspacePermissionKeys: mockWorkspacePermissionKeys, + workspacePermissionKeys: mockAppContextState.workspacePermissionKeys, }), })) +vi.mock('@/context/app-context-state', () => ({ + workspacePermissionKeysAtom: mockAppContextState.workspacePermissionKeysAtom, +})) + +vi.mock('jotai', () => ({ + useAtomValue: (atom: unknown) => { + if (atom === mockAppContextState.workspacePermissionKeysAtom) + return mockAppContextState.workspacePermissionKeys + + throw new Error('Unexpected atom') + }, +})) + // Mock the format time hook vi.mock('@/hooks/use-format-time-from-now', () => ({ useFormatTimeFromNow: () => ({ @@ -155,7 +171,7 @@ describe('MCPCard', () => { mockDeleteMCP.mockClear() mockUpdateMCP.mockResolvedValue({ result: 'success' }) mockDeleteMCP.mockResolvedValue({ result: 'success' }) - mockWorkspacePermissionKeys = ['mcp.manage'] + mockAppContextState.workspacePermissionKeys = ['mcp.manage'] }) describe('Rendering', () => { @@ -343,7 +359,7 @@ describe('MCPCard', () => { }) it('should not render operation dropdown when user lacks mcp.manage', () => { - mockWorkspacePermissionKeys = [] + mockAppContextState.workspacePermissionKeys = [] render(, { wrapper: createWrapper() }) diff --git a/web/app/components/tools/mcp/detail/__tests__/content.spec.tsx b/web/app/components/tools/mcp/detail/__tests__/content.spec.tsx index 7e843277067..5dc235a6683 100644 --- a/web/app/components/tools/mcp/detail/__tests__/content.spec.tsx +++ b/web/app/components/tools/mcp/detail/__tests__/content.spec.tsx @@ -107,16 +107,31 @@ vi.mock('../tool-item', () => ({ ), })) -// Mutable workspace permission state -let mockWorkspacePermissionKeys: string[] = ['mcp.manage'] +const mockAppContextState = vi.hoisted(() => ({ + workspacePermissionKeys: ['mcp.manage'] as string[], + workspacePermissionKeysAtom: Symbol('workspacePermissionKeysAtom'), +})) // Mock the app context vi.mock('@/context/app-context', () => ({ useSelector: (selector: (state: { workspacePermissionKeys: string[] }) => unknown) => selector({ - workspacePermissionKeys: mockWorkspacePermissionKeys, + workspacePermissionKeys: mockAppContextState.workspacePermissionKeys, }), })) +vi.mock('@/context/app-context-state', () => ({ + workspacePermissionKeysAtom: mockAppContextState.workspacePermissionKeysAtom, +})) + +vi.mock('jotai', () => ({ + useAtomValue: (atom: unknown) => { + if (atom === mockAppContextState.workspacePermissionKeysAtom) + return mockAppContextState.workspacePermissionKeys + + throw new Error('Unexpected atom') + }, +})) + // Mock the plugins service vi.mock('@/service/use-plugins', () => ({ useInstalledPluginList: () => ({ @@ -195,7 +210,7 @@ describe('MCPDetailContent', () => { mockIsFetching = false mockIsUpdating = false mockIsAuthorizing = false - mockWorkspacePermissionKeys = ['mcp.manage'] + mockAppContextState.workspacePermissionKeys = ['mcp.manage'] }) describe('Rendering', () => { @@ -232,7 +247,7 @@ describe('MCPDetailContent', () => { }) it('should render read-only detail when user lacks mcp.manage', () => { - mockWorkspacePermissionKeys = [] + mockAppContextState.workspacePermissionKeys = [] render(, { wrapper: createWrapper() }) @@ -466,7 +481,7 @@ describe('MCPDetailContent', () => { }) it('should disable authorize action when user lacks mcp.manage', () => { - mockWorkspacePermissionKeys = [] + mockAppContextState.workspacePermissionKeys = [] const detail = createMockDetail({ is_team_authorization: false }) render( , @@ -756,7 +771,7 @@ describe('MCPDetailContent', () => { }) it('should not run OAuth authorization when user lacks mcp.manage', async () => { - mockWorkspacePermissionKeys = [] + mockAppContextState.workspacePermissionKeys = [] mockAuthorizeMcp.mockResolvedValue({ authorization_url: 'https://oauth.example.com' }) const detail = createMockDetail({ is_team_authorization: false }) @@ -800,7 +815,7 @@ describe('MCPDetailContent', () => { }) it('should disable authorized button when user lacks mcp.manage', () => { - mockWorkspacePermissionKeys = [] + mockAppContextState.workspacePermissionKeys = [] const detail = createMockDetail({ is_team_authorization: true }) render( , diff --git a/web/app/components/tools/mcp/hooks/use-mcp-service-card.ts b/web/app/components/tools/mcp/hooks/use-mcp-service-card.ts index 81d3b84302a..074bf83ee0c 100644 --- a/web/app/components/tools/mcp/hooks/use-mcp-service-card.ts +++ b/web/app/components/tools/mcp/hooks/use-mcp-service-card.ts @@ -2,9 +2,10 @@ import type { AppDetailResponse } from '@/models/app' import type { AppSSO } from '@/types/app' import { useQuery, useQueryClient } from '@tanstack/react-query' +import { useAtomValue } from 'jotai' import { useCallback, useMemo, useState } from 'react' import { BlockEnum } from '@/app/components/workflow/types' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { userProfileIdAtom, workspacePermissionKeysAtom } from '@/context/app-context-state' import { fetchAppDetail } from '@/service/apps' import { useInvalidateMCPServerDetail, @@ -36,8 +37,8 @@ export const useMCPServiceCardState = ( const { mutateAsync: updateMCPServer } = useUpdateMCPServer() const { mutateAsync: refreshMCPServerCode, isPending: genLoading } = useRefreshMCPServerCode() const invalidateMCPServerDetail = useInvalidateMCPServerDetail() - const currentUserId = useAppContextWithSelector(state => state.userProfile?.id) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const currentUserId = useAtomValue(userProfileIdAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const canManageMCP = useMemo( () => getAppACLCapabilities(appInfo.permission_keys, { diff --git a/web/app/components/tools/provider/__tests__/custom-create-card.spec.tsx b/web/app/components/tools/provider/__tests__/custom-create-card.spec.tsx index 0092eebdd6c..948366f2eaa 100644 --- a/web/app/components/tools/provider/__tests__/custom-create-card.spec.tsx +++ b/web/app/components/tools/provider/__tests__/custom-create-card.spec.tsx @@ -4,14 +4,30 @@ import { beforeEach, describe, expect, it, vi } from 'vitest' import { AuthType } from '../../types' import CustomCreateCard, { NewCustomToolButton } from '../custom-create-card' -let mockWorkspacePermissionKeys: string[] = ['tool.manage'] +const mockAppContextState = vi.hoisted(() => ({ + workspacePermissionKeys: ['tool.manage'] as string[], + workspacePermissionKeysAtom: Symbol('workspacePermissionKeysAtom'), +})) vi.mock('@/context/app-context', () => ({ useSelector: (selector: (state: { workspacePermissionKeys: string[] }) => T): T => selector({ - workspacePermissionKeys: mockWorkspacePermissionKeys, + workspacePermissionKeys: mockAppContextState.workspacePermissionKeys, }), })) +vi.mock('@/context/app-context-state', () => ({ + workspacePermissionKeysAtom: mockAppContextState.workspacePermissionKeysAtom, +})) + +vi.mock('jotai', () => ({ + useAtomValue: (atom: unknown) => { + if (atom === mockAppContextState.workspacePermissionKeysAtom) + return mockAppContextState.workspacePermissionKeys + + throw new Error('Unexpected atom') + }, +})) + // Mock useLocale and useDocLink vi.mock('@/context/i18n', () => ({ useLocale: () => 'en-US', @@ -81,7 +97,7 @@ describe('CustomCreateCard', () => { beforeEach(() => { vi.clearAllMocks() - mockWorkspacePermissionKeys = ['tool.manage'] + mockAppContextState.workspacePermissionKeys = ['tool.manage'] mockModalVisible = false mockCreateCustomCollection.mockResolvedValue({}) }) @@ -94,7 +110,7 @@ describe('CustomCreateCard', () => { }) it('should not render anything when user does not have tool.manage', () => { - mockWorkspacePermissionKeys = [] + mockAppContextState.workspacePermissionKeys = [] const { container } = render() @@ -146,7 +162,7 @@ describe('CustomCreateCard', () => { }) it('should not render toolbar add button when user does not have tool.manage', () => { - mockWorkspacePermissionKeys = [] + mockAppContextState.workspacePermissionKeys = [] const { container } = render() diff --git a/web/app/components/tools/provider/__tests__/detail.spec.tsx b/web/app/components/tools/provider/__tests__/detail.spec.tsx index 670683d194c..497df1398c9 100644 --- a/web/app/components/tools/provider/__tests__/detail.spec.tsx +++ b/web/app/components/tools/provider/__tests__/detail.spec.tsx @@ -15,6 +15,7 @@ vi.mock('@/i18n-config/language', () => ({ const mockIsCurrentWorkspaceManager = vi.fn(() => true) const mockAppContextState = vi.hoisted(() => ({ workspacePermissionKeys: ['tool.manage', 'credential.use', 'credential.create', 'credential.manage'] as string[], + workspacePermissionKeysAtom: Symbol('workspacePermissionKeysAtom'), })) vi.mock('@/context/app-context', () => ({ useAppContext: () => ({ @@ -25,6 +26,19 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', () => ({ + workspacePermissionKeysAtom: mockAppContextState.workspacePermissionKeysAtom, +})) + +vi.mock('jotai', () => ({ + useAtomValue: (atom: unknown) => { + if (atom === mockAppContextState.workspacePermissionKeysAtom) + return mockAppContextState.workspacePermissionKeys + + throw new Error('Unexpected atom') + }, +})) + const mockSetShowModelModal = vi.fn() vi.mock('@/context/modal-context', () => ({ useModalContext: () => ({ diff --git a/web/app/components/workflow/block-selector/__tests__/tool-picker.spec.tsx b/web/app/components/workflow/block-selector/__tests__/tool-picker.spec.tsx index 85d3adccc68..ff7e6db519d 100644 --- a/web/app/components/workflow/block-selector/__tests__/tool-picker.spec.tsx +++ b/web/app/components/workflow/block-selector/__tests__/tool-picker.spec.tsx @@ -69,6 +69,18 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockWorkspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/hooks/use-theme', () => ({ default: vi.fn(), })) diff --git a/web/app/components/workflow/hooks/__tests__/use-node-plugin-installation.spec.ts b/web/app/components/workflow/hooks/__tests__/use-node-plugin-installation.spec.ts index 57652c54cce..9e80fb40f8a 100644 --- a/web/app/components/workflow/hooks/__tests__/use-node-plugin-installation.spec.ts +++ b/web/app/components/workflow/hooks/__tests__/use-node-plugin-installation.spec.ts @@ -21,6 +21,18 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockWorkspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/service/use-tools', () => ({ useAllBuiltInTools: (enabled: boolean) => mockBuiltInTools(enabled), useAllCustomTools: (enabled: boolean) => mockCustomTools(enabled), diff --git a/web/hooks/use-credential-permissions.spec.ts b/web/hooks/use-credential-permissions.spec.ts index 0f2cd554422..6d91b4b2299 100644 --- a/web/hooks/use-credential-permissions.spec.ts +++ b/web/hooks/use-credential-permissions.spec.ts @@ -1,13 +1,19 @@ import { renderHook } from '@testing-library/react' import { useCredentialPermissions } from './use-credential-permissions' -let mockWorkspacePermissionKeys: string[] | null = [] +let mockWorkspacePermissionKeys: string[] = [] -vi.mock('@/context/app-context', () => ({ - useSelector: (selector: (state: { workspacePermissionKeys: string[] | null }) => unknown) => selector({ +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ workspacePermissionKeys: mockWorkspacePermissionKeys, - }), -})) + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) describe('useCredentialPermissions', () => { beforeEach(() => { @@ -39,8 +45,8 @@ describe('useCredentialPermissions', () => { }) }) - it('should handle missing workspace permissions as no credential capabilities', () => { - mockWorkspacePermissionKeys = null + it('should handle empty workspace permissions as no credential capabilities', () => { + mockWorkspacePermissionKeys = [] const { result } = renderHook(() => useCredentialPermissions()) diff --git a/web/hooks/use-credential-permissions.ts b/web/hooks/use-credential-permissions.ts index 3f99c5e3345..c65f0e889db 100644 --- a/web/hooks/use-credential-permissions.ts +++ b/web/hooks/use-credential-permissions.ts @@ -1,8 +1,9 @@ -import { useSelector as useAppContextSelector } from '@/context/app-context' +import { useAtomValue } from 'jotai' +import { workspacePermissionKeysAtom } from '@/context/app-context-state' import { hasPermission } from '@/utils/permission' export const useCredentialPermissions = () => { - const workspacePermissionKeys = useAppContextSelector(state => state.workspacePermissionKeys) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) return { canUseCredential: hasPermission(workspacePermissionKeys, 'credential.use'), diff --git a/web/service/__tests__/use-plugins.spec.tsx b/web/service/__tests__/use-plugins.spec.tsx index 6d8374591a6..2c274da5a7b 100644 --- a/web/service/__tests__/use-plugins.spec.tsx +++ b/web/service/__tests__/use-plugins.spec.tsx @@ -19,9 +19,11 @@ import { const { mockGet, mockPost, + mockWorkspacePermissionKeysAtom, } = vi.hoisted(() => ({ mockGet: vi.fn(), mockPost: vi.fn(), + mockWorkspacePermissionKeysAtom: Symbol('workspacePermissionKeysAtom'), })) vi.mock('@/service/base', () => ({ @@ -37,12 +39,17 @@ vi.mock('@/app/components/plugins/install-plugin/hooks/use-refresh-plugin-list', }), })) -vi.mock('@/context/app-context', () => ({ - useAppContext: () => ({ - isCurrentWorkspaceManager: true, - isCurrentWorkspaceOwner: false, - workspacePermissionKeys: ['plugin.install'], - }), +vi.mock('@/context/app-context-state', () => ({ + workspacePermissionKeysAtom: mockWorkspacePermissionKeysAtom, +})) + +vi.mock('jotai', () => ({ + useAtomValue: (atom: unknown) => { + if (atom === mockWorkspacePermissionKeysAtom) + return ['plugin.install'] + + throw new Error('Unexpected atom') + }, })) vi.mock('../use-tools', () => ({ diff --git a/web/service/use-plugins.ts b/web/service/use-plugins.ts index 521f4093cec..07dc4db4f2e 100644 --- a/web/service/use-plugins.ts +++ b/web/service/use-plugins.ts @@ -40,12 +40,13 @@ import { useQueryClient, } from '@tanstack/react-query' import { cloneDeep } from 'es-toolkit/object' +import { useAtomValue } from 'jotai' import { useCallback, useEffect, useRef } from 'react' import { FormTypeEnum } from '@/app/components/base/form/types' import useRefreshPluginList from '@/app/components/plugins/install-plugin/hooks/use-refresh-plugin-list' import { getFormattedPlugin } from '@/app/components/plugins/marketplace/utils' import { PluginCategoryEnum, PluginSource, TaskStatus } from '@/app/components/plugins/types' -import { useAppContext } from '@/context/app-context' +import { workspacePermissionKeysAtom } from '@/context/app-context-state' import { fetchModelProviderModelList } from '@/service/common' import { fetchPluginInfoFromMarketPlace, uninstallPlugin } from '@/service/plugins' import { hasPermission } from '@/utils/permission' @@ -1233,7 +1234,7 @@ export const useFetchPluginsInMarketPlaceByInfo = (infos: MarketplacePluginInfoR export const usePluginTaskList = (category?: PluginCategoryEnum | string) => { const initializedRef = useRef(false) const queryClient = useQueryClient() - const { workspacePermissionKeys } = useAppContext() + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const canManagement = hasPermission(workspacePermissionKeys, 'plugin.install') const { refreshPluginList } = useRefreshPluginList() const query = useQuery({ From 3523da508f5b6970b447d346e2001967eb42e077 Mon Sep 17 00:00:00 2001 From: Yuzi <601709253@qq.com> Date: Wed, 8 Jul 2026 14:28:43 +0800 Subject: [PATCH 45/70] fix(web): guard invite-settings activate button against double-click (#38337) --- web/app/signin/invite-settings/page.tsx | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/web/app/signin/invite-settings/page.tsx b/web/app/signin/invite-settings/page.tsx index df48867a9b9..085443214d2 100644 --- a/web/app/signin/invite-settings/page.tsx +++ b/web/app/signin/invite-settings/page.tsx @@ -61,6 +61,7 @@ export default function InviteSettingsPage() { const token = decodeURIComponent(searchParams.get('invite_token') as string) const locale = useLocale() const [name, setName] = useState('') + const [isActivating, setIsActivating] = useState(false) const [language, setLanguage] = useState(() => getInitialLanguage(locale)) const [timezone, setTimezone] = useState(() => getBrowserTimezone() || 'America/Los_Angeles') const selectedLanguage = LANGUAGE_OPTIONS.find(item => item.value === language) @@ -93,6 +94,7 @@ export default function InviteSettingsPage() { toast.error(t('enterYourName', { ns: 'login' })) return } + setIsActivating(true) const body = requiresAccountSetup ? { token, @@ -118,8 +120,9 @@ export default function InviteSettingsPage() { } catch { recheck() + setIsActivating(false) } - }, [language, name, queryClient, recheck, requiresAccountSetup, searchParams, timezone, token, router, t]) + }, [isActivating, language, name, queryClient, recheck, requiresAccountSetup, searchParams, timezone, token, router, t]) if (!checkRes) return @@ -228,6 +231,8 @@ export default function InviteSettingsPage() { variant="primary" className="w-full" onClick={handleActivate} + loading={isActivating} + disabled={isActivating} > {`${t('join', { ns: 'login' })} ${checkRes?.data?.workspace_name}`} From 5a210cf03e7b67838b689dccd631dd59f57edd06 Mon Sep 17 00:00:00 2001 From: Stephen Zhou Date: Wed, 8 Jul 2026 15:02:47 +0800 Subject: [PATCH 46/70] refactor(web): migrate billing app context consumers (#38541) --- .../billing/billing-integration.test.tsx | 10 +++++++ .../billing/cloud-plan-payment-flow.test.tsx | 10 +++++++ .../education-verification-flow.test.tsx | 10 +++++++ .../billing/pricing-modal-flow.test.tsx | 10 +++++++ .../billing/self-hosted-plan-flow.test.tsx | 10 +++++++ web/__tests__/utils/mock-app-context-state.ts | 19 ++++++++++-- .../__tests__/index.spec.tsx | 15 +++++++++- .../billing/apps-full-in-dialog/index.tsx | 8 +++-- .../billing-page/__tests__/index.spec.tsx | 13 +++++++++ .../components/billing/billing-page/index.tsx | 6 ++-- .../billing/hooks/use-education-discount.ts | 5 ++-- .../billing/plan/__tests__/index.spec.tsx | 29 +++++++++++++++---- web/app/components/billing/plan/index.tsx | 8 +++-- .../billing/pricing/__tests__/dialog.spec.tsx | 18 +++++++++--- .../billing/pricing/__tests__/index.spec.tsx | 26 ++++++++++++----- web/app/components/billing/pricing/index.tsx | 5 ++-- .../cloud-plan-item/__tests__/index.spec.tsx | 29 +++++++++++++++---- .../pricing/plans/cloud-plan-item/index.tsx | 5 ++-- .../__tests__/index.spec.tsx | 21 ++++++++++++-- .../plans/self-hosted-plan-item/index.tsx | 5 ++-- web/context/app-context-state.ts | 12 ++++++++ 21 files changed, 230 insertions(+), 44 deletions(-) diff --git a/web/__tests__/billing/billing-integration.test.tsx b/web/__tests__/billing/billing-integration.test.tsx index d173392a420..5820e552b95 100644 --- a/web/__tests__/billing/billing-integration.test.tsx +++ b/web/__tests__/billing/billing-integration.test.tsx @@ -36,6 +36,16 @@ vi.mock('@/context/app-context', () => ({ useAppContext: () => mockAppCtx, })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppCtx) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/modal-context', () => ({ useModalContext: () => ({ setShowPricingModal: mockSetShowPricingModal, diff --git a/web/__tests__/billing/cloud-plan-payment-flow.test.tsx b/web/__tests__/billing/cloud-plan-payment-flow.test.tsx index d9e8d1d874a..e5e9ec4ed4e 100644 --- a/web/__tests__/billing/cloud-plan-payment-flow.test.tsx +++ b/web/__tests__/billing/cloud-plan-payment-flow.test.tsx @@ -28,6 +28,16 @@ vi.mock('@/context/app-context', () => ({ useAppContext: () => mockAppCtx, })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppCtx) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/i18n', () => ({ useGetLanguage: () => 'en-US', })) diff --git a/web/__tests__/billing/education-verification-flow.test.tsx b/web/__tests__/billing/education-verification-flow.test.tsx index 1b010964ce2..381bd37dafb 100644 --- a/web/__tests__/billing/education-verification-flow.test.tsx +++ b/web/__tests__/billing/education-verification-flow.test.tsx @@ -41,6 +41,16 @@ vi.mock('@/context/app-context', () => ({ useAppContext: () => mockAppCtx, })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppCtx) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/modal-context', () => ({ useModalContext: () => ({ setShowPricingModal: mockSetShowPricingModal, diff --git a/web/__tests__/billing/pricing-modal-flow.test.tsx b/web/__tests__/billing/pricing-modal-flow.test.tsx index dd1e710b1d6..2298105c651 100644 --- a/web/__tests__/billing/pricing-modal-flow.test.tsx +++ b/web/__tests__/billing/pricing-modal-flow.test.tsx @@ -28,6 +28,16 @@ vi.mock('@/context/app-context', () => ({ useAppContext: () => mockAppCtx, })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppCtx) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/i18n', () => ({ useGetLanguage: () => 'en-US', useGetPricingPageLanguage: () => 'en', diff --git a/web/__tests__/billing/self-hosted-plan-flow.test.tsx b/web/__tests__/billing/self-hosted-plan-flow.test.tsx index 08a3d0fed54..6fc2626ad63 100644 --- a/web/__tests__/billing/self-hosted-plan-flow.test.tsx +++ b/web/__tests__/billing/self-hosted-plan-flow.test.tsx @@ -24,6 +24,16 @@ vi.mock('@/context/app-context', () => ({ useAppContext: () => mockAppCtx, })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppCtx) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/i18n', () => ({ useGetLanguage: () => 'en-US', })) diff --git a/web/__tests__/utils/mock-app-context-state.ts b/web/__tests__/utils/mock-app-context-state.ts index 0c546de1471..5be87eee447 100644 --- a/web/__tests__/utils/mock-app-context-state.ts +++ b/web/__tests__/utils/mock-app-context-state.ts @@ -7,8 +7,8 @@ export type AppContextStateMockState = { id?: string name?: string email?: string - avatar?: string - avatar_url?: string + avatar?: string | null + avatar_url?: string | null is_password_set?: boolean } | null currentWorkspace?: { @@ -27,13 +27,16 @@ export type AppContextStateMockState = { type AppContextStateAtomKind = | 'userProfile' | 'userProfileId' + | 'userProfileEmail' | 'currentWorkspace' | 'currentWorkspaceId' | 'workspaceRoleFlags' + | 'isCurrentWorkspaceManager' | 'currentWorkspaceLoading' | 'workspacePermissionKeys' | 'workspacePermissionKeysLoading' | 'langGeniusVersionInfo' + | 'langGeniusCurrentVersion' type AppContextStateMockAtom = { [APP_CONTEXT_STATE_ATOM_KIND]: AppContextStateAtomKind @@ -101,13 +104,16 @@ export const createAppContextStateAtomMock = async ( ...actual, userProfileAtom: createMockAtom('userProfile'), userProfileIdAtom: createMockAtom('userProfileId'), + userProfileEmailAtom: createMockAtom('userProfileEmail'), currentWorkspaceAtom: createMockAtom('currentWorkspace'), currentWorkspaceIdAtom: createMockAtom('currentWorkspaceId'), workspaceRoleFlagsAtom: createMockAtom('workspaceRoleFlags'), + isCurrentWorkspaceManagerAtom: createMockAtom('isCurrentWorkspaceManager'), currentWorkspaceLoadingAtom: createMockAtom('currentWorkspaceLoading'), workspacePermissionKeysAtom: createMockAtom('workspacePermissionKeys'), workspacePermissionKeysLoadingAtom: createMockAtom('workspacePermissionKeysLoading'), langGeniusVersionInfoAtom: createMockAtom('langGeniusVersionInfo'), + langGeniusCurrentVersionAtom: createMockAtom('langGeniusCurrentVersion'), } } @@ -135,6 +141,9 @@ export const createAppContextStateJotaiMock = async ( if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'userProfileId') return userProfile.id + if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'userProfileEmail') + return userProfile.email + if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'currentWorkspace') return currentWorkspace @@ -150,6 +159,9 @@ export const createAppContextStateJotaiMock = async ( } } + if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'isCurrentWorkspaceManager') + return state.isCurrentWorkspaceManager ?? false + if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'currentWorkspaceLoading') return state.isLoadingCurrentWorkspace ?? false @@ -162,6 +174,9 @@ export const createAppContextStateJotaiMock = async ( if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'langGeniusVersionInfo') return state.langGeniusVersionInfo ?? defaultLangGeniusVersionInfo + if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'langGeniusCurrentVersion') + return (state.langGeniusVersionInfo ?? defaultLangGeniusVersionInfo).current_version + throw new Error(`Unsupported app context state atom: ${atom[APP_CONTEXT_STATE_ATOM_KIND]}`) }, } diff --git a/web/app/components/billing/apps-full-in-dialog/__tests__/index.spec.tsx b/web/app/components/billing/apps-full-in-dialog/__tests__/index.spec.tsx index ca0a355c547..72a35eb6843 100644 --- a/web/app/components/billing/apps-full-in-dialog/__tests__/index.spec.tsx +++ b/web/app/components/billing/apps-full-in-dialog/__tests__/index.spec.tsx @@ -11,6 +11,8 @@ import { useAppContext } from '@/context/app-context' import { baseProviderContextValue, useProviderContext } from '@/context/provider-context' import AppsFull from '../index' +let mockAppContextState: AppContextValue + vi.mock('@/config', async (importOriginal) => { const actual = await importOriginal() return { @@ -23,6 +25,16 @@ vi.mock('@/context/app-context', () => ({ useAppContext: vi.fn(), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppContextState) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/provider-context', async (importOriginal) => { const actual = await importOriginal() return { @@ -123,7 +135,8 @@ describe('AppsFull', () => { beforeEach(() => { vi.clearAllMocks() ;(useProviderContext as Mock).mockReturnValue(buildProviderContext()) - ;(useAppContext as Mock).mockReturnValue(buildAppContext()) + mockAppContextState = buildAppContext() + ;(useAppContext as Mock).mockReturnValue(mockAppContextState) ;(mailToSupport as Mock).mockReturnValue('mailto:support@example.com') }) diff --git a/web/app/components/billing/apps-full-in-dialog/index.tsx b/web/app/components/billing/apps-full-in-dialog/index.tsx index 486b2036317..51e76325a3d 100644 --- a/web/app/components/billing/apps-full-in-dialog/index.tsx +++ b/web/app/components/billing/apps-full-in-dialog/index.tsx @@ -4,11 +4,12 @@ import type { FC } from 'react' import { Button } from '@langgenius/dify-ui/button' import { cn } from '@langgenius/dify-ui/cn' import { MeterIndicator, MeterRoot, MeterTrack } from '@langgenius/dify-ui/meter' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useTranslation } from 'react-i18next' import { Plan } from '@/app/components/billing/type' import { mailToSupport } from '@/app/components/header/utils/util' -import { useAppContext } from '@/context/app-context' +import { langGeniusCurrentVersionAtom, userProfileEmailAtom } from '@/context/app-context-state' import { useProviderContext } from '@/context/provider-context' import UpgradeBtn from '../upgrade-btn' import s from './style.module.css' @@ -19,7 +20,8 @@ const AppsFull: FC<{ loc: string, className?: string }> = ({ }) => { const { t } = useTranslation() const { plan } = useProviderContext() - const { userProfile, langGeniusVersionInfo } = useAppContext() + const userProfileEmail = useAtomValue(userProfileEmailAtom) + const currentVersion = useAtomValue(langGeniusCurrentVersionAtom) const isTeam = plan.type === Plan.team const usage = plan.usage.buildApps const total = plan.total.buildApps @@ -54,7 +56,7 @@ const AppsFull: FC<{ loc: string, className?: string }> = ({ )} {plan.type !== Plan.sandbox && plan.type !== Plan.professional && ( diff --git a/web/app/components/billing/billing-page/__tests__/index.spec.tsx b/web/app/components/billing/billing-page/__tests__/index.spec.tsx index 07d6e62131a..25f93ff5c05 100644 --- a/web/app/components/billing/billing-page/__tests__/index.spec.tsx +++ b/web/app/components/billing/billing-page/__tests__/index.spec.tsx @@ -41,6 +41,19 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + isCurrentWorkspaceManager: isManager, + workspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/provider-context', () => ({ useProviderContext: () => ({ enableBilling, diff --git a/web/app/components/billing/billing-page/index.tsx b/web/app/components/billing/billing-page/index.tsx index d518810c668..a4454fbb370 100644 --- a/web/app/components/billing/billing-page/index.tsx +++ b/web/app/components/billing/billing-page/index.tsx @@ -1,8 +1,9 @@ 'use client' import type { FC } from 'react' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useTranslation } from 'react-i18next' -import { useAppContext } from '@/context/app-context' +import { isCurrentWorkspaceManagerAtom, workspacePermissionKeysAtom } from '@/context/app-context-state' import { useProviderContext } from '@/context/provider-context' import { useAsyncWindowOpen } from '@/hooks/use-async-window-open' import { useBillingUrl } from '@/service/use-billing' @@ -11,7 +12,8 @@ import PlanComp from '../plan' const Billing: FC = () => { const { t } = useTranslation() - const { isCurrentWorkspaceManager, workspacePermissionKeys } = useAppContext() + const isCurrentWorkspaceManager = useAtomValue(isCurrentWorkspaceManagerAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const { enableBilling } = useProviderContext() const canManageBillingSubscription = isCurrentWorkspaceManager && hasPermission(workspacePermissionKeys, BillingPermission.SubscriptionManage) const { data: billingUrl, isFetching, refetch } = useBillingUrl(enableBilling && canManageBillingSubscription) diff --git a/web/app/components/billing/hooks/use-education-discount.ts b/web/app/components/billing/hooks/use-education-discount.ts index 07977bf4f16..317be231003 100644 --- a/web/app/components/billing/hooks/use-education-discount.ts +++ b/web/app/components/billing/hooks/use-education-discount.ts @@ -1,15 +1,16 @@ 'use client' import { toast } from '@langgenius/dify-ui/toast' +import { useAtomValue } from 'jotai' import { useCallback, useState } from 'react' import { useTranslation } from 'react-i18next' -import { useAppContext } from '@/context/app-context' +import { workspacePermissionKeysAtom } from '@/context/app-context-state' import { fetchSubscriptionUrls } from '@/service/billing' import { BillingPermission, hasPermission } from '@/utils/permission' import { Plan } from '../type' export const useEducationDiscount = () => { const { t } = useTranslation() - const { workspacePermissionKeys } = useAppContext() + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const [isEducationDiscountLoading, setIsEducationDiscountLoading] = useState(false) const canManageBilling = hasPermission(workspacePermissionKeys, BillingPermission.Manage) diff --git a/web/app/components/billing/plan/__tests__/index.spec.tsx b/web/app/components/billing/plan/__tests__/index.spec.tsx index f062acba15e..9e0437ffd7f 100644 --- a/web/app/components/billing/plan/__tests__/index.spec.tsx +++ b/web/app/components/billing/plan/__tests__/index.spec.tsx @@ -21,11 +21,16 @@ vi.mock('@/next/navigation', () => ({ usePathname: () => currentPath, })) -vi.mock('@/config', () => ({ - get IS_CLOUD_EDITION() { - return mockConfig.isCloudEdition - }, -})) +vi.mock('@/config', async (importOriginal) => { + const actual = await importOriginal() + + return { + ...actual, + get IS_CLOUD_EDITION() { + return mockConfig.isCloudEdition + }, + } +}) const setShowAccountSettingModalMock = vi.fn() vi.mock('@/context/modal-context', () => ({ @@ -47,6 +52,20 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + userProfile: { email: 'user@example.com' }, + isCurrentWorkspaceManager, + workspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/app/education-apply/storage', () => ({ useSetEducationVerifying: () => setEducationVerifyingMock, })) diff --git a/web/app/components/billing/plan/index.tsx b/web/app/components/billing/plan/index.tsx index 775f6997f19..089841d0327 100644 --- a/web/app/components/billing/plan/index.tsx +++ b/web/app/components/billing/plan/index.tsx @@ -7,6 +7,7 @@ import { RiGroupLine, } from '@remixicon/react' import { useUnmountedRef } from 'ahooks' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useEffect } from 'react' import { useTranslation } from 'react-i18next' @@ -15,7 +16,7 @@ import UsageInfo from '@/app/components/billing/usage-info' import { useSetEducationVerifying } from '@/app/education-apply/storage' import VerifyStateModal from '@/app/education-apply/verify-state-modal' import { IS_CLOUD_EDITION } from '@/config' -import { useAppContext } from '@/context/app-context' +import { userProfileEmailAtom, workspacePermissionKeysAtom } from '@/context/app-context-state' import { useModalContextSelector } from '@/context/modal-context' import { useProviderContext } from '@/context/provider-context' import { usePathname, useRouter } from '@/next/navigation' @@ -41,7 +42,8 @@ const PlanComp: FC = ({ const { t } = useTranslation() const router = useRouter() const path = usePathname() - const { userProfile, workspacePermissionKeys } = useAppContext() + const userProfileEmail = useAtomValue(userProfileEmailAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const { plan, enableEducationPlan, allowRefreshEducationVerify, isEducationAccount } = useProviderContext() const isAboutToExpire = allowRefreshEducationVerify const { @@ -181,7 +183,7 @@ const PlanComp: FC = ({
    = {} vi.mock('@langgenius/dify-ui/dialog', () => ({ Dialog: ({ children, onOpenChange }: DialogProps) => { @@ -45,9 +45,19 @@ vi.mock('../footer', () => ({ })) vi.mock('@/context/app-context', () => ({ - useAppContext: vi.fn(), + useAppContext: () => mockAppCtx, })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppCtx) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/provider-context', () => ({ useProviderContext: vi.fn(), })) @@ -70,10 +80,10 @@ describe('Pricing dialog lifecycle', () => { beforeEach(() => { vi.clearAllMocks() latestOnOpenChange = undefined - ;(useAppContext as Mock).mockReturnValue({ + mockAppCtx = { isCurrentWorkspaceManager: true, workspacePermissionKeys: ['billing.manage'], - }) + } ;(useProviderContext as Mock).mockReturnValue({ plan: { type: Plan.sandbox, diff --git a/web/app/components/billing/pricing/__tests__/index.spec.tsx b/web/app/components/billing/pricing/__tests__/index.spec.tsx index 3a6d81aa0a2..1e5df7614d9 100644 --- a/web/app/components/billing/pricing/__tests__/index.spec.tsx +++ b/web/app/components/billing/pricing/__tests__/index.spec.tsx @@ -2,13 +2,13 @@ import type { Mock } from 'vitest' import type { UsagePlanInfo } from '../../type' import { fireEvent, render, screen } from '@testing-library/react' import * as React from 'react' -import { useAppContext } from '@/context/app-context' import { useGetPricingPageLanguage } from '@/context/i18n' import { useProviderContext } from '@/context/provider-context' import { Plan } from '../../type' import Pricing from '../index' let mockLanguage: string | null = 'en' +let mockAppCtx: Record = {} vi.mock('../plans/self-hosted-plan-item/list', () => ({ default: ({ plan }: { plan: string }) => ( @@ -28,9 +28,19 @@ vi.mock('@/next/link', () => ({ })) vi.mock('@/context/app-context', () => ({ - useAppContext: vi.fn(), + useAppContext: () => mockAppCtx, })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppCtx) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/provider-context', () => ({ useProviderContext: vi.fn(), })) @@ -53,10 +63,10 @@ describe('Pricing', () => { beforeEach(() => { vi.clearAllMocks() mockLanguage = 'en' - ;(useAppContext as Mock).mockReturnValue({ + mockAppCtx = { isCurrentWorkspaceManager: true, workspacePermissionKeys: ['billing.manage'], - }) + } ;(useProviderContext as Mock).mockReturnValue({ plan: { type: Plan.sandbox, @@ -79,10 +89,10 @@ describe('Pricing', () => { }) it('should default to yearly billing for education accounts', () => { - ;(useAppContext as Mock).mockReturnValue({ + mockAppCtx = { isCurrentWorkspaceManager: false, workspacePermissionKeys: ['billing.manage'], - }) + } ;(useProviderContext as Mock).mockReturnValue({ plan: { type: Plan.sandbox, @@ -99,10 +109,10 @@ describe('Pricing', () => { }) it('should not default to yearly billing when billing manage permission is missing', () => { - ;(useAppContext as Mock).mockReturnValue({ + mockAppCtx = { isCurrentWorkspaceManager: true, workspacePermissionKeys: [], - }) + } ;(useProviderContext as Mock).mockReturnValue({ plan: { type: Plan.sandbox, diff --git a/web/app/components/billing/pricing/index.tsx b/web/app/components/billing/pricing/index.tsx index 07dddc408bc..0640d63b92b 100644 --- a/web/app/components/billing/pricing/index.tsx +++ b/web/app/components/billing/pricing/index.tsx @@ -10,9 +10,10 @@ import { ScrollAreaThumb, ScrollAreaViewport, } from '@langgenius/dify-ui/scroll-area' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useState } from 'react' -import { useAppContext } from '@/context/app-context' +import { workspacePermissionKeysAtom } from '@/context/app-context-state' import { useGetPricingPageLanguage } from '@/context/i18n' import { useProviderContext } from '@/context/provider-context' import { BillingPermission, hasPermission } from '@/utils/permission' @@ -32,7 +33,7 @@ const Pricing: FC = ({ onCancel, }) => { const { plan, enableEducationPlan, isEducationAccount } = useProviderContext() - const { workspacePermissionKeys } = useAppContext() + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const canManageBilling = hasPermission(workspacePermissionKeys, BillingPermission.Manage) const shouldDefaultToYearly = canManageBilling && enableEducationPlan && isEducationAccount const [selectedPlanRange, setSelectedPlanRange] = React.useState() diff --git a/web/app/components/billing/pricing/plans/cloud-plan-item/__tests__/index.spec.tsx b/web/app/components/billing/pricing/plans/cloud-plan-item/__tests__/index.spec.tsx index c1cf7444d77..05e2e76adc7 100644 --- a/web/app/components/billing/pricing/plans/cloud-plan-item/__tests__/index.spec.tsx +++ b/web/app/components/billing/pricing/plans/cloud-plan-item/__tests__/index.spec.tsx @@ -12,10 +12,22 @@ import { Plan } from '../../../../type' import { PlanRange } from '../../../plan-switcher/plan-range-switcher' import CloudPlanItem from '../index' +let mockAppCtx: Record = {} + vi.mock('@/context/app-context', () => ({ useAppContext: vi.fn(), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppCtx) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/provider-context', () => ({ useProviderContext: vi.fn(), })) @@ -53,6 +65,11 @@ const mockFetchSubscriptionUrls = fetchSubscriptionUrls as Mock let assignedHref = '' const originalLocation = window.location +const mockAppContext = (state: Record) => { + mockAppCtx = state + mockUseAppContext.mockReturnValue(state) +} + const renderWithToastHost = (ui: React.ReactNode) => { return render( <> @@ -79,7 +96,7 @@ beforeAll(() => { beforeEach(() => { vi.clearAllMocks() toast.dismiss() - mockUseAppContext.mockReturnValue({ + mockAppContext({ isCurrentWorkspaceManager: true, workspacePermissionKeys: [ 'billing.view', @@ -183,7 +200,7 @@ describe('CloudPlanItem', () => { // Payment actions triggered from the CTA describe('Plan purchase flow', () => { it('should show toast when billing manage permission is missing for plan purchase', () => { - mockUseAppContext.mockReturnValue({ + mockAppContext({ isCurrentWorkspaceManager: true, workspacePermissionKeys: ['billing.subscription.manage'], }) @@ -206,7 +223,7 @@ describe('CloudPlanItem', () => { it('should open billing portal when upgrading current paid plan', async () => { const openWindow = vi.fn(async (cb: () => Promise) => await cb()) mockUseAsyncWindowOpen.mockReturnValue(openWindow) - mockUseAppContext.mockReturnValue({ + mockAppContext({ isCurrentWorkspaceManager: false, workspacePermissionKeys: ['billing.subscription.manage'], }) @@ -229,7 +246,7 @@ describe('CloudPlanItem', () => { }) it('should redirect to subscription url when selecting a new paid plan', async () => { - mockUseAppContext.mockReturnValue({ + mockAppContext({ isCurrentWorkspaceManager: false, workspacePermissionKeys: ['billing.manage'], }) @@ -316,7 +333,7 @@ describe('CloudPlanItem', () => { }) it('should show default CTA and hide warning when billing manage permission is missing', () => { - mockUseAppContext.mockReturnValue({ + mockAppContext({ isCurrentWorkspaceManager: true, workspacePermissionKeys: [], }) @@ -340,7 +357,7 @@ describe('CloudPlanItem', () => { }) it('should hide education unsupported warning when billing manage permission is missing', () => { - mockUseAppContext.mockReturnValue({ + mockAppContext({ isCurrentWorkspaceManager: true, workspacePermissionKeys: [], }) diff --git a/web/app/components/billing/pricing/plans/cloud-plan-item/index.tsx b/web/app/components/billing/pricing/plans/cloud-plan-item/index.tsx index 7da575a02bd..6dd974c12bd 100644 --- a/web/app/components/billing/pricing/plans/cloud-plan-item/index.tsx +++ b/web/app/components/billing/pricing/plans/cloud-plan-item/index.tsx @@ -10,10 +10,11 @@ import { DialogTitle, } from '@langgenius/dify-ui/dialog' import { toast } from '@langgenius/dify-ui/toast' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useMemo } from 'react' import { useTranslation } from 'react-i18next' -import { useAppContext } from '@/context/app-context' +import { workspacePermissionKeysAtom } from '@/context/app-context-state' import { useProviderContext } from '@/context/provider-context' import { useAsyncWindowOpen } from '@/hooks/use-async-window-open' import { fetchSubscriptionUrls } from '@/service/billing' @@ -56,7 +57,7 @@ const CloudPlanItem: FC = ({ const isCurrent = plan === currentPlan const isCurrentPaidPlan = isCurrent && !isFreePlan const isPlanDisabled = isCurrentPaidPlan ? false : planInfo.level <= ALL_PLANS[currentPlan].level - const { workspacePermissionKeys } = useAppContext() + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const canManageBilling = hasPermission(workspacePermissionKeys, BillingPermission.Manage) const canManageBillingSubscription = hasPermission(workspacePermissionKeys, BillingPermission.SubscriptionManage) const { enableEducationPlan, isEducationAccount } = useProviderContext() diff --git a/web/app/components/billing/pricing/plans/self-hosted-plan-item/__tests__/index.spec.tsx b/web/app/components/billing/pricing/plans/self-hosted-plan-item/__tests__/index.spec.tsx index c2923f98071..11c57d8ab67 100644 --- a/web/app/components/billing/pricing/plans/self-hosted-plan-item/__tests__/index.spec.tsx +++ b/web/app/components/billing/pricing/plans/self-hosted-plan-item/__tests__/index.spec.tsx @@ -7,6 +7,8 @@ import { contactSalesUrl, getStartedWithCommunityUrl, getWithPremiumUrl } from ' import { SelfHostedPlan } from '../../../../type' import SelfHostedPlanItem from '../index' +let mockAppCtx: Record = {} + vi.mock('../list', () => ({ default: ({ plan }: { plan: string }) => (
    @@ -20,6 +22,16 @@ vi.mock('@/context/app-context', () => ({ useAppContext: vi.fn(), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppCtx) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('../../../assets', () => ({ Community: () =>
    Community Icon
    , Premium: () =>
    Premium Icon
    , @@ -33,6 +45,11 @@ const mockUseAppContext = useAppContext as Mock let assignedHref = '' const originalLocation = window.location +const mockAppContext = (state: Record) => { + mockAppCtx = state + mockUseAppContext.mockReturnValue(state) +} + const renderWithToastHost = (ui: React.ReactNode) => { return render( <> @@ -59,7 +76,7 @@ beforeAll(() => { beforeEach(() => { vi.clearAllMocks() toast.dismiss() - mockUseAppContext.mockReturnValue({ + mockAppContext({ isCurrentWorkspaceManager: false, workspacePermissionKeys: ['billing.manage'], }) @@ -94,7 +111,7 @@ describe('SelfHostedPlanItem', () => { describe('CTA interactions', () => { it('should show toast when billing manage permission is missing', () => { - mockUseAppContext.mockReturnValue({ + mockAppContext({ isCurrentWorkspaceManager: true, workspacePermissionKeys: [], }) diff --git a/web/app/components/billing/pricing/plans/self-hosted-plan-item/index.tsx b/web/app/components/billing/pricing/plans/self-hosted-plan-item/index.tsx index 233f26ceb3f..b80fe3ff52d 100644 --- a/web/app/components/billing/pricing/plans/self-hosted-plan-item/index.tsx +++ b/web/app/components/billing/pricing/plans/self-hosted-plan-item/index.tsx @@ -2,11 +2,12 @@ import type { FC } from 'react' import { cn } from '@langgenius/dify-ui/cn' import { toast } from '@langgenius/dify-ui/toast' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useCallback } from 'react' import { useTranslation } from 'react-i18next' import { Azure, GoogleCloud } from '@/app/components/base/icons/src/public/billing' -import { useAppContext } from '@/context/app-context' +import { workspacePermissionKeysAtom } from '@/context/app-context-state' import { BillingPermission, hasPermission } from '@/utils/permission' import { contactSalesUrl, getStartedWithCommunityUrl, getWithPremiumUrl } from '../../../config' import { SelfHostedPlan } from '../../../type' @@ -52,7 +53,7 @@ const SelfHostedPlanItem: FC = ({ const isFreePlan = plan === SelfHostedPlan.community const isPremiumPlan = plan === SelfHostedPlan.premium const isEnterprisePlan = plan === SelfHostedPlan.enterprise - const { workspacePermissionKeys } = useAppContext() + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const canManageBilling = hasPermission(workspacePermissionKeys, BillingPermission.Manage) const handleGetPayUrl = useCallback(() => { diff --git a/web/context/app-context-state.ts b/web/context/app-context-state.ts index 04676f51099..2ad86521260 100644 --- a/web/context/app-context-state.ts +++ b/web/context/app-context-state.ts @@ -46,6 +46,10 @@ export const userProfileIdAtom = atom((get) => { return get(userProfileAtom).id }) +export const userProfileEmailAtom = atom((get) => { + return get(userProfileAtom).email +}) + const profileMetaAtom = atom((get) => { const accountProfileQuery = get(accountProfileQueryAtom) as SuspenseQueryResult @@ -81,6 +85,10 @@ export const isCurrentWorkspaceOwnerAtom = atom((get) => { return get(workspaceRoleFlagsAtom).isCurrentWorkspaceOwner }) +export const isCurrentWorkspaceManagerAtom = atom((get) => { + return get(workspaceRoleFlagsAtom).isCurrentWorkspaceManager +}) + const workspacePermissionKeysQueryAtom = atomWithQuery((get) => { const workspaceId = get(currentWorkspaceIdAtom) @@ -128,6 +136,10 @@ export const langGeniusVersionInfoAtom = atom((get) => { }) }) +export const langGeniusCurrentVersionAtom = atom((get) => { + return get(langGeniusVersionInfoAtom).current_version +}) + export const refreshUserProfileAtom = atom(null, (get) => { const queryClient = get(queryClientAtom) queryClient.invalidateQueries({ queryKey: userProfileQueryOptions().queryKey }) From 5d6131886053d3c8a0d7e35ecadf88b433e09a3a Mon Sep 17 00:00:00 2001 From: zyssyz123 <916125788@qq.com> Date: Wed, 8 Jul 2026 15:10:18 +0800 Subject: [PATCH 47/70] feat(api): use billing quota for credit pool (#38028) Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- api/services/billing_service.py | 115 ++++++++++-- api/services/credit_pool_service.py | 143 +++++++++++++-- .../services/test_billing_service.py | 126 ++++++++++++- .../services/test_credit_pool_service.py | 167 +++++++++++++++++- 4 files changed, 512 insertions(+), 39 deletions(-) diff --git a/api/services/billing_service.py b/api/services/billing_service.py index 2ee7179f432..4a829590f16 100644 --- a/api/services/billing_service.py +++ b/api/services/billing_service.py @@ -50,9 +50,26 @@ class QuotaReleaseResult(TypedDict): released: int +class QuotaBalanceResult(TypedDict): + available: int + reserved: int + quota: int + usage: int + + +class QuotaConsumeCappedResult(TypedDict): + deducted: int + available: int + reserved: int + quota: int + usage: int + + _quota_reserve_adapter = TypeAdapter(QuotaReserveResult) _quota_commit_adapter = TypeAdapter(QuotaCommitResult) _quota_release_adapter = TypeAdapter(QuotaReleaseResult) +_quota_balance_adapter = TypeAdapter(QuotaBalanceResult) +_quota_consume_capped_adapter = TypeAdapter(QuotaConsumeCappedResult) class _TenantFeatureQuota(TypedDict): @@ -176,6 +193,7 @@ class DismissNotificationDict(TypedDict): class BillingService: base_url = os.environ.get("BILLING_API_URL", "BILLING_API_URL") + quota_base_url = os.environ.get("BILLING_QUOTA_API_URL") or base_url secret_key = os.environ.get("BILLING_API_SECRET_KEY", "BILLING_API_SECRET_KEY") compliance_download_rate_limiter = RateLimiter("compliance_download_rate_limiter", 4, 60) @@ -215,12 +233,18 @@ class BillingService: def get_quota_info(cls, tenant_id: str) -> TenantFeatureQuotaInfo: params = {"tenant_id": tenant_id} return _tenant_feature_quota_info_adapter.validate_python( - cls._send_request("GET", "/quota/info", params=params) + cls._send_quota_request("GET", "/quota/info", params=params) ) @classmethod def quota_reserve( - cls, tenant_id: str, feature_key: str, request_id: str, amount: int = 1, meta: dict | None = None + cls, + tenant_id: str, + feature_key: str, + request_id: str, + amount: int = 1, + meta: dict | None = None, + bucket: str = "", ) -> QuotaReserveResult: """Reserve quota before task execution.""" payload: dict = { @@ -229,13 +253,21 @@ class BillingService: "request_id": request_id, "amount": amount, } + if bucket: + payload["bucket"] = bucket if meta: payload["meta"] = meta - return _quota_reserve_adapter.validate_python(cls._send_request("POST", "/quota/reserve", json=payload)) + return _quota_reserve_adapter.validate_python(cls._send_quota_request("POST", "/quota/reserve", json=payload)) @classmethod def quota_commit( - cls, tenant_id: str, feature_key: str, reservation_id: str, actual_amount: int, meta: dict | None = None + cls, + tenant_id: str, + feature_key: str, + reservation_id: str, + actual_amount: int, + meta: dict | None = None, + bucket: str = "", ) -> QuotaCommitResult: """Commit a reservation with actual consumption.""" payload: dict = { @@ -244,23 +276,57 @@ class BillingService: "reservation_id": reservation_id, "actual_amount": actual_amount, } + if bucket: + payload["bucket"] = bucket if meta: payload["meta"] = meta - return _quota_commit_adapter.validate_python(cls._send_request("POST", "/quota/commit", json=payload)) + return _quota_commit_adapter.validate_python(cls._send_quota_request("POST", "/quota/commit", json=payload)) @classmethod - def quota_release(cls, tenant_id: str, feature_key: str, reservation_id: str) -> QuotaReleaseResult: + def quota_release( + cls, tenant_id: str, feature_key: str, reservation_id: str, bucket: str = "" + ) -> QuotaReleaseResult: """Release a reservation (cancel, return frozen quota).""" - return _quota_release_adapter.validate_python( - cls._send_request( - "POST", - "/quota/release", - json={ - "tenant_id": tenant_id, - "feature_key": feature_key, - "reservation_id": reservation_id, - }, - ) + payload = { + "tenant_id": tenant_id, + "feature_key": feature_key, + "reservation_id": reservation_id, + } + if bucket: + payload["bucket"] = bucket + return _quota_release_adapter.validate_python(cls._send_quota_request("POST", "/quota/release", json=payload)) + + @classmethod + def quota_get_balance(cls, tenant_id: str, feature_key: str, bucket: str = "") -> QuotaBalanceResult: + """Get quota balance for a feature bucket.""" + params = {"tenant_id": tenant_id, "feature_key": feature_key} + if bucket: + params["bucket"] = bucket + return _quota_balance_adapter.validate_python(cls._send_quota_request("GET", "/quota/balance", params=params)) + + @classmethod + def quota_consume_capped( + cls, + tenant_id: str, + feature_key: str, + request_id: str, + amount: int, + meta: dict | None = None, + bucket: str = "", + ) -> QuotaConsumeCappedResult: + """Consume up to the available quota and return the actual deducted amount.""" + payload: dict = { + "tenant_id": tenant_id, + "feature_key": feature_key, + "request_id": request_id, + "amount": amount, + } + if bucket: + payload["bucket"] = bucket + if meta: + payload["meta"] = meta + return _quota_consume_capped_adapter.validate_python( + cls._send_quota_request("POST", "/quota/consume-capped", json=payload) ) @classmethod @@ -334,6 +400,12 @@ class BillingService: params = {"tenant_id": tenant_id, "feature_key": feature_key} return cls._send_request("GET", "/billing/tenant_feature_plan/usage", params=params) + @classmethod + def _send_quota_request( + cls, method: Literal["GET", "POST", "DELETE", "PUT"], endpoint: str, json=None, params=None + ): + return cls._send_request(method, endpoint, json=json, params=params, base_url=cls.quota_base_url) + @classmethod @retry( wait=wait_fixed(2), @@ -341,10 +413,17 @@ class BillingService: retry=retry_if_exception_type(httpx.RequestError), reraise=True, ) - def _send_request(cls, method: Literal["GET", "POST", "DELETE", "PUT"], endpoint: str, json=None, params=None): + def _send_request( + cls, + method: Literal["GET", "POST", "DELETE", "PUT"], + endpoint: str, + json=None, + params=None, + base_url: str | None = None, + ): headers = {"Content-Type": "application/json", "Billing-Api-Secret-Key": cls.secret_key} - url = f"{cls.base_url}{endpoint}" + url = f"{base_url or cls.base_url}{endpoint}" response = _http_client.request(method, url, json=json, params=params, headers=headers, follow_redirects=True) if method == "GET" and response.status_code != httpx.codes.OK: raise ValueError("Unable to retrieve billing information. Please try again later or contact support.") diff --git a/api/services/credit_pool_service.py b/api/services/credit_pool_service.py index afc49181185..837bf52c082 100644 --- a/api/services/credit_pool_service.py +++ b/api/services/credit_pool_service.py @@ -7,6 +7,8 @@ from piling up database transactions while preserving cross-tenant concurrency. import logging from collections.abc import Callable +from dataclasses import dataclass +from uuid import uuid4 from sqlalchemy import select from sqlalchemy.orm import Session @@ -19,11 +21,43 @@ from models.enums import ProviderQuotaType logger = logging.getLogger(__name__) +FEATURE_KEY_CREDIT_POOL = "credit_pool" CREDIT_POOL_TENANT_LOCK_TIMEOUT_SECONDS = 10 CREDIT_POOL_TENANT_LOCK_BLOCKING_TIMEOUT_SECONDS = 5 +@dataclass(frozen=True) +class CreditPoolBalance: + tenant_id: str + pool_type: str + quota_limit: int + quota_used: int + + @property + def remaining_credits(self) -> int: + if self.quota_limit == -1: + return -1 + return max(0, self.quota_limit - self.quota_used) + + def has_sufficient_credits(self, required_credits: int) -> bool: + return self.quota_limit == -1 or self.remaining_credits >= required_credits + + class CreditPoolService: + @staticmethod + def _normalize_pool_type(pool_type: str | ProviderQuotaType) -> str: + return pool_type.value if isinstance(pool_type, ProviderQuotaType) else str(pool_type) + + @staticmethod + def _use_billing_quota() -> bool: + return bool(dify_config.BILLING_ENABLED) + + @staticmethod + def _require_session(session: Session | None) -> Session: + if session is None: + raise ValueError("session is required when billing quota is disabled") + return session + @staticmethod def _get_tenant_lock_key(tenant_id: str) -> str: return f"credit_pool:tenant:{tenant_id}:deduct_lock" @@ -77,13 +111,36 @@ class CreditPoolService: return credit_pool @classmethod - def get_pool(cls, tenant_id: str, pool_type: str = "trial", *, session: Session) -> TenantCreditPool | None: + def get_pool( + cls, + tenant_id: str, + pool_type: str | ProviderQuotaType = "trial", + *, + session: Session | None = None, + ) -> TenantCreditPool | CreditPoolBalance | None: """get tenant credit pool""" + normalized_pool_type = cls._normalize_pool_type(pool_type) + if cls._use_billing_quota(): + from services.billing_service import BillingService + + balance = BillingService.quota_get_balance( + tenant_id=tenant_id, + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket=normalized_pool_type, + ) + return CreditPoolBalance( + tenant_id=tenant_id, + pool_type=normalized_pool_type, + quota_limit=balance["quota"], + quota_used=balance["usage"], + ) + + session = cls._require_session(session) return session.scalar( select(TenantCreditPool) .where( TenantCreditPool.tenant_id == tenant_id, - TenantCreditPool.pool_type == pool_type, + TenantCreditPool.pool_type == normalized_pool_type, ) .limit(1) ) @@ -93,31 +150,77 @@ class CreditPoolService: cls, tenant_id: str, credits_required: int, - pool_type: str = "trial", + pool_type: str | ProviderQuotaType = "trial", *, - session: Session, + session: Session | None = None, ) -> bool: """check if credits are available without deducting""" pool = cls.get_pool(tenant_id, pool_type, session=session) if not pool: return False - return pool.remaining_credits >= credits_required + return pool.has_sufficient_credits(credits_required) @classmethod def check_and_deduct_credits( cls, tenant_id: str, credits_required: int, - pool_type: str = "trial", + pool_type: str | ProviderQuotaType = "trial", *, - session: Session, + session: Session | None = None, ) -> int: """Deduct exactly the requested credits or raise without mutating the pool.""" if credits_required <= 0: return 0 + normalized_pool_type = cls._normalize_pool_type(pool_type) + + if cls._use_billing_quota(): + from services.billing_service import BillingService + + request_id = str(uuid4()) + result = BillingService.quota_reserve( + tenant_id=tenant_id, + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket=normalized_pool_type, + request_id=request_id, + amount=credits_required, + meta={"source": "credit_pool.check_and_deduct"}, + ) + reservation_id = result.get("reservation_id", "") + if not reservation_id: + raise QuotaExceededError("Insufficient credits remaining") + try: + BillingService.quota_commit( + tenant_id=tenant_id, + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket=normalized_pool_type, + reservation_id=reservation_id, + actual_amount=credits_required, + meta={"source": "credit_pool.check_and_deduct"}, + ) + except Exception: + try: + BillingService.quota_release( + tenant_id=tenant_id, + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket=normalized_pool_type, + reservation_id=reservation_id, + ) + except Exception: + logger.warning( + "Failed to release reserved credit pool quota, tenant_id=%s, pool_type=%s, reservation_id=%s", + tenant_id, + normalized_pool_type, + reservation_id, + exc_info=True, + ) + raise + return credits_required + + session = cls._require_session(session) def deduct() -> int: - pool = cls._get_locked_pool(session=session, tenant_id=tenant_id, pool_type=pool_type) + pool = cls._get_locked_pool(session=session, tenant_id=tenant_id, pool_type=normalized_pool_type) if not pool: raise QuotaExceededError("Credit pool not found") @@ -144,18 +247,34 @@ class CreditPoolService: cls, tenant_id: str, credits_required: int, - pool_type: str = "trial", + pool_type: str | ProviderQuotaType = "trial", *, - session: Session, + session: Session | None = None, ) -> int: """Deduct up to the available balance and return the actual deducted credits.""" if credits_required <= 0: return 0 + normalized_pool_type = cls._normalize_pool_type(pool_type) + + if cls._use_billing_quota(): + from services.billing_service import BillingService + + result = BillingService.quota_consume_capped( + tenant_id=tenant_id, + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket=normalized_pool_type, + request_id=str(uuid4()), + amount=credits_required, + meta={"source": "credit_pool.deduct_capped"}, + ) + return result["deducted"] + + session = cls._require_session(session) def deduct() -> int: - pool = cls._get_locked_pool(session=session, tenant_id=tenant_id, pool_type=pool_type) + pool = cls._get_locked_pool(session=session, tenant_id=tenant_id, pool_type=normalized_pool_type) if not pool: - logger.warning("Credit pool not found, tenant_id=%s, pool_type=%s", tenant_id, pool_type) + logger.warning("Credit pool not found, tenant_id=%s, pool_type=%s", tenant_id, normalized_pool_type) return 0 deducted_credits = min(credits_required, pool.remaining_credits) diff --git a/api/tests/unit_tests/services/test_billing_service.py b/api/tests/unit_tests/services/test_billing_service.py index dc691176114..2b1bdb5c5c4 100644 --- a/api/tests/unit_tests/services/test_billing_service.py +++ b/api/tests/unit_tests/services/test_billing_service.py @@ -73,6 +73,23 @@ class TestBillingServiceSendRequest: assert call_args[1]["headers"]["Billing-Api-Secret-Key"] == "test-secret-key" assert call_args[1]["headers"]["Content-Type"] == "application/json" + def test_send_request_with_base_url_override(self, mock_httpx_request, mock_billing_config): + """Quota APIs can use the new billing service without changing legacy billing calls.""" + # Arrange + expected_response = {"result": "success"} + mock_response = MagicMock() + mock_response.status_code = httpx.codes.OK + mock_response.json.return_value = expected_response + mock_httpx_request.return_value = mock_response + + # Act + result = BillingService._send_request("GET", "/quota/balance", base_url="https://quota.example.com") + + # Assert + assert result == expected_response + call_args = mock_httpx_request.call_args + assert call_args[0][1] == "https://quota.example.com/quota/balance" + @pytest.mark.parametrize( "status_code", [httpx.codes.NOT_FOUND, httpx.codes.INTERNAL_SERVER_ERROR, httpx.codes.BAD_REQUEST] ) @@ -393,6 +410,20 @@ class TestBillingServiceSubscriptionInfo: params={"tenant_id": tenant_id}, ) + 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: + mock_send_quota_request.return_value = {"quota": "200", "usage": "6", "available": "194", "reserved": "0"} + + result = BillingService.quota_get_balance(tenant_id, "credit_pool", bucket="trial") + + assert result == {"quota": 200, "usage": 6, "available": 194, "reserved": 0} + mock_send_quota_request.assert_called_once_with( + "GET", + "/quota/balance", + params={"tenant_id": tenant_id, "feature_key": "credit_pool", "bucket": "trial"}, + ) + def test_get_knowledge_rate_limit_with_defaults(self, mock_send_request): """Test knowledge rate limit retrieval with default values.""" # Arrange @@ -518,19 +549,20 @@ class TestBillingServiceUsageCalculation: assert result == expected_response mock_send_request.assert_called_once_with("GET", "/tenant-feature-usage/info", params={"tenant_id": tenant_id}) - def test_get_quota_info(self, mock_send_request): + def test_get_quota_info(self): """Test retrieval of quota info from new endpoint.""" # Arrange tenant_id = "tenant-123" expected_response = {"trigger_event": {"limit": 100, "usage": 30}, "api_rate_limit": {"limit": -1, "usage": 0}} - mock_send_request.return_value = expected_response + with patch.object(BillingService, "_send_quota_request") as mock_send_quota_request: + mock_send_quota_request.return_value = expected_response - # Act - result = BillingService.get_quota_info(tenant_id) + # Act + result = BillingService.get_quota_info(tenant_id) # Assert assert result == expected_response - mock_send_request.assert_called_once_with("GET", "/quota/info", params={"tenant_id": tenant_id}) + mock_send_quota_request.assert_called_once_with("GET", "/quota/info", params={"tenant_id": tenant_id}) def test_update_tenant_feature_plan_usage_positive_delta(self, mock_send_request): """Test updating tenant feature usage with positive delta (adding credits).""" @@ -614,7 +646,7 @@ class TestBillingServiceQuotaOperations: @pytest.fixture def mock_send_request(self): - with patch.object(BillingService, "_send_request") as mock: + with patch.object(BillingService, "_send_quota_request") as mock: yield mock def test_quota_reserve_success(self, mock_send_request): @@ -652,6 +684,16 @@ class TestBillingServiceQuotaOperations: call_json = mock_send_request.call_args[1]["json"] assert call_json["meta"] == {"source": "webhook"} + def test_quota_reserve_with_bucket(self, mock_send_request): + mock_send_request.return_value = {"reservation_id": "rid-2", "available": 98, "reserved": 1} + + BillingService.quota_reserve( + tenant_id="t1", feature_key="credit_pool", request_id="req-2", amount=1, bucket="trial" + ) + + call_json = mock_send_request.call_args[1]["json"] + assert call_json["bucket"] == "trial" + def test_quota_commit_success(self, mock_send_request): expected = {"available": 98, "reserved": 0, "refunded": 0} mock_send_request.return_value = expected @@ -696,6 +738,20 @@ class TestBillingServiceQuotaOperations: call_json = mock_send_request.call_args[1]["json"] assert call_json["meta"] == {"reason": "partial"} + def test_quota_commit_with_bucket(self, mock_send_request): + mock_send_request.return_value = {"available": 97, "reserved": 0, "refunded": 0} + + BillingService.quota_commit( + tenant_id="t1", + feature_key="credit_pool", + reservation_id="rid-1", + actual_amount=1, + bucket="paid", + ) + + call_json = mock_send_request.call_args[1]["json"] + assert call_json["bucket"] == "paid" + def test_quota_release_success(self, mock_send_request): expected = {"available": 100, "reserved": 0, "released": 1} mock_send_request.return_value = expected @@ -720,6 +776,64 @@ class TestBillingServiceQuotaOperations: assert result["released"] == 1 assert isinstance(result["released"], int) + def test_quota_release_with_bucket(self, mock_send_request): + mock_send_request.return_value = {"available": 100, "reserved": 0, "released": 1} + + BillingService.quota_release(tenant_id="t1", feature_key="credit_pool", reservation_id="rid-1", bucket="trial") + + call_json = mock_send_request.call_args[1]["json"] + assert call_json["bucket"] == "trial" + + def test_quota_consume_capped_success(self, mock_send_request): + mock_send_request.return_value = { + "deducted": "2", + "available": "8", + "reserved": "0", + "quota": "10", + "usage": "2", + } + + result = BillingService.quota_consume_capped( + tenant_id="t1", + feature_key="credit_pool", + request_id="req-1", + amount=5, + bucket="paid", + meta={"source": "test"}, + ) + + assert result == {"deducted": 2, "available": 8, "reserved": 0, "quota": 10, "usage": 2} + mock_send_request.assert_called_once_with( + "POST", + "/quota/consume-capped", + json={ + "tenant_id": "t1", + "feature_key": "credit_pool", + "request_id": "req-1", + "amount": 5, + "bucket": "paid", + "meta": {"source": "test"}, + }, + ) + + def test_send_quota_request_uses_quota_base_url(self): + with ( + patch.object(BillingService, "quota_base_url", "https://quota.example.com/v1"), + patch.object(BillingService, "_send_request") as mock_send_request, + ): + mock_send_request.return_value = {"ok": True} + + result = BillingService._send_quota_request("GET", "/quota/info", params={"tenant_id": "t1"}) + + assert result == {"ok": True} + mock_send_request.assert_called_once_with( + "GET", + "/quota/info", + json=None, + params={"tenant_id": "t1"}, + base_url="https://quota.example.com/v1", + ) + def test_get_quota_info_coerces_string_to_int(self, mock_send_request): """Test that TypeAdapter coerces string values to int for get_quota_info.""" mock_send_request.return_value = { diff --git a/api/tests/unit_tests/services/test_credit_pool_service.py b/api/tests/unit_tests/services/test_credit_pool_service.py index f31d067525a..8cafd3af590 100644 --- a/api/tests/unit_tests/services/test_credit_pool_service.py +++ b/api/tests/unit_tests/services/test_credit_pool_service.py @@ -1,5 +1,6 @@ +from collections.abc import Generator from types import SimpleNamespace -from unittest.mock import MagicMock, patch +from unittest.mock import ANY, MagicMock, patch from uuid import uuid4 import pytest @@ -13,6 +14,8 @@ from models.enums import ProviderQuotaType from services.credit_pool_service import ( CREDIT_POOL_TENANT_LOCK_BLOCKING_TIMEOUT_SECONDS, CREDIT_POOL_TENANT_LOCK_TIMEOUT_SECONDS, + FEATURE_KEY_CREDIT_POOL, + CreditPoolBalance, CreditPoolService, ) @@ -51,6 +54,12 @@ def _make_redis_lock() -> MagicMock: return lock +@pytest.fixture(autouse=True) +def _disable_billing_quota_by_default() -> Generator[None, None, None]: + with patch("services.credit_pool_service.dify_config.BILLING_ENABLED", False): + yield + + def test_get_pool_uses_provided_session() -> None: engine, tenant_id, _ = _create_engine_with_pool(quota_limit=10, quota_used=2) @@ -62,6 +71,13 @@ def test_get_pool_uses_provided_session() -> None: assert pool.quota_used == 2 +def test_credit_pool_balance_unlimited_remaining_and_sufficiency() -> None: + pool = CreditPoolBalance(tenant_id="tenant-1", pool_type="paid", quota_limit=-1, quota_used=999) + + assert pool.remaining_credits == -1 + assert pool.has_sufficient_credits(10_000) + + def test_check_and_deduct_credits_deducts_exact_amount_when_sufficient() -> None: engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=2) @@ -211,7 +227,7 @@ def test_check_and_deduct_credits_uses_tenant_redis_lock_before_db_deduction() - ) redis_lock.acquire.assert_called_once_with(blocking=True) redis_lock.release.assert_called_once_with() - get_locked_pool.assert_called_once_with(session=session, tenant_id=tenant_id, pool_type=ProviderQuotaType.TRIAL) + get_locked_pool.assert_called_once_with(session=session, tenant_id=tenant_id, pool_type="trial") def test_deduct_credits_capped_uses_tenant_redis_lock_before_db_deduction() -> None: @@ -240,7 +256,152 @@ def test_deduct_credits_capped_uses_tenant_redis_lock_before_db_deduction() -> N ) redis_lock.acquire.assert_called_once_with(blocking=True) redis_lock.release.assert_called_once_with() - get_locked_pool.assert_called_once_with(session=session, tenant_id=tenant_id, pool_type=ProviderQuotaType.PAID) + get_locked_pool.assert_called_once_with(session=session, tenant_id=tenant_id, pool_type="paid") + + +def test_get_pool_uses_billing_quota_balance_when_enabled() -> None: + tenant_id = "tenant-1" + with ( + patch("services.credit_pool_service.dify_config.BILLING_ENABLED", True), + patch("services.billing_service.BillingService.quota_get_balance") as quota_get_balance, + ): + quota_get_balance.return_value = {"quota": 1000, "usage": 250, "available": 750, "reserved": 0} + + pool = CreditPoolService.get_pool(tenant_id=tenant_id, pool_type=ProviderQuotaType.PAID) + + assert isinstance(pool, CreditPoolBalance) + assert pool.quota_limit == 1000 + assert pool.quota_used == 250 + assert pool.remaining_credits == 750 + quota_get_balance.assert_called_once_with( + tenant_id=tenant_id, + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket="paid", + ) + + +def test_check_and_deduct_credits_uses_billing_reserve_and_commit_when_enabled() -> None: + tenant_id = "tenant-1" + with ( + patch("services.credit_pool_service.dify_config.BILLING_ENABLED", True), + patch("services.billing_service.BillingService.quota_reserve") as quota_reserve, + patch("services.billing_service.BillingService.quota_commit") as quota_commit, + patch("services.billing_service.BillingService.quota_release") as quota_release, + ): + quota_reserve.return_value = {"reservation_id": "reservation-1", "available": 7, "reserved": 3} + + result = CreditPoolService.check_and_deduct_credits( + tenant_id=tenant_id, + credits_required=3, + pool_type=ProviderQuotaType.TRIAL, + ) + + assert result == 3 + quota_reserve.assert_called_once_with( + tenant_id=tenant_id, + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket="trial", + request_id=ANY, + amount=3, + meta={"source": "credit_pool.check_and_deduct"}, + ) + quota_commit.assert_called_once_with( + tenant_id=tenant_id, + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket="trial", + reservation_id="reservation-1", + actual_amount=3, + meta={"source": "credit_pool.check_and_deduct"}, + ) + quota_release.assert_not_called() + + +def test_check_and_deduct_credits_raises_when_billing_reserve_is_insufficient() -> None: + with ( + patch("services.credit_pool_service.dify_config.BILLING_ENABLED", True), + patch("services.billing_service.BillingService.quota_reserve") as quota_reserve, + ): + quota_reserve.return_value = {"reservation_id": "", "available": 1, "reserved": 0} + + with pytest.raises(QuotaExceededError, match="Insufficient credits remaining"): + CreditPoolService.check_and_deduct_credits(tenant_id="tenant-1", credits_required=3) + + +def test_check_and_deduct_credits_releases_billing_reservation_when_commit_fails() -> None: + with ( + patch("services.credit_pool_service.dify_config.BILLING_ENABLED", True), + patch("services.billing_service.BillingService.quota_reserve") as quota_reserve, + patch("services.billing_service.BillingService.quota_commit", side_effect=RuntimeError("commit failed")), + patch("services.billing_service.BillingService.quota_release") as quota_release, + ): + quota_reserve.return_value = {"reservation_id": "reservation-1", "available": 7, "reserved": 3} + + with pytest.raises(RuntimeError, match="commit failed"): + CreditPoolService.check_and_deduct_credits(tenant_id="tenant-1", credits_required=3) + + quota_release.assert_called_once_with( + tenant_id="tenant-1", + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket="trial", + reservation_id="reservation-1", + ) + + +def test_check_and_deduct_credits_logs_when_billing_release_fails() -> None: + with ( + patch("services.credit_pool_service.dify_config.BILLING_ENABLED", True), + patch("services.billing_service.BillingService.quota_reserve") as quota_reserve, + patch("services.billing_service.BillingService.quota_commit", side_effect=RuntimeError("commit failed")), + patch( + "services.billing_service.BillingService.quota_release", side_effect=RuntimeError("release failed") + ) as quota_release, + patch("services.credit_pool_service.logger.warning") as logger_warning, + ): + quota_reserve.return_value = {"reservation_id": "reservation-1", "available": 7, "reserved": 3} + + with pytest.raises(RuntimeError, match="commit failed"): + CreditPoolService.check_and_deduct_credits(tenant_id="tenant-1", credits_required=3) + + quota_release.assert_called_once_with( + tenant_id="tenant-1", + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket="trial", + reservation_id="reservation-1", + ) + logger_warning.assert_called_once() + assert logger_warning.call_args.args[3] == "reservation-1" + assert logger_warning.call_args.kwargs["exc_info"] is True + + +def test_deduct_credits_capped_uses_billing_consume_capped_when_enabled() -> None: + tenant_id = "tenant-1" + with ( + patch("services.credit_pool_service.dify_config.BILLING_ENABLED", True), + patch("services.billing_service.BillingService.quota_consume_capped") as quota_consume_capped, + ): + quota_consume_capped.return_value = { + "deducted": 2, + "available": 0, + "reserved": 0, + "quota": 10, + "usage": 10, + } + + result = CreditPoolService.deduct_credits_capped( + tenant_id=tenant_id, + credits_required=5, + pool_type=ProviderQuotaType.PAID, + ) + + assert result == 2 + quota_consume_capped.assert_called_once_with( + tenant_id=tenant_id, + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket="paid", + request_id=ANY, + amount=5, + meta={"source": "credit_pool.deduct_capped"}, + ) @pytest.mark.parametrize( From ae1e180b5460a91b501054a91e628a8b20f7d471 Mon Sep 17 00:00:00 2001 From: Xiyuan Chen <52963600+GareArc@users.noreply.github.com> Date: Wed, 8 Jul 2026 00:58:31 -0700 Subject: [PATCH 48/70] fix(cli): --insecure also skips TLS certificate verification (#38531) --- cli/src/auth/hosts.test.ts | 15 +++++ cli/src/auth/hosts.ts | 10 +++- cli/src/commands/_shared/authed-command.ts | 13 +++-- cli/src/commands/auth/login/index.ts | 4 +- cli/src/commands/auth/login/login.ts | 3 +- cli/src/commands/auth/logout/index.ts | 5 +- cli/src/http/client-tls.test.ts | 66 ++++++++++++++++++++++ cli/src/http/client.ts | 13 ++++- cli/src/http/proxy.test.ts | 24 ++++++++ cli/src/http/proxy.ts | 33 +++++++---- cli/src/http/types.ts | 2 + cli/src/util/host.ts | 11 ++++ cli/src/version/enforce.ts | 11 ++-- cli/src/version/probe.ts | 14 +++-- 14 files changed, 190 insertions(+), 34 deletions(-) create mode 100644 cli/src/http/client-tls.test.ts diff --git a/cli/src/auth/hosts.test.ts b/cli/src/auth/hosts.test.ts index 112538ccbb9..337b785a4c4 100644 --- a/cli/src/auth/hosts.test.ts +++ b/cli/src/auth/hosts.test.ts @@ -113,6 +113,21 @@ describe('Registry (pure)', () => { expect(active?.ctx.account.email).toBe('a@x') }) + it('resolveActive returns the active context with insecureTls', () => { + const reg = baseReg() + reg.upsert('h1', 'a@x', ctx('a@x')) + reg.setInsecureTls('h1', true) + reg.setHost('h1') + reg.setAccount('a@x') + expect(reg.resolveActive()?.insecureTls).toBe(true) + }) + + it('setInsecureTls is a no-op for an unknown host', () => { + const reg = baseReg() + reg.setInsecureTls('missing', true) + expect(reg.hosts.missing).toBeUndefined() + }) + it('resolveActive returns undefined for each missing pointer', () => { const reg = baseReg() expect(reg.resolveActive()).toBeUndefined() diff --git a/cli/src/auth/hosts.ts b/cli/src/auth/hosts.ts index 29305db951e..5df8cdabdbe 100644 --- a/cli/src/auth/hosts.ts +++ b/cli/src/auth/hosts.ts @@ -41,6 +41,7 @@ export type AccountContext = z.infer export const HostEntrySchema = z.object({ scheme: z.string().optional(), + insecure_tls: z.boolean().optional(), current_account: z.string().optional(), accounts: z.record(z.string(), AccountContextSchema).default({}), }) @@ -58,6 +59,7 @@ export type ActiveContext = { readonly email: string readonly ctx: AccountContext readonly scheme?: string + readonly insecureTls?: boolean } export function notLoggedInError(hint = 'run \'difyctl auth login\''): BaseError { @@ -104,7 +106,7 @@ export class Registry { const ctx = entry.accounts[email] if (ctx === undefined) return undefined - return { host, email, ctx, scheme: entry.scheme } + return { host, email, ctx, scheme: entry.scheme, insecureTls: entry.insecure_tls } } requireActive(hint?: string): ActiveContext { @@ -157,6 +159,12 @@ export class Registry { entry.scheme = scheme } + setInsecureTls(host: string, insecure: boolean): void { + const entry = this.data.hosts[host] + if (entry !== undefined) + entry.insecure_tls = insecure + } + activate(host: string, email: string, ctx: AccountContext): void { this.upsert(host, email, ctx) this.setHost(host) diff --git a/cli/src/commands/_shared/authed-command.ts b/cli/src/commands/_shared/authed-command.ts index 1b0d6803180..d40b9342041 100644 --- a/cli/src/commands/_shared/authed-command.ts +++ b/cli/src/commands/_shared/authed-command.ts @@ -13,7 +13,7 @@ import { formatErrorForCli } from '@/errors/format' import { createHttpClient } from '@/http/client' import { getTokenStore } from '@/store/manager' import { realStreams } from '@/sys/io/streams' -import { hostWithScheme, openAPIBase } from '@/util/host' +import { activeHostInfo, openAPIBase } from '@/util/host' import { enforceDifyVersion } from '@/version/enforce' import { versionInfo } from '@/version/info' import { maybeNudgeCompat } from '@/version/nudge' @@ -50,17 +50,17 @@ export async function buildAuthedContext( if (bearer === '') fail(cmd, opts, io) - const host = hostWithScheme(active.host, active.scheme) + const { host, insecure } = activeHostInfo(active) const retryAttempts = resolveRetryAttempts({ flag: opts.retryFlag, env: getEnv }) - const http = createHttpClient({ baseURL: openAPIBase(host), bearer, retryAttempts }) + const http = createHttpClient({ baseURL: openAPIBase(host), bearer, retryAttempts, insecure }) const cache = opts.withCache === true ? await loadAppInfoCache() : undefined // Hard gate: refuse a server too old for this difyctl (throws → exit 6). // Cached per host (1h) so most commands don't re-probe. Then the soft nudge // handles the "server too new" direction. - await enforceDifyVersion(host) - await runCompatNudge({ host, io }) + await enforceDifyVersion(host, { insecure }) + await runCompatNudge({ host, insecure, io }) return { reg, active, store, http, host, io, cache } } @@ -74,6 +74,7 @@ function fail(cmd: Pick, opts: AuthedContextOptions, io: IOStr // command flows through it without per-command wiring. async function runCompatNudge(opts: { readonly host: string + readonly insecure: boolean readonly io: IOStreams }): Promise { try { @@ -81,7 +82,7 @@ async function runCompatNudge(opts: { await maybeNudgeCompat(opts.host, { store, probe: async (host) => { - const http = createHttpClient({ baseURL: openAPIBase(host), timeoutMs: META_PROBE_TIMEOUT_MS, retryAttempts: 0 }) + const http = createHttpClient({ baseURL: openAPIBase(host), timeoutMs: META_PROBE_TIMEOUT_MS, retryAttempts: 0, insecure: opts.insecure }) return new MetaClient(http).serverVersion() }, emit: line => opts.io.err.write(line), diff --git a/cli/src/commands/auth/login/index.ts b/cli/src/commands/auth/login/index.ts index 214f38c3697..0b702262932 100644 --- a/cli/src/commands/auth/login/index.ts +++ b/cli/src/commands/auth/login/index.ts @@ -27,7 +27,7 @@ export default class Login extends DifyCommand { default: false, }), 'insecure': Flags.boolean({ - description: 'allow http:// hosts (local-dev only)', + description: 'allow http:// hosts and skip TLS certificate verification (local-dev only)', default: false, }), } @@ -40,7 +40,7 @@ export default class Login extends DifyCommand { noBrowser: flags['no-browser'], insecure: flags.insecure, verifyServer: async (host) => { - await enforceDifyVersion(host, { forceFresh: true }) + await enforceDifyVersion(host, { forceFresh: true, insecure: flags.insecure }) }, }) } diff --git a/cli/src/commands/auth/login/login.ts b/cli/src/commands/auth/login/login.ts index f0eb14dc8e1..7777f8c7c18 100644 --- a/cli/src/commands/auth/login/login.ts +++ b/cli/src/commands/auth/login/login.ts @@ -44,7 +44,7 @@ export async function runLogin(opts: LoginOptions): Promise { const host = await resolveLoginHost(opts, insecure) const label = opts.deviceLabel ?? defaultDeviceLabel() - const api = opts.api ?? new DeviceFlowApi(createHttpClient({ baseURL: openAPIBase(host) })) + const api = opts.api ?? new DeviceFlowApi(createHttpClient({ baseURL: openAPIBase(host), insecure })) const code = await api.requestCode({ device_label: label }) renderCodePrompt(opts.io.err, cs, code) @@ -88,6 +88,7 @@ export async function runLogin(opts: LoginOptions): Promise { reg.token_storage = storeBundle.mode reg.activate(display, email, ctx) applyScheme(reg, display, host) + reg.setInsecureTls(display, insecure) await reg.save() renderLoggedIn(opts.io.out, cs, host, success) diff --git a/cli/src/commands/auth/logout/index.ts b/cli/src/commands/auth/logout/index.ts index 6476b1726e8..786d525c748 100644 --- a/cli/src/commands/auth/logout/index.ts +++ b/cli/src/commands/auth/logout/index.ts @@ -6,7 +6,7 @@ import { createHttpClient } from '@/http/client' import { getTokenStore } from '@/store/manager' import { runWithSpinner } from '@/sys/io/spinner' import { realStreams } from '@/sys/io/streams' -import { hostWithScheme, openAPIBase } from '@/util/host' +import { activeHostInfo, openAPIBase } from '@/util/host' import { runLogout } from './logout.js' export default class Logout extends DifyCommand { @@ -32,7 +32,8 @@ export default class Logout extends DifyCommand { } catch { /* keyring locked — skip remote revocation, local cleanup still runs */ } if (bearer !== '') { - http = createHttpClient({ baseURL: openAPIBase(hostWithScheme(active.host, active.scheme)), bearer, retryAttempts: 0 }) + const { host, insecure } = activeHostInfo(active) + http = createHttpClient({ baseURL: openAPIBase(host), bearer, retryAttempts: 0, insecure }) } } diff --git a/cli/src/http/client-tls.test.ts b/cli/src/http/client-tls.test.ts new file mode 100644 index 00000000000..089b677fa97 --- /dev/null +++ b/cli/src/http/client-tls.test.ts @@ -0,0 +1,66 @@ +import type { Buffer } from 'node:buffer' +import type { AddressInfo } from 'node:net' +import { execFileSync } from 'node:child_process' +import { mkdtempSync, readFileSync, rmSync } from 'node:fs' +import * as https from 'node:https' +import { tmpdir } from 'node:os' +import { join } from 'node:path' +import { afterAll, beforeAll, describe, expect, it } from 'vitest' +import { createHttpClient } from './client.js' + +function generateSelfSignedCert(dir: string): { key: Buffer, cert: Buffer } { + const keyPath = join(dir, 'key.pem') + const certPath = join(dir, 'cert.pem') + execFileSync('openssl', [ + 'req', + '-x509', + '-newkey', + 'rsa:2048', + '-nodes', + '-keyout', + keyPath, + '-out', + certPath, + '-days', + '1', + '-subj', + '/CN=localhost', + ], { stdio: ['ignore', 'ignore', 'pipe'] }) + return { key: readFileSync(keyPath), cert: readFileSync(certPath) } +} + +// A real server, not a fetch mock, so this also covers Bun's native `tls` +// fetch option (ignored by Node, which only reads undici's `dispatcher`). +describe('createHttpClient against a real self-signed TLS server', () => { + let server: https.Server + let baseURL: string + let certDir: string + + beforeAll(async () => { + certDir = mkdtempSync(join(tmpdir(), 'difyctl-tls-test-')) + const { key, cert } = generateSelfSignedCert(certDir) + server = https.createServer({ key, cert }, (_req, res) => { + res.writeHead(200, { 'content-type': 'application/json' }) + res.end(JSON.stringify({ ok: true })) + }) + await new Promise(resolve => server.listen(0, resolve)) + const port = (server.address() as AddressInfo).port + baseURL = `https://localhost:${port}/` + }) + + afterAll(async () => { + await new Promise(resolve => server.close(() => resolve())) + rmSync(certDir, { recursive: true, force: true }) + }) + + it('rejects the self-signed cert by default', async () => { + const http = createHttpClient({ baseURL, retryAttempts: 0 }) + await expect(http.get('')).rejects.toBeDefined() + }) + + it('accepts the self-signed cert when insecure: true', async () => { + const http = createHttpClient({ baseURL, retryAttempts: 0, insecure: true }) + const res = await http.get<{ ok: boolean }>('') + expect(res.ok).toBe(true) + }) +}) diff --git a/cli/src/http/client.ts b/cli/src/http/client.ts index 940aee9152b..30f8ca27441 100644 --- a/cli/src/http/client.ts +++ b/cli/src/http/client.ts @@ -38,6 +38,7 @@ type ClientState = { readonly logger: HttpLogger | undefined readonly originalOptions: ClientOptions readonly dispatcher: ReturnType + readonly insecure: boolean } function toArray(value: T | T[] | undefined): T[] { @@ -74,7 +75,8 @@ function compileState(opts: ClientOptions): ClientState { hooks: { onRequest, onResponse, onRequestError, onResponseError }, logger: opts.logger, originalOptions: opts, - dispatcher: proxyDispatcher(), + dispatcher: proxyDispatcher({ insecure: opts.insecure }), + insecure: opts.insecure ?? false, } } @@ -171,9 +173,16 @@ async function execute( await runHooks(state.hooks.onRequest, ctx) - const init: RequestInit & { dispatcher?: unknown, verbose?: boolean } = { signal } + // Two runtimes, two options: Node's fetch reads undici's `dispatcher` (used + // below for TLS-skip + proxy routing); Bun's native fetch — what the compiled + // difyctl binary actually runs on — ignores `dispatcher` entirely and instead + // needs its own `tls` option. Set both; each runtime ignores the one it + // doesn't understand. + const init: RequestInit & { dispatcher?: unknown, tls?: { rejectUnauthorized: boolean }, verbose?: boolean } = { signal } if (state.dispatcher !== undefined) init.dispatcher = state.dispatcher + if (state.insecure) + init.tls = { rejectUnauthorized: false } if (isVerbose()) init.verbose = true diff --git a/cli/src/http/proxy.test.ts b/cli/src/http/proxy.test.ts index 80b36d0f295..2c1f386f320 100644 --- a/cli/src/http/proxy.test.ts +++ b/cli/src/http/proxy.test.ts @@ -52,4 +52,28 @@ describe('proxyDispatcher', () => { expect(proxyDispatcher()).toBe(first) await first?.close() }) + + it('builds a plain Agent with TLS verification disabled when insecure is set and no proxy env', async () => { + const { proxyDispatcher } = await import('./proxy.js') + const d = proxyDispatcher({ insecure: true }) + expect(d?.constructor.name).toBe('Agent') + await d?.close() + }) + + it('builds an EnvHttpProxyAgent with TLS verification disabled when insecure and a proxy are both set', async () => { + process.env.HTTP_PROXY = 'http://127.0.0.1:8888' + const { proxyDispatcher } = await import('./proxy.js') + const d = proxyDispatcher({ insecure: true }) + expect(d?.constructor.name).toBe('EnvHttpProxyAgent') + await d?.close() + }) + + it('re-resolves when the insecure flag changes for the same call site', async () => { + const { proxyDispatcher } = await import('./proxy.js') + const secure = proxyDispatcher({ insecure: false }) + expect(secure).toBeUndefined() + const insecure = proxyDispatcher({ insecure: true }) + expect(insecure?.constructor.name).toBe('Agent') + await insecure?.close() + }) }) diff --git a/cli/src/http/proxy.ts b/cli/src/http/proxy.ts index 5e58936f91a..bd9a745c922 100644 --- a/cli/src/http/proxy.ts +++ b/cli/src/http/proxy.ts @@ -1,4 +1,5 @@ -import { EnvHttpProxyAgent } from 'undici' +import type { Dispatcher } from 'undici' +import { Agent, EnvHttpProxyAgent } from 'undici' const PROXY_ENV_KEYS = ['HTTP_PROXY', 'http_proxy', 'HTTPS_PROXY', 'https_proxy'] as const @@ -6,18 +7,30 @@ export function hasProxyEnv(): boolean { return PROXY_ENV_KEYS.some(k => (process.env[k] ?? '') !== '') } -let resolved = false -let agent: EnvHttpProxyAgent | undefined +export type ProxyDispatcherOptions = { + // --insecure on a self-signed https:// host: skip certificate verification + // (local-dev only, same flag that allows plain http:// hosts). + readonly insecure?: boolean +} + +let resolvedKey: string | undefined +let agent: Dispatcher | undefined // Node's global fetch ignores HTTP_PROXY / HTTPS_PROXY / NO_PROXY. When a proxy // var is set, route requests through an EnvHttpProxyAgent (it also reads the -// lowercase variants and honours NO_PROXY); when none is set, return undefined so -// fetch keeps Node's default global dispatcher untouched. Resolved once per -// process — proxy env vars are fixed for a single CLI invocation. -export function proxyDispatcher(): EnvHttpProxyAgent | undefined { - if (!resolved) { - agent = hasProxyEnv() ? new EnvHttpProxyAgent() : undefined - resolved = true +// lowercase variants and honours NO_PROXY); when none is set and TLS verification +// isn't being skipped, return undefined so fetch keeps Node's default global +// dispatcher untouched. Resolved once per (proxy env, insecure) combination — +// both are fixed for a single CLI invocation. +export function proxyDispatcher(opts: ProxyDispatcherOptions = {}): Dispatcher | undefined { + const insecure = opts.insecure ?? false + const key = `${hasProxyEnv()}:${insecure}` + if (resolvedKey !== key) { + const tls = insecure ? { rejectUnauthorized: false } : undefined + agent = hasProxyEnv() + ? new EnvHttpProxyAgent(tls !== undefined ? { connect: tls, requestTls: tls, proxyTls: tls } : undefined) + : (tls !== undefined ? new Agent({ connect: tls }) : undefined) + resolvedKey = key } return agent } diff --git a/cli/src/http/types.ts b/cli/src/http/types.ts index d209e97460c..a022a253296 100644 --- a/cli/src/http/types.ts +++ b/cli/src/http/types.ts @@ -75,6 +75,8 @@ export type ClientOptions = { readonly retryAttempts?: number readonly logger?: HttpLogger readonly hooks?: Hooks + // Skip TLS certificate verification (local-dev only, self-signed hosts). + readonly insecure?: boolean } export type HttpClient = { diff --git a/cli/src/util/host.ts b/cli/src/util/host.ts index 1042b0875ea..453a68149ad 100644 --- a/cli/src/util/host.ts +++ b/cli/src/util/host.ts @@ -1,3 +1,4 @@ +import type { ActiveContext } from '@/auth/hosts' import { BaseError } from '@/errors/base' import { ErrorCode } from '@/errors/codes' @@ -44,6 +45,16 @@ export function hostWithScheme(host: string, scheme: string | undefined): string return `${proto}://${host}` } +// Every call site that builds an HTTP client for the active host needs both its +// scheme-qualified host and whether TLS verification is disabled for it — derive +// them together so neither is forgotten independently. +export function activeHostInfo(active: Pick): { host: string, insecure: boolean } { + return { + host: hostWithScheme(active.host, active.scheme), + insecure: active.insecureTls === true, + } +} + export function bareHost(raw: string): string { try { const u = new URL(raw) diff --git a/cli/src/version/enforce.ts b/cli/src/version/enforce.ts index d1b9fe87da9..8c995a912d1 100644 --- a/cli/src/version/enforce.ts +++ b/cli/src/version/enforce.ts @@ -16,15 +16,18 @@ const UPGRADE_HINT + '(https://docs.dify.ai/en/getting-started/install-self-hosted)' // /_version is unauthenticated; same timeout/no-retry budget as the auto-nudge probe. -const defaultProbe: ServerVersionProbe = async (host) => { - const http = createHttpClient({ baseURL: openAPIBase(host), timeoutMs: META_PROBE_TIMEOUT_MS, retryAttempts: 0 }) - return new MetaClient(http).serverVersion() +function buildDefaultProbe(insecure: boolean): ServerVersionProbe { + return async (host) => { + const http = createHttpClient({ baseURL: openAPIBase(host), timeoutMs: META_PROBE_TIMEOUT_MS, retryAttempts: 0, insecure }) + return new MetaClient(http).serverVersion() + } } export type EnforceOptions = { readonly probe?: ServerVersionProbe readonly store?: CompatStore readonly forceFresh?: boolean + readonly insecure?: boolean } /** @@ -45,7 +48,7 @@ export async function enforceDifyVersion( if (opts.forceFresh !== true && store.isFreshCompatible(host)) return undefined - const probe = opts.probe ?? defaultProbe + const probe = opts.probe ?? buildDefaultProbe(opts.insecure === true) let server: ServerVersionResponse try { server = await probe(host) diff --git a/cli/src/version/probe.ts b/cli/src/version/probe.ts index 266910e7c05..05ab3ea84e3 100644 --- a/cli/src/version/probe.ts +++ b/cli/src/version/probe.ts @@ -6,7 +6,7 @@ import { META_PROBE_TIMEOUT_MS, MetaClient } from '@/api/meta' import { Registry } from '@/auth/hosts' import { createHttpClient } from '@/http/client' import { arch, platform } from '@/sys/index' -import { hostWithScheme, openAPIBase } from '@/util/host' +import { activeHostInfo, openAPIBase } from '@/util/host' import { difyCompat, evaluateCompat } from './compat.js' import { versionInfo } from './info.js' @@ -51,9 +51,11 @@ const defaultLoadActive = async (): Promise => { return (await Registry.load()).resolveActive() } -const defaultProbe: MetaProbe = async (endpoint) => { - const http = createHttpClient({ baseURL: openAPIBase(endpoint), timeoutMs: META_PROBE_TIMEOUT_MS, retryAttempts: 0 }) - return new MetaClient(http).serverVersion() +function buildDefaultProbe(insecure: boolean): MetaProbe { + return async (endpoint) => { + const http = createHttpClient({ baseURL: openAPIBase(endpoint), timeoutMs: META_PROBE_TIMEOUT_MS, retryAttempts: 0, insecure }) + return new MetaClient(http).serverVersion() + } } function buildClientBlock(): ClientBlock { @@ -92,7 +94,6 @@ export async function runVersionProbe(opts: RunVersionProbeOptions): Promise Date: Wed, 8 Jul 2026 15:59:32 +0800 Subject: [PATCH 49/70] chore: remove superpowers artifacts (#38547) --- .../sdd/hitl-timeout-semantics-impl-report.md | 30 ------------------- 1 file changed, 30 deletions(-) delete mode 100644 .superpowers/sdd/hitl-timeout-semantics-impl-report.md diff --git a/.superpowers/sdd/hitl-timeout-semantics-impl-report.md b/.superpowers/sdd/hitl-timeout-semantics-impl-report.md deleted file mode 100644 index ba36a15d7dc..00000000000 --- a/.superpowers/sdd/hitl-timeout-semantics-impl-report.md +++ /dev/null @@ -1,30 +0,0 @@ -# HITL timeout semantics implementation report - -## What changed - -- Updated `api/core/workflow/nodes/human_input/callback.py` so `DifyHITLCallback` now preserves Dify's timeout split at the boundary: - - `HumanInputFormStatus.TIMEOUT` returns the graphon timeout branch via `Expired(selected_handle="__timeout__", ...)`. - - `HumanInputFormStatus.EXPIRED` is treated as an invalid resume state and raises `AssertionError`. - - `HumanInputFormStatus.WAITING` with a past global deadline is treated as an invalid resume state and raises `AssertionError`. - - `HumanInputFormStatus.WAITING` with only the node-level deadline expired still returns the timeout branch. -- Added `created_at` to `HumanInputFormEntity` and `_HumanInputFormEntityImpl` so the callback can compute the global deadline using Dify's shared `HUMAN_INPUT_GLOBAL_TIMEOUT_SECONDS` invariant. -- Kept the submitted and pause flows unchanged. -- Added focused unit coverage in `api/tests/unit_tests/core/workflow/test_human_input_callback.py` for: - - node timeout branch - - global expiration rejection - - waiting-form past node deadline timeout - - waiting-form past global deadline rejection - -## Verification - -- `uv run --project api pytest -o addopts='' api/tests/unit_tests/core/workflow/test_human_input_callback.py api/tests/unit_tests/core/workflow/nodes/human_input/test_human_input_form_filled_event.py -q` -- `git diff --check` - -## Result - -- The focused test set is expected to pass with the new `created_at` boundary in place. -- No unrelated files were modified. - -## Concerns - -- The callback now fails fast on invalid resume states by design. That is intentional, but any caller that previously relied on `EXPIRED` being mapped to the timeout branch will now see an assertion failure instead. From 5ffd4345dcb6fa2f70767219943b664e4c12c1aa Mon Sep 17 00:00:00 2001 From: Crazywoola <100913391+crazywoola@users.noreply.github.com> Date: Wed, 8 Jul 2026 16:07:14 +0800 Subject: [PATCH 50/70] fix: display errors for oauth page (#38546) Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- eslint-suppressions.json | 5 -- web/app/account/oauth/authorize/layout.tsx | 70 ++++++++-------------- 2 files changed, 25 insertions(+), 50 deletions(-) diff --git a/eslint-suppressions.json b/eslint-suppressions.json index 81c7678008c..e141792779d 100644 --- a/eslint-suppressions.json +++ b/eslint-suppressions.json @@ -292,11 +292,6 @@ "count": 1 } }, - "web/app/account/oauth/authorize/layout.tsx": { - "ts/no-explicit-any": { - "count": 1 - } - }, "web/app/account/oauth/authorize/page.tsx": { "ts/no-explicit-any": { "count": 1 diff --git a/web/app/account/oauth/authorize/layout.tsx b/web/app/account/oauth/authorize/layout.tsx index 053baeca925..339624fe027 100644 --- a/web/app/account/oauth/authorize/layout.tsx +++ b/web/app/account/oauth/authorize/layout.tsx @@ -1,60 +1,40 @@ 'use client' -import { cn } from '@langgenius/dify-ui/cn' -import { useQuery, useSuspenseQuery } from '@tanstack/react-query' +import type { ReactNode } from 'react' +import { useSuspenseQuery } from '@tanstack/react-query' -import Loading from '@/app/components/base/loading' import Header from '@/app/signin/_header' -import { AppContextProvider } from '@/context/app-context-provider' -import { isLegacyBase401, userProfileQueryOptions } from '@/features/account-profile/client' import { systemFeaturesQueryOptions } from '@/features/system-features/client' import useDocumentTitle from '@/hooks/use-document-title' -export default function SignInLayout({ children }: any) { +type Props = { + children: ReactNode +} + +const copyrightYear = new Date().getFullYear() + +export default function OAuthAuthorizeLayout({ children }: Props) { const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) useDocumentTitle('') - // Probe login state. 401 stays as `error` (not thrown) so this layout can render - // the signin/oauth UI for unauthenticated users; other errors bubble to error.tsx. - // (When unauthenticated, service/base.ts's auto-redirect to /signin still fires.) - const { isPending, data: userResp, error } = useQuery({ - ...userProfileQueryOptions(), - throwOnError: err => !isLegacyBase401(err), - }) - const isLoggedIn = !!userResp && !error - if (isPending) { - return ( -
    - -
    - ) - } return ( - <> -
    -
    -
    -
    -
    - {isLoggedIn - ? ( - - {children} - - ) - : children} -
    +
    +
    +
    +
    +
    + {children}
    - {systemFeatures.branding.enabled === false && ( -
    - © - {' '} - {new Date().getFullYear()} - {' '} - LangGenius, Inc. All rights reserved. -
    - )}
    + {systemFeatures.branding.enabled === false && ( +
    + © + {' '} + {copyrightYear} + {' '} + LangGenius, Inc. All rights reserved. +
    + )}
    - +
    ) } From dce45ef6ae4843466a2bc8c6cee54c7a257a895e Mon Sep 17 00:00:00 2001 From: Stephen Zhou Date: Wed, 8 Jul 2026 16:15:54 +0800 Subject: [PATCH 51/70] refactor(web): migrate account settings app context consumers (#38544) Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- eslint-suppressions.json | 32 ------------- .../header/account-dropdown-flow.test.tsx | 48 +++++++++++++------ web/__tests__/utils/mock-app-context-state.ts | 30 ++++++++++++ .../delete-account/components/check-email.tsx | 9 ++-- .../delete-account/components/feed-back.tsx | 9 ++-- .../account-dropdown/__tests__/index.spec.tsx | 31 +++++++++--- .../account-dropdown/default-menu-content.tsx | 7 ++- .../header/account-dropdown/index.tsx | 6 ++- .../main-nav-menu-content.tsx | 5 +- .../__tests__/access-rule-section.spec.tsx | 12 +++++ .../access-rules-page/access-rule-section.tsx | 5 +- .../__tests__/index.spec.tsx | 12 +++++ .../api-based-extension-page/index.tsx | 5 +- .../header/account-setting/index.tsx | 5 +- .../members-page/__tests__/index.spec.tsx | 29 ++++++++--- .../__tests__/invite-button.spec.tsx | 10 ++++ .../__tests__/dialog.spec.tsx | 20 +++++++- .../__tests__/index.spec.tsx | 17 ++++++- .../edit-workspace-modal/index.tsx | 9 ++-- .../account-setting/members-page/index.tsx | 12 +++-- .../members-page/invite-button.tsx | 7 +-- .../__tests__/index.spec.tsx | 17 ++++++- .../transfer-ownership-modal/index.tsx | 28 +++++++---- .../__tests__/model-list-item.spec.tsx | 12 +++++ .../__tests__/model-list.spec.tsx | 12 +++++ .../provider-added-card/index.tsx | 5 +- .../provider-added-card/model-list-item.tsx | 5 +- .../provider-added-card/model-list.tsx | 12 +++-- .../__tests__/index.spec.tsx | 12 +++++ .../system-model-selector/index.tsx | 5 +- .../permissions-page/__tests__/index.spec.tsx | 12 +++++ .../permissions-page/index.tsx | 5 +- .../role-list/__tests__/row-menu.spec.tsx | 12 +++++ .../permissions-page/role-list/row-menu.tsx | 5 +- .../preference-page/__tests__/index.spec.tsx | 13 +++++ .../account-setting/preference-page/index.tsx | 8 ++-- .../account-setting/update-setting-dialog.tsx | 5 +- .../header/env-nav/__tests__/index.spec.tsx | 32 ++++++++++--- web/app/components/header/env-nav/index.tsx | 5 +- .../main-nav/__tests__/index.spec.tsx | 15 ++++++ 40 files changed, 407 insertions(+), 133 deletions(-) diff --git a/eslint-suppressions.json b/eslint-suppressions.json index e141792779d..34cab985f05 100644 --- a/eslint-suppressions.json +++ b/eslint-suppressions.json @@ -279,11 +279,6 @@ "count": 1 } }, - "web/app/account/(commonLayout)/delete-account/components/check-email.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "web/app/account/(commonLayout)/delete-account/components/verify-email.tsx": { "no-restricted-imports": { "count": 1 @@ -3401,30 +3396,11 @@ "count": 2 } }, - "web/app/components/header/account-setting/members-page/edit-workspace-modal/index.tsx": { - "jsx-a11y/no-autofocus": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - } - }, "web/app/components/header/account-setting/members-page/invite-modal/index.tsx": { "jsx-a11y/no-autofocus": { "count": 1 } }, - "web/app/components/header/account-setting/members-page/transfer-ownership-modal/index.tsx": { - "erasable-syntax-only/enums": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - }, - "ts/no-explicit-any": { - "count": 2 - } - }, "web/app/components/header/account-setting/members-page/transfer-ownership-modal/member-selector.tsx": { "jsx-a11y/click-events-have-key-events": { "count": 1 @@ -3566,14 +3542,6 @@ "count": 2 } }, - "web/app/components/header/account-setting/model-provider-page/provider-added-card/model-list.tsx": { - "jsx-a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx-a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/header/account-setting/model-provider-page/provider-added-card/model-load-balancing-configs.tsx": { "jsx-a11y/click-events-have-key-events": { "count": 2 diff --git a/web/__tests__/header/account-dropdown-flow.test.tsx b/web/__tests__/header/account-dropdown-flow.test.tsx index fd651931b50..c949bb3f99a 100644 --- a/web/__tests__/header/account-dropdown-flow.test.tsx +++ b/web/__tests__/header/account-dropdown-flow.test.tsx @@ -6,11 +6,32 @@ import AccountDropdown from '@/app/components/header/account-dropdown' import { ACCOUNT_SETTING_TAB } from '@/app/components/header/account-setting/constants' const { + mockAppContextState, mockPush, mockLogout, mockResetUser, mockSetShowAccountSettingModal, } = vi.hoisted(() => ({ + mockAppContextState: { + userProfile: { + id: 'user-1', + name: 'Ada Lovelace', + email: 'ada@example.com', + avatar: '', + avatar_url: '', + is_password_set: true, + }, + langGeniusVersionInfo: { + current_env: 'CLOUD', + current_version: '1.0.0', + latest_version: '1.1.0', + release_date: '', + release_notes: 'https://example.com/releases/1.1.0', + version: '1.0.0', + can_auto_update: false, + }, + isCurrentWorkspaceOwner: false, + }, mockPush: vi.fn(), mockLogout: vi.fn(), mockResetUser: vi.fn(), @@ -28,21 +49,19 @@ vi.mock('react-i18next', () => ({ })) vi.mock('@/context/app-context', () => ({ - useAppContext: () => ({ - userProfile: { - name: 'Ada Lovelace', - email: 'ada@example.com', - avatar_url: '', - }, - langGeniusVersionInfo: { - current_version: '1.0.0', - latest_version: '1.1.0', - release_notes: 'https://example.com/releases/1.1.0', - }, - isCurrentWorkspaceOwner: false, - }), + useAppContext: () => mockAppContextState, })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppContextState) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/provider-context', () => ({ useProviderContext: () => ({ isEducationAccount: false, @@ -62,7 +81,8 @@ vi.mock('@/context/i18n', () => ({ useDocLink: () => (path: string) => `https://docs.example.com${path}`, })) -vi.mock('@/service/use-common', () => ({ +vi.mock('@/service/use-common', async importOriginal => ({ + ...await importOriginal(), useLogout: () => ({ mutateAsync: mockLogout, }), diff --git a/web/__tests__/utils/mock-app-context-state.ts b/web/__tests__/utils/mock-app-context-state.ts index 5be87eee447..d41dc9a2b47 100644 --- a/web/__tests__/utils/mock-app-context-state.ts +++ b/web/__tests__/utils/mock-app-context-state.ts @@ -13,6 +13,7 @@ export type AppContextStateMockState = { } | null currentWorkspace?: { id?: string + name?: string } | null isCurrentWorkspaceManager?: boolean isCurrentWorkspaceOwner?: boolean @@ -22,6 +23,8 @@ export type AppContextStateMockState = { isLoadingWorkspacePermissionKeys?: boolean workspacePermissionKeys?: string[] langGeniusVersionInfo?: LangGeniusVersionResponse + refreshUserProfile?: () => void + refreshCurrentWorkspace?: () => void } type AppContextStateAtomKind @@ -32,11 +35,14 @@ type AppContextStateAtomKind | 'currentWorkspaceId' | 'workspaceRoleFlags' | 'isCurrentWorkspaceManager' + | 'isCurrentWorkspaceOwner' | 'currentWorkspaceLoading' | 'workspacePermissionKeys' | 'workspacePermissionKeysLoading' | 'langGeniusVersionInfo' | 'langGeniusCurrentVersion' + | 'refreshUserProfile' + | 'refreshCurrentWorkspace' type AppContextStateMockAtom = { [APP_CONTEXT_STATE_ATOM_KIND]: AppContextStateAtomKind @@ -57,6 +63,7 @@ const defaultUserProfile = { const defaultCurrentWorkspace = { id: 'workspace-1', + name: 'Workspace', } const defaultLangGeniusVersionInfo = { @@ -109,11 +116,14 @@ export const createAppContextStateAtomMock = async ( currentWorkspaceIdAtom: createMockAtom('currentWorkspaceId'), workspaceRoleFlagsAtom: createMockAtom('workspaceRoleFlags'), isCurrentWorkspaceManagerAtom: createMockAtom('isCurrentWorkspaceManager'), + isCurrentWorkspaceOwnerAtom: createMockAtom('isCurrentWorkspaceOwner'), currentWorkspaceLoadingAtom: createMockAtom('currentWorkspaceLoading'), workspacePermissionKeysAtom: createMockAtom('workspacePermissionKeys'), workspacePermissionKeysLoadingAtom: createMockAtom('workspacePermissionKeysLoading'), langGeniusVersionInfoAtom: createMockAtom('langGeniusVersionInfo'), langGeniusCurrentVersionAtom: createMockAtom('langGeniusCurrentVersion'), + refreshUserProfileAtom: createMockAtom('refreshUserProfile'), + refreshCurrentWorkspaceAtom: createMockAtom('refreshCurrentWorkspace'), } } @@ -162,6 +172,9 @@ export const createAppContextStateJotaiMock = async ( if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'isCurrentWorkspaceManager') return state.isCurrentWorkspaceManager ?? false + if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'isCurrentWorkspaceOwner') + return state.isCurrentWorkspaceOwner ?? false + if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'currentWorkspaceLoading') return state.isLoadingCurrentWorkspace ?? false @@ -179,5 +192,22 @@ export const createAppContextStateJotaiMock = async ( throw new Error(`Unsupported app context state atom: ${atom[APP_CONTEXT_STATE_ATOM_KIND]}`) }, + useSetAtom: (atom: unknown) => { + if (!isAppContextStateMockAtom(atom)) + return actual.useSetAtom(atom as Parameters[0]) + + if (!appContextStateMockRegistry) + throw new Error('App context state atom mock is not initialized') + + const state = appContextStateMockRegistry.getState() + + if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'refreshUserProfile') + return state.refreshUserProfile ?? (() => {}) + + if (atom[APP_CONTEXT_STATE_ATOM_KIND] === 'refreshCurrentWorkspace') + return state.refreshCurrentWorkspace ?? (() => {}) + + throw new Error(`Unsupported app context state write atom: ${atom[APP_CONTEXT_STATE_ATOM_KIND]}`) + }, } } diff --git a/web/app/account/(commonLayout)/delete-account/components/check-email.tsx b/web/app/account/(commonLayout)/delete-account/components/check-email.tsx index 7e3351a10a2..b37a4995f02 100644 --- a/web/app/account/(commonLayout)/delete-account/components/check-email.tsx +++ b/web/app/account/(commonLayout)/delete-account/components/check-email.tsx @@ -1,9 +1,10 @@ 'use client' import { Button } from '@langgenius/dify-ui/button' +import { Input } from '@langgenius/dify-ui/input' +import { useAtomValue } from 'jotai' import { useCallback, useState } from 'react' import { useTranslation } from 'react-i18next' -import Input from '@/app/components/base/input' -import { useAppContext } from '@/context/app-context' +import { userProfileEmailAtom } from '@/context/app-context-state' import Link from '@/next/link' import { useSendDeleteAccountEmail } from '../state' @@ -14,7 +15,7 @@ type DeleteAccountProps = { export default function CheckEmail(props: DeleteAccountProps) { const { t } = useTranslation() - const { userProfile } = useAppContext() + const userProfileEmail = useAtomValue(userProfileEmailAtom) const [userInputEmail, setUserInputEmail] = useState('') const { isPending: isSendingEmail, mutateAsync: getDeleteEmailVerifyCode } = useSendDeleteAccountEmail() @@ -45,7 +46,7 @@ export default function CheckEmail(props: DeleteAccountProps) { }} />
    - +
    diff --git a/web/app/account/(commonLayout)/delete-account/components/feed-back.tsx b/web/app/account/(commonLayout)/delete-account/components/feed-back.tsx index ead73b2319c..697c6df095c 100644 --- a/web/app/account/(commonLayout)/delete-account/components/feed-back.tsx +++ b/web/app/account/(commonLayout)/delete-account/components/feed-back.tsx @@ -3,9 +3,10 @@ import { Button } from '@langgenius/dify-ui/button' import { Dialog, DialogContent, DialogTitle } from '@langgenius/dify-ui/dialog' import { Textarea } from '@langgenius/dify-ui/textarea' import { toast } from '@langgenius/dify-ui/toast' +import { useAtomValue } from 'jotai' import { useCallback, useState } from 'react' import { useTranslation } from 'react-i18next' -import { useAppContext } from '@/context/app-context' +import { userProfileEmailAtom } from '@/context/app-context-state' import { useRouter } from '@/next/navigation' import { useLogout } from '@/service/use-common' import { useDeleteAccountFeedback } from '../state' @@ -17,7 +18,7 @@ type DeleteAccountProps = { export default function FeedBack(props: DeleteAccountProps) { const { t } = useTranslation() - const { userProfile } = useAppContext() + const userProfileEmail = useAtomValue(userProfileEmailAtom) const router = useRouter() const [userFeedback, setUserFeedback] = useState('') const { isPending, mutateAsync: sendFeedback } = useDeleteAccountFeedback() @@ -35,12 +36,12 @@ export default function FeedBack(props: DeleteAccountProps) { const handleSubmit = useCallback(async () => { try { - await sendFeedback({ feedback: userFeedback, email: userProfile.email }) + await sendFeedback({ feedback: userFeedback, email: userProfileEmail }) props.onConfirm() await handleSuccess() } catch (error) { console.error(error) } - }, [handleSuccess, userFeedback, sendFeedback, userProfile, props]) + }, [handleSuccess, userFeedback, sendFeedback, userProfileEmail, props]) const handleSkip = useCallback(() => { props.onCancel() diff --git a/web/app/components/header/account-dropdown/__tests__/index.spec.tsx b/web/app/components/header/account-dropdown/__tests__/index.spec.tsx index 741580f6a5c..64065d65e15 100644 --- a/web/app/components/header/account-dropdown/__tests__/index.spec.tsx +++ b/web/app/components/header/account-dropdown/__tests__/index.spec.tsx @@ -43,6 +43,9 @@ vi.mock('@/app/components/base/theme-switcher', () => ({ const { mockSetTheme } = vi.hoisted(() => ({ mockSetTheme: vi.fn(), })) +const mockAppContextState = vi.hoisted(() => ({ + current: undefined as AppContextValue | undefined, +})) vi.mock('next-themes', () => ({ useTheme: () => ({ @@ -55,6 +58,16 @@ vi.mock('@/context/app-context', () => ({ useAppContext: vi.fn(), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppContextState.current ?? {}) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/provider-context', () => ({ useProviderContext: vi.fn(), })) @@ -63,7 +76,8 @@ vi.mock('@/context/modal-context', () => ({ useModalContext: vi.fn(), })) -vi.mock('@/service/use-common', () => ({ +vi.mock('@/service/use-common', async importOriginal => ({ + ...await importOriginal(), useLogout: vi.fn(), })) @@ -150,6 +164,11 @@ const baseAppContextValue: AppContextValue = { workspacePermissionKeys: [], } +const setAppContextValue = (value: AppContextValue) => { + mockAppContextState.current = value + vi.mocked(useAppContext).mockReturnValue(value) +} + describe('AccountDropdown', () => { const mockPush = vi.fn() const mockLogout = vi.fn() @@ -170,7 +189,7 @@ describe('AccountDropdown', () => { mockConfig.IS_CLOUD_EDITION = false mockEnv.env.NEXT_PUBLIC_SITE_ABOUT = 'show' - vi.mocked(useAppContext).mockReturnValue(baseAppContextValue) + setAppContextValue(baseAppContextValue) vi.mocked(useProviderContext).mockReturnValue({ isEducationAccount: false, plan: { type: Plan.sandbox }, @@ -267,7 +286,7 @@ describe('AccountDropdown', () => { it('should show Compliance in Cloud Edition for workspace owner', () => { // Arrange mockConfig.IS_CLOUD_EDITION = true - vi.mocked(useAppContext).mockReturnValue({ + setAppContextValue({ ...baseAppContextValue, userProfile: { ...baseAppContextValue.userProfile, name: 'User' }, isCurrentWorkspaceOwner: true, @@ -286,7 +305,7 @@ describe('AccountDropdown', () => { it('should hide Compliance in Cloud Edition when user is not workspace owner', () => { // Arrange mockConfig.IS_CLOUD_EDITION = true - vi.mocked(useAppContext).mockReturnValue({ + setAppContextValue({ ...baseAppContextValue, isCurrentWorkspaceOwner: false, }) @@ -382,7 +401,7 @@ describe('AccountDropdown', () => { describe('Version Indicators', () => { it('should show orange indicator when version is not latest', () => { // Arrange - vi.mocked(useAppContext).mockReturnValue({ + setAppContextValue({ ...baseAppContextValue, userProfile: { ...baseAppContextValue.userProfile, name: 'User' }, langGeniusVersionInfo: { @@ -402,7 +421,7 @@ describe('AccountDropdown', () => { it('should show green indicator when version is latest', () => { // Arrange - vi.mocked(useAppContext).mockReturnValue({ + setAppContextValue({ ...baseAppContextValue, userProfile: { ...baseAppContextValue.userProfile, name: 'User' }, langGeniusVersionInfo: { diff --git a/web/app/components/header/account-dropdown/default-menu-content.tsx b/web/app/components/header/account-dropdown/default-menu-content.tsx index bb0aded8e8c..ac5e981d971 100644 --- a/web/app/components/header/account-dropdown/default-menu-content.tsx +++ b/web/app/components/header/account-dropdown/default-menu-content.tsx @@ -10,12 +10,13 @@ import { } from '@langgenius/dify-ui/dropdown-menu' import { StatusDot } from '@langgenius/dify-ui/status-dot' import { useSuspenseQuery } from '@tanstack/react-query' +import { useAtomValue } from 'jotai' import { useTranslation } from 'react-i18next' import PremiumBadge from '@/app/components/base/premium-badge' import ThemeSwitcher from '@/app/components/base/theme-switcher' import { ACCOUNT_SETTING_TAB } from '@/app/components/header/account-setting/constants' import { IS_CLOUD_EDITION } from '@/config' -import { useAppContext } from '@/context/app-context' +import { isCurrentWorkspaceOwnerAtom, langGeniusVersionInfoAtom, userProfileAtom } from '@/context/app-context-state' import { useDocLink } from '@/context/i18n' import { useModalContext } from '@/context/modal-context' import { useProviderContext } from '@/context/provider-context' @@ -119,7 +120,9 @@ export function DefaultMenuContent({ const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) const { t } = useTranslation() const docLink = useDocLink() - const { userProfile, langGeniusVersionInfo, isCurrentWorkspaceOwner } = useAppContext() + const userProfile = useAtomValue(userProfileAtom) + const langGeniusVersionInfo = useAtomValue(langGeniusVersionInfoAtom) + const isCurrentWorkspaceOwner = useAtomValue(isCurrentWorkspaceOwnerAtom) const { isEducationAccount } = useProviderContext() const { setShowAccountSettingModal } = useModalContext() diff --git a/web/app/components/header/account-dropdown/index.tsx b/web/app/components/header/account-dropdown/index.tsx index dbadfc1f64c..518398c4c62 100644 --- a/web/app/components/header/account-dropdown/index.tsx +++ b/web/app/components/header/account-dropdown/index.tsx @@ -8,11 +8,12 @@ import { DropdownMenuContent, DropdownMenuTrigger, } from '@langgenius/dify-ui/dropdown-menu' +import { useAtomValue } from 'jotai' import { useState } from 'react' import { useTranslation } from 'react-i18next' import { resetUser } from '@/app/components/base/amplitude/utils' import { useSetEducationExpiredHasNoticed, useSetEducationReverifyHasNoticed, useSetEducationReverifyPrevExpireAt } from '@/app/education-apply/storage' -import { useAppContext } from '@/context/app-context' +import { langGeniusVersionInfoAtom, userProfileAtom } from '@/context/app-context-state' import { useRouter } from '@/next/navigation' import { useLogout } from '@/service/use-common' import AccountAbout from '../account-about' @@ -37,7 +38,8 @@ export default function AppSelector({ const [aboutVisible, setAboutVisible] = useState(false) const [isAccountMenuOpen, setIsAccountMenuOpen] = useState(false) const { t } = useTranslation() - const { userProfile, langGeniusVersionInfo } = useAppContext() + const userProfile = useAtomValue(userProfileAtom) + const langGeniusVersionInfo = useAtomValue(langGeniusVersionInfoAtom) const clearEducationReverifyPrevExpireAt = useSetEducationReverifyPrevExpireAt() const clearEducationReverifyHasNoticed = useSetEducationReverifyHasNoticed() const clearEducationExpiredHasNoticed = useSetEducationExpiredHasNoticed() diff --git a/web/app/components/header/account-dropdown/main-nav-menu-content.tsx b/web/app/components/header/account-dropdown/main-nav-menu-content.tsx index 3ae90caffcb..2faf2ad7f46 100644 --- a/web/app/components/header/account-dropdown/main-nav-menu-content.tsx +++ b/web/app/components/header/account-dropdown/main-nav-menu-content.tsx @@ -16,11 +16,12 @@ import { DropdownMenuSubContent, DropdownMenuSubTrigger, } from '@langgenius/dify-ui/dropdown-menu' +import { useAtomValue } from 'jotai' import { useTheme } from 'next-themes' import { useTranslation } from 'react-i18next' import PremiumBadge from '@/app/components/base/premium-badge' import { ACCOUNT_SETTING_TAB } from '@/app/components/header/account-setting/constants' -import { useAppContext } from '@/context/app-context' +import { userProfileAtom } from '@/context/app-context-state' import { useModalContext } from '@/context/modal-context' import { useProviderContext } from '@/context/provider-context' import Link from '@/next/link' @@ -87,7 +88,7 @@ export function MainNavMenuContent({ onLogout, }: MainNavMenuContentProps) { const { t } = useTranslation() - const { userProfile } = useAppContext() + const userProfile = useAtomValue(userProfileAtom) const { isEducationAccount } = useProviderContext() const { setShowAccountSettingModal } = useModalContext() diff --git a/web/app/components/header/account-setting/access-rules-page/__tests__/access-rule-section.spec.tsx b/web/app/components/header/account-setting/access-rules-page/__tests__/access-rule-section.spec.tsx index 35dd91459b6..f2a8fbc8f4d 100644 --- a/web/app/components/header/account-setting/access-rules-page/__tests__/access-rule-section.spec.tsx +++ b/web/app/components/header/account-setting/access-rules-page/__tests__/access-rule-section.spec.tsx @@ -15,6 +15,18 @@ vi.mock('@/context/app-context', () => ({ })), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mocks.workspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + const rule: AccessPolicyWithBindings = { policy: { id: 'app-rule-1', diff --git a/web/app/components/header/account-setting/access-rules-page/access-rule-section.tsx b/web/app/components/header/account-setting/access-rules-page/access-rule-section.tsx index fe1498a64f0..fc5bcbcc4b4 100644 --- a/web/app/components/header/account-setting/access-rules-page/access-rule-section.tsx +++ b/web/app/components/header/account-setting/access-rules-page/access-rule-section.tsx @@ -3,11 +3,12 @@ import type { AccessPolicyWithBindings } from '@/models/access-control' import { Button } from '@langgenius/dify-ui/button' import { cn } from '@langgenius/dify-ui/cn' +import { useAtomValue } from 'jotai' import { memo, useEffect, useRef, useState } from 'react' import { useTranslation } from 'react-i18next' import ActionButton from '@/app/components/base/action-button' import Loading from '@/app/components/base/loading' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { workspacePermissionKeysAtom } from '@/context/app-context-state' import { hasPermission } from '@/utils/permission' import AccessRuleRow from './access-rule-row' @@ -46,7 +47,7 @@ const AccessRuleSection = ({ const [expanded, setExpanded] = useState(defaultExpanded) const listRef = useRef(null) const anchorRef = useRef(null) - const workspacePermissionKeys = useAppContextWithSelector(s => s.workspacePermissionKeys) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const canManage = hasPermission(workspacePermissionKeys, 'workspace.role.manage') const ruleCount = totalCount ?? rules.length diff --git a/web/app/components/header/account-setting/api-based-extension-page/__tests__/index.spec.tsx b/web/app/components/header/account-setting/api-based-extension-page/__tests__/index.spec.tsx index af0536fd282..f5f475655e6 100644 --- a/web/app/components/header/account-setting/api-based-extension-page/__tests__/index.spec.tsx +++ b/web/app/components/header/account-setting/api-based-extension-page/__tests__/index.spec.tsx @@ -24,6 +24,18 @@ vi.mock('@/context/app-context', () => ({ })), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockWorkspacePermissionKeys.current, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/service/client', () => ({ consoleQuery: { apiBasedExtension: { diff --git a/web/app/components/header/account-setting/api-based-extension-page/index.tsx b/web/app/components/header/account-setting/api-based-extension-page/index.tsx index 78faacf711f..45d05fdf875 100644 --- a/web/app/components/header/account-setting/api-based-extension-page/index.tsx +++ b/web/app/components/header/account-setting/api-based-extension-page/index.tsx @@ -2,11 +2,12 @@ import type { ApiBasedExtensionResponse } from '@dify/contracts/api/console/api- import type { ReactNode } from 'react' import { Button } from '@langgenius/dify-ui/button' import { useQuery } from '@tanstack/react-query' +import { useAtomValue } from 'jotai' import { useMemo, useState } from 'react' import { useTranslation } from 'react-i18next' import { SearchInput } from '@/app/components/base/search-input' import { SkeletonContainer, SkeletonRectangle, SkeletonRow } from '@/app/components/base/skeleton' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { workspacePermissionKeysAtom } from '@/context/app-context-state' import { consoleQuery } from '@/service/client' import { hasPermission } from '@/utils/permission' import { Empty } from './empty' @@ -51,7 +52,7 @@ export function ApiBasedExtensionPage({ layout, }: ApiBasedExtensionPageProps = {}) { const { t } = useTranslation() - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const canManage = hasPermission(workspacePermissionKeys, 'api_extension.manage') const { data: apiBasedExtensions = [], isPending: isLoading } = useQuery(consoleQuery.apiBasedExtension.get.queryOptions()) const [dialogState, setDialogState] = useState(null) diff --git a/web/app/components/header/account-setting/index.tsx b/web/app/components/header/account-setting/index.tsx index 3d0be5fb7ef..2766957cd3a 100644 --- a/web/app/components/header/account-setting/index.tsx +++ b/web/app/components/header/account-setting/index.tsx @@ -4,6 +4,7 @@ import { Button } from '@langgenius/dify-ui/button' import { cn } from '@langgenius/dify-ui/cn' import { ScrollArea } from '@langgenius/dify-ui/scroll-area' import { useSuspenseQuery } from '@tanstack/react-query' +import { useAtomValue } from 'jotai' import { useCallback, useRef, useState } from 'react' import { useTranslation } from 'react-i18next' import BillingPage from '@/app/components/billing/billing-page' @@ -12,7 +13,7 @@ import { ACCOUNT_SETTING_TAB, } from '@/app/components/header/account-setting/constants' import MenuDialog from '@/app/components/header/account-setting/menu-dialog' -import { useAppContext } from '@/context/app-context' +import { workspacePermissionKeysAtom } from '@/context/app-context-state' import { useProviderContext } from '@/context/provider-context' import { systemFeaturesQueryOptions } from '@/features/system-features/client' import useBreakpoints, { MediaType } from '@/hooks/use-breakpoints' @@ -54,7 +55,7 @@ export default function AccountSetting({ const { t } = useTranslation() const { enableBilling, enableReplaceWebAppLogo } = useProviderContext() const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) - const { workspacePermissionKeys } = useAppContext() + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const isRbacEnabled = systemFeatures.rbac_enabled const canManageWorkspaceRoles = isRbacEnabled && hasPermission(workspacePermissionKeys, 'workspace.role.manage') const canViewBilling = enableBilling && hasPermission(workspacePermissionKeys, BillingPermission.View) diff --git a/web/app/components/header/account-setting/members-page/__tests__/index.spec.tsx b/web/app/components/header/account-setting/members-page/__tests__/index.spec.tsx index 29cae8e0a68..d93d0c10571 100644 --- a/web/app/components/header/account-setting/members-page/__tests__/index.spec.tsx +++ b/web/app/components/header/account-setting/members-page/__tests__/index.spec.tsx @@ -14,7 +14,19 @@ import { useUpdateRolesOfMember } from '@/service/access-control/use-member-role import { useMembers } from '@/service/use-common' import MembersPage from '../index' +const mockAppContextState = vi.hoisted(() => ({ + current: {} as Partial, +})) + vi.mock('@/context/app-context') +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppContextState.current) +}) +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) vi.mock('@/context/provider-context') vi.mock('@/hooks/use-format-time-from-now') vi.mock('@/service/access-control/use-member-roles') @@ -41,6 +53,11 @@ const createRole = (overrides: Partial): Role => ({ ...overrides, }) +const setAppContextValue = (value: AppContextValue) => { + mockAppContextState.current = value + vi.mocked(useAppContext).mockReturnValue(value) +} + vi.mock('../edit-workspace-modal', () => ({ default: ({ onCancel }: { onCancel: () => void }) => (
    @@ -181,7 +198,7 @@ describe('MembersPage', () => { beforeEach(() => { vi.clearAllMocks() - vi.mocked(useAppContext).mockReturnValue({ + setAppContextValue({ userProfile: { email: 'owner@example.com' }, currentWorkspace: { name: 'Test Workspace', role: 'owner' } as ICurrentWorkspace, isCurrentWorkspaceOwner: true, @@ -289,7 +306,7 @@ describe('MembersPage', () => { }) it('should hide manager controls for non-owner non-manager users', () => { - vi.mocked(useAppContext).mockReturnValue({ + setAppContextValue({ userProfile: { email: 'admin@example.com' }, currentWorkspace: { name: 'Test Workspace', role: 'admin' } as ICurrentWorkspace, isCurrentWorkspaceOwner: false, @@ -390,7 +407,7 @@ describe('MembersPage', () => { }) it('should show invite button when user is manager but not owner', () => { - vi.mocked(useAppContext).mockReturnValue({ + setAppContextValue({ userProfile: { email: 'admin@example.com' }, currentWorkspace: { name: 'Test Workspace', role: 'admin' } as ICurrentWorkspace, isCurrentWorkspaceOwner: false, @@ -405,7 +422,7 @@ describe('MembersPage', () => { }) it('should allow admins to operate other non-owner members only', () => { - vi.mocked(useAppContext).mockReturnValue({ + setAppContextValue({ userProfile: { email: 'admin@example.com' }, currentWorkspace: { name: 'Test Workspace', role: 'admin' } as ICurrentWorkspace, isCurrentWorkspaceOwner: false, @@ -483,7 +500,7 @@ describe('MembersPage', () => { }) it('should render role badge names from account roles', () => { - vi.mocked(useAppContext).mockReturnValue({ + setAppContextValue({ userProfile: { email: 'admin@example.com' }, currentWorkspace: { name: 'Test Workspace', role: 'admin' } as ICurrentWorkspace, isCurrentWorkspaceOwner: false, @@ -551,7 +568,7 @@ describe('MembersPage', () => { it('should not allow assigning roles from member details when target is current user', async () => { const user = userEvent.setup() - vi.mocked(useAppContext).mockReturnValue({ + setAppContextValue({ userProfile: { email: 'admin@example.com' }, currentWorkspace: { name: 'Test Workspace', role: 'admin' } as ICurrentWorkspace, isCurrentWorkspaceOwner: false, diff --git a/web/app/components/header/account-setting/members-page/__tests__/invite-button.spec.tsx b/web/app/components/header/account-setting/members-page/__tests__/invite-button.spec.tsx index 43fcce16f74..7ef5e69bacd 100644 --- a/web/app/components/header/account-setting/members-page/__tests__/invite-button.spec.tsx +++ b/web/app/components/header/account-setting/members-page/__tests__/invite-button.spec.tsx @@ -8,6 +8,16 @@ import { useWorkspacePermissions } from '@/service/use-workspace' import InviteButton from '../invite-button' vi.mock('@/context/app-context') +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + currentWorkspace: { id: 'workspace-id' }, + })) +}) +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) vi.mock('@/service/use-workspace') describe('InviteButton', () => { diff --git a/web/app/components/header/account-setting/members-page/edit-workspace-modal/__tests__/dialog.spec.tsx b/web/app/components/header/account-setting/members-page/edit-workspace-modal/__tests__/dialog.spec.tsx index e134f8f93c6..96cfd07cd54 100644 --- a/web/app/components/header/account-setting/members-page/edit-workspace-modal/__tests__/dialog.spec.tsx +++ b/web/app/components/header/account-setting/members-page/edit-workspace-modal/__tests__/dialog.spec.tsx @@ -1,4 +1,5 @@ import type { ReactNode } from 'react' +import type { AppContextValue } from '@/context/app-context' import { render } from '@testing-library/react' import { useAppContext } from '@/context/app-context' import EditWorkspaceModal from '../index' @@ -10,6 +11,9 @@ type DialogProps = { } let latestOnOpenChange: DialogProps['onOpenChange'] +const mockAppContextState = vi.hoisted(() => ({ + current: {} as Partial, +})) vi.mock('@langgenius/dify-ui/dialog', () => ({ Dialog: ({ children, onOpenChange }: DialogProps) => { @@ -29,14 +33,26 @@ vi.mock('@/context/app-context', () => ({ useAppContext: vi.fn(), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppContextState.current) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + describe('EditWorkspaceModal dialog lifecycle', () => { beforeEach(() => { vi.clearAllMocks() latestOnOpenChange = undefined - vi.mocked(useAppContext).mockReturnValue({ + const appContextValue = { currentWorkspace: { name: 'Test Workspace' }, isCurrentWorkspaceOwner: true, - } as never) + } as never + mockAppContextState.current = appContextValue + vi.mocked(useAppContext).mockReturnValue(appContextValue) }) it('should only call onCancel when the dialog requests closing', () => { diff --git a/web/app/components/header/account-setting/members-page/edit-workspace-modal/__tests__/index.spec.tsx b/web/app/components/header/account-setting/members-page/edit-workspace-modal/__tests__/index.spec.tsx index 1b85319613a..5ec4cb538d7 100644 --- a/web/app/components/header/account-setting/members-page/edit-workspace-modal/__tests__/index.spec.tsx +++ b/web/app/components/header/account-setting/members-page/edit-workspace-modal/__tests__/index.spec.tsx @@ -10,10 +10,21 @@ import EditWorkspaceModal from '../index' const toastMocks = vi.hoisted(() => ({ mockNotify: vi.fn(), })) +const mockAppContextState = vi.hoisted(() => ({ + current: {} as Partial, +})) const getSaveButton = () => screen.getByRole('button', { name: /operation\.(save|saving)/i }) vi.mock('@/context/app-context') +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppContextState.current) +}) +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) vi.mock('@/service/common') vi.mock('@langgenius/dify-ui/toast', () => ({ default: { @@ -34,10 +45,12 @@ describe('EditWorkspaceModal', () => { beforeEach(() => { vi.clearAllMocks() - vi.mocked(useAppContext).mockReturnValue({ + const appContextValue = { currentWorkspace: { name: 'Test Workspace' } as ICurrentWorkspace, isCurrentWorkspaceOwner: true, - } as unknown as AppContextValue) + } as unknown as AppContextValue + mockAppContextState.current = appContextValue + vi.mocked(useAppContext).mockReturnValue(appContextValue) }) afterEach(() => { diff --git a/web/app/components/header/account-setting/members-page/edit-workspace-modal/index.tsx b/web/app/components/header/account-setting/members-page/edit-workspace-modal/index.tsx index f9d84d414b4..7d17366df84 100644 --- a/web/app/components/header/account-setting/members-page/edit-workspace-modal/index.tsx +++ b/web/app/components/header/account-setting/members-page/edit-workspace-modal/index.tsx @@ -2,11 +2,12 @@ import { Button } from '@langgenius/dify-ui/button' import { cn } from '@langgenius/dify-ui/cn' import { Dialog, DialogCloseButton, DialogContent, DialogTitle } from '@langgenius/dify-ui/dialog' +import { Input } from '@langgenius/dify-ui/input' import { toast } from '@langgenius/dify-ui/toast' +import { useAtomValue } from 'jotai' import { useId, useMemo, useState } from 'react' import { useTranslation } from 'react-i18next' -import Input from '@/app/components/base/input' -import { useAppContext } from '@/context/app-context' +import { currentWorkspaceAtom, isCurrentWorkspaceOwnerAtom } from '@/context/app-context-state' import { updateWorkspaceInfo } from '@/service/common' type IEditWorkspaceModalProps = { @@ -14,7 +15,8 @@ type IEditWorkspaceModalProps = { } const EditWorkspaceModal = ({ onCancel }: IEditWorkspaceModalProps) => { const { t } = useTranslation() - const { currentWorkspace, isCurrentWorkspaceOwner } = useAppContext() + const currentWorkspace = useAtomValue(currentWorkspaceAtom) + const isCurrentWorkspaceOwner = useAtomValue(isCurrentWorkspaceOwnerAtom) const [name, setName] = useState(currentWorkspace.name) const [isSubmitting, setIsSubmitting] = useState(false) const inputId = useId() @@ -82,7 +84,6 @@ const EditWorkspaceModal = ({ onCancel }: IEditWorkspaceModalProps) => { { diff --git a/web/app/components/header/account-setting/members-page/index.tsx b/web/app/components/header/account-setting/members-page/index.tsx index 478952e5c6d..bd745507367 100644 --- a/web/app/components/header/account-setting/members-page/index.tsx +++ b/web/app/components/header/account-setting/members-page/index.tsx @@ -4,12 +4,13 @@ import type { InvitationResult, Member } from '@/models/common' import { toast } from '@langgenius/dify-ui/toast' import { Tooltip, TooltipContent, TooltipTrigger } from '@langgenius/dify-ui/tooltip' import { useSuspenseQuery } from '@tanstack/react-query' +import { useAtomValue } from 'jotai' import { useCallback, useState } from 'react' import { useTranslation } from 'react-i18next' import { NUM_INFINITE } from '@/app/components/billing/config' import { Plan } from '@/app/components/billing/type' import UpgradeBtn from '@/app/components/billing/upgrade-btn' -import { useAppContext } from '@/context/app-context' +import { currentWorkspaceAtom, isCurrentWorkspaceOwnerAtom, userProfileEmailAtom, workspacePermissionKeysAtom } from '@/context/app-context-state' import { useLocale } from '@/context/i18n' import { useProviderContext } from '@/context/provider-context' import { systemFeaturesQueryOptions } from '@/features/system-features/client' @@ -30,7 +31,10 @@ const MembersPage = () => { const locale = useLocale() const language = getAccessControlTemplateLanguage(locale) - const { userProfile, currentWorkspace, isCurrentWorkspaceOwner, workspacePermissionKeys } = useAppContext() + const userProfileEmail = useAtomValue(userProfileEmailAtom) + const currentWorkspace = useAtomValue(currentWorkspaceAtom) + const isCurrentWorkspaceOwner = useAtomValue(isCurrentWorkspaceOwnerAtom) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const { data, refetch } = useMembers(language) const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) const [inviteModalVisible, setInviteModalVisible] = useState(false) @@ -162,7 +166,7 @@ const MembersPage = () => { key={account.id} member={account} roles={account.roles} - isCurrentUser={userProfile.email === account.email} + isCurrentUser={userProfileEmail === account.email} canManage={canManageMembers} canTransferOwnership={isCurrentWorkspaceOwner && isAllowTransferWorkspace} allowMultipleRoles={systemFeatures.rbac_enabled} @@ -213,7 +217,7 @@ const MembersPage = () => { canAssignRoles={ canManageMembers && detailsMember.role !== 'owner' - && userProfile.email !== detailsMember.email + && userProfileEmail !== detailsMember.email } allowMultipleRoles={systemFeatures.rbac_enabled} onClose={handleCloseDetails} diff --git a/web/app/components/header/account-setting/members-page/invite-button.tsx b/web/app/components/header/account-setting/members-page/invite-button.tsx index e0a9408515a..e08415eedc9 100644 --- a/web/app/components/header/account-setting/members-page/invite-button.tsx +++ b/web/app/components/header/account-setting/members-page/invite-button.tsx @@ -1,9 +1,10 @@ import { Button } from '@langgenius/dify-ui/button' import { RiUserAddLine } from '@remixicon/react' import { useSuspenseQuery } from '@tanstack/react-query' +import { useAtomValue } from 'jotai' import { useTranslation } from 'react-i18next' import Loading from '@/app/components/base/loading' -import { useAppContext } from '@/context/app-context' +import { currentWorkspaceIdAtom } from '@/context/app-context-state' import { systemFeaturesQueryOptions } from '@/features/system-features/client' import { useWorkspacePermissions } from '@/service/use-workspace' @@ -14,9 +15,9 @@ type InviteButtonProps = { const InviteButton = (props: InviteButtonProps) => { const { t } = useTranslation() - const { currentWorkspace } = useAppContext() + const currentWorkspaceId = useAtomValue(currentWorkspaceIdAtom) const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) - const { data: workspacePermissions, isFetching: isFetchingWorkspacePermissions } = useWorkspacePermissions(currentWorkspace!.id, systemFeatures.branding.enabled) + const { data: workspacePermissions, isFetching: isFetchingWorkspacePermissions } = useWorkspacePermissions(currentWorkspaceId, systemFeatures.branding.enabled) if (systemFeatures.branding.enabled) { if (isFetchingWorkspacePermissions) { return diff --git a/web/app/components/header/account-setting/members-page/transfer-ownership-modal/__tests__/index.spec.tsx b/web/app/components/header/account-setting/members-page/transfer-ownership-modal/__tests__/index.spec.tsx index 07371740bd8..db02b773cc7 100644 --- a/web/app/components/header/account-setting/members-page/transfer-ownership-modal/__tests__/index.spec.tsx +++ b/web/app/components/header/account-setting/members-page/transfer-ownership-modal/__tests__/index.spec.tsx @@ -11,8 +11,19 @@ import TransferOwnershipModal from '../index' const toastMocks = vi.hoisted(() => ({ mockNotify: vi.fn(), })) +const mockAppContextState = vi.hoisted(() => ({ + current: {} as Partial, +})) vi.mock('@/context/app-context') +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppContextState.current) +}) +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) vi.mock('@/service/common') vi.mock('@/service/use-common') vi.mock('@langgenius/dify-ui/toast', () => ({ @@ -40,10 +51,12 @@ describe('TransferOwnershipModal', () => { beforeEach(() => { vi.clearAllMocks() - vi.mocked(useAppContext).mockReturnValue({ + const appContextValue = { currentWorkspace: { name: 'Test Workspace' } as ICurrentWorkspace, userProfile: { email: 'owner@example.com', id: 'owner-id' }, - } as unknown as AppContextValue) + } as unknown as AppContextValue + mockAppContextState.current = appContextValue + vi.mocked(useAppContext).mockReturnValue(appContextValue) vi.mocked(useMembers).mockReturnValue({ data: { accounts: [] }, diff --git a/web/app/components/header/account-setting/members-page/transfer-ownership-modal/index.tsx b/web/app/components/header/account-setting/members-page/transfer-ownership-modal/index.tsx index 7bfb3ee3f1f..ee198a7b555 100644 --- a/web/app/components/header/account-setting/members-page/transfer-ownership-modal/index.tsx +++ b/web/app/components/header/account-setting/members-page/transfer-ownership-modal/index.tsx @@ -1,11 +1,12 @@ import { Button } from '@langgenius/dify-ui/button' import { Dialog, DialogContent } from '@langgenius/dify-ui/dialog' +import { Input } from '@langgenius/dify-ui/input' import { toast } from '@langgenius/dify-ui/toast' +import { useAtomValue } from 'jotai' import * as React from 'react' import { useCallback, useState } from 'react' import { Trans, useTranslation } from 'react-i18next' -import Input from '@/app/components/base/input' -import { useAppContext } from '@/context/app-context' +import { currentWorkspaceAtom, userProfileAtom } from '@/context/app-context-state' import { ownershipTransfer, sendOwnerEmail, verifyOwnerEmail } from '@/service/common' import MemberSelector from './member-selector' @@ -13,15 +14,22 @@ type Props = Readonly<{ show: boolean onClose: () => void }> -enum STEP { - start = 'start', - verify = 'verify', - transfer = 'transfer', +const STEP = { + start: 'start', + verify: 'verify', + transfer: 'transfer', } +type Step = typeof STEP[keyof typeof STEP] + +const getErrorMessage = (error: unknown) => { + return error instanceof Error ? error.message : '' +} + const TransferOwnershipModal = ({ onClose, show }: Props) => { const { t } = useTranslation() - const { currentWorkspace, userProfile } = useAppContext() - const [step, setStep] = useState(STEP.start) + const currentWorkspace = useAtomValue(currentWorkspaceAtom) + const userProfile = useAtomValue(userProfileAtom) + const [step, setStep] = useState(STEP.start) const [code, setCode] = useState('') const [time, setTime] = useState(0) const [stepToken, setStepToken] = useState('') @@ -71,7 +79,7 @@ const TransferOwnershipModal = ({ onClose, show }: Props) => { } } catch (error) { - toast.error(`Error verifying email: ${error ? (error as any).message : ''}`) + toast.error(`Error verifying email: ${getErrorMessage(error)}`) } } const sendCodeToOriginEmail = async () => { @@ -96,7 +104,7 @@ const TransferOwnershipModal = ({ onClose, show }: Props) => { globalThis.location.reload() } catch (error) { - toast.error(`Error ownership transfer: ${error ? (error as any).message : ''}`) + toast.error(`Error ownership transfer: ${getErrorMessage(error)}`) } finally { setIsTransfer(false) diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/model-list-item.spec.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/model-list-item.spec.tsx index 6b6540ec261..065baa2496f 100644 --- a/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/model-list-item.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/model-list-item.spec.tsx @@ -22,6 +22,18 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockWorkspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/provider-context', () => ({ useProviderContext: () => ({ plan: { type: mockPlanType }, diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/model-list.spec.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/model-list.spec.tsx index c307407c0fd..1d5055a0741 100644 --- a/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/model-list.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/model-list.spec.tsx @@ -11,6 +11,18 @@ vi.mock('@/context/app-context', () => ({ selector({ workspacePermissionKeys: mockWorkspacePermissionKeys }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockWorkspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/modal-context', () => ({ useModalContextSelector: (selector: (state: { setShowModelLoadBalancingModal: typeof mockSetShowModelLoadBalancingModal }) => unknown) => selector({ setShowModelLoadBalancingModal: mockSetShowModelLoadBalancingModal }), diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/index.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/index.tsx index 90bd54c5e63..598254ceeda 100644 --- a/web/app/components/header/account-setting/model-provider-page/provider-added-card/index.tsx +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/index.tsx @@ -7,6 +7,7 @@ import type { PluginDetail } from '@/app/components/plugins/types' import { cn } from '@langgenius/dify-ui/cn' import { useQuery } from '@tanstack/react-query' +import { useAtomValue } from 'jotai' import { memo, useCallback } from 'react' import { useTranslation } from 'react-i18next' import { @@ -14,7 +15,7 @@ import { ManageCustomModelCredentials, } from '@/app/components/header/account-setting/model-provider-page/model-auth' import { IS_CE_EDITION } from '@/config' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { workspacePermissionKeysAtom } from '@/context/app-context-state' import { useProviderContextSelector } from '@/context/provider-context' import { useCredentialPermissions } from '@/hooks/use-credential-permissions' import { renderI18nObject } from '@/i18n-config' @@ -68,7 +69,7 @@ const ProviderAddedCard: FC = ({ })) const hasModelList = hasFetchedModelList && !!modelList.length const showCollapsedSection = !expanded || !hasFetchedModelList - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const showModelProvider = systemConfig.enabled && MODEL_PROVIDER_QUOTA_GET_PAID.includes(currentProviderName as ModelProviderQuotaGetPaid) && !IS_CE_EDITION const canConfigureModels = hasPermission(workspacePermissionKeys, 'plugin.model_config') const { canUseCredential, canCreateCredential, canManageCredential } = useCredentialPermissions() diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-list-item.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-list-item.tsx index ba890ab8c90..ea642daca49 100644 --- a/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-list-item.tsx +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-list-item.tsx @@ -4,12 +4,13 @@ import { Popover, PopoverContent, PopoverTrigger } from '@langgenius/dify-ui/pop import { Switch } from '@langgenius/dify-ui/switch' import { useQueryClient } from '@tanstack/react-query' import { useDebounceFn } from 'ahooks' +import { useAtomValue } from 'jotai' import { memo, useCallback } from 'react' import { useTranslation } from 'react-i18next' import Badge from '@/app/components/base/badge' import { Balance } from '@/app/components/base/icons/src/vender/line/financeAndECommerce' import { Plan } from '@/app/components/billing/type' -import { useAppContext } from '@/context/app-context' +import { workspacePermissionKeysAtom } from '@/context/app-context-state' import { useProviderContext, useProviderContextSelector } from '@/context/provider-context' import { consoleQuery } from '@/service/client' import { disableModel, enableModel } from '@/service/common' @@ -32,7 +33,7 @@ const ModelListItem = ({ model, provider, isConfigurable, onChange, onModifyLoad const { t } = useTranslation() const { plan } = useProviderContext() const modelLoadBalancingEnabled = useProviderContextSelector(state => state.modelLoadBalancingEnabled) - const { workspacePermissionKeys } = useAppContext() + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const canConfigureModels = hasPermission(workspacePermissionKeys, 'plugin.model_config') const queryClient = useQueryClient() const updateModelList = useUpdateModelList() diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-list.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-list.tsx index 0b818542e5c..b1a6ceccf27 100644 --- a/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-list.tsx +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-list.tsx @@ -4,13 +4,14 @@ import type { ModelItem, ModelProvider, } from '../declarations' +import { useAtomValue } from 'jotai' import { useCallback } from 'react' import { useTranslation } from 'react-i18next' import { AddCustomModel, ManageCustomModelCredentials, } from '@/app/components/header/account-setting/model-provider-page/model-auth' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { workspacePermissionKeysAtom } from '@/context/app-context-state' import { useModalContextSelector } from '@/context/modal-context' import { hasPermission } from '@/utils/permission' import { @@ -33,7 +34,7 @@ const ModelList: FC = ({ }) => { const { t } = useTranslation() const configurativeMethods = provider.configurate_methods.filter(method => method !== ConfigurationMethodEnum.fetchFromRemote) - const workspacePermissionKeys = useAppContextWithSelector(state => state.workspacePermissionKeys) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const canConfigureModels = hasPermission(workspacePermissionKeys, 'plugin.model_config') const isConfigurable = configurativeMethods.includes(ConfigurationMethodEnum.customizableModel) const setShowModelLoadBalancingModal = useModalContextSelector(state => state.setShowModelLoadBalancingModal) @@ -58,13 +59,14 @@ const ModelList: FC = ({ {t('modelProvider.modelsNum', { ns: 'common', num: models.length })} - onCollapse()} > {t('modelProvider.modelsNum', { ns: 'common', num: models.length })} - + { isConfigurable && canConfigureModels && ( diff --git a/web/app/components/header/account-setting/model-provider-page/system-model-selector/__tests__/index.spec.tsx b/web/app/components/header/account-setting/model-provider-page/system-model-selector/__tests__/index.spec.tsx index dc4d716b6ce..4bbf4ec7bb6 100644 --- a/web/app/components/header/account-setting/model-provider-page/system-model-selector/__tests__/index.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/system-model-selector/__tests__/index.spec.tsx @@ -40,6 +40,18 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockWorkspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/provider-context', () => ({ useProviderContext: () => ({ textGenerationModelList: [], diff --git a/web/app/components/header/account-setting/model-provider-page/system-model-selector/index.tsx b/web/app/components/header/account-setting/model-provider-page/system-model-selector/index.tsx index 6fe3fc917f2..ece50247ec8 100644 --- a/web/app/components/header/account-setting/model-provider-page/system-model-selector/index.tsx +++ b/web/app/components/header/account-setting/model-provider-page/system-model-selector/index.tsx @@ -12,10 +12,11 @@ import { DialogTitle, } from '@langgenius/dify-ui/dialog' import { toast } from '@langgenius/dify-ui/toast' +import { useAtomValue } from 'jotai' import { useState } from 'react' import { useTranslation } from 'react-i18next' import { Infotip } from '@/app/components/base/infotip' -import { useAppContext } from '@/context/app-context' +import { workspacePermissionKeysAtom } from '@/context/app-context-state' import { useProviderContext } from '@/context/provider-context' import { updateDefaultModel } from '@/service/common' import { hasPermission } from '@/utils/permission' @@ -68,7 +69,7 @@ const SystemModel: FC = ({ onOpenMarketplace, }) => { const { t } = useTranslation() - const { workspacePermissionKeys } = useAppContext() + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const { textGenerationModelList } = useProviderContext() const canManageSystemDefaultModel = hasPermission(workspacePermissionKeys, 'plugin.model_config') const updateModelList = useUpdateModelList() diff --git a/web/app/components/header/account-setting/permissions-page/__tests__/index.spec.tsx b/web/app/components/header/account-setting/permissions-page/__tests__/index.spec.tsx index b37c07d1054..c0c0efb4f95 100644 --- a/web/app/components/header/account-setting/permissions-page/__tests__/index.spec.tsx +++ b/web/app/components/header/account-setting/permissions-page/__tests__/index.spec.tsx @@ -26,6 +26,18 @@ vi.mock('@/context/app-context', () => ({ } as AppContextValue)), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mocks.workspacePermissionKeys, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/service/access-control/use-workspace-roles', () => ({ useCreateWorkspaceRole: vi.fn(), useUpdateWorkspaceRole: vi.fn(), diff --git a/web/app/components/header/account-setting/permissions-page/index.tsx b/web/app/components/header/account-setting/permissions-page/index.tsx index e5d3408106d..f1fa6572ca9 100644 --- a/web/app/components/header/account-setting/permissions-page/index.tsx +++ b/web/app/components/header/account-setting/permissions-page/index.tsx @@ -4,9 +4,10 @@ import type { RoleModalMode, submitRoleData } from './role-modal' import type { Role } from '@/models/access-control' import { Button } from '@langgenius/dify-ui/button' import { toast } from '@langgenius/dify-ui/toast' +import { useAtomValue } from 'jotai' import { useCallback, useEffect, useMemo, useRef, useState } from 'react' import { useTranslation } from 'react-i18next' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { workspacePermissionKeysAtom } from '@/context/app-context-state' import { useLocale } from '@/context/i18n' import { getAccessControlTemplateLanguage } from '@/i18n-config/language' import { useCreateWorkspaceRole, useUpdateWorkspaceRole } from '@/service/access-control/use-workspace-roles' @@ -32,7 +33,7 @@ const PermissionsPage = ({ containerRef }: PermissionsPageProps) => { const [modalState, setModalState] = useState(null) const anchorRef = useRef(null) - const workspacePermissionKeys = useAppContextWithSelector(s => s.workspacePermissionKeys) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const language = useMemo(() => getAccessControlTemplateLanguage(locale), [locale]) diff --git a/web/app/components/header/account-setting/permissions-page/role-list/__tests__/row-menu.spec.tsx b/web/app/components/header/account-setting/permissions-page/role-list/__tests__/row-menu.spec.tsx index 90558b49896..46412f3f16f 100644 --- a/web/app/components/header/account-setting/permissions-page/role-list/__tests__/row-menu.spec.tsx +++ b/web/app/components/header/account-setting/permissions-page/role-list/__tests__/row-menu.spec.tsx @@ -36,6 +36,18 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockWorkspacePermissionKeys.value, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/service/access-control/use-workspace-roles', () => ({ useCopyWorkspaceRole: () => ({ mutate: mockCopyRole, diff --git a/web/app/components/header/account-setting/permissions-page/role-list/row-menu.tsx b/web/app/components/header/account-setting/permissions-page/role-list/row-menu.tsx index 678357ea4bc..145a83c3467 100644 --- a/web/app/components/header/account-setting/permissions-page/role-list/row-menu.tsx +++ b/web/app/components/header/account-setting/permissions-page/role-list/row-menu.tsx @@ -18,10 +18,11 @@ import { DropdownMenuTrigger, } from '@langgenius/dify-ui/dropdown-menu' import { toast } from '@langgenius/dify-ui/toast' +import { useAtomValue } from 'jotai' import { useCallback, useState } from 'react' import { useTranslation } from 'react-i18next' import ActionButton from '@/app/components/base/action-button' -import { useSelector as useAppContextWithSelector } from '@/context/app-context' +import { workspacePermissionKeysAtom } from '@/context/app-context-state' import { useCopyWorkspaceRole, useDeleteWorkspaceRole } from '@/service/access-control/use-workspace-roles' import { hasPermission } from '@/utils/permission' import { CopyMembersConfirmDialog } from './copy-members-confirm-dialog' @@ -44,7 +45,7 @@ const RowMenu = ({ const [showDeleteConfirm, setShowDeleteConfirm] = useState(false) const [showCopyMembersConfirm, setShowCopyMembersConfirm] = useState(false) - const workspacePermissionKeys = useAppContextWithSelector(s => s.workspacePermissionKeys) + const workspacePermissionKeys = useAtomValue(workspacePermissionKeysAtom) const canManageRoles = hasPermission(workspacePermissionKeys, 'workspace.role.manage') const handleView = useCallback(() => onView?.(role), [onView, role]) diff --git a/web/app/components/header/account-setting/preference-page/__tests__/index.spec.tsx b/web/app/components/header/account-setting/preference-page/__tests__/index.spec.tsx index 55ee81481ad..db47bfe571d 100644 --- a/web/app/components/header/account-setting/preference-page/__tests__/index.spec.tsx +++ b/web/app/components/header/account-setting/preference-page/__tests__/index.spec.tsx @@ -75,6 +75,19 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + userProfile: mockUserProfile, + refreshUserProfile: mockMutateUserProfile, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/i18n', () => ({ useLocale: () => mockLocale, })) diff --git a/web/app/components/header/account-setting/preference-page/index.tsx b/web/app/components/header/account-setting/preference-page/index.tsx index 1574d254cd2..c8a62fb90c5 100644 --- a/web/app/components/header/account-setting/preference-page/index.tsx +++ b/web/app/components/header/account-setting/preference-page/index.tsx @@ -2,10 +2,11 @@ import type { Locale } from '@/i18n-config' import { Select, SelectContent, SelectItem, SelectItemIndicator, SelectItemText, SelectTrigger } from '@langgenius/dify-ui/select' import { toast } from '@langgenius/dify-ui/toast' +import { useAtomValue, useSetAtom } from 'jotai' import { useTheme } from 'next-themes' import { useState } from 'react' import { useTranslation } from 'react-i18next' -import { useAppContext } from '@/context/app-context' +import { refreshUserProfileAtom, userProfileAtom } from '@/context/app-context-state' import { useLocale } from '@/context/i18n' import { setLocaleOnClient } from '@/i18n-config' import { languages } from '@/i18n-config/language' @@ -35,7 +36,8 @@ const isThemeOption = (value: string): value is ThemeOption => { export default function PreferencePage() { const locale = useLocale() - const { userProfile, mutateUserProfile } = useAppContext() + const userProfile = useAtomValue(userProfileAtom) + const refreshUserProfile = useSetAtom(refreshUserProfileAtom) const [editing, setEditing] = useState(false) const { t } = useTranslation() const router = useRouter() @@ -77,7 +79,7 @@ export default function PreferencePage() { try { await updateUserProfile({ url, body: { [bodyKey]: item.value } }) toast.success(t('actionMsg.modifiedSuccessfully', { ns: 'common' })) - mutateUserProfile() + refreshUserProfile() } catch (e) { toast.error((e as Error).message) diff --git a/web/app/components/header/account-setting/update-setting-dialog.tsx b/web/app/components/header/account-setting/update-setting-dialog.tsx index 6e512b12af4..2d72c46cf3b 100644 --- a/web/app/components/header/account-setting/update-setting-dialog.tsx +++ b/web/app/components/header/account-setting/update-setting-dialog.tsx @@ -7,12 +7,13 @@ import { Button } from '@langgenius/dify-ui/button' import { cn } from '@langgenius/dify-ui/cn' import { Dialog, DialogCloseButton, DialogContent, DialogTitle, DialogTrigger } from '@langgenius/dify-ui/dialog' import { toast } from '@langgenius/dify-ui/toast' +import { useAtomValue } from 'jotai' import { useCallback, useMemo, useState } from 'react' import { useTranslation } from 'react-i18next' import { convertTimezoneToOffsetStr } from '@/app/components/base/date-and-time-picker/utils/dayjs' import { AUTO_UPDATE_MODE, AUTO_UPDATE_STRATEGY } from '@/app/components/plugins/reference-setting-modal/auto-update-setting/types' import { convertLocalSecondsToUTCDaySeconds, convertUTCDaySecondsToLocalSeconds, dayjsToTimeOfDay, timeOfDayToDayjs } from '@/app/components/plugins/reference-setting-modal/auto-update-setting/utils' -import { useAppContext } from '@/context/app-context' +import { userProfileAtom } from '@/context/app-context-state' import { useMutationPluginAutoUpgradeSettings, usePluginAutoUpgradeSettings } from '@/service/use-plugins' import UpdateSettingDialogForm from './update-setting-dialog-form' @@ -26,7 +27,7 @@ const UpdateSettingDialog = ({ disabled = false, }: Props) => { const { t } = useTranslation() - const { userProfile } = useAppContext() + const userProfile = useAtomValue(userProfileAtom) const timezone = userProfile.timezone || 'UTC' const { data: autoUpgradeSetting, diff --git a/web/app/components/header/env-nav/__tests__/index.spec.tsx b/web/app/components/header/env-nav/__tests__/index.spec.tsx index 81076e773fe..98832145dc3 100644 --- a/web/app/components/header/env-nav/__tests__/index.spec.tsx +++ b/web/app/components/header/env-nav/__tests__/index.spec.tsx @@ -4,10 +4,24 @@ import { vi } from 'vitest' import { useAppContext } from '@/context/app-context' import EnvNav from '../index' +const mockAppContextState = vi.hoisted(() => ({ + current: {} as Partial, +})) + vi.mock('@/context/app-context', () => ({ useAppContext: vi.fn(), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppContextState.current) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + describe('EnvNav', () => { const mockUseAppContext = vi.mocked(useAppContext) @@ -16,33 +30,39 @@ describe('EnvNav', () => { }) it('should render null when environment is PRODUCTION', () => { - mockUseAppContext.mockReturnValue({ + const appContextValue = { langGeniusVersionInfo: { current_env: 'PRODUCTION', }, - } as unknown as AppContextValue) + } as unknown as AppContextValue + mockAppContextState.current = appContextValue + mockUseAppContext.mockReturnValue(appContextValue) const { container } = render() expect(container.firstChild).toBeNull() }) it('should render TESTING tag and icon when environment is TESTING', () => { - mockUseAppContext.mockReturnValue({ + const appContextValue = { langGeniusVersionInfo: { current_env: 'TESTING', }, - } as unknown as AppContextValue) + } as unknown as AppContextValue + mockAppContextState.current = appContextValue + mockUseAppContext.mockReturnValue(appContextValue) render() expect(screen.getByText('common.environment.testing')).toBeInTheDocument() }) it('should render DEVELOPMENT tag and icon when environment is DEVELOPMENT', () => { - mockUseAppContext.mockReturnValue({ + const appContextValue = { langGeniusVersionInfo: { current_env: 'DEVELOPMENT', }, - } as unknown as AppContextValue) + } as unknown as AppContextValue + mockAppContextState.current = appContextValue + mockUseAppContext.mockReturnValue(appContextValue) render() expect( diff --git a/web/app/components/header/env-nav/index.tsx b/web/app/components/header/env-nav/index.tsx index 6323f3428e9..a27daa7638a 100644 --- a/web/app/components/header/env-nav/index.tsx +++ b/web/app/components/header/env-nav/index.tsx @@ -1,9 +1,10 @@ 'use client' +import { useAtomValue } from 'jotai' import { useTranslation } from 'react-i18next' import { TerminalSquare } from '@/app/components/base/icons/src/vender/solid/development' import { Beaker02 } from '@/app/components/base/icons/src/vender/solid/education' -import { useAppContext } from '@/context/app-context' +import { langGeniusVersionInfoAtom } from '@/context/app-context-state' const headerEnvClassName: { [k: string]: string } = { DEVELOPMENT: 'bg-[#FEC84B] border-[#FDB022] text-[#93370D]', @@ -12,7 +13,7 @@ const headerEnvClassName: { [k: string]: string } = { const EnvNav = () => { const { t } = useTranslation() - const { langGeniusVersionInfo } = useAppContext() + const langGeniusVersionInfo = useAtomValue(langGeniusVersionInfoAtom) const showEnvTag = langGeniusVersionInfo.current_env === 'TESTING' || langGeniusVersionInfo.current_env === 'DEVELOPMENT' if (!showEnvTag) diff --git a/web/app/components/main-nav/__tests__/index.spec.tsx b/web/app/components/main-nav/__tests__/index.spec.tsx index c4175b44418..f9c9f6ca999 100644 --- a/web/app/components/main-nav/__tests__/index.spec.tsx +++ b/web/app/components/main-nav/__tests__/index.spec.tsx @@ -30,6 +30,9 @@ const { mockIsAgentV2Enabled, mockSwitchWorkspace, mockToastSuccess } = vi.hoist mockToastSuccess: vi.fn(), mockIsAgentV2Enabled: vi.fn(() => true), })) +const mockAppContextState = vi.hoisted(() => ({ + current: undefined as AppContextValue | undefined, +})) vi.mock('@/features/agent-v2/feature-flag', () => ({ isAgentV2Enabled: () => mockIsAgentV2Enabled(), @@ -40,6 +43,16 @@ vi.mock('@/context/app-context', () => ({ useSelector: vi.fn(), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppContextState.current ?? {}) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/context/provider-context', () => ({ useProviderContext: vi.fn(), })) @@ -240,6 +253,7 @@ const renderMainNav = ( const queryClient = createTestQueryClient() const getMockAppContext = useAppContext as Mock const currentAppContext = getMockAppContext() as AppContextValue + mockAppContextState.current = currentAppContext queryClient.setQueryData(consoleQuery.workspaces.current.post.queryKey(), currentAppContext.currentWorkspace) queryClient.setQueryData(consoleQuery.workspaces.get.queryKey(), { workspaces: mockWorkspaces }) const resolvedSystemFeatures = { @@ -287,6 +301,7 @@ describe('MainNav', () => { refresh: vi.fn(), }) ;(useAppContext as Mock).mockReturnValue(appContextValue) + mockAppContextState.current = appContextValue ;(useAppContextSelector as Mock).mockImplementation((selector: (state: AppContextValue) => unknown) => selector((useAppContext as Mock)() as AppContextValue)) ;(useProviderContext as Mock).mockReturnValue({ enableBilling: true, From a278d21741652434d9b05086405a7fd420a72cc6 Mon Sep 17 00:00:00 2001 From: Asuka Minato Date: Wed, 8 Jul 2026 17:24:31 +0900 Subject: [PATCH 52/70] test: more caplog (#38452) Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- .../telemetry/test_enterprise_trace.py | 24 ++++++++++++++----- 1 file changed, 18 insertions(+), 6 deletions(-) diff --git a/api/tests/unit_tests/enterprise/telemetry/test_enterprise_trace.py b/api/tests/unit_tests/enterprise/telemetry/test_enterprise_trace.py index 24c905c75c2..29ddf73e24d 100644 --- a/api/tests/unit_tests/enterprise/telemetry/test_enterprise_trace.py +++ b/api/tests/unit_tests/enterprise/telemetry/test_enterprise_trace.py @@ -3,8 +3,9 @@ from __future__ import annotations import json +import logging from datetime import UTC, datetime -from typing import Any +from typing import Any, cast from unittest.mock import MagicMock, patch import pytest @@ -475,13 +476,24 @@ class TestWorkflowTrace: assert span_call[1]["start_time"] == _T0 assert span_call[1]["end_time"] == _T1 - def test_emits_companion_log_with_event_name(self, trace_handler: EnterpriseOtelTrace, mock_exporter): - with patch("enterprise.telemetry.enterprise_trace.emit_telemetry_log") as mock_log: + def test_emits_companion_log_with_event_name( + self, + trace_handler: EnterpriseOtelTrace, + caplog: pytest.LogCaptureFixture, + ): + with caplog.at_level(logging.INFO, logger="dify.telemetry"): trace_handler._workflow_trace(make_workflow_info()) - mock_log.assert_called_once() - assert mock_log.call_args[1]["event_name"] == EnterpriseTelemetryEvent.WORKFLOW_RUN - assert mock_log.call_args[1]["tenant_id"] == "tenant-abc" + records = [record for record in caplog.records if record.name == "dify.telemetry"] + + assert len(records) == 1 + record = records[0] + + attrs = cast(dict[str, Any], record.__dict__["attributes"]) + + assert attrs["dify.event.name"] == EnterpriseTelemetryEvent.WORKFLOW_RUN + assert attrs["dify.event.signal"] == "span_detail" + assert record.__dict__["tenant_id"] == "tenant-abc" def test_companion_log_includes_content_when_enabled(self, trace_handler: EnterpriseOtelTrace, mock_exporter): mock_exporter.include_content = True From eca2d419b222bfa8a9face7042dacebf149ccb86 Mon Sep 17 00:00:00 2001 From: Stephen Zhou Date: Wed, 8 Jul 2026 17:01:30 +0800 Subject: [PATCH 53/70] refactor(web): migrate workflow app context consumers (#38552) --- .../__tests__/selection-contextmenu.spec.tsx | 12 +++++++ .../workflow/comment/comment-icon.spec.tsx | 26 ++++++++++++-- .../workflow/comment/comment-icon.tsx | 7 ++-- .../workflow/comment/comment-input.spec.tsx | 23 +++++++++--- .../workflow/comment/comment-input.tsx | 5 +-- .../workflow/comment/thread.spec.tsx | 23 +++++++++--- .../components/workflow/comment/thread.tsx | 13 +++---- .../__tests__/header-in-restoring.spec.tsx | 16 +++++++++ .../header/__tests__/header-layouts.spec.tsx | 16 +++++++++ .../workflow/header/header-in-restoring.tsx | 5 +-- .../workflow/header/online-users.tsx | 6 ++-- .../__tests__/use-workflow-comment.spec.ts | 25 +++++++++---- .../workflow/hooks/use-workflow-comment.ts | 5 +-- .../nodes/_base/__tests__/node.spec.tsx | 19 ++++++++++ .../_base/components/workflow-panel/index.tsx | 5 +-- .../components/workflow/nodes/_base/node.tsx | 5 +-- .../__tests__/email-configure-modal.spec.tsx | 23 ++++++++---- .../__tests__/method-item.spec.tsx | 23 ++++++++---- .../__tests__/test-email-sender.spec.tsx | 36 +++++++++++++------ .../delivery-method/email-configure-modal.tsx | 5 +-- .../delivery-method/method-item.tsx | 5 +-- .../recipient/__tests__/index.spec.tsx | 19 +++++++--- .../delivery-method/recipient/index.tsx | 10 +++--- .../delivery-method/test-email-sender.tsx | 14 ++++---- .../__tests__/integration.spec.tsx | 15 ++++++++ .../components/dataset-list.tsx | 7 ++-- web/app/components/workflow/operator/hooks.ts | 5 +-- .../comments-panel/__tests__/index.spec.tsx | 15 +++++++- .../workflow/panel/comments-panel/index.tsx | 9 ++--- .../__tests__/index.spec.tsx | 18 +++++++++- .../panel/version-history-panel/index.tsx | 5 +-- .../workflow/selection-contextmenu.tsx | 5 +-- 32 files changed, 328 insertions(+), 97 deletions(-) diff --git a/web/app/components/workflow/__tests__/selection-contextmenu.spec.tsx b/web/app/components/workflow/__tests__/selection-contextmenu.spec.tsx index 3ae8851a23f..422a602c2d8 100644 --- a/web/app/components/workflow/__tests__/selection-contextmenu.spec.tsx +++ b/web/app/components/workflow/__tests__/selection-contextmenu.spec.tsx @@ -29,6 +29,18 @@ vi.mock('@/context/app-context', () => ({ }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + workspacePermissionKeys: mockWorkspacePermissionKeys.value, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/app/components/snippets/hooks/use-create-snippet', async () => { const React = await vi.importActual('react') diff --git a/web/app/components/workflow/comment/comment-icon.spec.tsx b/web/app/components/workflow/comment/comment-icon.spec.tsx index df579442fce..017c83b272b 100644 --- a/web/app/components/workflow/comment/comment-icon.spec.tsx +++ b/web/app/components/workflow/comment/comment-icon.spec.tsx @@ -6,6 +6,13 @@ import { CommentIcon } from './comment-icon' type Position = { x: number, y: number } let mockUserId = 'user-1' +const mockAppContextState = vi.hoisted(() => ({ + userProfile: { + id: 'user-1', + name: 'User', + avatar_url: 'avatar', + }, +})) const mockFlowToScreenPosition = vi.fn((position: Position) => position) const mockScreenToFlowPosition = vi.fn((position: Position) => position) @@ -25,13 +32,28 @@ vi.mock('reactflow', () => ({ vi.mock('@/context/app-context', () => ({ useAppContext: () => ({ userProfile: { + ...mockAppContextState.userProfile, id: mockUserId, - name: 'User', - avatar_url: 'avatar', }, }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => ({ + ...mockAppContextState, + userProfile: { + ...mockAppContextState.userProfile, + id: mockUserId, + }, + })) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/app/components/base/user-avatar-list', () => ({ UserAvatarList: ({ users }: { users: Array<{ id: string }> }) => (
    {users.map(user => user.id).join(',')}
    diff --git a/web/app/components/workflow/comment/comment-icon.tsx b/web/app/components/workflow/comment/comment-icon.tsx index 658ad6a2c6b..e46576e145d 100644 --- a/web/app/components/workflow/comment/comment-icon.tsx +++ b/web/app/components/workflow/comment/comment-icon.tsx @@ -2,10 +2,11 @@ import type { FC, PointerEvent as ReactPointerEvent } from 'react' import type { WorkflowCommentList } from '@/app/components/workflow/comment/types' +import { useAtomValue } from 'jotai' import { memo, useCallback, useMemo, useRef, useState } from 'react' import { useReactFlow, useViewport } from 'reactflow' import { UserAvatarList } from '@/app/components/base/user-avatar-list' -import { useAppContext } from '@/context/app-context' +import { userProfileIdAtom } from '@/context/app-context-state' import CommentPreview from './comment-preview' type CommentIconProps = { @@ -18,8 +19,8 @@ type CommentIconProps = { export const CommentIcon: FC = memo(({ comment, onClick, isActive = false, onPositionUpdate }) => { const { flowToScreenPosition, screenToFlowPosition } = useReactFlow() const viewport = useViewport() - const { userProfile } = useAppContext() - const isAuthor = comment.created_by_account?.id === userProfile?.id + const currentUserId = useAtomValue(userProfileIdAtom) + const isAuthor = comment.created_by_account?.id === currentUserId const [showPreview, setShowPreview] = useState(false) const [dragPosition, setDragPosition] = useState<{ x: number, y: number } | null>(null) const [isDragging, setIsDragging] = useState(false) diff --git a/web/app/components/workflow/comment/comment-input.spec.tsx b/web/app/components/workflow/comment/comment-input.spec.tsx index 8796b4ec364..1696f85fa89 100644 --- a/web/app/components/workflow/comment/comment-input.spec.tsx +++ b/web/app/components/workflow/comment/comment-input.spec.tsx @@ -18,6 +18,13 @@ const stableT = (key: string, options?: { ns?: string }) => ( ) let mentionInputProps: MentionInputProps | null = null +const mockAppContextState = vi.hoisted(() => ({ + userProfile: { + id: 'user-1', + name: 'Alice', + avatar_url: 'avatar', + }, +})) vi.mock('react-i18next', () => ({ useTranslation: () => ({ @@ -27,14 +34,20 @@ vi.mock('react-i18next', () => ({ vi.mock('@/context/app-context', () => ({ useAppContext: () => ({ - userProfile: { - id: 'user-1', - name: 'Alice', - avatar_url: 'avatar', - }, + userProfile: mockAppContextState.userProfile, }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppContextState) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@langgenius/dify-ui/avatar', () => ({ Avatar: ({ name }: { name: string }) =>
    {name}
    , default: ({ name }: { name: string }) =>
    {name}
    , diff --git a/web/app/components/workflow/comment/comment-input.tsx b/web/app/components/workflow/comment/comment-input.tsx index 963111aaa26..2d616e3a798 100644 --- a/web/app/components/workflow/comment/comment-input.tsx +++ b/web/app/components/workflow/comment/comment-input.tsx @@ -1,9 +1,10 @@ import type { FC, PointerEvent as ReactPointerEvent } from 'react' import { Avatar } from '@langgenius/dify-ui/avatar' import { cn } from '@langgenius/dify-ui/cn' +import { useAtomValue } from 'jotai' import { memo, useCallback, useEffect, useRef, useState } from 'react' import { useTranslation } from 'react-i18next' -import { useAppContext } from '@/context/app-context' +import { userProfileAtom } from '@/context/app-context-state' import { MentionInput } from './mention-input' type CommentInputProps = { @@ -30,7 +31,7 @@ export const CommentInput: FC = memo(({ }) => { const [content, setContent] = useState('') const { t } = useTranslation() - const { userProfile } = useAppContext() + const userProfile = useAtomValue(userProfileAtom) const dragStateRef = useRef<{ pointerId: number | null startPointerX: number diff --git a/web/app/components/workflow/comment/thread.spec.tsx b/web/app/components/workflow/comment/thread.spec.tsx index 996a8e8d344..771443e37ff 100644 --- a/web/app/components/workflow/comment/thread.spec.tsx +++ b/web/app/components/workflow/comment/thread.spec.tsx @@ -4,6 +4,13 @@ import { CommentThread } from './thread' const mockSetCommentPreviewHovering = vi.hoisted(() => vi.fn()) const mockFlowToScreenPosition = vi.hoisted(() => vi.fn(({ x, y }: { x: number, y: number }) => ({ x, y }))) +const mockAppContextState = vi.hoisted(() => ({ + userProfile: { + id: 'user-1', + name: 'Alice', + avatar_url: 'alice.png', + }, +})) const storeState = vi.hoisted(() => ({ mentionableUsersCache: { @@ -33,14 +40,20 @@ vi.mock('@/hooks/use-format-time-from-now', () => ({ vi.mock('@/context/app-context', () => ({ useAppContext: () => ({ - userProfile: { - id: 'user-1', - name: 'Alice', - avatar_url: 'alice.png', - }, + userProfile: mockAppContextState.userProfile, }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppContextState) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('reactflow', () => ({ useReactFlow: () => ({ flowToScreenPosition: mockFlowToScreenPosition, diff --git a/web/app/components/workflow/comment/thread.tsx b/web/app/components/workflow/comment/thread.tsx index 97aadf04b39..1d4b9545449 100644 --- a/web/app/components/workflow/comment/thread.tsx +++ b/web/app/components/workflow/comment/thread.tsx @@ -11,13 +11,14 @@ import { } from '@langgenius/dify-ui/dropdown-menu' import { Tooltip, TooltipContent, TooltipTrigger } from '@langgenius/dify-ui/tooltip' import { RiArrowDownSLine, RiArrowUpSLine, RiCheckboxCircleFill, RiCheckboxCircleLine, RiCloseLine, RiDeleteBinLine, RiMoreFill } from '@remixicon/react' +import { useAtomValue } from 'jotai' import { memo, useCallback, useEffect, useMemo, useRef, useState } from 'react' import { useTranslation } from 'react-i18next' import { useReactFlow, useViewport } from 'reactflow' import Divider from '@/app/components/base/divider' import InlineDeleteConfirm from '@/app/components/base/inline-delete-confirm' import { getUserColor } from '@/app/components/workflow/collaboration/utils/user-color' -import { useAppContext } from '@/context/app-context' +import { userProfileAtom, userProfileIdAtom } from '@/context/app-context-state' import { useFormatTimeFromNow } from '@/hooks/use-format-time-from-now' import { useParams } from '@/next/navigation' import { useStore } from '../store' @@ -52,8 +53,7 @@ const ThreadMessage: FC<{ className?: string }> = ({ authorId, authorName, avatarUrl, createdAt, content, mentionableNames, className }) => { const { formatTimeFromNow } = useFormatTimeFromNow() - const { userProfile } = useAppContext() - const currentUserId = userProfile?.id + const currentUserId = useAtomValue(userProfileIdAtom) const isCurrentUser = authorId === currentUserId const userColor = isCurrentUser ? undefined : getUserColor(authorId) @@ -175,7 +175,8 @@ export const CommentThread: FC = memo(({ const appId = params.appId as string const { flowToScreenPosition } = useReactFlow() const viewport = useViewport() - const { userProfile } = useAppContext() + const userProfile = useAtomValue(userProfileAtom) + const currentUserId = userProfile.id const { t } = useTranslation() const [replyContent, setReplyContent] = useState('') const [editingCommentContent, setEditingCommentContent] = useState('') @@ -369,7 +370,7 @@ export const CommentThread: FC = memo(({ }, [editingReply, onReplyEdit]) const replies = comment.replies || [] - const isOwnComment = comment.created_by_account?.id === userProfile?.id + const isOwnComment = comment.created_by_account?.id === currentUserId const messageListRef = useRef(null) const previousReplyCountRef = useRef(undefined) const previousCommentIdRef = useRef(undefined) @@ -604,7 +605,7 @@ export const CommentThread: FC = memo(({
    {replies.map((reply) => { const isReplyEditing = editingReply?.id === reply.id - const isOwnReply = reply.created_by_account?.id === userProfile?.id + const isOwnReply = reply.created_by_account?.id === currentUserId return (
    ({ + userProfile: { + id: '', + name: '', + }, +})) + +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppContextState) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) vi.mock('@/context/provider-context', () => ({ useProviderContext: () => ({ diff --git a/web/app/components/workflow/header/__tests__/header-layouts.spec.tsx b/web/app/components/workflow/header/__tests__/header-layouts.spec.tsx index 9b76bc8e46b..a93aaa2a4cd 100644 --- a/web/app/components/workflow/header/__tests__/header-layouts.spec.tsx +++ b/web/app/components/workflow/header/__tests__/header-layouts.spec.tsx @@ -25,6 +25,22 @@ const mockViewHistory = vi.fn() let mockNodesReadOnly = false let mockTheme: 'light' | 'dark' = 'light' +const mockAppContextState = vi.hoisted(() => ({ + userProfile: { + id: '', + name: '', + }, +})) + +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppContextState) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) vi.mock('reactflow', () => ({ useNodes: () => mockUseNodes(), diff --git a/web/app/components/workflow/header/header-in-restoring.tsx b/web/app/components/workflow/header/header-in-restoring.tsx index b098bed6018..cb1ab7c7d64 100644 --- a/web/app/components/workflow/header/header-in-restoring.tsx +++ b/web/app/components/workflow/header/header-in-restoring.tsx @@ -2,6 +2,7 @@ import { Button } from '@langgenius/dify-ui/button' import { cn } from '@langgenius/dify-ui/cn' import { toast } from '@langgenius/dify-ui/toast' import { RiHistoryLine } from '@remixicon/react' +import { useAtomValue } from 'jotai' import { useCallback, useState, @@ -9,7 +10,7 @@ import { import { useTranslation } from 'react-i18next' import { PlanUpgradeModal } from '@/app/components/billing/plan-upgrade-modal' import { Plan } from '@/app/components/billing/type' -import { useSelector as useAppContextSelector } from '@/context/app-context' +import { userProfileAtom } from '@/context/app-context-state' import { useProviderContext } from '@/context/provider-context' import useTheme from '@/hooks/use-theme' import { useInvalidAllLastRun, useResetWorkflowVersionHistory, useRestoreWorkflow } from '@/service/use-workflow' @@ -39,7 +40,7 @@ const HeaderInRestoring = ({ const [isRestorePlanUpgradeModalOpen, setIsRestorePlanUpgradeModalOpen] = useState(false) const { plan, enableBilling } = useProviderContext() const workflowStore = useWorkflowStore() - const userProfile = useAppContextSelector(s => s.userProfile) + const userProfile = useAtomValue(userProfileAtom) const configsMap = useHooksStore(s => s.configsMap) const invalidAllLastRun = useInvalidAllLastRun(configsMap?.flowType, configsMap?.flowId) const { diff --git a/web/app/components/workflow/header/online-users.tsx b/web/app/components/workflow/header/online-users.tsx index efef3868193..2094f69d526 100644 --- a/web/app/components/workflow/header/online-users.tsx +++ b/web/app/components/workflow/header/online-users.tsx @@ -9,10 +9,11 @@ import { PopoverTrigger, } from '@langgenius/dify-ui/popover' import { Tooltip, TooltipContent, TooltipTrigger } from '@langgenius/dify-ui/tooltip' +import { useAtomValue } from 'jotai' import { useEffect, useState } from 'react' import { useTranslation } from 'react-i18next' import { useReactFlow } from 'reactflow' -import { useAppContext } from '@/context/app-context' +import { userProfileIdAtom } from '@/context/app-context-state' import { getAvatar } from '@/service/common' import { useCollaboration } from '../collaboration/hooks/use-collaboration' import { getUserColor } from '../collaboration/utils/user-color' @@ -54,12 +55,11 @@ const OnlineUsers = () => { const { t } = useTranslation() const appId = useStore(s => s.appId) const { onlineUsers, cursors, isEnabled: isCollaborationEnabled } = useCollaboration(appId as string) - const { userProfile } = useAppContext() + const currentUserId = useAtomValue(userProfileIdAtom) const reactFlow = useReactFlow() const [dropdownOpen, setDropdownOpen] = useState(false) const avatarUrls = useAvatarUrls(onlineUsers || []) - const currentUserId = userProfile?.id const fallbackUsername = t('comments.fallback.user', { ns: 'workflow' }) const currentUserSuffix = t('members.you', { ns: 'common' }) diff --git a/web/app/components/workflow/hooks/__tests__/use-workflow-comment.spec.ts b/web/app/components/workflow/hooks/__tests__/use-workflow-comment.spec.ts index 52cffa1d7ba..548ec5eef25 100644 --- a/web/app/components/workflow/hooks/__tests__/use-workflow-comment.spec.ts +++ b/web/app/components/workflow/hooks/__tests__/use-workflow-comment.spec.ts @@ -28,6 +28,14 @@ const commentsUpdateState = vi.hoisted(() => ({ const globalFeatureState = vi.hoisted(() => ({ enableCollaboration: true, })) +const mockAppContextState = vi.hoisted(() => ({ + userProfile: { + id: 'user-1', + name: 'Alice', + email: 'alice@example.com', + avatar_url: 'alice.png', + }, +})) vi.mock('reactflow', () => ({ useReactFlow: () => ({ @@ -43,15 +51,20 @@ vi.mock('@/next/navigation', () => ({ vi.mock('@/context/app-context', () => ({ useAppContext: () => ({ - userProfile: { - id: 'user-1', - name: 'Alice', - email: 'alice@example.com', - avatar_url: 'alice.png', - }, + userProfile: mockAppContextState.userProfile, }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppContextState) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('@/service/client', () => ({ consoleClient: { systemFeatures: { diff --git a/web/app/components/workflow/hooks/use-workflow-comment.ts b/web/app/components/workflow/hooks/use-workflow-comment.ts index 6373dbed6aa..04b17449771 100644 --- a/web/app/components/workflow/hooks/use-workflow-comment.ts +++ b/web/app/components/workflow/hooks/use-workflow-comment.ts @@ -1,9 +1,10 @@ import type { UserProfile, WorkflowCommentDetail, WorkflowCommentList } from '@/app/components/workflow/comment/types' import { useSuspenseQuery } from '@tanstack/react-query' +import { useAtomValue } from 'jotai' import { useCallback, useEffect, useMemo, useRef } from 'react' import { useReactFlow } from 'reactflow' import { collaborationManager } from '@/app/components/workflow/collaboration/core/collaboration-manager' -import { useAppContext } from '@/context/app-context' +import { userProfileAtom } from '@/context/app-context-state' import { systemFeaturesQueryOptions } from '@/features/system-features/client' import { useParams } from '@/next/navigation' import { consoleClient } from '@/service/client' @@ -67,7 +68,7 @@ export const useWorkflowComment = () => { () => new Map(mentionableUsers.map(user => [user.id, user])), [mentionableUsers], ) - const { userProfile } = useAppContext() + const userProfile = useAtomValue(userProfileAtom) const { data: isCollaborationEnabled } = useSuspenseQuery({ ...systemFeaturesQueryOptions(), select: s => s.enable_collaboration_mode, diff --git a/web/app/components/workflow/nodes/_base/__tests__/node.spec.tsx b/web/app/components/workflow/nodes/_base/__tests__/node.spec.tsx index b6f5c32d025..51a3f8813d8 100644 --- a/web/app/components/workflow/nodes/_base/__tests__/node.spec.tsx +++ b/web/app/components/workflow/nodes/_base/__tests__/node.spec.tsx @@ -11,6 +11,25 @@ const mockHandleNodeIterationChildSizeChange = vi.fn() const mockHandleNodeLoopChildSizeChange = vi.fn() const mockUseNodeResizeObserver = vi.fn() const mockUseCollaboration = vi.fn() +const mockAppContextState = vi.hoisted(() => ({ + userProfile: { + id: 'user-1', + name: 'User', + email: 'user@example.com', + avatar: '', + avatar_url: '', + }, +})) + +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppContextState) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) vi.mock('@/app/components/workflow/hooks', () => ({ useNodesReadOnly: () => ({ nodesReadOnly: false }), diff --git a/web/app/components/workflow/nodes/_base/components/workflow-panel/index.tsx b/web/app/components/workflow/nodes/_base/components/workflow-panel/index.tsx index f30c428890e..515f307db56 100644 --- a/web/app/components/workflow/nodes/_base/components/workflow-panel/index.tsx +++ b/web/app/components/workflow/nodes/_base/components/workflow-panel/index.tsx @@ -13,6 +13,7 @@ import { RiPlayLargeLine, } from '@remixicon/react' import { debounce } from 'es-toolkit/compat' +import { useAtomValue } from 'jotai' import * as React from 'react' import { cloneElement, @@ -69,7 +70,7 @@ import { hasRetryNode, isSupportCustomRunForm, } from '@/app/components/workflow/utils' -import { useAppContext } from '@/context/app-context' +import { userProfileAtom } from '@/context/app-context-state' import { useAllBuiltInTools } from '@/service/use-tools' import { useAllTriggerPlugins } from '@/service/use-triggers' import { FlowType } from '@/types/common' @@ -114,7 +115,7 @@ const BasePanel: FC = ({ const { t } = useTranslation() const language = useLanguage() const appId = useStore(s => s.appId) - const { userProfile } = useAppContext() + const userProfile = useAtomValue(userProfileAtom) const { isConnected, nodePanelPresence } = useCollaboration(appId as string) const { showMessageLogModal } = useAppStore(useShallow(state => ({ showMessageLogModal: state.showMessageLogModal, diff --git a/web/app/components/workflow/nodes/_base/node.tsx b/web/app/components/workflow/nodes/_base/node.tsx index 9fd1fb6f97e..80eeadad016 100644 --- a/web/app/components/workflow/nodes/_base/node.tsx +++ b/web/app/components/workflow/nodes/_base/node.tsx @@ -4,6 +4,7 @@ import type { } from 'react' import type { NodeProps } from '@/app/components/workflow/types' import { cn } from '@langgenius/dify-ui/cn' +import { useAtomValue } from 'jotai' import { cloneElement, memo, @@ -28,7 +29,7 @@ import { NodeRunningStatus, } from '@/app/components/workflow/types' import { hasErrorHandleNode, hasRetryNode } from '@/app/components/workflow/utils' -import { useAppContext } from '@/context/app-context' +import { userProfileAtom } from '@/context/app-context-state' import { selectWorkflowNode } from '../../utils/node-navigation' import AddVariablePopupWithPosition from './components/add-variable-popup-with-position' import EntryNodeContainer, { StartNodeTypeEnum } from './components/entry-node-container' @@ -76,7 +77,7 @@ const BaseNode: FC = ({ const { handleNodeIterationChildSizeChange } = useNodeIterationInteractions() const { handleNodeLoopChildSizeChange } = useNodeLoopInteractions() const toolIcon = useToolIcon(data) - const { userProfile } = useAppContext() + const userProfile = useAtomValue(userProfileAtom) const appId = useStore(s => s.appId) const { nodePanelPresence } = useCollaboration(appId as string) const controlMode = useStore(s => s.controlMode) diff --git a/web/app/components/workflow/nodes/human-input/components/delivery-method/__tests__/email-configure-modal.spec.tsx b/web/app/components/workflow/nodes/human-input/components/delivery-method/__tests__/email-configure-modal.spec.tsx index d9875c75397..7ba7347869f 100644 --- a/web/app/components/workflow/nodes/human-input/components/delivery-method/__tests__/email-configure-modal.spec.tsx +++ b/web/app/components/workflow/nodes/human-input/components/delivery-method/__tests__/email-configure-modal.spec.tsx @@ -3,7 +3,11 @@ import { fireEvent, render, screen } from '@testing-library/react' import EmailConfigureModal from '../email-configure-modal' const mockToastError = vi.hoisted(() => vi.fn()) -const mockUseAppContextSelector = vi.hoisted(() => vi.fn()) +const mockAppContextState = vi.hoisted(() => ({ + userProfile: { + email: 'owner@example.com', + }, +})) vi.mock('@langgenius/dify-ui/toast', () => ({ toast: { @@ -13,9 +17,19 @@ vi.mock('@langgenius/dify-ui/toast', () => ({ vi.mock('@/context/app-context', () => ({ useSelector: (selector: (state: { userProfile: { email: string } }) => string) => - mockUseAppContextSelector(selector), + selector({ userProfile: mockAppContextState.userProfile }), })) +vi.mock('@/context/app-context-state', async (importOriginal) => { + const { createAppContextStateAtomMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateAtomMock(importOriginal, () => mockAppContextState) +}) + +vi.mock('jotai', async (importOriginal) => { + const { createAppContextStateJotaiMock } = await import('@/__tests__/utils/mock-app-context-state') + return createAppContextStateJotaiMock(importOriginal) +}) + vi.mock('../mail-body-input', () => ({ default: ({ value, onChange }: { value: string, onChange: (value: string) => void }) => (