Compare commits

...
14 Commits
Author SHA1 Message Date
6f8ed69ee1 fix: fix mcp output_schema is optional (#39453)
Co-authored-by: yunlu.wen <yunlu.wen@dify.ai>
2026-07-28 02:31:18 +00:00
Yunlu WenGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
59b879d2df chore: bump version to 1.16.1 (#39653)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-28 01:09:25 +00:00
Xiyuan ChenandGitHub a57b0b9b58 feat(plugin): allow disabling the tenant plugin model providers cache (#39632) 2026-07-27 23:53:59 +00:00
Xiyuan ChenandGitHub d8506efed6 fix(cli): decouple release script tests from the live compat window (#39658) 2026-07-27 23:50:05 +00:00
JingyiandGitHub b25b28cc76 fix(workflow): align block icon vector sizes (#39657) 2026-07-27 23:43:03 +00:00
yyhGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
f68cabe226 fix(web): align tour trigger DOM order (#39654)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-27 13:00:16 +00:00
Asuka MinatoGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>Byron.wang
f73f83c5d6 test: use SQLite sessions in core app (#39075)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: Byron.wang <byron@dify.ai>
2026-07-27 10:45:16 +00:00
Asuka MinatoGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>Byron.wang
29801d6a65 test: use SQLite sessions in core tools (#39077)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: Byron.wang <byron@dify.ai>
2026-07-27 10:45:11 +00:00
Asuka MinatoGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>Byron.wang
75016e8bfe test: use SQLite sessions in core ops (#39072)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: Byron.wang <byron@dify.ai>
2026-07-27 10:41:17 +00:00
yyhandGitHub b6dea7b2ba fix(workflow): preserve latest collaboration session (#39646) 2026-07-27 10:37:45 +00:00
yyhandGitHub b2b1cd7e97 chore: update workspace dependencies (#39641) 2026-07-27 10:27:31 +00:00
Asuka MinatoGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>Byron.wang
ca63e01d45 test: use SQLite sessions in controllers service api (#39069)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: Byron.wang <byron@dify.ai>
2026-07-27 10:26:23 +00:00
yyhandGitHub 271b6e1f5c fix(ui): unify combobox trigger focus rings (#39643) 2026-07-27 10:21:32 +00:00
Asuka MinatoGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>Byron.wang
609ddf9e01 test: use SQLite sessions in core memory (#39066)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: Byron.wang <byron@dify.ai>
2026-07-27 10:15:08 +00:00
54 changed files with 3773 additions and 3424 deletions
+1
View File
@@ -666,6 +666,7 @@ PLUGIN_REMOTE_INSTALL_PORT=5003
PLUGIN_REMOTE_INSTALL_HOST=localhost
PLUGIN_MAX_PACKAGE_SIZE=15728640
PLUGIN_MODEL_SCHEMA_CACHE_TTL=3600
PLUGIN_MODEL_PROVIDERS_CACHE_ENABLED=true
PLUGIN_MODEL_PROVIDERS_CACHE_TTL=86400
# Comma-separated marketplace plugin IDs whose latest versions are installed for newly registered users.
# Example: langgenius/openai,langgenius/gemini
+6
View File
@@ -266,6 +266,12 @@ class PluginConfig(BaseSettings):
default=60 * 60,
)
PLUGIN_MODEL_PROVIDERS_CACHE_ENABLED: bool = Field(
description="Whether tenant plugin model providers are cached in Redis. Disable when plugins are installed "
"by a system other than this one, which cannot invalidate the cache when a tenant's plugins change.",
default=True,
)
PLUGIN_MODEL_PROVIDERS_CACHE_TTL: PositiveInt = Field(
description="TTL in seconds for caching tenant plugin model providers in Redis",
default=60 * 60 * 24,
-2
View File
@@ -23,7 +23,6 @@ from .knowledge import retrieval as _knowledge_retrieval
from .plugin import agent_config as _agent_config
from .plugin import agent_drive as _agent_drive
from .plugin import plugin as _plugin
from .workspace import plugin_model_providers as _plugin_model_providers
from .workspace import workspace as _workspace
api.add_namespace(inner_api_ns)
@@ -36,7 +35,6 @@ __all__ = [
"_knowledge_retrieval",
"_mail",
"_plugin",
"_plugin_model_providers",
"_runtime_credentials",
"_workspace",
"api",
@@ -1,39 +0,0 @@
from flask_restx import Resource
from pydantic import BaseModel, ConfigDict, Field
from controllers.common.schema import register_schema_model
from controllers.console.wraps import setup_required
from controllers.inner_api import inner_api_ns
from controllers.inner_api.wraps import enterprise_inner_api_only
from core.plugin.plugin_service import PluginService
class InvalidatePluginModelProvidersCachePayload(BaseModel):
model_config = ConfigDict(extra="forbid")
tenant_ids: list[str] = Field(default_factory=list, description="Workspace ids whose cache should be invalidated")
register_schema_model(inner_api_ns, InvalidatePluginModelProvidersCachePayload)
@inner_api_ns.route("/enterprise/workspace/plugin-model-providers/invalidate")
class EnterprisePluginModelProvidersCacheInvalidate(Resource):
@setup_required
@enterprise_inner_api_only
@inner_api_ns.doc(
"enterprise_invalidate_plugin_model_providers_cache",
responses={
200: "Cache invalidated",
400: "Invalid request",
401: "Unauthorized - invalid API key",
},
)
@inner_api_ns.expect(inner_api_ns.models[InvalidatePluginModelProvidersCachePayload.__name__])
def post(self):
args = InvalidatePluginModelProvidersCachePayload.model_validate(inner_api_ns.payload or {})
for tenant_id in args.tenant_ids:
PluginService.invalidate_plugin_model_providers_cache(tenant_id)
return {"result": "success"}, 200
+11 -4
View File
@@ -434,14 +434,18 @@ class PluginService:
exc_info=True,
)
@classmethod
def _fetch_plugin_model_providers_uncached(
cls, tenant_id: str, client: PluginModelClient | None
) -> tuple[ProviderEntity, ...]:
model_client = client or PluginModelClient()
return tuple(cls._to_provider_entity(provider) for provider in model_client.fetch_model_providers(tenant_id))
@classmethod
def _fetch_and_cache_plugin_model_providers(
cls, tenant_id: str, client: PluginModelClient | None, *, refresh_generation: int | None
) -> tuple[ProviderEntity, ...]:
model_client = client or PluginModelClient()
providers = tuple(
cls._to_provider_entity(provider) for provider in model_client.fetch_model_providers(tenant_id)
)
providers = cls._fetch_plugin_model_providers_uncached(tenant_id, client)
generation = cls._load_plugin_model_providers_generation(tenant_id)
if generation is not None and generation == refresh_generation:
cls._store_cached_plugin_model_providers(tenant_id, generation, providers)
@@ -471,6 +475,9 @@ class PluginService:
are intentionally owned by this service so tenant isolation and cache
expiry are handled in one place.
"""
if not dify_config.PLUGIN_MODEL_PROVIDERS_CACHE_ENABLED:
return cls._fetch_plugin_model_providers_uncached(tenant_id, client)
deadline = time.monotonic() + cls.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_TIMEOUT
while True:
+2
View File
@@ -107,6 +107,8 @@ class MCPTool(Tool):
if self.entity.output_schema and result.structuredContent:
for k, v in result.structuredContent.items():
yield self.create_variable_message(k, v)
elif result.structuredContent:
yield self.create_json_message(result.structuredContent)
def _process_text_content(self, content: TextContent) -> Generator[ToolInvokeMessage, None, None]:
"""Process text content and yield appropriate messages."""
+1 -1
View File
@@ -1,6 +1,6 @@
[project]
name = "dify-api"
version = "1.16.0"
version = "1.16.1"
requires-python = "~=3.12.0"
dependencies = [
@@ -1,64 +0,0 @@
import inspect
from unittest.mock import call, patch
import pytest
from flask import Flask
from pydantic import ValidationError
from controllers.inner_api.workspace.plugin_model_providers import (
EnterprisePluginModelProvidersCacheInvalidate,
InvalidatePluginModelProvidersCachePayload,
)
class TestInvalidatePluginModelProvidersCachePayload:
def test_valid_payload(self):
payload = InvalidatePluginModelProvidersCachePayload.model_validate(
{"tenant_ids": ["tenant-alpha", "tenant-beta"]}
)
assert payload.tenant_ids == ["tenant-alpha", "tenant-beta"]
def test_missing_tenant_ids_defaults_to_empty(self):
payload = InvalidatePluginModelProvidersCachePayload.model_validate({})
assert payload.tenant_ids == []
def test_unknown_field_rejected(self):
with pytest.raises(ValidationError):
InvalidatePluginModelProvidersCachePayload.model_validate({"tenant_ids": ["tenant-alpha"], "generation": 7})
class TestEnterprisePluginModelProvidersCacheInvalidate:
@pytest.fixture
def api_instance(self):
return EnterprisePluginModelProvidersCacheInvalidate()
def _post(self, api_instance, app: Flask, payload):
unwrapped_post = inspect.unwrap(api_instance.post)
with app.test_request_context():
with patch("controllers.inner_api.workspace.plugin_model_providers.inner_api_ns") as mock_ns:
mock_ns.payload = payload
return unwrapped_post(api_instance)
@patch("controllers.inner_api.workspace.plugin_model_providers.PluginService")
def test_post_invalidates_once_per_tenant(self, mock_plugin_service, api_instance, app: Flask):
result = self._post(api_instance, app, {"tenant_ids": ["tenant-alpha", "tenant-beta"]})
assert result == ({"result": "success"}, 200)
assert mock_plugin_service.invalidate_plugin_model_providers_cache.call_args_list == [
call("tenant-alpha"),
call("tenant-beta"),
]
@patch("controllers.inner_api.workspace.plugin_model_providers.PluginService")
def test_post_with_empty_list_is_a_no_op(self, mock_plugin_service, api_instance, app: Flask):
result = self._post(api_instance, app, {"tenant_ids": []})
assert result == ({"result": "success"}, 200)
mock_plugin_service.invalidate_plugin_model_providers_cache.assert_not_called()
@patch("controllers.inner_api.workspace.plugin_model_providers.PluginService")
def test_post_with_missing_payload_is_a_no_op(self, mock_plugin_service, api_instance, app: Flask):
result = self._post(api_instance, app, None)
assert result == ({"result": "success"}, 200)
mock_plugin_service.invalidate_plugin_model_providers_cache.assert_not_called()
@@ -1,17 +1,31 @@
"""
Unit tests for Service API File Preview endpoint
"""Unit tests for the Service API file-preview endpoint.
Ownership checks run against persisted message, file, app, and upload rows so the
tests exercise the same SQLAlchemy statements and tenant boundary as production.
Storage remains mocked because it is the external I/O boundary of the endpoint.
"""
import logging
import uuid
from collections.abc import Iterator
from dataclasses import dataclass
from datetime import datetime
from decimal import Decimal
from typing import Protocol, cast
from unittest.mock import Mock, patch
from uuid import uuid4
import pytest
from sqlalchemy import event
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session
from controllers.service_api.app.error import FileAccessDeniedError, FileNotFoundError
from controllers.service_api.app.file_preview import FilePreviewApi
from models.model import App, EndUser, Message, MessageFile, UploadFile
from extensions.storage.storage_type import StorageType
from graphon.file import FileTransferMethod, FileType
from models.base import TypeBase
from models.enums import ConversationFromSource, CreatorUserRole
from models.model import App, AppMode, Message, MessageFile, UploadFile
class _FilePreviewLogRecord(Protocol):
@@ -20,367 +34,252 @@ class _FilePreviewLogRecord(Protocol):
error: str
@dataclass(frozen=True)
class _Database:
"""Expose the real test session through the interface used by the controller."""
session: Session
@dataclass(frozen=True)
class _PreviewRecords:
app: App
message: Message
message_file: MessageFile
upload_file: UploadFile
@pytest.fixture
def database(sqlite_engine: Engine) -> Iterator[_Database]:
"""Create only the tables required by file ownership validation."""
models = (App, Message, MessageFile, UploadFile)
tables = [TypeBase.metadata.tables[model.__tablename__] for model in models]
TypeBase.metadata.create_all(sqlite_engine, tables=tables)
with Session(sqlite_engine, expire_on_commit=False) as session:
yield _Database(session)
@pytest.fixture
def file_preview_api() -> FilePreviewApi:
"""Create the resource instance under test."""
return FilePreviewApi()
def _upload_file(*, tenant_id: str, file_id: str | None = None) -> UploadFile:
upload_file = UploadFile(
tenant_id=tenant_id,
storage_type=StorageType.LOCAL,
key="storage/key/test_file.jpg",
name="test_file.jpg",
size=1024,
extension="jpg",
mime_type="image/jpeg",
created_by_role=CreatorUserRole.ACCOUNT,
created_by=str(uuid4()),
created_at=datetime(2026, 1, 1),
used=True,
)
if file_id is not None:
upload_file.id = file_id
return upload_file
def _persist_preview_records(
session: Session,
*,
app_id: str | None = None,
app_tenant_id: str | None = None,
upload_tenant_id: str | None = None,
) -> _PreviewRecords:
app_id = app_id or str(uuid4())
app_tenant_id = app_tenant_id or str(uuid4())
upload_file = _upload_file(tenant_id=upload_tenant_id or app_tenant_id)
app = App(
id=app_id,
tenant_id=app_tenant_id,
name="Preview app",
description="",
mode=AppMode.CHAT,
icon_type=None,
icon="",
icon_background=None,
enable_site=True,
enable_api=True,
)
message = Message(
id=str(uuid4()),
app_id=app_id,
conversation_id=str(uuid4()),
_inputs={},
query="preview",
message={},
message_unit_price=Decimal(0),
answer="answer",
answer_unit_price=Decimal(0),
currency="USD",
from_source=ConversationFromSource.API,
)
message_file = MessageFile(
message_id=message.id,
type=FileType.IMAGE,
transfer_method=FileTransferMethod.LOCAL_FILE,
created_by_role=CreatorUserRole.ACCOUNT,
created_by=str(uuid4()),
upload_file_id=upload_file.id,
)
session.add_all([app, message, message_file, upload_file])
session.commit()
return _PreviewRecords(app=app, message=message, message_file=message_file, upload_file=upload_file)
class TestFilePreviewApi:
"""Test suite for FilePreviewApi"""
"""Exercise ownership validation and response construction."""
@pytest.fixture
def file_preview_api(self):
"""Create FilePreviewApi instance for testing"""
return FilePreviewApi()
def test_validate_file_ownership_success(self, file_preview_api: FilePreviewApi, database: _Database):
records = _persist_preview_records(database.session)
@pytest.fixture
def mock_app(self):
"""Mock App model"""
app = Mock(spec=App)
app.id = str(uuid.uuid4())
app.tenant_id = str(uuid.uuid4())
return app
with patch("controllers.service_api.app.file_preview.db", database):
message_file, upload_file = file_preview_api._validate_file_ownership(
records.upload_file.id, records.app.id
)
@pytest.fixture
def mock_end_user(self):
"""Mock EndUser model"""
end_user = Mock(spec=EndUser)
end_user.id = str(uuid.uuid4())
return end_user
assert message_file.id == records.message_file.id
assert upload_file.id == records.upload_file.id
assert upload_file.tenant_id == records.app.tenant_id
@pytest.fixture
def mock_upload_file(self):
"""Mock UploadFile model"""
upload_file = Mock(spec=UploadFile)
upload_file.id = str(uuid.uuid4())
upload_file.name = "test_file.jpg"
upload_file.extension = "jpg"
upload_file.mime_type = "image/jpeg"
upload_file.size = 1024
upload_file.key = "storage/key/test_file.jpg"
upload_file.tenant_id = str(uuid.uuid4())
return upload_file
def test_validate_file_ownership_file_not_found(self, file_preview_api: FilePreviewApi, database: _Database):
with patch("controllers.service_api.app.file_preview.db", database):
with pytest.raises(FileNotFoundError, match="File not found in message context"):
file_preview_api._validate_file_ownership(str(uuid4()), str(uuid4()))
@pytest.fixture
def mock_message_file(self):
"""Mock MessageFile model"""
message_file = Mock(spec=MessageFile)
message_file.id = str(uuid.uuid4())
message_file.upload_file_id = str(uuid.uuid4())
message_file.message_id = str(uuid.uuid4())
return message_file
def test_validate_file_ownership_access_denied(self, file_preview_api: FilePreviewApi, database: _Database):
records = _persist_preview_records(database.session)
@pytest.fixture
def mock_message(self):
"""Mock Message model"""
message = Mock(spec=Message)
message.id = str(uuid.uuid4())
message.app_id = str(uuid.uuid4())
return message
with patch("controllers.service_api.app.file_preview.db", database):
with pytest.raises(FileAccessDeniedError, match="not owned by requesting app"):
file_preview_api._validate_file_ownership(records.upload_file.id, str(uuid4()))
def test_validate_file_ownership_success(
self, file_preview_api: FilePreviewApi, mock_app, mock_upload_file, mock_message_file, mock_message
):
"""Test successful file ownership validation"""
file_id = str(uuid.uuid4())
app_id = mock_app.id
def test_validate_file_ownership_upload_file_not_found(self, file_preview_api: FilePreviewApi, database: _Database):
records = _persist_preview_records(database.session)
database.session.delete(records.upload_file)
database.session.commit()
# Set up the mocks
mock_upload_file.tenant_id = mock_app.tenant_id
mock_message.app_id = app_id
mock_message_file.upload_file_id = file_id
mock_message_file.message_id = mock_message.id
with patch("controllers.service_api.app.file_preview.db", database):
with pytest.raises(FileNotFoundError, match="Upload file record not found"):
file_preview_api._validate_file_ownership(records.upload_file.id, records.app.id)
with patch("controllers.service_api.app.file_preview.db") as mock_db:
# Mock scalar() for MessageFile and Message queries
mock_db.session.scalar.side_effect = [
mock_message_file, # MessageFile query
mock_message, # Message query
]
# Mock get() for UploadFile and App PK lookups
mock_db.session.get.side_effect = [
mock_upload_file, # UploadFile query
mock_app, # App query for tenant validation
]
def test_validate_file_ownership_tenant_mismatch(self, file_preview_api: FilePreviewApi, database: _Database):
records = _persist_preview_records(database.session, upload_tenant_id=str(uuid4()))
# Execute the method
result_message_file, result_upload_file = file_preview_api._validate_file_ownership(file_id, app_id)
# Assertions
assert result_message_file == mock_message_file
assert result_upload_file == mock_upload_file
def test_validate_file_ownership_file_not_found(self, file_preview_api: FilePreviewApi):
"""Test file ownership validation when MessageFile not found"""
file_id = str(uuid.uuid4())
app_id = str(uuid.uuid4())
with patch("controllers.service_api.app.file_preview.db") as mock_db:
# Mock MessageFile not found via scalar()
mock_db.session.scalar.return_value = None
# Execute and assert exception
with pytest.raises(FileNotFoundError) as exc_info:
file_preview_api._validate_file_ownership(file_id, app_id)
assert "File not found in message context" in str(exc_info.value)
def test_validate_file_ownership_access_denied(self, file_preview_api: FilePreviewApi, mock_message_file):
"""Test file ownership validation when Message not owned by app"""
file_id = str(uuid.uuid4())
app_id = str(uuid.uuid4())
with patch("controllers.service_api.app.file_preview.db") as mock_db:
# Mock MessageFile found but Message not owned by app via scalar()
mock_db.session.scalar.side_effect = [
mock_message_file, # MessageFile query - found
None, # Message query - not found (access denied)
]
# Execute and assert exception
with pytest.raises(FileAccessDeniedError) as exc_info:
file_preview_api._validate_file_ownership(file_id, app_id)
assert "not owned by requesting app" in str(exc_info.value)
def test_validate_file_ownership_upload_file_not_found(
self, file_preview_api: FilePreviewApi, mock_message_file, mock_message
):
"""Test file ownership validation when UploadFile not found"""
file_id = str(uuid.uuid4())
app_id = str(uuid.uuid4())
with patch("controllers.service_api.app.file_preview.db") as mock_db:
# Mock scalar() for MessageFile and Message
mock_db.session.scalar.side_effect = [
mock_message_file, # MessageFile query - found
mock_message, # Message query - found
]
# Mock get() for UploadFile - not found
mock_db.session.get.return_value = None
# Execute and assert exception
with pytest.raises(FileNotFoundError) as exc_info:
file_preview_api._validate_file_ownership(file_id, app_id)
assert "Upload file record not found" in str(exc_info.value)
def test_validate_file_ownership_tenant_mismatch(
self, file_preview_api: FilePreviewApi, mock_app, mock_upload_file, mock_message_file, mock_message
):
"""Test file ownership validation with tenant mismatch"""
file_id = str(uuid.uuid4())
app_id = mock_app.id
# Set up tenant mismatch
mock_upload_file.tenant_id = "different_tenant_id"
mock_app.tenant_id = "app_tenant_id"
mock_message.app_id = app_id
mock_message_file.upload_file_id = file_id
mock_message_file.message_id = mock_message.id
with patch("controllers.service_api.app.file_preview.db") as mock_db:
# Mock scalar() for MessageFile and Message queries
mock_db.session.scalar.side_effect = [
mock_message_file, # MessageFile query
mock_message, # Message query
]
# Mock get() for UploadFile and App PK lookups
mock_db.session.get.side_effect = [
mock_upload_file, # UploadFile query
mock_app, # App query for tenant validation
]
# Execute and assert exception
with pytest.raises(FileAccessDeniedError) as exc_info:
file_preview_api._validate_file_ownership(file_id, app_id)
assert "tenant mismatch" in str(exc_info.value)
with patch("controllers.service_api.app.file_preview.db", database):
with pytest.raises(FileAccessDeniedError, match="tenant mismatch"):
file_preview_api._validate_file_ownership(records.upload_file.id, records.app.id)
def test_validate_file_ownership_invalid_input(self, file_preview_api: FilePreviewApi):
"""Test file ownership validation with invalid input"""
# Test with empty file_id
with pytest.raises(FileAccessDeniedError) as exc_info:
with pytest.raises(FileAccessDeniedError, match="Invalid file or app identifier"):
file_preview_api._validate_file_ownership("", "app_id")
assert "Invalid file or app identifier" in str(exc_info.value)
# Test with empty app_id
with pytest.raises(FileAccessDeniedError) as exc_info:
with pytest.raises(FileAccessDeniedError, match="Invalid file or app identifier"):
file_preview_api._validate_file_ownership("file_id", "")
assert "Invalid file or app identifier" in str(exc_info.value)
def test_build_file_response_basic(self, file_preview_api: FilePreviewApi, mock_upload_file):
"""Test basic file response building"""
mock_generator = Mock()
@pytest.mark.parametrize(
("as_attachment", "mime_type", "name", "extension", "size"),
[
(False, "image/jpeg", "test_file.jpg", "jpg", 1024),
(True, "image/jpeg", "test_file.jpg", "jpg", 1024),
(False, "text/html", "unsafe.html", "html", 1024),
(False, "video/mp4", "test_file.mp4", "mp4", 1024),
(False, "image/jpeg", "test_file.jpg", "jpg", 0),
],
)
def test_build_file_response(
self,
file_preview_api: FilePreviewApi,
as_attachment: bool,
mime_type: str,
name: str,
extension: str,
size: int,
):
upload_file = _upload_file(tenant_id=str(uuid4()))
upload_file.mime_type = mime_type
upload_file.name = name
upload_file.extension = extension
upload_file.size = size
response = file_preview_api._build_file_response(mock_generator, mock_upload_file, False)
response = file_preview_api._build_file_response(Mock(), upload_file, as_attachment)
# Check response properties
assert response.mimetype == mock_upload_file.mime_type
assert response.direct_passthrough is True
assert response.headers["Content-Length"] == str(mock_upload_file.size)
assert "Cache-Control" in response.headers
def test_build_file_response_as_attachment(self, file_preview_api: FilePreviewApi, mock_upload_file):
"""Test file response building with attachment flag"""
mock_generator = Mock()
response = file_preview_api._build_file_response(mock_generator, mock_upload_file, True)
# Check attachment-specific headers
assert "attachment" in response.headers["Content-Disposition"]
assert mock_upload_file.name in response.headers["Content-Disposition"]
assert response.headers["Content-Type"] == "application/octet-stream"
def test_build_file_response_html_forces_attachment(self, file_preview_api: FilePreviewApi, mock_upload_file):
"""Test HTML files are forced to download"""
mock_generator = Mock()
mock_upload_file.mime_type = "text/html"
mock_upload_file.name = "unsafe.html"
mock_upload_file.extension = "html"
response = file_preview_api._build_file_response(mock_generator, mock_upload_file, False)
assert "attachment" in response.headers["Content-Disposition"]
assert response.headers["Content-Type"] == "application/octet-stream"
assert response.headers["X-Content-Type-Options"] == "nosniff"
def test_build_file_response_audio_video(self, file_preview_api: FilePreviewApi, mock_upload_file):
"""Test file response building for audio/video files"""
mock_generator = Mock()
mock_upload_file.mime_type = "video/mp4"
response = file_preview_api._build_file_response(mock_generator, mock_upload_file, False)
# Check Range support for media files
assert response.headers["Accept-Ranges"] == "bytes"
def test_build_file_response_no_size(self, file_preview_api: FilePreviewApi, mock_upload_file):
"""Test file response building when size is unknown"""
mock_generator = Mock()
mock_upload_file.size = 0 # Unknown size
response = file_preview_api._build_file_response(mock_generator, mock_upload_file, False)
# Content-Length should not be set when size is unknown
assert "Content-Length" not in response.headers
assert ("Content-Length" in response.headers) is bool(size)
if as_attachment or mime_type == "text/html":
assert "attachment" in response.headers["Content-Disposition"]
assert response.headers["Content-Type"] == "application/octet-stream"
else:
assert response.mimetype == mime_type
if mime_type == "text/html":
assert response.headers["X-Content-Type-Options"] == "nosniff"
if mime_type.startswith("video/"):
assert response.headers["Accept-Ranges"] == "bytes"
@patch("controllers.service_api.app.file_preview.storage")
def test_get_method_integration(
self,
mock_storage,
file_preview_api: FilePreviewApi,
mock_app,
mock_end_user,
mock_upload_file,
mock_message_file,
mock_message,
def test_components_use_validated_file(
self, mock_storage: Mock, file_preview_api: FilePreviewApi, database: _Database
):
"""Test the full GET method integration (without decorator)"""
file_id = str(uuid.uuid4())
app_id = mock_app.id
records = _persist_preview_records(database.session)
generator = Mock()
# Set up mocks
mock_upload_file.tenant_id = mock_app.tenant_id
mock_message.app_id = app_id
mock_message_file.upload_file_id = file_id
mock_message_file.message_id = mock_message.id
with patch("controllers.service_api.app.file_preview.db", database):
message_file, upload_file = file_preview_api._validate_file_ownership(
records.upload_file.id, records.app.id
)
response = file_preview_api._build_file_response(generator, upload_file, False)
mock_generator = Mock()
mock_storage.load.return_value = mock_generator
with patch("controllers.service_api.app.file_preview.db") as mock_db:
# Mock scalar() for MessageFile and Message queries
mock_db.session.scalar.side_effect = [
mock_message_file, # MessageFile query
mock_message, # Message query
]
# Mock get() for UploadFile and App PK lookups
mock_db.session.get.side_effect = [
mock_upload_file, # UploadFile query
mock_app, # App query for tenant validation
]
# Test the core logic directly without Flask decorators
# Validate file ownership
result_message_file, result_upload_file = file_preview_api._validate_file_ownership(file_id, app_id)
assert result_message_file == mock_message_file
assert result_upload_file == mock_upload_file
# Test file response building
response = file_preview_api._build_file_response(mock_generator, mock_upload_file, False)
assert response is not None
# Verify storage was called correctly
mock_storage.load.assert_not_called() # Since we're testing components separately
assert message_file.id == records.message_file.id
assert response.mimetype == "image/jpeg"
mock_storage.load.assert_not_called()
@patch("controllers.service_api.app.file_preview.storage")
def test_storage_error_handling(
self,
mock_storage,
file_preview_api: FilePreviewApi,
mock_app,
mock_upload_file,
mock_message_file,
mock_message,
def test_storage_error_remains_external(
self, mock_storage: Mock, file_preview_api: FilePreviewApi, database: _Database
):
"""Test storage error handling in the core logic"""
file_id = str(uuid.uuid4())
app_id = mock_app.id
records = _persist_preview_records(database.session)
mock_storage.load.side_effect = OSError("Storage error")
# Set up mocks
mock_upload_file.tenant_id = mock_app.tenant_id
mock_message.app_id = app_id
mock_message_file.upload_file_id = file_id
mock_message_file.message_id = mock_message.id
with patch("controllers.service_api.app.file_preview.db", database):
_, upload_file = file_preview_api._validate_file_ownership(records.upload_file.id, records.app.id)
# Mock storage error
mock_storage.load.side_effect = Exception("Storage error")
with patch("controllers.service_api.app.file_preview.db") as mock_db:
# Mock scalar() for MessageFile and Message queries
mock_db.session.scalar.side_effect = [
mock_message_file, # MessageFile query
mock_message, # Message query
]
# Mock get() for UploadFile and App PK lookups
mock_db.session.get.side_effect = [
mock_upload_file, # UploadFile query
mock_app, # App query for tenant validation
]
# First validate file ownership works
result_message_file, result_upload_file = file_preview_api._validate_file_ownership(file_id, app_id)
assert result_message_file == mock_message_file
assert result_upload_file == mock_upload_file
# Test storage error handling
with pytest.raises(Exception) as exc_info:
mock_storage.load(mock_upload_file.key, stream=True)
assert "Storage error" in str(exc_info.value)
with pytest.raises(OSError, match="Storage error"):
mock_storage.load(upload_file.key, stream=True)
def test_validate_file_ownership_unexpected_error_logging(
self, file_preview_api: FilePreviewApi, caplog: pytest.LogCaptureFixture
self,
file_preview_api: FilePreviewApi,
database: _Database,
sqlite_engine: Engine,
caplog: pytest.LogCaptureFixture,
):
"""Test that unexpected errors are logged properly"""
file_id = str(uuid.uuid4())
app_id = str(uuid.uuid4())
file_id = str(uuid4())
app_id = str(uuid4())
with patch("controllers.service_api.app.file_preview.db") as mock_db:
# Mock database scalar to raise unexpected exception
mock_db.session.scalar.side_effect = Exception("Unexpected database error")
def fail_statement(*_args: object) -> None:
raise RuntimeError("Unexpected database error")
# Execute and assert exception
with caplog.at_level(logging.ERROR, logger="controllers.service_api.app.file_preview"):
with pytest.raises(FileAccessDeniedError) as exc_info:
file_preview_api._validate_file_ownership(file_id, app_id)
event.listen(sqlite_engine, "before_cursor_execute", fail_statement)
try:
with patch("controllers.service_api.app.file_preview.db", database):
with caplog.at_level(logging.ERROR, logger="controllers.service_api.app.file_preview"):
with pytest.raises(FileAccessDeniedError, match="File access validation failed"):
file_preview_api._validate_file_ownership(file_id, app_id)
finally:
event.remove(sqlite_engine, "before_cursor_execute", fail_statement)
# Verify error message
assert "File access validation failed" in str(exc_info.value)
# Verify logging was called with the structured context fields. The ``extra`` keys
# are attached to the LogRecord as attributes, so they are not in ``caplog.text``.
assert len(caplog.records) == 1
log_record = caplog.records[0]
assert log_record.getMessage() == "Unexpected error during file ownership validation"
record = cast(_FilePreviewLogRecord, log_record)
assert record.file_id == file_id
assert record.app_id == app_id
assert record.error == "Unexpected database error"
assert len(caplog.records) == 1
log_record = caplog.records[0]
assert log_record.getMessage() == "Unexpected error during file ownership validation"
record = cast(_FilePreviewLogRecord, log_record)
assert record.file_id == file_id
assert record.app_id == app_id
assert record.error == "Unexpected database error"
@@ -1,24 +1,95 @@
"""Unit tests for the message cycle manager optimization."""
import logging
from collections.abc import Iterator
from dataclasses import dataclass
from types import SimpleNamespace
from unittest.mock import Mock, patch
import pytest
from flask import Flask, current_app
from sqlalchemy import Engine, event, select
from sqlalchemy.orm import Session, sessionmaker
from core.app.entities.queue_entities import QueueAnnotationReplyEvent, QueueRetrieverResourcesEvent
from core.app.entities.task_entities import MessageStreamResponse, StreamEvent, TaskStateMetadata
from core.app.task_pipeline import message_cycle_manager as message_cycle_manager_module
from core.app.task_pipeline.message_cycle_manager import MessageCycleManager
from core.rag.entities import RetrievalSourceMetadata
from models.model import App, AppMode
from graphon.file import FileTransferMethod, FileType
from models import model as model_module
from models.base import TypeBase
from models.enums import ConversationFromSource, CreatorUserRole, MessageFileBelongsTo
from models.model import App, AppMode, Conversation, MessageFile
def _patch_create_session(mock_session):
session_cm = Mock()
session_cm.__enter__ = Mock(return_value=mock_session)
session_cm.__exit__ = Mock(return_value=False)
return patch("core.app.task_pipeline.message_cycle_manager.session_factory.create_session", return_value=session_cm)
@dataclass(frozen=True)
class _SQLiteDb:
engine: Engine
session: Session
@pytest.fixture
def cycle_db(sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch) -> Iterator[Session]:
"""Bind request-owned and cycle-manager-owned sessions to isolated SQLite."""
TypeBase.metadata.create_all(
sqlite_engine,
tables=[App.__table__, Conversation.__table__, MessageFile.__table__],
)
owned_session_factory = sessionmaker(bind=sqlite_engine, expire_on_commit=False)
with owned_session_factory() as request_session:
sqlite_db = _SQLiteDb(engine=sqlite_engine, session=request_session)
monkeypatch.setattr(message_cycle_manager_module, "db", sqlite_db)
monkeypatch.setattr(model_module, "db", sqlite_db)
monkeypatch.setattr(message_cycle_manager_module.session_factory, "create_session", owned_session_factory)
yield request_session
def _app(*, app_id: str = "app-id", tenant_id: str = "tenant-1") -> App:
return App(
id=app_id,
tenant_id=tenant_id,
name="Test App",
description="",
mode=AppMode.CHAT,
enable_site=True,
enable_api=True,
max_active_requests=0,
)
def _conversation(*, conversation_id: str = "conv-1", app_id: str = "app-id") -> Conversation:
conversation = Conversation(
app_id=app_id,
mode=AppMode.CHAT,
name="",
status="normal",
from_source=ConversationFromSource.API,
inputs={},
)
conversation.id = conversation_id
return conversation
def _message_file(
*,
file_id: str = "file-1",
message_id: str = "test-message-id",
belongs_to: MessageFileBelongsTo | None = MessageFileBelongsTo.ASSISTANT,
url: str | None = "http://example.com/image.png",
file_type: FileType = FileType.IMAGE,
) -> MessageFile:
message_file = MessageFile(
message_id=message_id,
type=file_type,
transfer_method=FileTransferMethod.TOOL_FILE,
created_by_role=CreatorUserRole.ACCOUNT,
created_by="account-id",
belongs_to=belongs_to,
url=url,
)
message_file.id = file_id
return message_file
class TestMessageCycleManagerOptimization:
@@ -37,30 +108,22 @@ class TestMessageCycleManagerOptimization:
task_state = Mock()
return MessageCycleManager(application_generate_entity=mock_application_generate_entity, task_state=task_state)
def test_get_message_event_type_with_assistant_file(self, message_cycle_manager):
def test_get_message_event_type_with_assistant_file(self, message_cycle_manager, cycle_db: Session):
"""Test get_message_event_type returns MESSAGE_FILE when message has assistant-generated files.
This ensures that AI-generated images (belongs_to='assistant') trigger the MESSAGE_FILE event,
allowing the frontend to properly display generated image files with url field.
"""
with patch("core.app.task_pipeline.message_cycle_manager.session_factory") as mock_session_factory:
# Setup mock session and message file
mock_session = Mock()
mock_session_factory.create_session.return_value.__enter__.return_value = mock_session
cycle_db.add(_message_file())
cycle_db.commit()
mock_message_file = Mock()
mock_message_file.belongs_to = "assistant"
mock_session.scalar.return_value = mock_message_file
with current_app.app_context():
result = message_cycle_manager.get_message_event_type("test-message-id")
# Execute
with current_app.app_context():
result = message_cycle_manager.get_message_event_type("test-message-id")
assert result == StreamEvent.MESSAGE_FILE
assert "test-message-id" in message_cycle_manager._message_has_file
# Assert
assert result == StreamEvent.MESSAGE_FILE
mock_session.scalar.assert_called_once()
def test_get_message_event_type_with_user_file(self, message_cycle_manager):
def test_get_message_event_type_with_user_file(self, message_cycle_manager, cycle_db: Session):
"""Test get_message_event_type returns MESSAGE when message only has user-uploaded files.
This is a regression test for the issue where user-uploaded images (belongs_to='user')
@@ -68,90 +131,81 @@ class TestMessageCycleManagerOptimization:
resulting in broken images in the chat UI. The query filters for belongs_to='assistant',
so when only user files exist, the database query returns None, resulting in MESSAGE event type.
"""
with patch("core.app.task_pipeline.message_cycle_manager.session_factory") as mock_session_factory:
# Setup mock session and message file
mock_session = Mock()
mock_session_factory.create_session.return_value.__enter__.return_value = mock_session
cycle_db.add(_message_file(belongs_to=MessageFileBelongsTo.USER))
cycle_db.commit()
# When querying for assistant files with only user files present, return None
# (simulates database query with belongs_to='assistant' filter returning no results)
mock_session.scalar.return_value = None
with current_app.app_context():
result = message_cycle_manager.get_message_event_type("test-message-id")
# Execute
with current_app.app_context():
result = message_cycle_manager.get_message_event_type("test-message-id")
assert result == StreamEvent.MESSAGE
assert "test-message-id" not in message_cycle_manager._message_has_file
# Assert
assert result == StreamEvent.MESSAGE
mock_session.scalar.assert_called_once()
def test_get_message_event_type_without_message_file(self, message_cycle_manager):
def test_get_message_event_type_without_message_file(self, message_cycle_manager, cycle_db: Session):
"""Test get_message_event_type returns MESSAGE when message has no files."""
with patch("core.app.task_pipeline.message_cycle_manager.session_factory") as mock_session_factory:
# Setup mock session and no message file
mock_session = Mock()
mock_session_factory.create_session.return_value.__enter__.return_value = mock_session
# Current implementation uses session.scalar(select(...))
mock_session.scalar.return_value = None
assert list(cycle_db.scalars(select(MessageFile)).all()) == []
# Execute
with current_app.app_context():
result = message_cycle_manager.get_message_event_type("test-message-id")
with current_app.app_context():
result = message_cycle_manager.get_message_event_type("test-message-id")
# Assert
assert result == StreamEvent.MESSAGE
mock_session.scalar.assert_called_once()
assert result == StreamEvent.MESSAGE
def test_get_message_event_type_uses_cache_without_query(self, message_cycle_manager):
def test_get_message_event_type_uses_cache_without_query(
self, message_cycle_manager, cycle_db: Session, sqlite_engine: Engine
):
"""Return MESSAGE_FILE directly from in-memory cache without opening a DB session."""
message_cycle_manager._message_has_file.add("cached-message")
statements: list[str] = []
with patch("core.app.task_pipeline.message_cycle_manager.session_factory") as mock_session_factory:
def record_statement(_conn, _cursor, statement, _parameters, _context, _executemany) -> None:
statements.append(statement)
event.listen(sqlite_engine, "before_cursor_execute", record_statement)
try:
result = message_cycle_manager.get_message_event_type("cached-message")
finally:
event.remove(sqlite_engine, "before_cursor_execute", record_statement)
assert result == StreamEvent.MESSAGE_FILE
mock_session_factory.create_session.assert_not_called()
assert statements == []
def test_message_to_stream_response_with_precomputed_event_type(self, message_cycle_manager):
def test_message_to_stream_response_with_precomputed_event_type(self, message_cycle_manager, cycle_db: Session):
"""MessageCycleManager.message_to_stream_response expects a valid event_type; callers should precompute it."""
with patch("core.app.task_pipeline.message_cycle_manager.session_factory") as mock_session_factory:
# Setup mock session and message file
mock_session = Mock()
mock_session_factory.create_session.return_value.__enter__.return_value = mock_session
cycle_db.add(_message_file())
cycle_db.commit()
mock_message_file = Mock()
mock_message_file.belongs_to = "assistant"
mock_session.scalar.return_value = mock_message_file
with current_app.app_context():
event_type = message_cycle_manager.get_message_event_type("test-message-id")
result = message_cycle_manager.message_to_stream_response(
answer="Hello world", message_id="test-message-id", event_type=event_type
)
# Execute: compute event type once, then pass to message_to_stream_response
with current_app.app_context():
event_type = message_cycle_manager.get_message_event_type("test-message-id")
result = message_cycle_manager.message_to_stream_response(
answer="Hello world", message_id="test-message-id", event_type=event_type
)
assert isinstance(result, MessageStreamResponse)
assert result.answer == "Hello world"
assert result.id == "test-message-id"
assert result.event == StreamEvent.MESSAGE_FILE
# Assert
assert isinstance(result, MessageStreamResponse)
assert result.answer == "Hello world"
assert result.id == "test-message-id"
assert result.event == StreamEvent.MESSAGE_FILE
mock_session.scalar.assert_called_once()
def test_message_to_stream_response_with_event_type_skips_query(self, message_cycle_manager):
def test_message_to_stream_response_with_event_type_skips_query(
self, message_cycle_manager, cycle_db: Session, sqlite_engine: Engine
):
"""Test that message_to_stream_response skips database query when event_type is provided."""
with patch("core.app.task_pipeline.message_cycle_manager.session_factory") as mock_session_factory:
# Execute with event_type provided
statements: list[str] = []
def record_statement(_conn, _cursor, statement, _parameters, _context, _executemany) -> None:
statements.append(statement)
event.listen(sqlite_engine, "before_cursor_execute", record_statement)
try:
result = message_cycle_manager.message_to_stream_response(
answer="Hello world", message_id="test-message-id", event_type=StreamEvent.MESSAGE
)
finally:
event.remove(sqlite_engine, "before_cursor_execute", record_statement)
# Assert
assert isinstance(result, MessageStreamResponse)
assert result.answer == "Hello world"
assert result.id == "test-message-id"
assert result.event == StreamEvent.MESSAGE
# Should not open a session when event_type is provided
mock_session_factory.create_session.assert_not_called()
assert isinstance(result, MessageStreamResponse)
assert result.answer == "Hello world"
assert result.id == "test-message-id"
assert result.event == StreamEvent.MESSAGE
assert statements == []
def test_message_to_stream_response_with_from_variable_selector(self, message_cycle_manager):
"""Test message_to_stream_response with from_variable_selector parameter."""
@@ -168,40 +222,32 @@ class TestMessageCycleManagerOptimization:
assert result.from_variable_selector == ["var1", "var2"]
assert result.event == StreamEvent.MESSAGE
def test_optimization_usage_example(self, message_cycle_manager):
def test_optimization_usage_example(self, message_cycle_manager, cycle_db: Session, sqlite_engine: Engine):
"""Test the optimization pattern that should be used by callers."""
# Step 1: Get event type once (this queries database)
with patch("core.app.task_pipeline.message_cycle_manager.session_factory") as mock_session_factory:
mock_session = Mock()
mock_session_factory.create_session.return_value.__enter__.return_value = mock_session
# Current implementation uses session.scalar(select(...))
mock_session.scalar.return_value = None # No files
statements: list[str] = []
def record_statement(_conn, _cursor, statement, _parameters, _context, _executemany) -> None:
statements.append(statement)
event.listen(sqlite_engine, "before_cursor_execute", record_statement)
try:
with current_app.app_context():
event_type = message_cycle_manager.get_message_event_type("test-message-id")
# Should open session once
mock_session_factory.create_session.assert_called_once()
assert event_type == StreamEvent.MESSAGE
# Step 2: Use event_type for multiple calls (no additional queries)
with patch("core.app.task_pipeline.message_cycle_manager.session_factory") as mock_session_factory:
mock_session_factory.create_session.return_value.__enter__.return_value = Mock()
chunk1_response = message_cycle_manager.message_to_stream_response(
answer="Chunk 1", message_id="test-message-id", event_type=event_type
)
chunk2_response = message_cycle_manager.message_to_stream_response(
answer="Chunk 2", message_id="test-message-id", event_type=event_type
)
finally:
event.remove(sqlite_engine, "before_cursor_execute", record_statement)
# Should not open session again when event_type provided
mock_session_factory.create_session.assert_not_called()
assert chunk1_response.event == StreamEvent.MESSAGE
assert chunk2_response.event == StreamEvent.MESSAGE
assert chunk1_response.answer == "Chunk 1"
assert chunk2_response.answer == "Chunk 2"
assert event_type == StreamEvent.MESSAGE
assert len([statement for statement in statements if statement.lstrip().upper().startswith("SELECT")]) == 1
assert chunk1_response.event == StreamEvent.MESSAGE
assert chunk2_response.event == StreamEvent.MESSAGE
assert chunk1_response.answer == "Chunk 1"
assert chunk2_response.answer == "Chunk 2"
def test_generate_conversation_name_returns_none_for_completion(self, message_cycle_manager):
"""Return None when completion entities are used for conversation naming.
@@ -269,51 +315,38 @@ class TestMessageCycleManagerOptimization:
assert message_cycle_manager._application_generate_entity.is_new_conversation is False
mock_timer.assert_not_called()
def test_generate_conversation_name_worker_returns_when_conversation_missing(self, message_cycle_manager):
def test_generate_conversation_name_worker_returns_when_conversation_missing(
self, message_cycle_manager, cycle_db: Session
):
"""Return early when the conversation cannot be found."""
flask_app = Flask(__name__)
db_session = Mock()
db_session.scalar.return_value = None
assert list(cycle_db.scalars(select(Conversation)).all()) == []
with _patch_create_session(db_session):
message_cycle_manager._generate_conversation_name_worker(flask_app, "conv-missing", "hello")
message_cycle_manager._generate_conversation_name_worker(flask_app, "conv-missing", "hello")
db_session.commit.assert_not_called()
assert list(cycle_db.scalars(select(Conversation)).all()) == []
def test_generate_conversation_name_worker_returns_when_app_missing(self, message_cycle_manager):
def test_generate_conversation_name_worker_returns_when_app_missing(self, message_cycle_manager, cycle_db: Session):
"""Return early when non-completion conversation has no app relation."""
flask_app = Flask(__name__)
conversation = SimpleNamespace(mode=AppMode.CHAT, app=None, app_id="app-id")
db_session = Mock()
db_session.scalar.return_value = conversation
db_session.get.return_value = None
conversation = _conversation()
cycle_db.add(conversation)
cycle_db.commit()
with _patch_create_session(db_session):
message_cycle_manager._generate_conversation_name_worker(flask_app, "conv-1", "hello")
message_cycle_manager._generate_conversation_name_worker(flask_app, "conv-1", "hello")
db_session.commit.assert_not_called()
assert cycle_db.get(Conversation, "conv-1").name == ""
assert cycle_db.get(App, "app-id") is None
def test_generate_conversation_name_worker_uses_cached_name(self, message_cycle_manager):
def test_generate_conversation_name_worker_uses_cached_name(
self, message_cycle_manager, cycle_db: Session, sqlite_engine: Engine
):
"""Use cached conversation name when present and avoid LLM call."""
flask_app = Flask(__name__)
class ConversationWithPoisonedApp:
mode = AppMode.CHAT
app_id = "app-id"
name = ""
@property
def app(self):
raise AssertionError("conversation.app must not open an implicit session")
conversation = ConversationWithPoisonedApp()
app_model = SimpleNamespace(tenant_id="tenant-1")
db_session = Mock()
db_session.scalar.return_value = conversation
db_session.get.return_value = app_model
cycle_db.add_all([_app(), _conversation()])
cycle_db.commit()
with (
_patch_create_session(db_session) as create_session,
patch("core.app.task_pipeline.message_cycle_manager.redis_client") as mock_redis,
patch("core.app.task_pipeline.message_cycle_manager.LLMGenerator") as mock_llm_generator,
):
@@ -321,27 +354,23 @@ class TestMessageCycleManagerOptimization:
message_cycle_manager._generate_conversation_name_worker(flask_app, "conv-1", "hello")
assert cycle_db.in_transaction() is False
with Session(sqlite_engine) as verification_session:
conversation = verification_session.get(Conversation, "conv-1")
assert conversation is not None
assert conversation.name == "cached-title"
create_session.assert_called_once_with()
db_session.get.assert_called_once_with(App, "app-id")
db_session.commit.assert_called_once()
mock_llm_generator.generate_conversation_name.assert_not_called()
mock_redis.setex.assert_not_called()
def test_generate_conversation_name_worker_generates_and_caches_name(self, message_cycle_manager):
def test_generate_conversation_name_worker_generates_and_caches_name(
self, message_cycle_manager, cycle_db: Session, sqlite_engine: Engine
):
"""Generate conversation name and write it to redis cache on cache miss."""
flask_app = Flask(__name__)
conversation = SimpleNamespace(
mode=AppMode.CHAT,
app=SimpleNamespace(tenant_id="tenant-1"),
app_id="app-id",
name="",
)
db_session = Mock()
db_session.scalar.return_value = conversation
cycle_db.add_all([_app(), _conversation()])
cycle_db.commit()
with (
_patch_create_session(db_session),
patch("core.app.task_pipeline.message_cycle_manager.redis_client") as mock_redis,
patch("core.app.task_pipeline.message_cycle_manager.LLMGenerator") as mock_llm_generator,
):
@@ -350,27 +379,27 @@ class TestMessageCycleManagerOptimization:
message_cycle_manager._generate_conversation_name_worker(flask_app, "conv-1", "hello")
assert cycle_db.in_transaction() is False
with Session(sqlite_engine) as verification_session:
conversation = verification_session.get(Conversation, "conv-1")
assert conversation is not None
assert conversation.name == "generated-title"
db_session.commit.assert_called_once()
mock_redis.setex.assert_called_once()
def test_generate_conversation_name_worker_falls_back_when_generation_fails(
self, message_cycle_manager, caplog: pytest.LogCaptureFixture
self,
message_cycle_manager,
cycle_db: Session,
sqlite_engine: Engine,
caplog: pytest.LogCaptureFixture,
):
"""Fallback to truncated query when LLM generation fails."""
flask_app = Flask(__name__)
conversation = SimpleNamespace(
mode=AppMode.CHAT,
app=SimpleNamespace(tenant_id="tenant-1"),
app_id="app-id",
name="",
)
db_session = Mock()
db_session.scalar.return_value = conversation
cycle_db.add_all([_app(), _conversation()])
cycle_db.commit()
long_query = "q" * 60
with (
_patch_create_session(db_session),
patch("core.app.task_pipeline.message_cycle_manager.redis_client") as mock_redis,
patch("core.app.task_pipeline.message_cycle_manager.LLMGenerator") as mock_llm_generator,
patch("core.app.task_pipeline.message_cycle_manager.dify_config") as mock_dify_config,
@@ -382,8 +411,11 @@ class TestMessageCycleManagerOptimization:
with caplog.at_level(logging.ERROR, logger="core.app.task_pipeline.message_cycle_manager"):
message_cycle_manager._generate_conversation_name_worker(flask_app, "conv-1", long_query)
assert cycle_db.in_transaction() is False
with Session(sqlite_engine) as verification_session:
conversation = verification_session.get(Conversation, "conv-1")
assert conversation is not None
assert conversation.name == (long_query[:47] + "...")
db_session.commit.assert_called_once()
assert any(record.levelno == logging.ERROR for record in caplog.records)
def test_handle_annotation_reply_sets_metadata(self, message_cycle_manager):
@@ -454,33 +486,25 @@ class TestMessageCycleManagerOptimization:
assert message_cycle_manager._task_state.metadata.retriever_resources[0].position == 1
assert message_cycle_manager._task_state.metadata.retriever_resources[1].position == 2
def test_message_file_to_stream_response_builds_signed_url(self, message_cycle_manager):
def test_message_file_to_stream_response_builds_signed_url(self, message_cycle_manager, cycle_db: Session):
"""Build a stream response with a signed tool file URL.
Args: message_cycle_manager with mocked Session/db and sign_tool_file.
Args: message_cycle_manager with a persisted MessageFile and mocked sign_tool_file.
Returns: MessageStreamResponse with signed url and belongs_to normalized to user.
Side effects: Calls sign_tool_file for tool file ids.
"""
message_cycle_manager._application_generate_entity.task_id = "task-1"
message_file = SimpleNamespace(
id="file-1",
type="image",
belongs_to=None,
url="tool://file.verylongextension",
message_id="msg-1",
cycle_db.add(
_message_file(
file_id="file-1",
message_id="msg-1",
belongs_to=None,
url="tool://file.verylongextension",
)
)
cycle_db.commit()
session = Mock()
session.scalar.return_value = message_file
with (
patch("core.app.task_pipeline.message_cycle_manager.Session") as mock_session_cls,
patch("core.app.task_pipeline.message_cycle_manager.sign_tool_file") as mock_sign,
patch("core.app.task_pipeline.message_cycle_manager.db") as mock_db,
):
mock_db.engine = Mock()
mock_session_cls.return_value.__enter__.return_value = session
with patch("core.app.task_pipeline.message_cycle_manager.sign_tool_file") as mock_sign:
mock_sign.return_value = "signed-url"
response = message_cycle_manager.message_file_to_stream_response(SimpleNamespace(message_file_id="file-1"))
@@ -514,56 +538,42 @@ class TestMessageCycleManagerOptimization:
assert len(message_cycle_manager._task_state.metadata.retriever_resources) == 1
assert message_cycle_manager._task_state.metadata.retriever_resources[0].position == 1
def test_message_file_to_stream_response_uses_http_url_directly(self, message_cycle_manager):
def test_message_file_to_stream_response_uses_http_url_directly(self, message_cycle_manager, cycle_db: Session):
"""Use original URL when message file URL is already HTTP."""
message_cycle_manager._application_generate_entity.task_id = "task-http"
message_file = SimpleNamespace(
id="file-http",
type="image",
belongs_to="assistant",
url="http://example.com/pic.png",
message_id="msg-http",
)
session = Mock()
session.scalar.return_value = message_file
with (
patch("core.app.task_pipeline.message_cycle_manager.Session") as mock_session_cls,
patch("core.app.task_pipeline.message_cycle_manager.db") as mock_db,
):
mock_db.engine = Mock()
mock_session_cls.return_value.__enter__.return_value = session
response = message_cycle_manager.message_file_to_stream_response(
SimpleNamespace(message_file_id="file-http")
cycle_db.add(
_message_file(
file_id="file-http",
message_id="msg-http",
belongs_to=MessageFileBelongsTo.ASSISTANT,
url="http://example.com/pic.png",
)
)
cycle_db.commit()
response = message_cycle_manager.message_file_to_stream_response(SimpleNamespace(message_file_id="file-http"))
assert response is not None
assert response.url == "http://example.com/pic.png"
assert "msg-http" in message_cycle_manager._message_has_file
def test_message_file_to_stream_response_defaults_extension_to_bin_without_dot(self, message_cycle_manager):
def test_message_file_to_stream_response_defaults_extension_to_bin_without_dot(
self, message_cycle_manager, cycle_db: Session
):
"""Default tool file extension to .bin when URL has no extension part."""
message_cycle_manager._application_generate_entity.task_id = "task-bin"
message_file = SimpleNamespace(
id="file-bin",
type="file",
belongs_to="assistant",
url="tool-file-id",
message_id="msg-bin",
cycle_db.add(
_message_file(
file_id="file-bin",
message_id="msg-bin",
belongs_to=MessageFileBelongsTo.ASSISTANT,
url="tool-file-id",
file_type=FileType.CUSTOM,
)
)
cycle_db.commit()
session = Mock()
session.scalar.return_value = message_file
with (
patch("core.app.task_pipeline.message_cycle_manager.Session") as mock_session_cls,
patch("core.app.task_pipeline.message_cycle_manager.sign_tool_file") as mock_sign,
patch("core.app.task_pipeline.message_cycle_manager.db") as mock_db,
):
mock_db.engine = Mock()
mock_session_cls.return_value.__enter__.return_value = session
with patch("core.app.task_pipeline.message_cycle_manager.sign_tool_file") as mock_sign:
mock_sign.return_value = "signed-bin-url"
response = message_cycle_manager.message_file_to_stream_response(
@@ -574,19 +584,13 @@ class TestMessageCycleManagerOptimization:
assert response.url == "signed-bin-url"
mock_sign.assert_called_once_with(tool_file_id="tool-file-id", extension=".bin")
def test_message_file_to_stream_response_returns_none_when_file_missing(self, message_cycle_manager):
def test_message_file_to_stream_response_returns_none_when_file_missing(
self, message_cycle_manager, cycle_db: Session
):
"""Return None when message file lookup does not find a record."""
session = Mock()
session.scalar.return_value = None
assert list(cycle_db.scalars(select(MessageFile)).all()) == []
with (
patch("core.app.task_pipeline.message_cycle_manager.Session") as mock_session_cls,
patch("core.app.task_pipeline.message_cycle_manager.db") as mock_db,
):
mock_db.engine = Mock()
mock_session_cls.return_value.__enter__.return_value = session
response = message_cycle_manager.message_file_to_stream_response(SimpleNamespace(message_file_id="missing"))
response = message_cycle_manager.message_file_to_stream_response(SimpleNamespace(message_file_id="missing"))
assert response is None
@@ -1,11 +1,19 @@
"""Comprehensive unit tests for core/memory/token_buffer_memory.py"""
"""Comprehensive SQLite-backed tests for token-buffer memory."""
from collections.abc import Iterator
from dataclasses import dataclass
from datetime import UTC, datetime, timedelta
from decimal import Decimal
from unittest.mock import MagicMock, patch
from uuid import uuid4
import pytest
from sqlalchemy import Engine, event
from sqlalchemy.orm import Session
from core.memory import token_buffer_memory as memory_module
from core.memory.token_buffer_memory import TokenBufferMemory
from graphon.file import FileTransferMethod, FileType
from graphon.model_runtime.entities import (
AssistantPromptMessage,
ImagePromptMessageContent,
@@ -13,13 +21,44 @@ from graphon.model_runtime.entities import (
TextPromptMessageContent,
UserPromptMessage,
)
from models.model import AppMode
from models.base import TypeBase
from models.enums import ConversationFromSource, CreatorUserRole, MessageFileBelongsTo
from models.model import AppMode, Message, MessageFile
from models.workflow import Workflow, WorkflowType
# ---------------------------------------------------------------------------
# Helpers / shared fixtures
# ---------------------------------------------------------------------------
@dataclass(frozen=True)
class Database:
"""Typed SQLite binding plus executed SQL for query-count assertions."""
engine: Engine
session: Session
statements: list[tuple[str, object]]
@pytest.fixture
def database(sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch) -> Iterator[Database]:
TypeBase.metadata.create_all(
sqlite_engine,
tables=[Message.__table__, MessageFile.__table__, Workflow.__table__],
)
statements: list[tuple[str, object]] = []
def record_statement(_connection, _cursor, statement, parameters, _context, _executemany) -> None:
statements.append((statement, parameters))
event.listen(sqlite_engine, "before_cursor_execute", record_statement)
with Session(sqlite_engine, expire_on_commit=False) as session:
database = Database(engine=sqlite_engine, session=session, statements=statements)
monkeypatch.setattr(memory_module, "db", database)
yield database
event.remove(sqlite_engine, "before_cursor_execute", record_statement)
def _make_conversation(mode: AppMode = AppMode.CHAT) -> MagicMock:
"""Return a minimal Conversation mock."""
conv = MagicMock()
@@ -47,6 +86,73 @@ def _make_message(answer: str = "hello", answer_tokens: int = 5) -> MagicMock:
return msg
def _persist_message(
database: Database,
conversation_id: str,
*,
query: str = "user query",
answer: str = "hello",
answer_tokens: int = 5,
created_at: datetime | None = None,
workflow_run_id: str | None = None,
) -> Message:
message = Message(
id=str(uuid4()),
app_id="app-1",
conversation_id=conversation_id,
_inputs={},
query=query,
message={},
message_unit_price=Decimal(0),
answer=answer,
answer_tokens=answer_tokens,
answer_unit_price=Decimal(0),
currency="USD",
from_source=ConversationFromSource.API,
workflow_run_id=workflow_run_id,
created_at=created_at or datetime.now(UTC).replace(tzinfo=None),
)
database.session.add(message)
database.session.commit()
return message
def _persist_message_file(
database: Database,
message: Message,
*,
belongs_to: MessageFileBelongsTo | None,
) -> MessageFile:
message_file = MessageFile(
message_id=message.id,
type=FileType.IMAGE,
transfer_method=FileTransferMethod.REMOTE_URL,
created_by_role=CreatorUserRole.ACCOUNT,
created_by="account-1",
belongs_to=belongs_to,
url="https://example.com/image.png",
)
database.session.add(message_file)
database.session.commit()
return message_file
def _persist_workflow(database: Database, *, workflow_id: str) -> Workflow:
workflow = Workflow(
id=workflow_id,
tenant_id="tenant-1",
app_id="app-1",
type=WorkflowType.CHAT,
version="1",
graph="{}",
features="{}",
created_by="account-1",
)
database.session.add(workflow)
database.session.commit()
return workflow
# ===========================================================================
# Tests for __init__ and workflow_run_repo property
# ===========================================================================
@@ -61,25 +167,25 @@ class TestInit:
assert mem.model_instance is mi
assert mem._workflow_run_repo is None
def test_workflow_run_repo_is_created_lazily(self):
def test_workflow_run_repo_is_created_lazily(self, database: Database):
conv = _make_conversation()
mi = _make_model_instance()
mem = TokenBufferMemory(conversation=conv, model_instance=mi)
mock_repo = MagicMock()
with (
patch("core.memory.token_buffer_memory.sessionmaker") as mock_sm,
patch("core.memory.token_buffer_memory.db") as mock_db,
patch(
"core.memory.token_buffer_memory.DifyAPIRepositoryFactory.create_api_workflow_run_repository",
return_value=mock_repo,
),
):
mock_db.engine = MagicMock()
with patch(
"core.memory.token_buffer_memory.DifyAPIRepositoryFactory.create_api_workflow_run_repository",
return_value=mock_repo,
) as repository_factory:
repo = mem.workflow_run_repo
assert repo is mock_repo
assert mem._workflow_run_repo is mock_repo
session_factory = repository_factory.call_args.args[0]
with session_factory() as session:
assert isinstance(session, Session)
assert session.get_bind() is database.engine
def test_workflow_run_repo_cached_after_first_access(self):
conv = _make_conversation()
mi = _make_model_instance()
@@ -410,7 +516,7 @@ class TestBuildPromptMessageWithFiles:
)
@pytest.mark.parametrize("mode", [AppMode.ADVANCED_CHAT, AppMode.WORKFLOW])
def test_workflow_mode_workflow_not_found_raises(self, mode):
def test_workflow_mode_workflow_not_found_raises(self, mode, database: Database):
"""Raises ValueError when Workflow lookup returns None."""
conv = _make_conversation(mode)
conv.app = MagicMock()
@@ -422,22 +528,17 @@ class TestBuildPromptMessageWithFiles:
mem._workflow_run_repo = MagicMock()
mem._workflow_run_repo.get_workflow_run_by_id.return_value = mock_workflow_run
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
):
mock_db.session.scalar.return_value = None # workflow not found
with pytest.raises(ValueError, match="Workflow not found"):
mem._build_prompt_message_with_files(
message_files=[],
text_content="text",
message=_make_message(),
app_record=MagicMock(),
is_user_message=True,
)
with pytest.raises(ValueError, match="Workflow not found"):
mem._build_prompt_message_with_files(
message_files=[],
text_content="text",
message=_make_message(),
app_record=MagicMock(),
is_user_message=True,
)
@pytest.mark.parametrize("mode", [AppMode.ADVANCED_CHAT, AppMode.WORKFLOW])
def test_workflow_mode_success_no_files_user(self, mode):
def test_workflow_mode_success_no_files_user(self, mode, database: Database):
"""Happy path: workflow mode, no message files → plain UserPromptMessage."""
conv = _make_conversation(mode)
conv.app = MagicMock()
@@ -445,22 +546,16 @@ class TestBuildPromptMessageWithFiles:
mock_workflow_run = MagicMock()
mock_workflow_run.workflow_id = str(uuid4())
mock_workflow = MagicMock()
mock_workflow.features_dict = {}
workflow = _persist_workflow(database, workflow_id=mock_workflow_run.workflow_id)
mem = TokenBufferMemory(conversation=conv, model_instance=_make_model_instance())
mem._workflow_run_repo = MagicMock()
mem._workflow_run_repo.get_workflow_run_by_id.return_value = mock_workflow_run
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch(
"core.memory.token_buffer_memory.FileUploadConfigManager.convert",
return_value=None,
),
with patch(
"core.memory.token_buffer_memory.FileUploadConfigManager.convert",
return_value=None,
):
mock_db.session.scalar.return_value = mock_workflow
result = mem._build_prompt_message_with_files(
message_files=[],
text_content="wf text",
@@ -471,6 +566,7 @@ class TestBuildPromptMessageWithFiles:
assert isinstance(result, UserPromptMessage)
assert result.content == "wf text"
assert database.session.get(Workflow, workflow.id) is workflow
# ------------------------------------------------------------------
# Invalid mode
@@ -498,417 +594,140 @@ class TestBuildPromptMessageWithFiles:
class TestGetHistoryPromptMessages:
"""Tests for get_history_prompt_messages."""
"""Tests for persisted history retrieval, file batching, and pruning."""
def _make_memory(self, mode: AppMode = AppMode.CHAT) -> TokenBufferMemory:
conv = _make_conversation(mode)
conv.app = MagicMock()
return TokenBufferMemory(conversation=conv, model_instance=_make_model_instance())
def test_returns_empty_when_no_messages(self):
def test_returns_empty_when_no_messages(self, database: Database) -> None:
assert self._make_memory().get_history_prompt_messages() == []
def test_skips_newest_message_without_answer(self, database: Database) -> None:
mem = self._make_memory()
with patch("core.memory.token_buffer_memory.db") as mock_db:
mock_db.session.scalars.return_value.all.return_value = []
result = mem.get_history_prompt_messages()
assert result == []
message = _persist_message(database, mem.conversation.id, answer="", answer_tokens=0)
def test_skips_first_message_without_answer(self):
"""The newest message (index 0 after extraction) without answer and tokens==0 is skipped."""
assert mem.get_history_prompt_messages() == []
assert database.session.get(Message, message.id) is message
def test_message_with_answer_returns_user_and_assistant_prompts(self, database: Database) -> None:
mem = self._make_memory()
_persist_message(database, mem.conversation.id, query="My query", answer="My answer", answer_tokens=10)
msg_no_answer = _make_message(answer="", answer_tokens=0)
msg_no_answer.parent_message_id = None # ensures extract_thread_messages returns it
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch(
"core.memory.token_buffer_memory.extract_thread_messages",
return_value=[msg_no_answer],
),
):
mock_db.session.scalars.return_value.all.side_effect = [
[msg_no_answer], # first call: messages query
[], # second call: user files query (never hit, but safe)
]
result = mem.get_history_prompt_messages()
assert result == []
def test_message_with_answer_not_skipped(self):
"""A message with a non-empty answer is NOT popped."""
mem = self._make_memory()
msg = _make_message(answer="some answer", answer_tokens=10)
msg.parent_message_id = None
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch(
"core.memory.token_buffer_memory.extract_thread_messages",
return_value=[msg],
),
patch(
"core.memory.token_buffer_memory.FileUploadConfigManager.convert",
return_value=None,
),
):
# user files query → empty; assistant files query → empty
mock_db.session.scalars.return_value.all.return_value = []
result = mem.get_history_prompt_messages()
assert len(result) == 2 # one user + one assistant
def test_message_limit_default_is_500(self):
"""When message_limit is None the stmt is limited to 500."""
mem = self._make_memory()
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch("core.memory.token_buffer_memory.select") as mock_select,
patch("core.memory.token_buffer_memory.extract_thread_messages", return_value=[]),
):
mock_stmt = MagicMock()
mock_select.return_value.where.return_value.order_by.return_value = mock_stmt
mock_stmt.limit.return_value = mock_stmt
mock_db.session.scalars.return_value.all.return_value = []
mem.get_history_prompt_messages(message_limit=None)
mock_stmt.limit.assert_called_with(500)
def test_message_limit_clipped_to_500(self):
"""A message_limit > 500 is clamped to 500."""
mem = self._make_memory()
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch("core.memory.token_buffer_memory.select") as mock_select,
patch("core.memory.token_buffer_memory.extract_thread_messages", return_value=[]),
):
mock_stmt = MagicMock()
mock_select.return_value.where.return_value.order_by.return_value = mock_stmt
mock_stmt.limit.return_value = mock_stmt
mock_db.session.scalars.return_value.all.return_value = []
mem.get_history_prompt_messages(message_limit=9999)
mock_stmt.limit.assert_called_with(500)
def test_message_limit_positive_used(self):
"""A positive message_limit < 500 is used as-is."""
mem = self._make_memory()
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch("core.memory.token_buffer_memory.select") as mock_select,
patch("core.memory.token_buffer_memory.extract_thread_messages", return_value=[]),
):
mock_stmt = MagicMock()
mock_select.return_value.where.return_value.order_by.return_value = mock_stmt
mock_stmt.limit.return_value = mock_stmt
mock_db.session.scalars.return_value.all.return_value = []
mem.get_history_prompt_messages(message_limit=10)
mock_stmt.limit.assert_called_with(10)
def test_message_limit_zero_uses_default(self):
"""message_limit=0 triggers the else branch → default 500."""
mem = self._make_memory()
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch("core.memory.token_buffer_memory.select") as mock_select,
patch("core.memory.token_buffer_memory.extract_thread_messages", return_value=[]),
):
mock_stmt = MagicMock()
mock_select.return_value.where.return_value.order_by.return_value = mock_stmt
mock_stmt.limit.return_value = mock_stmt
mock_db.session.scalars.return_value.all.return_value = []
mem.get_history_prompt_messages(message_limit=0)
mock_stmt.limit.assert_called_with(500)
def test_user_files_cause_build_with_files_call(self):
"""When user_files is non-empty _build_prompt_message_with_files is invoked."""
mem = self._make_memory()
msg = _make_message()
msg.parent_message_id = None
mock_user_file = MagicMock()
mock_user_file.message_id = msg.id # must match so batched grouping keys it to this message
mock_user_prompt = UserPromptMessage(content="from build")
mock_assistant_prompt = AssistantPromptMessage(content="answer")
call_count = {"n": 0}
def scalars_side_effect(stmt):
r = MagicMock()
if call_count["n"] == 0:
# messages query
r.all.return_value = [msg]
elif call_count["n"] == 1:
# user files
r.all.return_value = [mock_user_file]
else:
# assistant files
r.all.return_value = []
call_count["n"] += 1
return r
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch(
"core.memory.token_buffer_memory.extract_thread_messages",
return_value=[msg],
),
patch.object(
mem,
"_build_prompt_message_with_files",
side_effect=[mock_user_prompt, mock_assistant_prompt],
) as mock_build,
patch(
"core.memory.token_buffer_memory.FileUploadConfigManager.convert",
return_value=None,
),
):
mock_db.session.scalars.side_effect = scalars_side_effect
result = mem.get_history_prompt_messages()
assert mock_build.call_count >= 1
# First call should be user message
first_call_kwargs = mock_build.call_args_list[0][1]
assert first_call_kwargs["is_user_message"] is True
def test_assistant_files_cause_build_with_files_call(self):
"""When assistant_files is non-empty, build is called with is_user_message=False."""
mem = self._make_memory()
msg = _make_message()
msg.parent_message_id = None
mock_assistant_file = MagicMock()
mock_assistant_file.message_id = msg.id # must match so batched grouping keys it to this message
mock_user_prompt = UserPromptMessage(content="query")
mock_assistant_prompt = AssistantPromptMessage(content="built")
call_count = {"n": 0}
def scalars_side_effect(stmt):
r = MagicMock()
if call_count["n"] == 0:
r.all.return_value = [msg]
elif call_count["n"] == 1:
r.all.return_value = [] # no user files
else:
r.all.return_value = [mock_assistant_file]
call_count["n"] += 1
return r
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch(
"core.memory.token_buffer_memory.extract_thread_messages",
return_value=[msg],
),
patch.object(
mem,
"_build_prompt_message_with_files",
return_value=mock_assistant_prompt,
) as mock_build,
):
mock_db.session.scalars.side_effect = scalars_side_effect
result = mem.get_history_prompt_messages()
mock_build.assert_called_once()
call_kwargs = mock_build.call_args[1]
assert call_kwargs["is_user_message"] is False
def test_message_files_loaded_with_constant_query_count(self):
"""Regression guard against N+1: message files must be batch-loaded.
Regardless of the number of messages in the thread, file loading must use a
constant number of queries (1 messages query + 2 batched file queries),
never 2 queries per message.
"""
mem = self._make_memory()
messages = [_make_message() for _ in range(5)]
for m in messages:
m.parent_message_id = None
scalars_calls = {"n": 0}
def scalars_side_effect(stmt):
r = MagicMock()
# First call returns the thread messages; the batched file queries return none.
r.all.return_value = messages if scalars_calls["n"] == 0 else []
scalars_calls["n"] += 1
return r
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch("core.memory.token_buffer_memory.extract_thread_messages", return_value=messages),
patch("core.memory.token_buffer_memory.FileUploadConfigManager.convert", return_value=None),
):
mock_db.session.scalars.side_effect = scalars_side_effect
mem.get_history_prompt_messages()
# 1 (messages) + 2 (batched user/assistant files) = 3, independent of message count.
# Before this fix it would have been 1 + 2 * 5 = 11 (an N+1 pattern).
assert scalars_calls["n"] == 3
def test_token_pruning_removes_oldest_messages(self):
"""If tokens exceed limit, oldest messages are removed until within limit."""
conv = _make_conversation()
conv.app = MagicMock()
# Model returns tokens that decrease only after removing pairs
token_values = [3000, 1500] # first call over limit, second within
mi = MagicMock()
mi.get_llm_num_tokens.side_effect = token_values
mem = TokenBufferMemory(conversation=conv, model_instance=mi)
msg = _make_message()
msg.parent_message_id = None
call_count = {"n": 0}
def scalars_side_effect(stmt):
r = MagicMock()
if call_count["n"] == 0:
r.all.return_value = [msg]
else:
r.all.return_value = []
call_count["n"] += 1
return r
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch(
"core.memory.token_buffer_memory.extract_thread_messages",
return_value=[msg],
),
patch(
"core.memory.token_buffer_memory.FileUploadConfigManager.convert",
return_value=None,
),
):
mock_db.session.scalars.side_effect = scalars_side_effect
result = mem.get_history_prompt_messages(max_token_limit=2000)
# After pruning, we should have fewer than the 2 initial messages
assert len(result) <= 1
def test_token_pruning_stops_at_single_message(self):
"""Pruning stops when only 1 message remains (to prevent empty list)."""
conv = _make_conversation()
conv.app = MagicMock()
# Always over limit
mi = MagicMock()
mi.get_llm_num_tokens.return_value = 99999
mem = TokenBufferMemory(conversation=conv, model_instance=mi)
msg = _make_message()
msg.parent_message_id = None
call_count = {"n": 0}
def scalars_side_effect(stmt):
r = MagicMock()
if call_count["n"] == 0:
r.all.return_value = [msg]
else:
r.all.return_value = []
call_count["n"] += 1
return r
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch(
"core.memory.token_buffer_memory.extract_thread_messages",
return_value=[msg],
),
patch(
"core.memory.token_buffer_memory.FileUploadConfigManager.convert",
return_value=None,
),
):
mock_db.session.scalars.side_effect = scalars_side_effect
result = mem.get_history_prompt_messages(max_token_limit=1)
# At least 1 message should remain
assert len(result) >= 1
def test_no_pruning_when_within_limit(self):
"""When tokens ≤ limit, no pruning occurs."""
mem = self._make_memory()
mem.model_instance.get_llm_num_tokens.return_value = 50 # well under default 2000
msg = _make_message()
msg.parent_message_id = None
call_count = {"n": 0}
def scalars_side_effect(stmt):
r = MagicMock()
if call_count["n"] == 0:
r.all.return_value = [msg]
else:
r.all.return_value = []
call_count["n"] += 1
return r
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch(
"core.memory.token_buffer_memory.extract_thread_messages",
return_value=[msg],
),
patch(
"core.memory.token_buffer_memory.FileUploadConfigManager.convert",
return_value=None,
),
):
mock_db.session.scalars.side_effect = scalars_side_effect
result = mem.get_history_prompt_messages(max_token_limit=2000)
assert len(result) == 2 # user + assistant
def test_plain_user_and_assistant_messages_returned(self):
"""Without files, plain UserPromptMessage and AssistantPromptMessage appear."""
mem = self._make_memory()
msg = _make_message(answer="My answer")
msg.query = "My query"
msg.parent_message_id = None
call_count = {"n": 0}
def scalars_side_effect(stmt):
r = MagicMock()
if call_count["n"] == 0:
r.all.return_value = [msg]
else:
r.all.return_value = []
call_count["n"] += 1
return r
with (
patch("core.memory.token_buffer_memory.db") as mock_db,
patch(
"core.memory.token_buffer_memory.extract_thread_messages",
return_value=[msg],
),
patch(
"core.memory.token_buffer_memory.FileUploadConfigManager.convert",
return_value=None,
),
):
mock_db.session.scalars.side_effect = scalars_side_effect
result = mem.get_history_prompt_messages()
result = mem.get_history_prompt_messages()
assert len(result) == 2
user_msg, ai_msg = result
assert isinstance(user_msg, UserPromptMessage)
assert user_msg.content == "My query"
assert isinstance(ai_msg, AssistantPromptMessage)
assert ai_msg.content == "My answer"
assert isinstance(result[0], UserPromptMessage)
assert result[0].content == "My query"
assert isinstance(result[1], AssistantPromptMessage)
assert result[1].content == "My answer"
def test_history_is_conversation_scoped(self, database: Database) -> None:
mem = self._make_memory()
_persist_message(database, mem.conversation.id, answer="visible")
_persist_message(database, "other-conversation", answer="hidden")
result = mem.get_history_prompt_messages()
assert [prompt.content for prompt in result] == ["user query", "visible"]
@pytest.mark.parametrize(
("message_limit", "expected_limit"),
[(None, 500), (9999, 500), (10, 10), (0, 500)],
)
def test_message_limit_is_applied_to_executable_query(
self,
database: Database,
message_limit: int | None,
expected_limit: int,
) -> None:
mem = self._make_memory()
before = len(database.statements)
mem.get_history_prompt_messages(message_limit=message_limit)
statements = database.statements[before:]
assert len(statements) == 1
sql, parameters = statements[0]
assert "LIMIT" in sql
assert expected_limit in parameters
@pytest.mark.parametrize(
("belongs_to", "is_user_message"),
[
(MessageFileBelongsTo.USER, True),
(None, True),
(MessageFileBelongsTo.ASSISTANT, False),
],
)
def test_message_files_use_persisted_ownership(
self,
database: Database,
belongs_to: MessageFileBelongsTo | None,
is_user_message: bool,
) -> None:
mem = self._make_memory()
message = _persist_message(database, mem.conversation.id)
message_file = _persist_message_file(database, message, belongs_to=belongs_to)
built_prompt = (
UserPromptMessage(content="built user")
if is_user_message
else AssistantPromptMessage(content="built assistant")
)
with patch.object(mem, "_build_prompt_message_with_files", return_value=built_prompt) as build_prompt:
result = mem.get_history_prompt_messages()
build_prompt.assert_called_once()
assert build_prompt.call_args.kwargs["message_files"] == [message_file]
assert build_prompt.call_args.kwargs["is_user_message"] is is_user_message
assert built_prompt in result
def test_message_files_are_batch_loaded_with_constant_query_count(self, database: Database) -> None:
mem = self._make_memory()
base_time = datetime.now(UTC).replace(tzinfo=None)
messages = [
_persist_message(
database,
mem.conversation.id,
query=f"query-{index}",
answer=f"answer-{index}",
created_at=base_time + timedelta(seconds=index),
)
for index in range(5)
]
before = len(database.statements)
with patch("core.memory.token_buffer_memory.extract_thread_messages", return_value=messages):
result = mem.get_history_prompt_messages()
selects = [sql for sql, _ in database.statements[before:] if sql.lstrip().upper().startswith("SELECT")]
assert len(selects) == 3
assert len(result) == 10
@pytest.mark.parametrize(
("token_values", "max_token_limit", "expected_length"),
[
([3000, 1500], 2000, 1),
([99999, 99999], 1, 1),
([50], 2000, 2),
],
)
def test_token_pruning_uses_persisted_history(
self,
database: Database,
token_values: list[int],
max_token_limit: int,
expected_length: int,
) -> None:
mem = self._make_memory()
mem.model_instance.get_llm_num_tokens.side_effect = token_values
_persist_message(database, mem.conversation.id)
result = mem.get_history_prompt_messages(max_token_limit=max_token_limit)
assert len(result) == expected_length
# ===========================================================================
@@ -7,36 +7,212 @@ Covers:
- TraceTask._get_user_id_from_metadata
"""
from unittest.mock import MagicMock, patch
import uuid
from collections.abc import Iterator
from contextlib import contextmanager
from unittest.mock import PropertyMock, patch
import pytest
from sqlalchemy import Engine, event
from sqlalchemy.orm import Session
from core.tools.entities.tool_entities import ApiProviderSchemaType
from extensions.ext_database import db
from graphon.model_runtime.entities.model_entities import ModelType
from models.account import Tenant
from models.base import TypeBase
from models.model import App, AppMode, IconType
from models.provider import Provider, ProviderCredential, ProviderModel, ProviderModelCredential, ProviderType
from models.tools import ApiToolProvider, BuiltinToolProvider, MCPToolProvider, WorkflowToolProvider
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_db_and_session_patches(scalar_side_effect=None, scalar_return_value=None):
"""Return (mock_db, cm, session) ready to patch 'core.ops.ops_trace_manager.db'
and 'core.ops.ops_trace_manager.Session'.
@pytest.fixture
def orm_session(sqlite_engine: Engine) -> Iterator[Session]:
models = (
App,
Tenant,
Provider,
ProviderCredential,
ProviderModel,
ProviderModelCredential,
BuiltinToolProvider,
ApiToolProvider,
WorkflowToolProvider,
MCPToolProvider,
)
tables = [model.metadata.tables[model.__tablename__] for model in models]
TypeBase.metadata.create_all(sqlite_engine, tables=tables)
Provide either scalar_side_effect (list, for multiple calls) or
scalar_return_value (single value).
"""
mock_db = MagicMock()
mock_db.engine = MagicMock()
with patch.object(type(db), "engine", new_callable=PropertyMock, return_value=sqlite_engine):
with Session(sqlite_engine, expire_on_commit=False) as session:
yield session
session = MagicMock()
if scalar_side_effect is not None:
session.scalar.side_effect = scalar_side_effect
def _persist_app(session: Session, *, tenant_id: str, name: str = "MyApp") -> App:
app = App(
id=str(uuid.uuid4()),
tenant_id=tenant_id,
name=name,
mode=AppMode.WORKFLOW,
icon_type=IconType.EMOJI,
icon="workflow",
icon_background="#FFFFFF",
enable_site=True,
enable_api=False,
)
session.add(app)
session.commit()
return app
def _persist_tenant(session: Session, *, name: str = "MyWorkspace") -> Tenant:
tenant = Tenant(name=name)
session.add(tenant)
session.commit()
return tenant
def _persist_tool_provider(
session: Session, provider_type: str
) -> BuiltinToolProvider | ApiToolProvider | WorkflowToolProvider | MCPToolProvider:
tenant_id = str(uuid.uuid4())
user_id = str(uuid.uuid4())
if provider_type in {"builtin", "plugin"}:
provider = BuiltinToolProvider(
name="CredentialA",
tenant_id=tenant_id,
user_id=user_id,
provider="test/provider",
)
elif provider_type == "api":
provider = ApiToolProvider(
name="CredentialA",
icon="icon.svg",
schema="{}",
schema_type_str=ApiProviderSchemaType.OPENAPI,
user_id=user_id,
tenant_id=tenant_id,
description="API provider",
tools_str="[]",
credentials_str="{}",
)
elif provider_type == "workflow":
provider = WorkflowToolProvider(
name="CredentialA",
label="CredentialA",
icon="icon.svg",
app_id=str(uuid.uuid4()),
version="1",
user_id=user_id,
tenant_id=tenant_id,
description="Workflow provider",
)
elif provider_type == "mcp":
provider = MCPToolProvider(
name="CredentialA",
server_identifier="credential-a",
server_url="https://example.com/mcp",
server_url_hash="credential-a-hash",
icon="icon.svg",
tenant_id=tenant_id,
user_id=user_id,
)
else:
session.scalar.return_value = scalar_return_value
raise ValueError(f"unsupported provider type: {provider_type}")
cm = MagicMock()
cm.__enter__ = MagicMock(return_value=session)
cm.__exit__ = MagicMock(return_value=False)
session.add(provider)
session.commit()
return provider
return mock_db, cm, session
def _persist_provider_credential(
session: Session,
*,
tenant_id: str,
credential_name: str = "ProvCredName",
) -> ProviderCredential:
credential = ProviderCredential(
tenant_id=tenant_id,
provider_name="openai",
credential_name=credential_name,
encrypted_config="{}",
)
session.add(credential)
session.commit()
return credential
def _persist_model_credential(
session: Session,
*,
tenant_id: str,
credential_name: str = "ModelCredName",
) -> ProviderModelCredential:
credential = ProviderModelCredential(
tenant_id=tenant_id,
provider_name="openai",
model_name="gpt-4",
model_type=ModelType.LLM,
credential_name=credential_name,
encrypted_config="{}",
)
session.add(credential)
session.commit()
return credential
def _persist_provider(
session: Session,
*,
tenant_id: str,
credential_id: str | None,
) -> Provider:
provider = Provider(
tenant_id=tenant_id,
provider_name="openai",
provider_type=ProviderType.CUSTOM,
credential_id=credential_id,
)
session.add(provider)
session.commit()
return provider
def _persist_provider_model(
session: Session,
*,
tenant_id: str,
credential_id: str | None,
) -> ProviderModel:
model = ProviderModel(
tenant_id=tenant_id,
provider_name="openai",
model_name="gpt-4",
model_type=ModelType.LLM,
credential_id=credential_id,
)
session.add(model)
session.commit()
return model
@contextmanager
def _raise_on_table(engine: Engine, table_name: str) -> Iterator[None]:
"""Raise only when SQL targets the named table, leaving other real lookups intact."""
def fail_target_query(_conn, _cursor, statement, _parameters, _context, _executemany):
if f"FROM {table_name}" in statement:
raise RuntimeError(f"forced failure for {table_name}")
event.listen(engine, "before_cursor_execute", fail_target_query)
try:
yield
finally:
event.remove(engine, "before_cursor_execute", fail_target_query)
# ---------------------------------------------------------------------------
@@ -47,62 +223,42 @@ def _make_db_and_session_patches(scalar_side_effect=None, scalar_return_value=No
class TestLookupAppAndWorkspaceNames:
"""Tests for _lookup_app_and_workspace_names(app_id, tenant_id)."""
def test_both_found(self):
def test_both_found(self, orm_session: Session):
"""Returns (app_name, workspace_name) when both records exist."""
from core.ops.ops_trace_manager import _lookup_app_and_workspace_names
mock_db, cm, _session = _make_db_and_session_patches(scalar_side_effect=["MyApp", "MyWorkspace"])
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", return_value=cm),
):
app_name, workspace_name = _lookup_app_and_workspace_names("app-123", "tenant-456")
tenant = _persist_tenant(orm_session)
app = _persist_app(orm_session, tenant_id=tenant.id)
app_name, workspace_name = _lookup_app_and_workspace_names(app.id, tenant.id)
assert app_name == "MyApp"
assert workspace_name == "MyWorkspace"
def test_app_only_found(self):
def test_app_only_found(self, orm_session: Session):
"""Returns (app_name, '') when tenant record is absent."""
from core.ops.ops_trace_manager import _lookup_app_and_workspace_names
mock_db, cm, _session = _make_db_and_session_patches(scalar_side_effect=["MyApp", None])
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", return_value=cm),
):
app_name, workspace_name = _lookup_app_and_workspace_names("app-123", "tenant-456")
app = _persist_app(orm_session, tenant_id=str(uuid.uuid4()))
app_name, workspace_name = _lookup_app_and_workspace_names(app.id, str(uuid.uuid4()))
assert app_name == "MyApp"
assert workspace_name == ""
def test_tenant_only_found(self):
def test_tenant_only_found(self, orm_session: Session):
"""Returns ('', workspace_name) when app record is absent."""
from core.ops.ops_trace_manager import _lookup_app_and_workspace_names
mock_db, cm, _session = _make_db_and_session_patches(scalar_side_effect=[None, "MyWorkspace"])
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", return_value=cm),
):
app_name, workspace_name = _lookup_app_and_workspace_names("app-123", "tenant-456")
tenant = _persist_tenant(orm_session)
app_name, workspace_name = _lookup_app_and_workspace_names(str(uuid.uuid4()), tenant.id)
assert app_name == ""
assert workspace_name == "MyWorkspace"
def test_neither_found(self):
def test_neither_found(self, orm_session: Session):
"""Returns ('', '') when both DB lookups return None."""
from core.ops.ops_trace_manager import _lookup_app_and_workspace_names
mock_db, cm, _session = _make_db_and_session_patches(scalar_side_effect=[None, None])
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", return_value=cm),
):
app_name, workspace_name = _lookup_app_and_workspace_names("app-123", "tenant-456")
app_name, workspace_name = _lookup_app_and_workspace_names(str(uuid.uuid4()), str(uuid.uuid4()))
assert app_name == ""
assert workspace_name == ""
@@ -111,50 +267,30 @@ class TestLookupAppAndWorkspaceNames:
"""Returns ('', '') immediately when both IDs are None — no DB access."""
from core.ops.ops_trace_manager import _lookup_app_and_workspace_names
mock_db = MagicMock()
mock_session_cls = MagicMock()
app_name, workspace_name = _lookup_app_and_workspace_names(None, None)
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", mock_session_cls),
):
app_name, workspace_name = _lookup_app_and_workspace_names(None, None)
mock_session_cls.assert_not_called()
assert app_name == ""
assert workspace_name == ""
def test_app_id_none_only_queries_tenant(self):
def test_app_id_none_only_queries_tenant(self, orm_session: Session):
"""When app_id is None, only the tenant query is issued."""
from core.ops.ops_trace_manager import _lookup_app_and_workspace_names
mock_db, cm, session = _make_db_and_session_patches(scalar_return_value="OnlyWorkspace")
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", return_value=cm),
):
app_name, workspace_name = _lookup_app_and_workspace_names(None, "tenant-456")
tenant = _persist_tenant(orm_session, name="OnlyWorkspace")
app_name, workspace_name = _lookup_app_and_workspace_names(None, tenant.id)
assert app_name == ""
assert workspace_name == "OnlyWorkspace"
assert session.scalar.call_count == 1
def test_tenant_id_none_only_queries_app(self):
def test_tenant_id_none_only_queries_app(self, orm_session: Session):
"""When tenant_id is None, only the app query is issued."""
from core.ops.ops_trace_manager import _lookup_app_and_workspace_names
mock_db, cm, session = _make_db_and_session_patches(scalar_return_value="OnlyApp")
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", return_value=cm),
):
app_name, workspace_name = _lookup_app_and_workspace_names("app-123", None)
app = _persist_app(orm_session, tenant_id=str(uuid.uuid4()), name="OnlyApp")
app_name, workspace_name = _lookup_app_and_workspace_names(app.id, None)
assert app_name == "OnlyApp"
assert workspace_name == ""
assert session.scalar.call_count == 1
# ---------------------------------------------------------------------------
@@ -166,32 +302,20 @@ class TestLookupCredentialName:
"""Tests for _lookup_credential_name(credential_id, provider_type)."""
@pytest.mark.parametrize("provider_type", ["builtin", "plugin", "api", "workflow", "mcp"])
def test_known_provider_types_return_name(self, provider_type):
def test_known_provider_types_return_name(self, provider_type: str, orm_session: Session):
"""Each valid provider_type results in a DB query and returns the credential name."""
from core.ops.ops_trace_manager import _lookup_credential_name
mock_db, cm, session = _make_db_and_session_patches(scalar_return_value="CredentialA")
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", return_value=cm),
):
result = _lookup_credential_name("cred-123", provider_type)
provider = _persist_tool_provider(orm_session, provider_type)
result = _lookup_credential_name(provider.id, provider_type)
assert result == "CredentialA"
session.scalar.assert_called_once()
def test_credential_not_found_returns_empty_string(self):
def test_credential_not_found_returns_empty_string(self, orm_session: Session):
"""Returns '' when DB yields None for the given credential_id."""
from core.ops.ops_trace_manager import _lookup_credential_name
mock_db, cm, _session = _make_db_and_session_patches(scalar_return_value=None)
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", return_value=cm),
):
result = _lookup_credential_name("cred-999", "api")
result = _lookup_credential_name(str(uuid.uuid4()), "api")
assert result == ""
@@ -199,48 +323,24 @@ class TestLookupCredentialName:
"""Returns '' immediately for an unrecognised provider_type — no DB access."""
from core.ops.ops_trace_manager import _lookup_credential_name
mock_db = MagicMock()
mock_session_cls = MagicMock()
result = _lookup_credential_name(str(uuid.uuid4()), "unknown_type")
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", mock_session_cls),
):
result = _lookup_credential_name("cred-123", "unknown_type")
mock_session_cls.assert_not_called()
assert result == ""
def test_none_credential_id_returns_empty_string_without_db(self):
"""Returns '' immediately when credential_id is None — no DB access."""
from core.ops.ops_trace_manager import _lookup_credential_name
mock_db = MagicMock()
mock_session_cls = MagicMock()
result = _lookup_credential_name(None, "api")
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", mock_session_cls),
):
result = _lookup_credential_name(None, "api")
mock_session_cls.assert_not_called()
assert result == ""
def test_none_provider_type_returns_empty_string_without_db(self):
"""Returns '' immediately when provider_type is None — no DB access."""
from core.ops.ops_trace_manager import _lookup_credential_name
mock_db = MagicMock()
mock_session_cls = MagicMock()
result = _lookup_credential_name(str(uuid.uuid4()), None)
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", mock_session_cls),
):
result = _lookup_credential_name("cred-123", None)
mock_session_cls.assert_not_called()
assert result == ""
def test_builtin_and_plugin_map_to_same_model(self):
@@ -281,106 +381,78 @@ class TestLookupCredentialName:
class TestLookupLlmCredentialInfo:
"""Tests for _lookup_llm_credential_info(tenant_id, provider, model, model_type)."""
def _provider_record(self, credential_id: str | None = None) -> MagicMock:
record = MagicMock()
record.credential_id = credential_id
return record
def _model_record(self, credential_id: str | None = None) -> MagicMock:
record = MagicMock()
record.credential_id = credential_id
return record
def test_model_level_credential_found(self):
def test_model_level_credential_found(self, orm_session: Session):
"""Returns model-level credential_id and name when ProviderModel has a credential."""
from core.ops.ops_trace_manager import _lookup_llm_credential_info
provider_record = self._provider_record(credential_id=None)
model_record = self._model_record(credential_id="model-cred-id")
tenant_id = str(uuid.uuid4())
model_credential = _persist_model_credential(orm_session, tenant_id=tenant_id)
_persist_provider(orm_session, tenant_id=tenant_id, credential_id=None)
_persist_provider_model(orm_session, tenant_id=tenant_id, credential_id=model_credential.id)
# scalar calls: (1) Provider, (2) ProviderModel, (3) ProviderModelCredential.credential_name
mock_db, cm, _session = _make_db_and_session_patches(
scalar_side_effect=[provider_record, model_record, "ModelCredName"]
decoy_tenant_id = str(uuid.uuid4())
decoy_credential = _persist_model_credential(
orm_session,
tenant_id=decoy_tenant_id,
credential_name="WrongTenantCredential",
)
_persist_provider(orm_session, tenant_id=decoy_tenant_id, credential_id=None)
_persist_provider_model(orm_session, tenant_id=decoy_tenant_id, credential_id=decoy_credential.id)
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", return_value=cm),
):
cred_id, cred_name = _lookup_llm_credential_info("tenant-1", "openai", "gpt-4")
cred_id, cred_name = _lookup_llm_credential_info(tenant_id, "openai", "gpt-4")
assert cred_id == "model-cred-id"
assert cred_id == model_credential.id
assert cred_name == "ModelCredName"
def test_provider_level_fallback_when_no_model_credential(self):
def test_provider_level_fallback_when_no_model_credential(self, orm_session: Session):
"""Falls back to provider-level credential when ProviderModel has no credential_id."""
from core.ops.ops_trace_manager import _lookup_llm_credential_info
provider_record = self._provider_record(credential_id="prov-cred-id")
model_record = self._model_record(credential_id=None)
tenant_id = str(uuid.uuid4())
provider_credential = _persist_provider_credential(orm_session, tenant_id=tenant_id)
_persist_provider(orm_session, tenant_id=tenant_id, credential_id=provider_credential.id)
_persist_provider_model(orm_session, tenant_id=tenant_id, credential_id=None)
# scalar calls: (1) Provider, (2) ProviderModel (no cred), (3) ProviderCredential.credential_name
mock_db, cm, _session = _make_db_and_session_patches(
scalar_side_effect=[provider_record, model_record, "ProvCredName"]
)
cred_id, cred_name = _lookup_llm_credential_info(tenant_id, "openai", "gpt-4")
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", return_value=cm),
):
cred_id, cred_name = _lookup_llm_credential_info("tenant-1", "openai", "gpt-4")
assert cred_id == "prov-cred-id"
assert cred_id == provider_credential.id
assert cred_name == "ProvCredName"
def test_provider_level_fallback_when_no_model_record(self):
def test_provider_level_fallback_when_no_model_record(self, orm_session: Session):
"""Falls back to provider-level credential when no ProviderModel row exists."""
from core.ops.ops_trace_manager import _lookup_llm_credential_info
provider_record = self._provider_record(credential_id="prov-cred-id")
tenant_id = str(uuid.uuid4())
provider_credential = _persist_provider_credential(orm_session, tenant_id=tenant_id)
_persist_provider(orm_session, tenant_id=tenant_id, credential_id=provider_credential.id)
# scalar calls: (1) Provider, (2) ProviderModel → None, (3) ProviderCredential.credential_name
mock_db, cm, _session = _make_db_and_session_patches(scalar_side_effect=[provider_record, None, "ProvCredName"])
cred_id, cred_name = _lookup_llm_credential_info(tenant_id, "openai", "gpt-4")
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", return_value=cm),
):
cred_id, cred_name = _lookup_llm_credential_info("tenant-1", "openai", "gpt-4")
assert cred_id == "prov-cred-id"
assert cred_id == provider_credential.id
assert cred_name == "ProvCredName"
def test_no_model_arg_uses_provider_level_only(self):
def test_no_model_arg_uses_provider_level_only(self, orm_session: Session):
"""When model is None, skips ProviderModel query and uses provider credential."""
from core.ops.ops_trace_manager import _lookup_llm_credential_info
provider_record = self._provider_record(credential_id="prov-cred-id")
tenant_id = str(uuid.uuid4())
provider_credential = _persist_provider_credential(orm_session, tenant_id=tenant_id)
_persist_provider(orm_session, tenant_id=tenant_id, credential_id=provider_credential.id)
# scalar calls: (1) Provider, (2) ProviderCredential.credential_name — no ProviderModel
mock_db, cm, session = _make_db_and_session_patches(scalar_side_effect=[provider_record, "ProvCredName"])
cred_id, cred_name = _lookup_llm_credential_info(tenant_id, "openai", None)
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", return_value=cm),
):
cred_id, cred_name = _lookup_llm_credential_info("tenant-1", "openai", None)
assert cred_id == "prov-cred-id"
assert cred_id == provider_credential.id
assert cred_name == "ProvCredName"
assert session.scalar.call_count == 2
def test_provider_not_found_returns_none_and_empty(self):
def test_provider_not_found_returns_none_and_empty(self, orm_session: Session):
"""Returns (None, '') when Provider record does not exist."""
from core.ops.ops_trace_manager import _lookup_llm_credential_info
mock_db, cm, _session = _make_db_and_session_patches(scalar_return_value=None)
other_tenant_id = str(uuid.uuid4())
_persist_provider(orm_session, tenant_id=other_tenant_id, credential_id=None)
tenant_id = str(uuid.uuid4())
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", return_value=cm),
):
cred_id, cred_name = _lookup_llm_credential_info("tenant-1", "openai", "gpt-4")
cred_id, cred_name = _lookup_llm_credential_info(tenant_id, "openai", "gpt-4")
assert cred_id is None
assert cred_name == ""
@@ -389,16 +461,8 @@ class TestLookupLlmCredentialInfo:
"""Returns (None, '') immediately when tenant_id is None — no DB access."""
from core.ops.ops_trace_manager import _lookup_llm_credential_info
mock_db = MagicMock()
mock_session_cls = MagicMock()
cred_id, cred_name = _lookup_llm_credential_info(None, "openai", "gpt-4")
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", mock_session_cls),
):
cred_id, cred_name = _lookup_llm_credential_info(None, "openai", "gpt-4")
mock_session_cls.assert_not_called()
assert cred_id is None
assert cred_name == ""
@@ -406,69 +470,46 @@ class TestLookupLlmCredentialInfo:
"""Returns (None, '') immediately when provider is None — no DB access."""
from core.ops.ops_trace_manager import _lookup_llm_credential_info
mock_db = MagicMock()
mock_session_cls = MagicMock()
cred_id, cred_name = _lookup_llm_credential_info(str(uuid.uuid4()), None, "gpt-4")
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", mock_session_cls),
):
cred_id, cred_name = _lookup_llm_credential_info("tenant-1", None, "gpt-4")
mock_session_cls.assert_not_called()
assert cred_id is None
assert cred_name == ""
def test_db_error_on_outer_query_returns_none_and_empty(self):
def test_db_error_on_outer_query_returns_none_and_empty(self, orm_session: Session, sqlite_engine: Engine):
"""Returns (None, '') and logs a warning when the outer DB query raises."""
from core.ops.ops_trace_manager import _lookup_llm_credential_info
mock_db, cm, session = _make_db_and_session_patches()
session.scalar.side_effect = Exception("DB connection failed")
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", return_value=cm),
):
cred_id, cred_name = _lookup_llm_credential_info("tenant-1", "openai", "gpt-4")
with _raise_on_table(sqlite_engine, "providers"):
cred_id, cred_name = _lookup_llm_credential_info(str(uuid.uuid4()), "openai", "gpt-4")
assert cred_id is None
assert cred_name == ""
def test_credential_name_lookup_failure_returns_id_with_empty_name(self):
def test_credential_name_lookup_failure_returns_id_with_empty_name(
self, orm_session: Session, sqlite_engine: Engine
):
"""When credential name sub-query fails, returns cred_id but '' for name."""
from core.ops.ops_trace_manager import _lookup_llm_credential_info
provider_record = self._provider_record(credential_id="prov-cred-id")
tenant_id = str(uuid.uuid4())
provider_credential = _persist_provider_credential(orm_session, tenant_id=tenant_id)
_persist_provider(orm_session, tenant_id=tenant_id, credential_id=provider_credential.id)
# Provider found, no model record, then name lookup raises
mock_db, cm, _session = _make_db_and_session_patches(
scalar_side_effect=[provider_record, None, Exception("deleted")]
)
with _raise_on_table(sqlite_engine, "provider_credentials"):
cred_id, cred_name = _lookup_llm_credential_info(tenant_id, "openai", "gpt-4")
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", return_value=cm),
):
cred_id, cred_name = _lookup_llm_credential_info("tenant-1", "openai", "gpt-4")
assert cred_id == "prov-cred-id"
assert cred_id == provider_credential.id
assert cred_name == ""
def test_no_credential_on_provider_or_model_returns_none_id(self):
def test_no_credential_on_provider_or_model_returns_none_id(self, orm_session: Session):
"""Returns (None, '') when neither provider nor model has a credential_id."""
from core.ops.ops_trace_manager import _lookup_llm_credential_info
provider_record = self._provider_record(credential_id=None)
model_record = self._model_record(credential_id=None)
tenant_id = str(uuid.uuid4())
_persist_provider(orm_session, tenant_id=tenant_id, credential_id=None)
_persist_provider_model(orm_session, tenant_id=tenant_id, credential_id=None)
mock_db, cm, _session = _make_db_and_session_patches(scalar_side_effect=[provider_record, model_record])
with (
patch("core.ops.ops_trace_manager.db", mock_db),
patch("core.ops.ops_trace_manager.Session", return_value=cm),
):
cred_id, cred_name = _lookup_llm_credential_info("tenant-1", "openai", "gpt-4")
cred_id, cred_name = _lookup_llm_credential_info(tenant_id, "openai", "gpt-4")
assert cred_id is None
assert cred_name == ""
@@ -1,16 +1,16 @@
"""Unit tests for workflow-as-tool behavior.
StubSession/StubScalars emulate SQLAlchemy session/scalars with minimal methods
(`scalar`, `scalars`, `expunge`, `commit`, `refresh`, context manager) to keep
database access mocked and predictable in tests.
"""
"""Unit tests for workflow-as-tool behavior with real SQLite ORM boundaries."""
import json
import uuid
from collections.abc import Iterator
from dataclasses import dataclass
from types import SimpleNamespace
from typing import Any
from unittest.mock import MagicMock, Mock, patch
import pytest
from sqlalchemy import Engine, inspect
from sqlalchemy.orm import Session, sessionmaker
from core.app.entities.app_invoke_entities import InvokeFrom
from core.tools.__base.tool_runtime import ToolRuntime
@@ -23,74 +23,142 @@ from core.tools.entities.tool_entities import (
ToolProviderType,
)
from core.tools.errors import ToolInvokeError
from core.tools.workflow_as_tool import tool as workflow_tool_module
from core.tools.workflow_as_tool.tool import WorkflowTool
from graphon.file import FILE_MODEL_IDENTITY, FileTransferMethod, FileType
from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole
from models.base import TypeBase
from models.enums import EndUserType
from models.model import App, AppMode, EndUser
from models.workflow import Workflow, WorkflowType
TENANT_ID = "00000000-0000-0000-0000-000000000001"
OTHER_TENANT_ID = "00000000-0000-0000-0000-000000000002"
APP_ID = "00000000-0000-0000-0000-000000000003"
ACCOUNT_ID = "00000000-0000-0000-0000-000000000004"
END_USER_ID = "00000000-0000-0000-0000-000000000005"
CREATOR_ID = "00000000-0000-0000-0000-000000000006"
class StubScalars:
"""Minimal stub for SQLAlchemy scalar results."""
_value: Any
def __init__(self, value: Any) -> None:
self._value = value
def first(self) -> Any:
return self._value
@dataclass(frozen=True)
class SqliteToolDb:
engine: Engine
session_maker: sessionmaker[Session]
caller_session: Session
class StubSession:
"""Minimal stub for session_factory-created sessions."""
@pytest.fixture
def sqlite_tool_db(
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
) -> Iterator[SqliteToolDb]:
"""Bind service-owned sessions and Account tenant reloads to SQLite."""
models = (App, Workflow, EndUser, Account, Tenant, TenantAccountJoin)
TypeBase.metadata.create_all(sqlite_engine, tables=[model.__table__ for model in models])
session_maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False)
monkeypatch.setattr(workflow_tool_module.session_factory, "create_session", session_maker)
scalar_results: list[Any]
scalars_results: list[Any]
expunge_calls: list[object]
from models import account as account_module
def __init__(self, *, scalar_results: list[Any] | None = None, scalars_results: list[Any] | None = None) -> None:
self.scalar_results = list(scalar_results or [])
self.scalars_results = list(scalars_results or [])
self.expunge_calls: list[object] = []
def scalar(self, _stmt: Any) -> Any:
return self.scalar_results.pop(0)
def scalars(self, _stmt: Any) -> StubScalars:
return StubScalars(self.scalars_results.pop(0))
def expunge(self, value: Any) -> None:
self.expunge_calls.append(value)
def begin(self) -> "StubSession":
return self
def commit(self) -> None:
pass
def refresh(self, _value: Any) -> None:
pass
def close(self) -> None:
pass
def __enter__(self) -> "StubSession":
return self
def __exit__(self, exc_type: Any, exc: Any, tb: Any) -> bool:
return False
monkeypatch.setattr(account_module, "db", SimpleNamespace(engine=sqlite_engine))
with session_maker() as caller_session:
yield SqliteToolDb(engine=sqlite_engine, session_maker=session_maker, caller_session=caller_session)
def _build_tool() -> WorkflowTool:
def _persist_tenant(db: SqliteToolDb, *, tenant_id: str = TENANT_ID) -> Tenant:
tenant = Tenant(name="Tenant")
tenant.id = tenant_id
db.caller_session.add(tenant)
db.caller_session.commit()
return tenant
def _persist_account(db: SqliteToolDb, *, tenant_id: str = TENANT_ID) -> Account:
account = Account(name="Account", email="account@example.com")
account.id = ACCOUNT_ID
join = TenantAccountJoin(
tenant_id=tenant_id,
account_id=account.id,
current=True,
role=TenantAccountRole.NORMAL,
)
db.caller_session.add_all([account, join])
db.caller_session.commit()
return account
def _persist_end_user(
db: SqliteToolDb,
*,
end_user_id: str = END_USER_ID,
tenant_id: str = TENANT_ID,
) -> EndUser:
end_user = EndUser(
id=end_user_id,
tenant_id=tenant_id,
app_id=APP_ID,
type=EndUserType.SERVICE_API,
name="End user",
session_id="end-user-session",
)
db.caller_session.add(end_user)
db.caller_session.commit()
return end_user
def _persist_app(db: SqliteToolDb) -> App:
app = App(
id=APP_ID,
tenant_id=TENANT_ID,
name="Workflow app",
description="",
mode=AppMode.WORKFLOW,
icon_type=None,
icon="",
icon_background=None,
app_model_config_id=None,
workflow_id=None,
enable_site=False,
enable_api=True,
max_active_requests=None,
created_by=CREATOR_ID,
)
db.caller_session.add(app)
db.caller_session.commit()
return app
def _persist_workflow(db: SqliteToolDb, *, version: str, workflow_id: str | None = None) -> Workflow:
workflow = Workflow.new(
tenant_id=TENANT_ID,
app_id=APP_ID,
type=WorkflowType.WORKFLOW.value,
version=version,
graph=json.dumps({"nodes": [], "edges": []}),
features="{}",
created_by=CREATOR_ID,
environment_variables=[],
conversation_variables=[],
rag_pipeline_variables=[],
)
workflow.id = workflow_id or str(uuid.uuid4())
db.caller_session.add(workflow)
db.caller_session.commit()
return workflow
def _build_tool(*, tenant_id: str = "test_tool", workflow_app_id: str = "app-1", version: str = "1") -> WorkflowTool:
entity = ToolEntity(
identity=ToolIdentity(author="test", name="test tool", label=I18nObject(en_US="test tool"), provider="test"),
parameters=[],
description=None,
has_runtime_parameters=False,
)
runtime = ToolRuntime(tenant_id="test_tool", invoke_from=InvokeFrom.EXPLORE)
runtime = ToolRuntime(tenant_id=tenant_id, invoke_from=InvokeFrom.EXPLORE)
return WorkflowTool(
workflow_app_id="app-1",
workflow_app_id=workflow_app_id,
workflow_as_tool_id="wf-tool-1",
version="1",
version=version,
workflow_entities={},
workflow_call_depth=1,
entity=entity,
@@ -98,7 +166,10 @@ def _build_tool() -> WorkflowTool:
)
def test_workflow_tool_should_raise_tool_invoke_error_when_result_has_error_field(monkeypatch: pytest.MonkeyPatch):
def test_workflow_tool_should_raise_tool_invoke_error_when_result_has_error_field(
monkeypatch: pytest.MonkeyPatch,
sqlite_tool_db: SqliteToolDb,
):
"""Ensure that WorkflowTool will throw a `ToolInvokeError` exception when
`WorkflowAppGenerator.generate` returns a result with `error` key inside
the `data` element.
@@ -122,11 +193,14 @@ def test_workflow_tool_should_raise_tool_invoke_error_when_result_has_error_fiel
with pytest.raises(ToolInvokeError) as exc_info:
# WorkflowTool always returns a generator, so we need to iterate to
# actually `run` the tool.
list(tool.invoke(MagicMock(), "test_user", {}))
list(tool.invoke(sqlite_tool_db.caller_session, "test_user", {}))
assert exc_info.value.args == ("oops",)
def test_workflow_tool_does_not_use_pause_state_config(monkeypatch: pytest.MonkeyPatch):
def test_workflow_tool_does_not_use_pause_state_config(
monkeypatch: pytest.MonkeyPatch,
sqlite_tool_db: SqliteToolDb,
):
"""Ensure pause_state_config is passed as None."""
tool = _build_tool()
@@ -140,14 +214,17 @@ def test_workflow_tool_does_not_use_pause_state_config(monkeypatch: pytest.Monke
monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock)
monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None)
list(tool.invoke(MagicMock(), "test_user", {}))
list(tool.invoke(sqlite_tool_db.caller_session, "test_user", {}))
call_kwargs = generate_mock.call_args.kwargs
assert "pause_state_config" in call_kwargs
assert call_kwargs["pause_state_config"] is None
def test_workflow_tool_passes_parent_trace_context_from_runtime(monkeypatch: pytest.MonkeyPatch):
def test_workflow_tool_passes_parent_trace_context_from_runtime(
monkeypatch: pytest.MonkeyPatch,
sqlite_tool_db: SqliteToolDb,
):
"""Ensure nested workflow runtime metadata is forwarded as parent trace context."""
tool = _build_tool()
tool.set_parent_trace_context(
@@ -165,7 +242,7 @@ def test_workflow_tool_passes_parent_trace_context_from_runtime(monkeypatch: pyt
monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock)
monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None)
list(tool.invoke(MagicMock(), "test_user", {}))
list(tool.invoke(sqlite_tool_db.caller_session, "test_user", {}))
call_kwargs = generate_mock.call_args.kwargs
assert call_kwargs["args"]["parent_trace_context"].model_dump() == {
@@ -174,7 +251,10 @@ def test_workflow_tool_passes_parent_trace_context_from_runtime(monkeypatch: pyt
}
def test_workflow_tool_passes_parent_trace_session_id(monkeypatch: pytest.MonkeyPatch):
def test_workflow_tool_passes_parent_trace_session_id(
monkeypatch: pytest.MonkeyPatch,
sqlite_tool_db: SqliteToolDb,
):
"""Ensure nested workflows inherit the parent observability session ID."""
tool = _build_tool()
tool.entity.parameters = [
@@ -197,14 +277,17 @@ def test_workflow_tool_passes_parent_trace_session_id(monkeypatch: pytest.Monkey
monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock)
monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None)
list(tool.invoke(MagicMock(), "test_user", {"trace_session_id": "user-input-session"}))
list(tool.invoke(sqlite_tool_db.caller_session, "test_user", {"trace_session_id": "user-input-session"}))
call_kwargs = generate_mock.call_args.kwargs
assert call_kwargs["args"]["inputs"]["trace_session_id"] == "user-input-session"
assert call_kwargs["args"]["trace_session_id"] == "session-1"
def test_workflow_tool_keeps_user_inputs_named_like_trace_runtime_keys(monkeypatch: pytest.MonkeyPatch):
def test_workflow_tool_keeps_user_inputs_named_like_trace_runtime_keys(
monkeypatch: pytest.MonkeyPatch,
sqlite_tool_db: SqliteToolDb,
):
"""Ensure private trace context does not overwrite same-named workflow inputs."""
tool = _build_tool()
tool.entity.parameters = [
@@ -238,7 +321,7 @@ def test_workflow_tool_keeps_user_inputs_named_like_trace_runtime_keys(monkeypat
list(
tool.invoke(
MagicMock(),
sqlite_tool_db.caller_session,
"test_user",
{
"outer_workflow_run_id": "user-workflow-input",
@@ -256,7 +339,10 @@ def test_workflow_tool_keeps_user_inputs_named_like_trace_runtime_keys(monkeypat
}
def test_workflow_tool_can_clear_parent_trace_context(monkeypatch: pytest.MonkeyPatch):
def test_workflow_tool_can_clear_parent_trace_context(
monkeypatch: pytest.MonkeyPatch,
sqlite_tool_db: SqliteToolDb,
):
"""Ensure reused WorkflowTool instances do not keep stale parent trace context."""
tool = _build_tool()
tool.set_parent_trace_context(
@@ -275,13 +361,16 @@ def test_workflow_tool_can_clear_parent_trace_context(monkeypatch: pytest.Monkey
monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock)
monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None)
list(tool.invoke(MagicMock(), "test_user", {}))
list(tool.invoke(sqlite_tool_db.caller_session, "test_user", {}))
call_kwargs = generate_mock.call_args.kwargs
assert "parent_trace_context" not in call_kwargs["args"]
def test_workflow_tool_can_clear_trace_session_id(monkeypatch: pytest.MonkeyPatch):
def test_workflow_tool_can_clear_trace_session_id(
monkeypatch: pytest.MonkeyPatch,
sqlite_tool_db: SqliteToolDb,
):
"""Ensure reused WorkflowTool instances do not keep stale trace session IDs."""
tool = _build_tool()
tool.set_trace_session_id("session-1")
@@ -297,7 +386,7 @@ def test_workflow_tool_can_clear_trace_session_id(monkeypatch: pytest.MonkeyPatc
monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock)
monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None)
list(tool.invoke(MagicMock(), "test_user", {}))
list(tool.invoke(sqlite_tool_db.caller_session, "test_user", {}))
call_kwargs = generate_mock.call_args.kwargs
assert "trace_session_id" not in call_kwargs["args"]
@@ -315,6 +404,7 @@ def test_workflow_tool_can_clear_trace_session_id(monkeypatch: pytest.MonkeyPatc
def test_workflow_tool_omits_parent_trace_context_when_runtime_is_incomplete(
monkeypatch: pytest.MonkeyPatch,
runtime_parameters: dict[str, Any],
sqlite_tool_db: SqliteToolDb,
):
"""Ensure incomplete runtime metadata does not leak parent trace context into generator args."""
tool = _build_tool()
@@ -330,13 +420,16 @@ def test_workflow_tool_omits_parent_trace_context_when_runtime_is_incomplete(
monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock)
monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None)
list(tool.invoke(MagicMock(), "test_user", {}))
list(tool.invoke(sqlite_tool_db.caller_session, "test_user", {}))
call_kwargs = generate_mock.call_args.kwargs
assert "parent_trace_context" not in call_kwargs["args"]
def test_workflow_tool_should_generate_variable_messages_for_outputs(monkeypatch: pytest.MonkeyPatch):
def test_workflow_tool_should_generate_variable_messages_for_outputs(
monkeypatch: pytest.MonkeyPatch,
sqlite_tool_db: SqliteToolDb,
):
"""Test that WorkflowTool should generate variable messages when there are outputs"""
tool = _build_tool()
@@ -359,7 +452,7 @@ def test_workflow_tool_should_generate_variable_messages_for_outputs(monkeypatch
monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None)
# Execute tool invocation
messages = list(tool.invoke(MagicMock(), "test_user", {}))
messages = list(tool.invoke(sqlite_tool_db.caller_session, "test_user", {}))
# Verify variable messages
variable_messages = [msg for msg in messages if msg.type == ToolInvokeMessage.MessageType.VARIABLE]
@@ -382,7 +475,10 @@ def test_workflow_tool_should_generate_variable_messages_for_outputs(monkeypatch
assert json_messages[0].message.json_object == mock_outputs
def test_workflow_tool_should_handle_empty_outputs(monkeypatch: pytest.MonkeyPatch):
def test_workflow_tool_should_handle_empty_outputs(
monkeypatch: pytest.MonkeyPatch,
sqlite_tool_db: SqliteToolDb,
):
"""Test that WorkflowTool should handle empty outputs correctly"""
tool = _build_tool()
@@ -402,7 +498,7 @@ def test_workflow_tool_should_handle_empty_outputs(monkeypatch: pytest.MonkeyPat
monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None)
# Execute tool invocation
messages = list(tool.invoke(MagicMock(), "test_user", {}))
messages = list(tool.invoke(sqlite_tool_db.caller_session, "test_user", {}))
# Verify generated messages
# Should contain: 0 variable messages + 1 text message + 1 JSON message = 2 messages
@@ -458,41 +554,32 @@ def test_create_file_message_should_include_file_marker():
assert message.meta == {"file": file_obj}
def test_resolve_user_from_database_falls_back_to_end_user(monkeypatch: pytest.MonkeyPatch):
def test_resolve_user_from_database_falls_back_to_end_user(sqlite_tool_db: SqliteToolDb):
"""Ensure worker context can resolve EndUser when Account is missing."""
tenant = SimpleNamespace(id="tenant_id")
end_user = SimpleNamespace(id="end_user_id", tenant_id="tenant_id")
# Monkeypatch session factory to return our stub session
stub_session = StubSession(scalar_results=[tenant, None, end_user])
monkeypatch.setattr(
"core.tools.workflow_as_tool.tool.session_factory.create_session",
lambda: stub_session,
_persist_tenant(sqlite_tool_db)
end_user = _persist_end_user(sqlite_tool_db)
other_tenant_end_user = _persist_end_user(
sqlite_tool_db,
end_user_id="00000000-0000-0000-0000-000000000007",
tenant_id=OTHER_TENANT_ID,
)
tool = _build_tool()
tool = _build_tool(tenant_id=TENANT_ID)
tool.runtime.invoke_from = InvokeFrom.SERVICE_API
tool.runtime.tenant_id = "tenant_id"
resolved_user = tool._resolve_user_from_database(user_id=end_user.id)
assert resolved_user is end_user
assert stub_session.expunge_calls == [end_user]
assert isinstance(resolved_user, EndUser)
assert resolved_user.id == end_user.id
assert resolved_user.tenant_id == TENANT_ID
assert inspect(resolved_user).detached is True
assert tool._resolve_user_from_database(user_id=other_tenant_end_user.id) is None
def test_resolve_user_from_database_returns_none_when_no_tenant(monkeypatch: pytest.MonkeyPatch):
def test_resolve_user_from_database_returns_none_when_no_tenant(sqlite_tool_db: SqliteToolDb):
"""Return None if tenant cannot be found in worker context."""
# Monkeypatch session factory to return our stub session with no tenant
monkeypatch.setattr(
"core.tools.workflow_as_tool.tool.session_factory.create_session",
lambda: StubSession(scalar_results=[None]),
)
tool = _build_tool()
tool = _build_tool(tenant_id=OTHER_TENANT_ID)
tool.runtime.invoke_from = InvokeFrom.SERVICE_API
tool.runtime.tenant_id = "missing_tenant"
resolved_user = tool._resolve_user_from_database(user_id="any")
@@ -544,7 +631,10 @@ def test_extract_usage_from_nested():
assert nested == {"total_tokens": 3}
def test_invoke_raises_when_user_not_found(monkeypatch: pytest.MonkeyPatch):
def test_invoke_raises_when_user_not_found(
monkeypatch: pytest.MonkeyPatch,
sqlite_tool_db: SqliteToolDb,
):
"""Raise ToolInvokeError when user resolution fails."""
tool = _build_tool()
monkeypatch.setattr(tool, "_get_app", lambda *args, **kwargs: None)
@@ -552,58 +642,45 @@ def test_invoke_raises_when_user_not_found(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(tool, "_resolve_user", lambda *args, **kwargs: None)
with pytest.raises(ToolInvokeError, match="User not found"):
list(tool.invoke(MagicMock(), "missing", {}))
list(tool.invoke(sqlite_tool_db.caller_session, "missing", {}))
def test_resolve_user_from_database_returns_account(monkeypatch: pytest.MonkeyPatch):
def test_resolve_user_from_database_returns_account(sqlite_tool_db: SqliteToolDb):
"""Resolve Account and set tenant in worker context."""
tenant = SimpleNamespace(id="tenant_id")
account = SimpleNamespace(id="account_id", current_tenant=None)
set_current_tenant = Mock(side_effect=lambda tenant, *, session: setattr(account, "current_tenant", tenant))
account.set_current_tenant_with_session = set_current_tenant
session = StubSession(scalar_results=[tenant, account])
tenant = _persist_tenant(sqlite_tool_db)
account = _persist_account(sqlite_tool_db)
tool = _build_tool(tenant_id=TENANT_ID)
monkeypatch.setattr("core.tools.workflow_as_tool.tool.session_factory.create_session", lambda: session)
tool = _build_tool()
tool.runtime.tenant_id = "tenant_id"
resolved = tool._resolve_user_from_database(user_id="account_id")
assert resolved is account
assert account.current_tenant is tenant
set_current_tenant.assert_called_once_with(tenant, session=session)
assert session.expunge_calls == [account]
resolved = tool._resolve_user_from_database(user_id=account.id)
assert isinstance(resolved, Account)
assert resolved.id == account.id
assert resolved.current_tenant_id == tenant.id
assert inspect(resolved).detached is True
def test_get_workflow_and_get_app_db_branches(monkeypatch: pytest.MonkeyPatch):
def test_get_workflow_and_get_app_db_branches(sqlite_tool_db: SqliteToolDb):
"""Cover workflow/app retrieval branches and error cases."""
tool = _build_tool()
latest_workflow = SimpleNamespace(id="wf-latest")
specific_workflow = SimpleNamespace(id="wf-v1")
app = SimpleNamespace(id="app-1")
sessions = iter(
[
StubSession(scalar_results=[], scalars_results=[latest_workflow]),
StubSession(scalar_results=[specific_workflow], scalars_results=[]),
StubSession(scalar_results=[app], scalars_results=[]),
]
)
monkeypatch.setattr(
"core.tools.workflow_as_tool.tool.session_factory.create_session",
lambda: next(sessions),
)
app = _persist_app(sqlite_tool_db)
specific_workflow = _persist_workflow(sqlite_tool_db, version="1")
latest_workflow = _persist_workflow(sqlite_tool_db, version="2")
_persist_workflow(sqlite_tool_db, version=Workflow.VERSION_DRAFT)
tool = _build_tool(tenant_id=TENANT_ID, workflow_app_id=APP_ID)
assert tool._get_workflow("app-1", "") is latest_workflow
assert tool._get_workflow("app-1", "1") is specific_workflow
assert tool._get_app("app-1") is app
latest = tool._get_workflow(APP_ID, "")
specific = tool._get_workflow(APP_ID, "1")
resolved_app = tool._get_app(APP_ID)
assert latest.id == latest_workflow.id
assert specific.id == specific_workflow.id
assert resolved_app.id == app.id
assert inspect(latest).detached is True
assert inspect(specific).detached is True
assert inspect(resolved_app).detached is True
monkeypatch.setattr(
"core.tools.workflow_as_tool.tool.session_factory.create_session",
lambda: StubSession(scalar_results=[None, None], scalars_results=[None]),
)
with pytest.raises(ValueError, match="workflow not found"):
tool._get_workflow("app-1", "1")
tool._get_workflow(APP_ID, "missing")
with pytest.raises(ValueError, match="app not found"):
tool._get_app("app-1")
tool._get_app("00000000-0000-0000-0000-000000000099")
def _setup_transform_args_tool(monkeypatch: pytest.MonkeyPatch) -> WorkflowTool:
@@ -722,7 +799,10 @@ def test_transform_args_normalizes_optional_files_parameter(
assert files == []
def test_workflow_tool_invocation_normalizes_optional_files_parameter(monkeypatch: pytest.MonkeyPatch):
def test_workflow_tool_invocation_normalizes_optional_files_parameter(
monkeypatch: pytest.MonkeyPatch,
sqlite_tool_db: SqliteToolDb,
):
"""Ensure casted empty FILES values do not reach workflow input validation as [None]."""
tool = _build_tool()
images_param = ToolParameter.get_simple_instance(
@@ -741,7 +821,7 @@ def test_workflow_tool_invocation_normalizes_optional_files_parameter(monkeypatc
generate_mock = MagicMock(return_value={"data": {}})
monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock)
list(tool.invoke(MagicMock(), "test_user", {"images": None}))
list(tool.invoke(sqlite_tool_db.caller_session, "test_user", {"images": None}))
call_kwargs = generate_mock.call_args.kwargs
assert call_kwargs["args"]["inputs"]["images"] == []
@@ -246,6 +246,26 @@ class TestPluginModelProviderCache:
call([cache_key]),
]
def test_fetch_plugin_model_providers_bypasses_redis_when_cache_disabled(self) -> None:
"""With the cache disabled the daemon is the only source, and Redis is never touched."""
with patch(f"{MODULE}.redis_client") as redis_client, patch(f"{MODULE}.dify_config") as config:
config.PLUGIN_MODEL_PROVIDERS_CACHE_ENABLED = False
client = Mock()
client.fetch_model_providers.return_value = [_build_plugin_model_provider()]
from core.plugin.plugin_service import PluginService
first = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1", client=client)
second = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1", client=client)
assert [provider.provider for provider in first] == ["langgenius/openai/openai"]
assert [provider.provider for provider in second] == ["langgenius/openai/openai"]
assert client.fetch_model_providers.call_count == 2
redis_client.get.assert_not_called()
redis_client.mget.assert_not_called()
redis_client.setex.assert_not_called()
redis_client.lock.assert_not_called()
def test_fetch_plugin_model_providers_refetches_when_cache_read_fails(self) -> None:
"""Redis read failures do not block provider discovery for the tenant."""
with patch(f"{MODULE}.redis_client") as redis_client:
@@ -135,6 +135,19 @@ class TestMCPToolInvoke:
values = {m.message.variable_name: m.message.variable_value for m in var_msgs}
assert values == {"a": 1, "b": "x"}
def test_invoke_yields_json_when_structured_content_has_no_output_schema(self, orm_session: Session) -> None:
tool = _make_mcp_tool()
result = CallToolResult(content=[], structuredContent={"a": 1, "b": "x"})
with patch.object(tool, "invoke_remote_mcp_tool", return_value=result):
messages = list(tool._invoke(session=orm_session, user_id="test_user", tool_parameters={}))
assert len(messages) == 1
msg = messages[0]
assert msg.type == ToolInvokeMessage.MessageType.JSON
assert isinstance(msg.message, ToolInvokeMessage.JsonMessage)
assert msg.message.json_object == {"a": 1, "b": "x"}
class TestMCPToolUsageExtraction:
"""Test usage metadata extraction from MCP tool results."""
Generated
+2 -2
View File
@@ -1281,7 +1281,7 @@ wheels = [
[[package]]
name = "dify-agent"
version = "1.16.0"
version = "1.16.1"
source = { editable = "../dify-agent" }
dependencies = [
{ name = "httpx" },
@@ -1331,7 +1331,7 @@ docs = [
[[package]]
name = "dify-api"
version = "1.16.0"
version = "1.16.1"
source = { virtual = "." }
dependencies = [
{ name = "aliyun-log-python-sdk" },
+1 -1
View File
@@ -71,7 +71,7 @@
"channel": "alpha",
"compat": {
"minDify": "1.16.0",
"maxDify": "1.16.0"
"maxDify": "1.16.1"
},
"release": {
"tagPrefix": "difyctl-v",
+8 -2
View File
@@ -114,9 +114,15 @@ function die(msg) {
process.exit(1)
}
// Tests point this at a fixture manifest so their assertions stay fixed while
// the real version and compat window move with every release. The name is
// mirrored in test/fixtures/pkg-manifest.ts rather than imported from here,
// because this file's shebang breaks the Windows test runner.
const PKG_PATH_ENV = 'DIFYCTL_PKG_PATH'
function loadPkg() {
const pkgUrl = new URL('../package.json', import.meta.url)
const pkg = JSON.parse(readFileSync(pkgUrl, 'utf8'))
const pkgPath = process.env[PKG_PATH_ENV] || new URL('../package.json', import.meta.url)
const pkg = JSON.parse(readFileSync(pkgPath, 'utf8'))
if (!pkg.difyctl?.release) die('cli/package.json missing difyctl.release')
return {
version: pkg.version,
+36 -16
View File
@@ -1,12 +1,19 @@
import { execFileSync } from 'node:child_process'
import { fileURLToPath } from 'node:url'
import { describe, expect, it } from 'vitest'
import { FIXTURE_COMPAT, pkgManifestEnv } from '../test/fixtures/pkg-manifest'
const SCRIPT = fileURLToPath(new URL('./release-naming.mjs', import.meta.url))
function run(args: string[]): { code: number; stdout: string; stderr: string } {
function run(
args: string[],
env: Record<string, string> = {},
): { code: number; stdout: string; stderr: string } {
try {
const stdout = execFileSync('node', [SCRIPT, ...args], { encoding: 'utf8' })
const stdout = execFileSync('node', [SCRIPT, ...args], {
encoding: 'utf8',
env: { ...process.env, ...env },
})
return { code: 0, stdout, stderr: '' }
} catch (e) {
const err = e as { status?: number; stdout?: string; stderr?: string }
@@ -14,45 +21,50 @@ function run(args: string[]): { code: number; stdout: string; stderr: string } {
}
}
describe('release-naming compat-check (compat 1.16.0..1.16.0)', () => {
describe('release-naming compat-check', () => {
const { minDify, maxDify } = FIXTURE_COMPAT // 2.0.0 .. 2.5.0
const pkgEnv = pkgManifestEnv()
const compatCheck = (difyVersion?: string) =>
run(difyVersion === undefined ? ['compat-check'] : ['compat-check', difyVersion], pkgEnv).code
it('accepts a version inside the window', () => {
expect(run(['compat-check', '1.16.0']).code).toBe(0)
expect(compatCheck('2.3.0')).toBe(0)
})
it('accepts the inclusive lower bound', () => {
expect(run(['compat-check', '1.16.0']).code).toBe(0)
expect(compatCheck(minDify)).toBe(0)
})
it('accepts the inclusive upper bound', () => {
expect(run(['compat-check', '1.16.0']).code).toBe(0)
expect(compatCheck(maxDify)).toBe(0)
})
it('accepts a v-prefixed tag', () => {
expect(run(['compat-check', 'v1.16.0']).code).toBe(0)
expect(compatCheck('v2.3.0')).toBe(0)
})
it('rejects a version below the lower bound', () => {
expect(run(['compat-check', '1.15.9']).code).not.toBe(0)
expect(compatCheck('1.9.9')).not.toBe(0)
})
it('rejects a version above the upper bound', () => {
expect(run(['compat-check', '1.16.1']).code).not.toBe(0)
expect(compatCheck('2.5.1')).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('treats a prerelease of the lower bound as below it', () => {
expect(compatCheck(`${minDify}-rc1`)).not.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 on the bound', () => {
expect(compatCheck(`${maxDify}+build123`)).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('ignores build metadata when out of range', () => {
expect(compatCheck('2.5.1+build123')).not.toBe(0)
})
it('requires a version argument', () => {
expect(run(['compat-check']).code).not.toBe(0)
expect(compatCheck()).not.toBe(0)
})
})
@@ -67,6 +79,14 @@ describe('release-naming github-env', () => {
for (const key of ['version', 'channel', 'prerelease', 'minDify', 'maxDify', 'tagPrefix'])
expect(stdout).toMatch(new RegExp(`^${key}=`, 'm'))
})
// The only assertion against the live manifest: the window must exist and be
// well-formed, whatever release it currently points at.
it('emits a well-formed compat window from the real cli/package.json', () => {
const { stdout } = run(['github-env'])
expect(stdout).toMatch(/^minDify=\d+\.\d+\.\d+$/m)
expect(stdout).toMatch(/^maxDify=\d+\.\d+\.\d+$/m)
})
})
describe('release-naming edge channel', () => {
+8 -2
View File
@@ -4,14 +4,20 @@ import { tmpdir } from 'node:os'
import { join } from 'node:path'
import { fileURLToPath } from 'node:url'
import { describe, expect, it } from 'vitest'
import { FIXTURE_COMPAT, pkgManifestEnv } from '../test/fixtures/pkg-manifest'
const SCRIPT = fileURLToPath(new URL('./release-r2-edge.mjs', import.meta.url))
const PKG_ENV = pkgManifestEnv()
function run(args: string[]): { code: number; stdout: string; stderr: string } {
try {
return {
code: 0,
stdout: execFileSync('node', [SCRIPT, ...args], { encoding: 'utf8' }),
stdout: execFileSync('node', [SCRIPT, ...args], {
encoding: 'utf8',
env: { ...process.env, ...PKG_ENV },
}),
stderr: '',
}
} catch (e) {
@@ -108,7 +114,7 @@ describe('release-r2-edge manifest', () => {
it('carries the compat window from package.json', () => {
const { json } = buildManifest()
expect(json.compat).toEqual({ minDify: '1.16.0', maxDify: '1.16.0' })
expect(json.compat).toEqual(FIXTURE_COMPAT)
})
it('lists all 5 targets with asset name + sha256 from the checksums file', () => {
+57
View File
@@ -0,0 +1,57 @@
import { mkdtempSync, writeFileSync } from 'node:fs'
import { tmpdir } from 'node:os'
import { join } from 'node:path'
// Mirrors PKG_PATH_ENV in scripts/release-naming.mjs, which cannot be imported
// here: its shebang breaks the Windows test runner. Divergence is self-
// reporting, not silent — the script would fall back to the real
// cli/package.json and every fixture-window assertion would fail.
const PKG_PATH_ENV = 'DIFYCTL_PKG_PATH'
// release-naming.mjs and release-r2-edge.mjs read their data from
// cli/package.json. Tests spawn them against this fixture instead, so
// assertions can name exact versions without tracking the live release.
// Deliberately far from any real Dify version, and min != max so "inside the
// window" is a case distinct from either bound.
export const FIXTURE_COMPAT = { minDify: '2.0.0', maxDify: '2.5.0' }
export const FIXTURE_TARGET_IDS = [
'linux-x64',
'linux-arm64',
'darwin-x64',
'darwin-arm64',
'windows-x64',
] as const
const FIXTURE_RELEASE = {
tagPrefix: 'difyctl-v',
binName: 'difyctl',
checksumsSuffix: '-checksums.txt',
targets: FIXTURE_TARGET_IDS.map((id) => ({
id,
bunTarget: `bun-${id}`,
exe: id.startsWith('windows'),
})),
}
export type PkgManifestOverrides = {
version?: string
channel?: string
compat?: { minDify: string; maxDify: string }
}
// Returns the env additions that point a spawned script at the fixture.
export function pkgManifestEnv(overrides: PkgManifestOverrides = {}): Record<string, string> {
const manifest = {
version: overrides.version ?? '0.2.0-alpha',
difyctl: {
channel: overrides.channel ?? 'alpha',
compat: overrides.compat ?? FIXTURE_COMPAT,
release: FIXTURE_RELEASE,
},
}
const path = join(mkdtempSync(join(tmpdir(), 'difyctl-pkg-')), 'package.json')
writeFileSync(path, JSON.stringify(manifest))
return { [PKG_PATH_ENV]: path }
}
+1 -1
View File
@@ -1,6 +1,6 @@
[project]
name = "dify-agent"
version = "1.16.0"
version = "1.16.1"
description = "Add your description here"
readme = "README.md"
requires-python = ">=3.12,<4.0"
+1 -1
View File
@@ -581,7 +581,7 @@ wheels = [
[[package]]
name = "dify-agent"
version = "1.16.0"
version = "1.16.1"
source = { editable = "." }
dependencies = [
{ name = "httpx" },
+7 -7
View File
@@ -220,7 +220,7 @@ services:
# API service
api:
<<: *shared-api-worker-config
image: langgenius/dify-api:1.16.0
image: langgenius/dify-api:1.16.1
environment:
MODE: api
SENTRY_DSN: ${API_SENTRY_DSN:-}
@@ -271,7 +271,7 @@ services:
# WebSocket service for workflow collaboration.
api_websocket:
<<: *shared-api-worker-config
image: langgenius/dify-api:1.16.0
image: langgenius/dify-api:1.16.1
profiles:
- collaboration
environment:
@@ -297,7 +297,7 @@ services:
# The Celery worker for processing all queues (dataset, workflow, mail, etc.)
worker:
<<: *shared-worker-config
image: langgenius/dify-api:1.16.0
image: langgenius/dify-api:1.16.1
environment:
MODE: worker
SENTRY_DSN: ${API_SENTRY_DSN:-}
@@ -347,7 +347,7 @@ services:
# Celery beat for scheduling periodic tasks.
worker_beat:
<<: *shared-worker-beat-config
image: langgenius/dify-api:1.16.0
image: langgenius/dify-api:1.16.1
environment:
MODE: beat
depends_on:
@@ -380,7 +380,7 @@ services:
# Frontend web application.
web:
image: langgenius/dify-web:1.16.0
image: langgenius/dify-web:1.16.1
restart: always
env_file:
- path: ./envs/core-services/web.env
@@ -542,7 +542,7 @@ services:
# on port 3128, which only allows agent_backend /agent-stub/ and the Dify API
# /files/* endpoints (see ssrf_proxy/squid-agent.conf.template).
local_sandbox:
image: langgenius/dify-agent-local-sandbox:1.16.0
image: langgenius/dify-agent-local-sandbox:1.16.1
restart: always
env_file:
- path: ./envs/core-services/local-sandbox.env
@@ -651,7 +651,7 @@ services:
# Dify Agent backend service.
agent_backend:
image: langgenius/dify-agent-backend:1.16.0
image: langgenius/dify-agent-backend:1.16.1
restart: always
env_file:
- path: ./envs/core-services/dify-agent.env
+7 -7
View File
@@ -226,7 +226,7 @@ services:
# API service
api:
<<: *shared-api-worker-config
image: langgenius/dify-api:1.16.0
image: langgenius/dify-api:1.16.1
environment:
MODE: api
SENTRY_DSN: ${API_SENTRY_DSN:-}
@@ -277,7 +277,7 @@ services:
# WebSocket service for workflow collaboration.
api_websocket:
<<: *shared-api-worker-config
image: langgenius/dify-api:1.16.0
image: langgenius/dify-api:1.16.1
profiles:
- collaboration
environment:
@@ -303,7 +303,7 @@ services:
# The Celery worker for processing all queues (dataset, workflow, mail, etc.)
worker:
<<: *shared-worker-config
image: langgenius/dify-api:1.16.0
image: langgenius/dify-api:1.16.1
environment:
MODE: worker
SENTRY_DSN: ${API_SENTRY_DSN:-}
@@ -353,7 +353,7 @@ services:
# Celery beat for scheduling periodic tasks.
worker_beat:
<<: *shared-worker-beat-config
image: langgenius/dify-api:1.16.0
image: langgenius/dify-api:1.16.1
environment:
MODE: beat
depends_on:
@@ -386,7 +386,7 @@ services:
# Frontend web application.
web:
image: langgenius/dify-web:1.16.0
image: langgenius/dify-web:1.16.1
restart: always
env_file:
- path: ./envs/core-services/web.env
@@ -548,7 +548,7 @@ services:
# on port 3128, which only allows agent_backend /agent-stub/ and the Dify API
# /files/* endpoints (see ssrf_proxy/squid-agent.conf.template).
local_sandbox:
image: langgenius/dify-agent-local-sandbox:1.16.0
image: langgenius/dify-agent-local-sandbox:1.16.1
restart: always
env_file:
- path: ./envs/core-services/local-sandbox.env
@@ -657,7 +657,7 @@ services:
# Dify Agent backend service.
agent_backend:
image: langgenius/dify-agent-backend:1.16.0
image: langgenius/dify-agent-backend:1.16.1
restart: always
env_file:
- path: ./envs/core-services/dify-agent.env
@@ -73,6 +73,7 @@ SSRF_PROXY_HTTPS_URL=http://ssrf_proxy:3128
PGDATA=/var/lib/postgresql/data/pgdata
PLUGIN_MAX_PACKAGE_SIZE=52428800
PLUGIN_MODEL_SCHEMA_CACHE_TTL=3600
PLUGIN_MODEL_PROVIDERS_CACHE_ENABLED=true
PLUGIN_MODEL_PROVIDERS_CACHE_TTL=86400
# Comma-separated marketplace plugin IDs whose latest versions are installed for newly registered users.
# Example: langgenius/openai,langgenius/gemini
+11
View File
@@ -31,6 +31,17 @@ const createApiResponse = ({
status: () => status,
statusText: () => statusText,
text: async () => body,
timing: () => ({
connectEnd: -1,
connectStart: -1,
domainLookupEnd: -1,
domainLookupStart: -1,
requestStart: -1,
responseEnd: -1,
responseStart: -1,
secureConnectionStart: -1,
startTime: -1,
}),
url: () => url,
[Symbol.asyncDispose]: async () => {},
}
-29
View File
@@ -4,26 +4,11 @@
"count": 2
}
},
"cli/test/e2e/helpers/retry.ts": {
"no-throw-literal": {
"count": 1
}
},
"cli/test/e2e/suites/output/json-yaml-output.e2e.ts": {
"prefer-const": {
"count": 1
}
},
"e2e/support/web-server.ts": {
"no-throw-literal": {
"count": 1
}
},
"packages/dev-proxy/src/cli.spec.ts": {
"no-throw-literal": {
"count": 1
}
},
"packages/migrate-no-unchecked-indexed-access/src/no-unchecked-indexed-access/migrate.ts": {
"no-console": {
"count": 11
@@ -2141,11 +2126,6 @@
"count": 1
}
},
"web/app/components/base/voice-input/recorder.ts": {
"no-throw-literal": {
"count": 1
}
},
"web/app/components/billing/plan/assets/index.tsx": {
"no-barrel-files/no-barrel-files": {
"count": 4
@@ -3225,9 +3205,6 @@
}
},
"web/app/components/plugins/marketplace/hooks.ts": {
"@tanstack/query/exhaustive-deps": {
"count": 1
},
"no-restricted-imports": {
"count": 1
}
@@ -6349,9 +6326,6 @@
}
},
"web/service/use-pipeline.ts": {
"@tanstack/query/exhaustive-deps": {
"count": 1
},
"no-restricted-imports": {
"count": 1
}
@@ -6370,9 +6344,6 @@
}
},
"web/service/use-workflow.ts": {
"@tanstack/query/exhaustive-deps": {
"count": 1
},
"no-restricted-imports": {
"count": 1
},
+1 -1
View File
@@ -56,5 +56,5 @@
"engines": {
"node": "^22.22.1"
},
"packageManager": "pnpm@11.15.0"
"packageManager": "pnpm@11.17.0"
}
+1 -1
View File
@@ -60,7 +60,7 @@ const comboboxTriggerVariants = cva(
[
'group/combobox-trigger flex w-full min-w-0 items-center border-0 bg-components-input-bg-normal text-start text-components-input-text-filled outline-hidden transition-colors',
'hover:bg-state-base-hover-alt focus-visible:bg-state-base-hover-alt data-popup-open:bg-state-base-hover-alt',
'focus-visible:inset-ring-1 focus-visible:inset-ring-components-input-border-active',
'focus-visible:ring-2 focus-visible:ring-state-accent-solid',
'data-placeholder:text-components-input-text-placeholder',
'data-readonly:cursor-default data-readonly:bg-transparent data-readonly:hover:bg-transparent',
'data-disabled:cursor-not-allowed data-disabled:bg-components-input-bg-disabled data-disabled:text-components-input-text-filled-disabled data-disabled:hover:bg-components-input-bg-disabled',
+1989 -1585
View File
File diff suppressed because it is too large Load Diff
+58 -58
View File
@@ -39,32 +39,32 @@ overrides:
postcss-selector-parser@>=6.0.0 <6.1.3: 6.1.4
postcss-selector-parser@>=7.0.0 <7.1.3: 7.1.4
postcss@<8.5.10: ^8.5.10
rollup@>=4.0.0 <4.59.0: 4.62.2
rollup@>=4.0.0 <4.59.0: 4.62.3
safer-buffer: npm:@nolyfill/safer-buffer@^1.0.44
side-channel: npm:@nolyfill/side-channel@^1.0.44
solid-js: 1.9.14
string-width: ~8.2.2
tar@<=7.5.15: ^7.5.16
vite: npm:@voidzero-dev/vite-plus-core@0.2.5
vite: npm:@voidzero-dev/vite-plus-core@0.2.6
ws@>=8.0.0 <8.20.1: ^8.21.1
yaml@>=2.0.0 <2.8.3: 2.9.0
yauzl@<3.2.1: 3.2.1
catalog:
'@amplitude/analytics-browser': 2.45.2
'@amplitude/plugin-session-replay-browser': 1.33.4
'@amplitude/analytics-browser': 2.45.4
'@amplitude/plugin-session-replay-browser': 1.33.6
'@base-ui/react': 1.6.0
'@chromatic-com/storybook': 5.2.1
'@cucumber/cucumber': 13.1.1
'@cucumber/cucumber': 13.2.0
'@egoist/tailwindcss-icons': 1.9.2
'@emoji-mart/data': 1.2.1
'@eslint-community/eslint-plugin-eslint-comments': 4.7.2
'@eslint-react/eslint-plugin': 5.17.1
'@eslint-react/eslint-plugin': 5.18.0
'@eslint/markdown': 8.0.3
'@floating-ui/react': 0.27.20
'@formatjs/intl-localematcher': 0.8.13
'@heroicons/react': 2.2.0
'@hey-api/openapi-ts': 0.98.2
'@hono/node-server': 2.0.10
'@hono/node-server': 2.0.12
'@iconify-json/heroicons': 1.2.3
'@iconify-json/ri': 1.2.10
'@lexical/code': 0.47.0
@@ -77,41 +77,41 @@ catalog:
'@mdx-js/loader': 3.1.1
'@mdx-js/react': 3.1.1
'@mdx-js/rollup': 3.1.1
'@mediabunny/mp3-encoder': 1.50.9
'@mediabunny/mp3-encoder': 1.51.0
'@monaco-editor/react': 4.7.0
'@napi-rs/keyring': 1.3.0
'@next/mdx': 16.2.11
'@orpc/client': 1.14.8
'@orpc/contract': 1.14.8
'@orpc/openapi-client': 1.14.8
'@orpc/tanstack-query': 1.14.8
'@playwright/test': 1.61.1
'@next/mdx': 16.2.12
'@orpc/client': 1.14.10
'@orpc/contract': 1.14.10
'@orpc/openapi-client': 1.14.10
'@orpc/tanstack-query': 1.14.10
'@playwright/test': 1.62.0
'@remixicon/react': 4.9.0
'@rgrove/parse-xml': 4.2.2
'@sentry/react': 10.66.0
'@storybook/addon-a11y': 10.5.2
'@storybook/addon-docs': 10.5.2
'@storybook/addon-links': 10.5.2
'@storybook/addon-onboarding': 10.5.2
'@storybook/addon-themes': 10.5.2
'@storybook/addon-vitest': 10.5.2
'@storybook/nextjs-vite': 10.5.2
'@storybook/react': 10.5.2
'@storybook/react-vite': 10.5.2
'@rgrove/parse-xml': 4.2.3
'@sentry/react': 10.68.0
'@storybook/addon-a11y': 10.5.4
'@storybook/addon-docs': 10.5.4
'@storybook/addon-links': 10.5.4
'@storybook/addon-onboarding': 10.5.4
'@storybook/addon-themes': 10.5.4
'@storybook/addon-vitest': 10.5.4
'@storybook/nextjs-vite': 10.5.4
'@storybook/react': 10.5.4
'@storybook/react-vite': 10.5.4
'@streamdown/math': 1.0.2
'@svgdotjs/svg.js': 3.2.6
'@svgdotjs/svg.js': 3.2.7
'@t3-oss/env-core': 0.13.11
'@t3-oss/env-nextjs': 0.13.11
'@tailwindcss/postcss': 4.3.3
'@tailwindcss/typography': 0.5.20
'@tailwindcss/vite': 4.3.3
'@tanstack/eslint-plugin-query': 5.101.2
'@tanstack/eslint-plugin-query': 5.101.4
'@tanstack/form-core': 1.33.2
'@tanstack/query-core': 5.101.2
'@tanstack/query-core': 5.101.4
'@tanstack/react-form': 1.33.2
'@tanstack/react-hotkeys': 0.10.0
'@tanstack/react-query': 5.101.2
'@tanstack/react-virtual': 3.14.6
'@tanstack/react-query': 5.101.4
'@tanstack/react-virtual': 3.14.8
'@testing-library/dom': 10.4.1
'@testing-library/jest-dom': 6.9.1
'@testing-library/react': 16.3.2
@@ -127,9 +127,9 @@ catalog:
'@types/react': 19.2.17
'@types/react-dom': 19.2.3
'@types/sortablejs': 1.15.9
'@typescript-eslint/parser': 8.64.0
'@typescript-eslint/parser': 8.65.0
'@typescript/native': npm:typescript@7.0.2
'@vitejs/plugin-react': 6.0.3
'@vitejs/plugin-react': 6.0.4
'@vitejs/plugin-rsc': 0.5.30
'@vitest/browser': 4.1.10
'@vitest/browser-playwright': 4.1.10
@@ -143,7 +143,7 @@ catalog:
cli-table3: 0.6.5
clsx: 2.1.1
code-inspector-plugin: 1.6.6
concurrently: ^10.0.3
concurrently: 10.0.4
copy-to-clipboard: 4.0.2
cron-parser: 5.6.2
dayjs: 1.11.21
@@ -156,32 +156,32 @@ catalog:
embla-carousel-fade: 8.6.0
embla-carousel-react: 8.6.0
emoji-mart: 5.6.0
es-toolkit: 1.49.0
eslint: 10.7.0
es-toolkit: 1.50.0
eslint: 10.8.0
eslint-markdown: 0.12.1
eslint-plugin-antfu: 3.2.3
eslint-plugin-command: 3.5.3
eslint-plugin-erasable-syntax-only: 0.4.2
eslint-plugin-hyoban: 0.14.1
eslint-plugin-jsdoc: 63.1.0
eslint-plugin-jsdoc: 63.3.1
eslint-plugin-jsonc: 3.3.0
eslint-plugin-markdown-preferences: 0.41.1
eslint-plugin-n: 18.2.2
eslint-plugin-no-barrel-files: 1.3.1
eslint-plugin-perfectionist: 5.10.0
eslint-plugin-pnpm: 1.6.1
eslint-plugin-pnpm: 1.7.0
eslint-plugin-regexp: 3.1.1
eslint-plugin-storybook: 10.5.2
eslint-plugin-toml: 1.4.0
eslint-plugin-storybook: 10.5.4
eslint-plugin-toml: 1.5.0
eslint-plugin-unicorn: 71.1.0
eslint-plugin-yml: 3.6.0
eventsource-parser: 3.1.0
fast-deep-equal: 3.1.3
foxact: 0.3.8
fuse.js: 7.5.0
happy-dom: 20.11.0
happy-dom: 20.11.1
hast-util-to-jsx-runtime: 2.3.6
hono: 4.12.31
hono: 4.12.32
html-entities: 2.6.0
html-to-image: 1.11.13
i18next: 26.3.6
@@ -189,39 +189,39 @@ catalog:
iconify-import-svg: 0.2.0
immer: 11.1.15
jotai: 2.20.2
jotai-effect: 2.3.1
jotai-effect: 2.4.1
jotai-scope: 0.11.0
jotai-tanstack-query: 0.11.0
js-cookie: 3.0.8
js-yaml: 5.2.1
js-yaml: 5.2.2
jsonschema: 1.5.0
katex: 0.17.0
knip: 6.27.0
knip: 6.29.0
ky: 2.0.2
lexical: 0.47.0
lockfile: 1.0.4
loro-crdt: 1.13.7
mediabunny: 1.50.9
loro-crdt: 1.13.8
mediabunny: 1.51.0
mermaid: 11.16.0
mime: 4.1.0
mitt: 3.0.1
motion: 12.42.2
negotiator: 1.0.0
next: 16.2.11
next: 16.2.12
next-themes: 0.4.6
nuqs: 2.9.1
nuqs: 2.9.2
open: 11.0.0
ora: 9.4.1
picocolors: 1.1.1
pinyin-pro: 3.28.1
playwright: 1.61.1
postcss: 8.5.19
pinyin-pro: 3.28.2
playwright: 1.62.0
postcss: 8.5.23
qrcode.react: 4.2.0
qs: 6.15.3
react: 19.2.8
react-dom: 19.2.8
react-easy-crop: 6.2.2
react-i18next: 17.0.10
react-easy-crop: 6.2.3
react-i18next: 17.0.11
react-papaparse: 4.4.0
react-pdf-highlighter: 8.0.0-rc.0
react-server-dom-webpack: 19.2.8
@@ -237,7 +237,7 @@ catalog:
socket.io-client: 4.8.3
sortablejs: 1.15.7
std-semver: 1.0.8
storybook: 10.5.2
storybook: 10.5.4
streamdown: 2.5.0
string-ts: 2.3.1
tailwind-merge: 3.6.0
@@ -246,14 +246,14 @@ catalog:
tsx: 4.23.1
typescript: npm:@typescript/typescript6@6.0.2
uglify-js: 3.19.3
undici: 7.28.0
undici: 7.29.0
unist-util-visit: 5.1.0
use-context-selector: 2.0.0
uuid: 14.0.1
vinext: 1.0.0-beta.2
vite: npm:@voidzero-dev/vite-plus-core@0.2.5
vinext: 1.0.0-beta.4
vite: npm:@voidzero-dev/vite-plus-core@0.2.6
vite-plugin-inspect: 12.0.2
vite-plus: 0.2.5
vite-plus: 0.2.6
vitest: 4.1.10
vitest-browser-react: 2.2.0
vitest-canvas-mock: 1.1.4
@@ -388,12 +388,10 @@ const ProviderConfigModal: FC<Props> = ({
isRequired
value={(config as ArizeConfig).api_key}
onChange={handleConfigChange('api_key')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: 'API Key',
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: 'API Key',
})!}
/>
<Field
label="Space ID"
@@ -401,12 +399,10 @@ const ProviderConfigModal: FC<Props> = ({
isRequired
value={(config as ArizeConfig).space_id}
onChange={handleConfigChange('space_id')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: 'Space ID',
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: 'Space ID',
})!}
/>
<Field
label={t(($) => $[`${I18N_PREFIX}.project`], { ns: 'app' })!}
@@ -414,12 +410,10 @@ const ProviderConfigModal: FC<Props> = ({
isRequired
value={(config as ArizeConfig).project}
onChange={handleConfigChange('project')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.project`], { ns: 'app' }),
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.project`], { ns: 'app' }),
})!}
/>
<Field
label="Endpoint"
@@ -438,12 +432,10 @@ const ProviderConfigModal: FC<Props> = ({
isRequired
value={(config as PhoenixConfig).api_key}
onChange={handleConfigChange('api_key')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: 'API Key',
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: 'API Key',
})!}
/>
<Field
label={t(($) => $[`${I18N_PREFIX}.project`], { ns: 'app' })!}
@@ -451,12 +443,10 @@ const ProviderConfigModal: FC<Props> = ({
isRequired
value={(config as PhoenixConfig).project}
onChange={handleConfigChange('project')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.project`], { ns: 'app' }),
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.project`], { ns: 'app' }),
})!}
/>
<Field
label="Endpoint"
@@ -475,12 +465,10 @@ const ProviderConfigModal: FC<Props> = ({
isRequired
value={(config as AliyunConfig).license_key}
onChange={handleConfigChange('license_key')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: 'License Key',
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: 'License Key',
})!}
/>
<Field
label="Endpoint"
@@ -505,9 +493,10 @@ const ProviderConfigModal: FC<Props> = ({
isRequired
value={(config as TencentConfig).token}
onChange={handleConfigChange('token')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], { ns: 'app', key: 'Token' })!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: 'Token',
})!}
/>
<Field
label="Endpoint"
@@ -535,12 +524,10 @@ const ProviderConfigModal: FC<Props> = ({
isRequired
value={(config as WeaveConfig).api_key}
onChange={handleConfigChange('api_key')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: 'API Key',
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: 'API Key',
})!}
/>
<Field
label={t(($) => $[`${I18N_PREFIX}.project`], { ns: 'app' })!}
@@ -548,21 +535,20 @@ const ProviderConfigModal: FC<Props> = ({
isRequired
value={(config as WeaveConfig).project}
onChange={handleConfigChange('project')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.project`], { ns: 'app' }),
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.project`], { ns: 'app' }),
})!}
/>
<Field
label="Entity"
labelClassName="text-sm!"
value={(config as WeaveConfig).entity}
onChange={handleConfigChange('entity')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], { ns: 'app', key: 'Entity' })!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: 'Entity',
})!}
/>
<Field
label="Endpoint"
@@ -588,12 +574,10 @@ const ProviderConfigModal: FC<Props> = ({
isRequired
value={(config as LangSmithConfig).api_key}
onChange={handleConfigChange('api_key')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: 'API Key',
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: 'API Key',
})!}
/>
<Field
label={t(($) => $[`${I18N_PREFIX}.project`], { ns: 'app' })!}
@@ -601,12 +585,10 @@ const ProviderConfigModal: FC<Props> = ({
isRequired
value={(config as LangSmithConfig).project}
onChange={handleConfigChange('project')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.project`], { ns: 'app' }),
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.project`], { ns: 'app' }),
})!}
/>
<Field
label="Endpoint"
@@ -625,12 +607,10 @@ const ProviderConfigModal: FC<Props> = ({
value={(config as LangFuseConfig).secret_key}
isRequired
onChange={handleConfigChange('secret_key')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.secretKey`], { ns: 'app' }),
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.secretKey`], { ns: 'app' }),
})!}
/>
<Field
label={t(($) => $[`${I18N_PREFIX}.publicKey`], { ns: 'app' })!}
@@ -638,12 +618,10 @@ const ProviderConfigModal: FC<Props> = ({
isRequired
value={(config as LangFuseConfig).public_key}
onChange={handleConfigChange('public_key')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.publicKey`], { ns: 'app' }),
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.publicKey`], { ns: 'app' }),
})!}
/>
<Field
label="Host"
@@ -662,24 +640,20 @@ const ProviderConfigModal: FC<Props> = ({
labelClassName="text-sm!"
value={(config as OpikConfig).api_key}
onChange={handleConfigChange('api_key')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: 'API Key',
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: 'API Key',
})!}
/>
<Field
label={t(($) => $[`${I18N_PREFIX}.project`], { ns: 'app' })!}
labelClassName="text-sm!"
value={(config as OpikConfig).project}
onChange={handleConfigChange('project')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.project`], { ns: 'app' }),
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.project`], { ns: 'app' }),
})!}
/>
<Field
label="Workspace"
@@ -713,36 +687,30 @@ const ProviderConfigModal: FC<Props> = ({
isRequired
value={(config as MLflowConfig).experiment_id}
onChange={handleConfigChange('experiment_id')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.experimentId`], { ns: 'app' }),
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.experimentId`], { ns: 'app' }),
})!}
/>
<Field
label={t(($) => $[`${I18N_PREFIX}.username`], { ns: 'app' })!}
labelClassName="text-sm!"
value={(config as MLflowConfig).username}
onChange={handleConfigChange('username')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.username`], { ns: 'app' }),
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.username`], { ns: 'app' }),
})!}
/>
<Field
label={t(($) => $[`${I18N_PREFIX}.password`], { ns: 'app' })!}
labelClassName="text-sm!"
value={(config as MLflowConfig).password}
onChange={handleConfigChange('password')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.password`], { ns: 'app' }),
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.password`], { ns: 'app' }),
})!}
/>
</>
)}
@@ -753,12 +721,10 @@ const ProviderConfigModal: FC<Props> = ({
labelClassName="text-sm!"
value={(config as DatabricksConfig).experiment_id}
onChange={handleConfigChange('experiment_id')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.experimentId`], { ns: 'app' }),
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.experimentId`], { ns: 'app' }),
})!}
isRequired
/>
<Field
@@ -766,12 +732,10 @@ const ProviderConfigModal: FC<Props> = ({
labelClassName="text-sm!"
value={(config as DatabricksConfig).host}
onChange={handleConfigChange('host')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.databricksHost`], { ns: 'app' }),
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.databricksHost`], { ns: 'app' }),
})!}
isRequired
/>
<Field
@@ -779,36 +743,30 @@ const ProviderConfigModal: FC<Props> = ({
labelClassName="text-sm!"
value={(config as DatabricksConfig).client_id}
onChange={handleConfigChange('client_id')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.clientId`], { ns: 'app' }),
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.clientId`], { ns: 'app' }),
})!}
/>
<Field
label={t(($) => $[`${I18N_PREFIX}.clientSecret`], { ns: 'app' })!}
labelClassName="text-sm!"
value={(config as DatabricksConfig).client_secret}
onChange={handleConfigChange('client_secret')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.clientSecret`], { ns: 'app' }),
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.clientSecret`], { ns: 'app' }),
})!}
/>
<Field
label={t(($) => $[`${I18N_PREFIX}.personalAccessToken`], { ns: 'app' })!}
labelClassName="text-sm!"
value={(config as DatabricksConfig).personal_access_token}
onChange={handleConfigChange('personal_access_token')}
placeholder={
t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.personalAccessToken`], { ns: 'app' }),
})!
}
placeholder={t(($) => $[`${I18N_PREFIX}.placeholder`], {
ns: 'app',
key: t(($) => $[`${I18N_PREFIX}.personalAccessToken`], { ns: 'app' }),
})!}
/>
</>
)}
@@ -884,12 +842,10 @@ const ProviderConfigModal: FC<Props> = ({
<AlertDialogContent>
<div className="flex flex-col gap-2 px-6 pt-6 pb-4">
<AlertDialogTitle className="w-full truncate title-2xl-semi-bold text-text-primary">
{
t(($) => $[`${I18N_PREFIX}.removeConfirmTitle`], {
ns: 'app',
key: t(($) => $[`tracing.${type}.title`], { ns: 'app' }),
})!
}
{t(($) => $[`${I18N_PREFIX}.removeConfirmTitle`], {
ns: 'app',
key: t(($) => $[`tracing.${type}.title`], { ns: 'app' }),
})!}
</AlertDialogTitle>
<AlertDialogDescription className="w-full system-md-regular wrap-break-word whitespace-pre-wrap text-text-tertiary">
{t(($) => $[`${I18N_PREFIX}.removeConfirmContent`], { ns: 'app' })}
@@ -118,7 +118,7 @@ export default function AddMemberOrGroupDialog() {
aria-label={t(($) => $['operation.add'], { ns: 'common' })}
icon={false}
size="small"
className="h-6 w-auto min-w-[52px] shrink-0 rounded-md border-0 bg-transparent px-2 py-0 text-xs font-medium text-components-button-secondary-accent-text hover:bg-state-accent-hover focus-visible:bg-state-accent-hover focus-visible:ring-2 focus-visible:ring-state-accent-solid data-popup-open:bg-state-accent-hover"
className="h-6 w-auto min-w-[52px] shrink-0 rounded-md border-0 bg-transparent px-2 py-0 text-xs font-medium text-components-button-secondary-accent-text hover:bg-state-accent-hover focus-visible:bg-state-accent-hover data-popup-open:bg-state-accent-hover"
>
<span className="inline-flex min-w-0 items-center justify-center gap-x-0.5 whitespace-nowrap">
<span className="i-ri-add-circle-fill size-4 shrink-0" aria-hidden="true" />
@@ -145,7 +145,7 @@ export function DocumentPicker({ datasetId, value, parentMode, onChange }: Props
aria-label={value?.name || t(($) => $['operation.search'], { ns: 'common' })}
icon={false}
className={cn(
'ml-1 flex size-auto rounded-lg border-0 bg-transparent px-2 py-1 hover:bg-state-base-hover focus-visible:bg-state-base-hover focus-visible:ring-1 focus-visible:ring-components-input-border-active data-popup-open:bg-state-base-hover',
'ml-1 flex size-auto rounded-lg border-0 bg-transparent px-2 py-1 hover:bg-state-base-hover focus-visible:bg-state-base-hover data-popup-open:bg-state-base-hover',
)}
>
<ComboboxValue>
@@ -143,9 +143,9 @@ export const ParentChildOptions: FC<ParentChildOptionsProps> = ({
<div className="flex gap-3">
<DelimiterInput
value={parentChildConfig.parent.delimiter}
tooltip={
t(($) => $['stepTwo.parentChildDelimiterTip'], { ns: 'datasetCreation' })!
}
tooltip={t(($) => $['stepTwo.parentChildDelimiterTip'], {
ns: 'datasetCreation',
})!}
onChange={(e) => onParentDelimiterChange(e.target.value)}
/>
<MaxLengthInput
@@ -178,9 +178,9 @@ export const ParentChildOptions: FC<ParentChildOptionsProps> = ({
<div className="mt-1 flex gap-3">
<DelimiterInput
value={parentChildConfig.child.delimiter}
tooltip={
t(($) => $['stepTwo.parentChildChunkDelimiterTip'], { ns: 'datasetCreation' })!
}
tooltip={t(($) => $['stepTwo.parentChildChunkDelimiterTip'], {
ns: 'datasetCreation',
})!}
onChange={(e) => onChildDelimiterChange(e.target.value)}
/>
<MaxLengthInput
@@ -144,7 +144,7 @@ function ModelSelector({
<ComboboxTrigger
aria-label={t(($) => $['detailPanel.configureModel'], { ns: 'plugin' })}
icon={false}
className="block h-auto w-full border-0 bg-transparent p-0 text-left hover:bg-transparent focus-visible:bg-transparent focus-visible:ring-0 data-popup-open:bg-transparent"
className="block h-auto w-full border-0 bg-transparent p-0 text-left hover:bg-transparent focus-visible:bg-transparent data-popup-open:bg-transparent"
disabled={readonly}
>
<ModelSelectorTrigger
@@ -655,6 +655,21 @@ describe('MainNav', () => {
expect(helpButton.parentElement).toHaveClass('shrink-0', 'rounded-full', 'p-1')
})
it('orders the Step-by-step Tour before the account and help actions', async () => {
localStorage.setItem(STEP_BY_STEP_TOUR_SHELL_MODE_STORAGE_KEY, 'collapsed')
renderMainNav()
const tourTrigger = await screen.findByRole('button', { name: 'Open step-by-step tour' })
const accountButton = screen.getByRole('button', { name: 'common.account.account' })
const helpButton = screen.getByRole('button', { name: 'common.mainNav.help.openMenu' })
expect(tourTrigger.compareDocumentPosition(accountButton)).toBe(
Node.DOCUMENT_POSITION_FOLLOWING,
)
expect(accountButton.compareDocumentPosition(helpButton)).toBe(Node.DOCUMENT_POSITION_FOLLOWING)
})
it('keeps the global navigation account section expanded on home routes', () => {
localStorage.setItem(DETAIL_SIDEBAR_STORAGE_KEY, 'collapse')
mockPathname = '/'
+1 -1
View File
@@ -133,6 +133,7 @@ export function MainNav({ className }: MainNavProps) {
)}
</div>
<div className="relative w-60 shrink-0">
<StepByStepTourMount className="absolute -top-7 left-2.5 h-8 w-[183px] overflow-visible" />
<div className="flex w-60 items-center justify-between bg-gradient-to-b from-background-body-transparent to-background-body to-50% py-3 pr-1 pl-3 backdrop-blur-[2px]">
<div className="flex min-w-0 items-center gap-1 overflow-hidden">
<AccountSection />
@@ -141,7 +142,6 @@ export function MainNav({ className }: MainNavProps) {
<HelpMenu />
</div>
</div>
<StepByStepTourMount className="absolute -top-7 left-2.5 h-8 w-[183px] overflow-visible" />
</div>
</aside>
)
@@ -130,7 +130,7 @@ export function AppPicker({
<ComboboxTrigger
aria-label={t(($) => $['appSelector.label'], { ns: 'app' })}
icon={false}
className="block h-auto w-full border-0 bg-transparent p-0 text-left hover:bg-transparent focus-visible:bg-transparent focus-visible:ring-0 data-popup-open:bg-transparent"
className="block h-auto w-full border-0 bg-transparent p-0 text-left hover:bg-transparent focus-visible:bg-transparent data-popup-open:bg-transparent"
>
{trigger}
</ComboboxTrigger>
@@ -217,11 +217,9 @@ export default function ConfigCredential({ positionCenter, credential, onChange,
onChange={(e) =>
setTempCredential({ ...tempCredential, api_key_header: e.target.value })
}
placeholder={
t(($) => $['createTool.authMethod.types.apiKeyPlaceholder'], {
ns: 'tools',
})!
}
placeholder={t(($) => $['createTool.authMethod.types.apiKeyPlaceholder'], {
ns: 'tools',
})!}
/>
</div>
<div>
@@ -233,11 +231,12 @@ export default function ConfigCredential({ positionCenter, credential, onChange,
onChange={(e) =>
setTempCredential({ ...tempCredential, api_key_value: e.target.value })
}
placeholder={
t(($) => $['createTool.authMethod.types.apiValuePlaceholder'], {
placeholder={t(
($) => $['createTool.authMethod.types.apiValuePlaceholder'],
{
ns: 'tools',
})!
}
},
)!}
/>
</div>
</>
@@ -265,11 +264,12 @@ export default function ConfigCredential({ positionCenter, credential, onChange,
api_key_query_param: e.target.value,
})
}
placeholder={
t(($) => $['createTool.authMethod.types.queryParamPlaceholder'], {
placeholder={t(
($) => $['createTool.authMethod.types.queryParamPlaceholder'],
{
ns: 'tools',
})!
}
},
)!}
/>
</div>
<div>
@@ -281,11 +281,12 @@ export default function ConfigCredential({ positionCenter, credential, onChange,
onChange={(e) =>
setTempCredential({ ...tempCredential, api_key_value: e.target.value })
}
placeholder={
t(($) => $['createTool.authMethod.types.apiValuePlaceholder'], {
placeholder={t(
($) => $['createTool.authMethod.types.apiValuePlaceholder'],
{
ns: 'tools',
})!
}
},
)!}
/>
</div>
</>
@@ -256,9 +256,9 @@ const EditCustomCollectionModal: FC<Props> = ({
/>
<Input
className="h-10 grow"
placeholder={
t(($) => $['createTool.toolNamePlaceHolder'], { ns: 'tools' })!
}
placeholder={t(($) => $['createTool.toolNamePlaceHolder'], {
ns: 'tools',
})!}
value={customCollection.provider}
onChange={(e) => {
const newCollection = produce(customCollection, (draft) => {
@@ -298,9 +298,9 @@ const EditCustomCollectionModal: FC<Props> = ({
className="h-[240px] resize-none"
value={schema}
onValueChange={(value) => setSchema(value)}
placeholder={
t(($) => $['createTool.schemaPlaceHolder'], { ns: 'tools' })!
}
placeholder={t(($) => $['createTool.schemaPlaceHolder'], {
ns: 'tools',
})!}
/>
</div>
@@ -355,11 +355,12 @@ export function WorkflowToolDrawer({
<input
type="text"
className="w-full appearance-none bg-transparent text-[13px] leading-[18px] font-normal text-text-secondary caret-primary-600 outline-hidden placeholder:text-text-quaternary"
placeholder={
t(($) => $['createTool.toolInput.descriptionPlaceholder'], {
placeholder={t(
($) => $['createTool.toolInput.descriptionPlaceholder'],
{
ns: 'tools',
})!
}
},
)!}
value={item.description}
onChange={(e) =>
handleParameterChange('description', e.target.value, index)
+12 -15
View File
@@ -31,17 +31,24 @@ import {
import { API_PREFIX } from '@/config'
import { BlockEnum } from './types'
type BlockIconSize = 'xs' | 'sm' | 'md'
type BlockIconProps = {
type: BlockEnum
size?: string
size?: BlockIconSize
className?: string
toolIcon?: string | { content: string; background: string }
}
const ICON_CONTAINER_CLASSNAME_SIZE_MAP: Record<string, string> = {
const ICON_CONTAINER_CLASSNAME_SIZE_MAP: Record<BlockIconSize, string> = {
xs: 'w-4 h-4 rounded-[5px] shadow-xs',
sm: 'w-5 h-5 rounded-md shadow-xs',
md: 'w-6 h-6 rounded-lg shadow-md',
}
const ICON_CLASSNAME_SIZE_MAP: Record<BlockIconSize, string> = {
xs: 'size-3',
sm: 'size-3.5',
md: 'size-4',
}
const DEFAULT_ICON_MAP: Record<BlockEnum, React.ComponentType<{ className: string }>> = {
[BlockEnum.Start]: Home,
@@ -144,7 +151,7 @@ const BlockIcon: FC<BlockIconProps> = ({ type, size = 'sm', className, toolIcon
>
<span
aria-hidden
className={cn('i-custom-vender-workflow-user-input', size === 'xs' ? 'size-4' : 'size-4')}
className={cn('i-custom-vender-workflow-user-input', ICON_CLASSNAME_SIZE_MAP[size])}
/>
</div>
)
@@ -163,7 +170,7 @@ const BlockIcon: FC<BlockIconProps> = ({ type, size = 'sm', className, toolIcon
aria-hidden
className={cn(
'i-custom-vender-workflow-start-placeholder text-text-primary opacity-30',
size === 'xs' ? 'size-3' : 'size-3.5',
ICON_CLASSNAME_SIZE_MAP[size],
)}
/>
</div>
@@ -180,17 +187,7 @@ const BlockIcon: FC<BlockIconProps> = ({ type, size = 'sm', className, toolIcon
className,
)}
>
{showDefaultIcon &&
getIcon(
type,
type === BlockEnum.TriggerSchedule || type === BlockEnum.TriggerWebhook
? size === 'xs'
? 'w-4 h-4'
: 'w-4.5 h-4.5'
: size === 'xs'
? 'w-3 h-3'
: 'w-3.5 h-3.5',
)}
{showDefaultIcon && getIcon(type, ICON_CLASSNAME_SIZE_MAP[size])}
{!showDefaultIcon && (
<>
{typeof resolvedToolIcon === 'string' ? (
@@ -16,6 +16,15 @@ const createMockSocket = (): Socket =>
off: vi.fn(),
}) as unknown as Socket
const createDeferred = <T>() => {
let resolve!: (value: T) => void
const promise = new Promise<T>((resolvePromise) => {
resolve = resolvePromise
})
return { promise, resolve }
}
const loadCollaborationModules = async () => {
const [{ CollaborationManager }, { webSocketClient }] = await Promise.all([
import('../collaboration-manager'),
@@ -40,10 +49,11 @@ describe('CollaborationManager CRDT runtime loading', () => {
expect(manager.isConnected()).toBe(false)
})
it('does not create connection state when the runtime fails to load', async () => {
it('does not create connection state when the runtime fails to load and allows a retry', async () => {
const { CollaborationManager, webSocketClient } = await loadCollaborationModules()
const manager = new CollaborationManager()
const runtimeError = new Error('runtime-load-failed')
const retryError = new Error('runtime-retry-failed')
const loadRuntimeSpy = vi
.spyOn(
manager as unknown as {
@@ -51,7 +61,8 @@ describe('CollaborationManager CRDT runtime loading', () => {
},
'loadCrdtRuntime',
)
.mockRejectedValue(runtimeError)
.mockRejectedValueOnce(runtimeError)
.mockRejectedValueOnce(retryError)
const connectSpy = vi.spyOn(webSocketClient, 'connect')
await expect(manager.connect('app-runtime-failure')).rejects.toBe(runtimeError)
@@ -59,11 +70,22 @@ describe('CollaborationManager CRDT runtime loading', () => {
expect(loadRuntimeSpy).toHaveBeenCalledTimes(1)
expect(connectSpy).not.toHaveBeenCalled()
expect(manager.isConnected()).toBe(false)
await expect(manager.connect('app-runtime-failure')).rejects.toBe(retryError)
expect(loadRuntimeSpy).toHaveBeenCalledTimes(2)
expect(connectSpy).not.toHaveBeenCalled()
})
it('initializes one session for concurrent consumers of the same app', async () => {
const { CollaborationManager, webSocketClient } = await loadCollaborationModules()
const manager = new CollaborationManager()
const loadRuntimeSpy = vi.spyOn(
manager as unknown as {
loadCrdtRuntime: () => Promise<(typeof import('../crdt-runtime'))['crdtRuntime']>
},
'loadCrdtRuntime',
)
const socket = createMockSocket()
const connectSpy = vi.spyOn(webSocketClient, 'connect').mockReturnValue(socket)
const disconnectSpy = vi
@@ -76,6 +98,7 @@ describe('CollaborationManager CRDT runtime loading', () => {
])
expect(firstConnectionId).not.toBe(secondConnectionId)
expect(loadRuntimeSpy).toHaveBeenCalledTimes(1)
expect(connectSpy).toHaveBeenCalledTimes(1)
expect(loroModuleState.evaluations).toBe(1)
@@ -85,4 +108,61 @@ describe('CollaborationManager CRDT runtime loading', () => {
manager.disconnect(secondConnectionId)
expect(disconnectSpy).toHaveBeenCalledWith('app-concurrent')
})
it('keeps the latest app session when the app changes during runtime loading', async () => {
const { CollaborationManager, webSocketClient } = await loadCollaborationModules()
const { crdtRuntime } = await import('../crdt-runtime')
const manager = new CollaborationManager()
const runtime = createDeferred<typeof crdtRuntime>()
vi.spyOn(
manager as unknown as {
loadCrdtRuntime: () => Promise<typeof crdtRuntime>
},
'loadCrdtRuntime',
).mockReturnValue(runtime.promise)
const socket = createMockSocket()
const connectSpy = vi.spyOn(webSocketClient, 'connect').mockReturnValue(socket)
const disconnectSpy = vi
.spyOn(webSocketClient, 'disconnect')
.mockImplementation(() => undefined)
const firstConnection = manager.connect('app-first')
const secondConnection = manager.connect('app-second')
runtime.resolve(crdtRuntime)
const [firstConnectionId, secondConnectionId] = await Promise.all([
firstConnection,
secondConnection,
])
expect(connectSpy).toHaveBeenCalledTimes(1)
expect(connectSpy).toHaveBeenCalledWith('app-second')
manager.disconnect(firstConnectionId)
expect(disconnectSpy).not.toHaveBeenCalled()
manager.disconnect(secondConnectionId)
expect(disconnectSpy).toHaveBeenCalledWith('app-second')
})
it('does not initialize a pending session after the manager is destroyed', async () => {
const { CollaborationManager, webSocketClient } = await loadCollaborationModules()
const { crdtRuntime } = await import('../crdt-runtime')
const manager = new CollaborationManager()
const runtime = createDeferred<typeof crdtRuntime>()
vi.spyOn(
manager as unknown as {
loadCrdtRuntime: () => Promise<typeof crdtRuntime>
},
'loadCrdtRuntime',
).mockReturnValue(runtime.promise)
const connectSpy = vi.spyOn(webSocketClient, 'connect')
const connection = manager.connect('app-destroyed')
manager.destroy()
runtime.resolve(crdtRuntime)
await connection
expect(connectSpy).not.toHaveBeenCalled()
})
})
@@ -907,13 +907,17 @@ describe('CollaborationManager socket and subscription behavior', () => {
expect(secondConnectionId).toBeTruthy()
expect(disconnectSpy).not.toHaveBeenCalled()
await manager.connect('app-2', reactFlowStore)
const thirdConnectionId = await manager.connect('app-2', reactFlowStore)
expect(disconnectSpy).toHaveBeenCalledWith('app-1')
expect(internals.currentAppId).toBe('app-2')
internals.isLeader = true
manager.disconnect(secondConnectionId)
manager.disconnect(firstConnectionId)
expect(disconnectSpy).not.toHaveBeenCalledWith('app-2')
expect(internals.currentAppId).toBe('app-2')
manager.disconnect(thirdConnectionId)
expect(disconnectSpy).toHaveBeenCalledWith('app-2')
expect(eventEmitSpy).toHaveBeenCalledWith('leaderChange', false)
expect(internals.currentAppId).toBeNull()
@@ -156,6 +156,9 @@ const toUint8Array = (value: unknown): Uint8Array | null => {
export class CollaborationManager {
private crdtRuntime: CrdtRuntime | null = null
private crdtRuntimePromise: Promise<CrdtRuntime> | null = null
private targetAppId: string | null = null
private connectGeneration = 0
private doc: LoroDoc | null = null
private undoManager: UndoManager | null = null
private provider: CRDTProvider | null = null
@@ -538,6 +541,20 @@ export class CollaborationManager {
return crdtRuntime
}
private async ensureCrdtRuntime(): Promise<void> {
if (this.crdtRuntime) return
const runtimePromise = this.crdtRuntimePromise ?? this.loadCrdtRuntime()
this.crdtRuntimePromise = runtimePromise
try {
this.crdtRuntime = await runtimePromise
} catch (error) {
if (this.crdtRuntimePromise === runtimePromise) this.crdtRuntimePromise = null
throw error
}
}
private getCrdtRuntime(): CrdtRuntime {
if (!this.crdtRuntime) throw new Error('CRDT runtime not initialized')
return this.crdtRuntime
@@ -639,21 +656,31 @@ export class CollaborationManager {
}
async connect(appId: string, reactFlowStore?: ReactFlowStore): Promise<string> {
this.crdtRuntime ??= await this.loadCrdtRuntime()
const connectionId = Math.random().toString(36).substring(2, 11)
if (this.targetAppId !== appId) {
this.targetAppId = appId
this.connectGeneration += 1
}
const connectGeneration = this.connectGeneration
this.activeConnections.add(connectionId)
if (!this.crdtRuntime) await this.ensureCrdtRuntime()
if (connectGeneration !== this.connectGeneration || this.targetAppId !== appId)
return connectionId
if (this.currentAppId === appId && this.doc) {
// Already connected to the same app, only update store if provided and we don't have one
if (reactFlowStore && !this.reactFlowStore) this.reactFlowStore = reactFlowStore
this.activeConnections.add(connectionId)
return connectionId
}
// Only disconnect if switching to a different app
if (this.currentAppId && this.currentAppId !== appId) this.forceDisconnect()
if (this.currentAppId && this.currentAppId !== appId)
this.forceDisconnect({ preserveConnectIntent: true })
this.activeConnections.add(connectionId)
this.hasEstablishedConnection = false
this.currentAppId = appId
@@ -685,7 +712,7 @@ export class CollaborationManager {
}
disconnect = (connectionId?: string): void => {
if (connectionId) this.activeConnections.delete(connectionId)
if (connectionId && !this.activeConnections.delete(connectionId)) return
// Only disconnect when no more connections
if (this.activeConnections.size === 0) this.forceDisconnect()
@@ -699,7 +726,9 @@ export class CollaborationManager {
this.pendingWorkflowSyncRequests.clear()
}
private forceDisconnect = (): void => {
private forceDisconnect = ({
preserveConnectIntent = false,
}: { preserveConnectIntent?: boolean } = {}): void => {
if (this.currentAppId) webSocketClient.disconnect(this.currentAppId)
this.clearInitialSyncRetry()
@@ -740,6 +769,10 @@ export class CollaborationManager {
if (wasLeader) this.eventEmitter.emit('leaderChange', false)
this.activeConnections.clear()
if (!preserveConnectIntent) {
this.targetAppId = null
this.connectGeneration += 1
}
this.eventEmitter.removeAllListeners()
}
@@ -68,9 +68,9 @@ const FileTypeItem: FC<Props> = ({
<TagInput
items={customFileTypes}
onChange={onCustomFileTypesChange}
placeholder={
t(($) => $['variableConfig.file.custom.createPlaceholder'], { ns: 'appDebug' })!
}
placeholder={t(($) => $['variableConfig.file.custom.createPlaceholder'], {
ns: 'appDebug',
})!}
/>
</div>
</div>
@@ -177,11 +177,12 @@ const AddExtractParameter: FC<Props> = ({ type, payload, onSave, onCancel }) =>
<Input
value={param.name}
onChange={(e) => handleParamChange('name')(e.target.value)}
placeholder={
t(($) => $[`${i18nPrefix}.addExtractParameterContent.namePlaceholder`], {
placeholder={t(
($) => $[`${i18nPrefix}.addExtractParameterContent.namePlaceholder`],
{
ns: 'workflow',
})!
}
},
)!}
/>
</Field>
<Field
@@ -224,12 +225,10 @@ const AddExtractParameter: FC<Props> = ({ type, payload, onSave, onCancel }) =>
)}
value={param.description}
onValueChange={(value) => handleParamChange('description')(value)}
placeholder={
t(
($) => $[`${i18nPrefix}.addExtractParameterContent.descriptionPlaceholder`],
{ ns: 'workflow' },
)!
}
placeholder={t(
($) => $[`${i18nPrefix}.addExtractParameterContent.descriptionPlaceholder`],
{ ns: 'workflow' },
)!}
/>
</Field>
<Field
@@ -193,7 +193,7 @@ export function SourceAppPicker({
<ComboboxTrigger
aria-label={t(($) => $['versions.sourceAppOption'])}
icon={false}
className="block h-auto w-full border-0 bg-transparent p-0 text-left hover:bg-transparent focus-visible:bg-transparent focus-visible:ring-0 data-open:bg-transparent"
className="block h-auto w-full border-0 bg-transparent p-0 text-left hover:bg-transparent focus-visible:bg-transparent data-open:bg-transparent"
>
<SourceAppTrigger app={value} />
</ComboboxTrigger>
@@ -128,7 +128,7 @@ export function AccessSubjectAddButton({
onClick={() => {
if (open) closeMenu()
}}
className="h-6 w-auto min-w-[52px] shrink-0 rounded-md border-0 bg-transparent px-2 py-0 text-xs font-medium text-components-button-secondary-accent-text hover:bg-state-accent-hover focus-visible:bg-state-accent-hover focus-visible:ring-2 focus-visible:ring-state-accent-solid data-popup-open:bg-state-accent-hover"
className="h-6 w-auto min-w-[52px] shrink-0 rounded-md border-0 bg-transparent px-2 py-0 text-xs font-medium text-components-button-secondary-accent-text hover:bg-state-accent-hover focus-visible:bg-state-accent-hover data-popup-open:bg-state-accent-hover"
>
<span className="inline-flex min-w-0 items-center justify-center gap-x-0.5 whitespace-nowrap">
<span className="i-ri-add-circle-fill size-4 shrink-0" aria-hidden="true" />
@@ -86,7 +86,7 @@ export const TagFilter = ({
aria-label={triggerLabel}
icon={false}
className={cn(
'flex h-8 max-w-60 min-w-28 cursor-pointer items-center gap-1 rounded-lg border-[0.5px] border-transparent bg-components-input-bg-normal px-2 py-0 text-left whitespace-nowrap select-none hover:bg-components-input-bg-normal focus-visible:bg-components-input-bg-normal focus-visible:ring-2 focus-visible:ring-state-accent-solid data-popup-open:bg-components-input-bg-normal',
'flex h-8 max-w-60 min-w-28 cursor-pointer items-center gap-1 rounded-lg border-[0.5px] border-transparent bg-components-input-bg-normal px-2 py-0 text-left whitespace-nowrap select-none hover:bg-components-input-bg-normal focus-visible:bg-components-input-bg-normal data-popup-open:bg-components-input-bg-normal',
!!value.length && 'pr-6 shadow-xs',
triggerClassName,
)}
@@ -257,7 +257,7 @@ export const TagSelector = ({
disabled={!canManageTags && !canBindOrUnbindTags}
aria-label={triggerLabel}
className={cn(
'block h-auto w-full rounded-lg border-0 bg-transparent p-0 text-left hover:bg-transparent focus:outline-hidden focus-visible:bg-transparent focus-visible:inset-ring-2 focus-visible:inset-ring-state-accent-solid data-popup-open:bg-state-base-hover data-popup-open:hover:bg-state-base-hover',
'block h-auto w-full rounded-lg border-0 bg-transparent p-0 text-left hover:bg-transparent focus:outline-hidden focus-visible:bg-transparent data-popup-open:bg-state-base-hover data-popup-open:hover:bg-state-base-hover',
)}
icon={false}
>
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "dify-web",
"version": "1.16.0",
"version": "1.16.1",
"private": true,
"type": "module",
"imports": {