Compare commits
13
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f1c32e1bc8 | ||
|
|
a57b0b9b58 | ||
|
|
d8506efed6 | ||
|
|
b25b28cc76 | ||
|
|
f68cabe226 | ||
|
|
f73f83c5d6 | ||
|
|
29801d6a65 | ||
|
|
75016e8bfe | ||
|
|
b6dea7b2ba | ||
|
|
b2b1cd7e97 | ||
|
|
ca63e01d45 | ||
|
|
271b6e1f5c | ||
|
|
609ddf9e01 |
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
@@ -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:
|
||||
|
||||
@@ -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"
|
||||
|
||||
+235
-231
@@ -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:
|
||||
|
||||
@@ -1,402 +0,0 @@
|
||||
import { execFileSync, spawnSync } from 'node:child_process'
|
||||
import { chmodSync, mkdtempSync, readFileSync, rmSync, writeFileSync } from 'node:fs'
|
||||
import { tmpdir } from 'node:os'
|
||||
import { join } from 'node:path'
|
||||
import { fileURLToPath } from 'node:url'
|
||||
import { describe, expect, it } from 'vitest'
|
||||
|
||||
const SCRIPT = fileURLToPath(new URL('./install-cli.sh', import.meta.url))
|
||||
|
||||
function pickAsset(target: string, releaseJson: string): string {
|
||||
return execFileSync('sh', ['-c', `. "${SCRIPT}"; pick_asset "$1"`, 'sh', target], {
|
||||
input: releaseJson,
|
||||
encoding: 'utf8',
|
||||
env: { ...process.env, DIFYCTL_INSTALL_LIB: '1' },
|
||||
}).trim()
|
||||
}
|
||||
|
||||
function assetVersion(name: string, target: string): string {
|
||||
return execFileSync('sh', ['-c', `. "${SCRIPT}"; asset_version "$1" "$2"`, 'sh', name, target], {
|
||||
encoding: 'utf8',
|
||||
env: { ...process.env, DIFYCTL_INSTALL_LIB: '1' },
|
||||
}).trim()
|
||||
}
|
||||
|
||||
// Stubs the only network primitive (fetch_json) so resolution logic runs fully
|
||||
// offline. Routes by URL; release bodies come from env (TAG_<tag-with-._->_>),
|
||||
// the latest release from LATEST_JSON, the listing from LIST_JSON. A missing
|
||||
// fixture returns 22 to mimic `curl -f` on a 4xx.
|
||||
/* oxlint-disable no-template-curly-in-string -- shell parameter expansions, not JS template literals */
|
||||
const FETCH_STUB = [
|
||||
'fetch_json() {',
|
||||
' case "$1" in',
|
||||
' *"/releases/latest") [ -n "${LATEST_JSON:-}" ] || return 22; printf "%s" "$LATEST_JSON" ;;',
|
||||
' *"/releases?per_page=100") [ -n "${LIST_JSON:-}" ] || return 22; printf "%s" "$LIST_JSON" ;;',
|
||||
' *"/releases/tags/"*)',
|
||||
' _t=${1##*/releases/tags/};',
|
||||
' _k=$(printf "TAG_%s" "$_t" | tr ".-" "__");',
|
||||
' eval "_v=\\${$_k:-}";',
|
||||
' [ -n "$_v" ] || return 22;',
|
||||
' printf "%s" "$_v" ;;',
|
||||
' *) return 22 ;;',
|
||||
' esac',
|
||||
'}',
|
||||
].join('\n')
|
||||
/* oxlint-enable no-template-curly-in-string */
|
||||
|
||||
function runLib(
|
||||
program: string,
|
||||
env: Record<string, string> = {},
|
||||
): { code: number; stdout: string; stderr: string } {
|
||||
const full = `. "${SCRIPT}"\n${FETCH_STUB}\n${program}`
|
||||
const r = spawnSync('sh', ['-c', full], {
|
||||
encoding: 'utf8',
|
||||
env: {
|
||||
...process.env,
|
||||
DIFYCTL_INSTALL_LIB: '1',
|
||||
DIFY_VERSION: '',
|
||||
DIFYCTL_VERSION: '',
|
||||
...env,
|
||||
},
|
||||
})
|
||||
return { code: r.status ?? 1, stdout: (r.stdout ?? '').trim(), stderr: r.stderr ?? '' }
|
||||
}
|
||||
|
||||
// Like runLib but with a caller-supplied fetch_json stub, so we can drive the
|
||||
// real rate_limit_hint / maybe_ratelimit_exit / fetch_hit_ratelimit (which the
|
||||
// script defines) by writing a classified reason to FETCH_ERR_FILE.
|
||||
function runLibStub(
|
||||
stub: string,
|
||||
program: string,
|
||||
env: Record<string, string> = {},
|
||||
): { code: number; stderr: string } {
|
||||
const full = `. "${SCRIPT}"\n${stub}\n${program}`
|
||||
const r = spawnSync('sh', ['-c', full], {
|
||||
encoding: 'utf8',
|
||||
env: {
|
||||
...process.env,
|
||||
DIFYCTL_INSTALL_LIB: '1',
|
||||
DIFY_VERSION: '',
|
||||
DIFYCTL_VERSION: '',
|
||||
...env,
|
||||
},
|
||||
})
|
||||
return { code: r.status ?? 1, stderr: r.stderr ?? '' }
|
||||
}
|
||||
|
||||
// A fetch_json that always fails with the given classification, mimicking the
|
||||
// real one writing to FETCH_ERR_FILE from inside a command-substitution subshell.
|
||||
function failStub(reason: string): string {
|
||||
return `fetch_json() { printf '%s' '${reason}' > "$FETCH_ERR_FILE"; return 1; }`
|
||||
}
|
||||
|
||||
const REL_1142 = JSON.stringify({
|
||||
tag_name: '1.14.2',
|
||||
assets: [{ name: 'difyctl-v0.2.0-linux-x64' }, { name: 'difyctl-v0.2.0-checksums.txt' }],
|
||||
})
|
||||
const REL_1150 = JSON.stringify({
|
||||
tag_name: '1.15.0',
|
||||
assets: [{ name: 'difyctl-v0.3.0-linux-x64' }],
|
||||
})
|
||||
const LIST_NEWEST_FIRST = JSON.stringify({
|
||||
releases: [{ tag_name: '1.15.0' }, { tag_name: '1.14.2' }],
|
||||
})
|
||||
|
||||
const RELEASE = JSON.stringify({
|
||||
tag_name: '1.14.2',
|
||||
name: 'Dify 1.14.2',
|
||||
assets: [
|
||||
{ name: 'difyctl-v0.1.0-rc.1-linux-x64' },
|
||||
{ name: 'difyctl-v0.2.0-linux-x64' },
|
||||
{ name: 'difyctl-v0.2.0-linux-arm64' },
|
||||
{ name: 'difyctl-v0.2.0-darwin-arm64' },
|
||||
{ name: 'difyctl-v0.2.0-windows-x64.exe' },
|
||||
{ name: 'difyctl-v0.2.0-checksums.txt' },
|
||||
{ name: 'some-other-asset.zip' },
|
||||
],
|
||||
})
|
||||
|
||||
describe('install-cli pick_asset', () => {
|
||||
it('picks the highest difyctl version for a linux target', () => {
|
||||
expect(pickAsset('linux-x64', RELEASE)).toBe('difyctl-v0.2.0-linux-x64')
|
||||
})
|
||||
|
||||
it('matches the windows .exe asset', () => {
|
||||
expect(pickAsset('windows-x64', RELEASE)).toBe('difyctl-v0.2.0-windows-x64.exe')
|
||||
})
|
||||
|
||||
it('matches an arm64 target exactly (no x64 bleed-through)', () => {
|
||||
expect(pickAsset('darwin-arm64', RELEASE)).toBe('difyctl-v0.2.0-darwin-arm64')
|
||||
})
|
||||
|
||||
it('excludes the checksums asset', () => {
|
||||
expect(pickAsset('linux-x64', RELEASE)).not.toContain('checksums')
|
||||
})
|
||||
|
||||
it('yields empty when no asset matches the target', () => {
|
||||
expect(pickAsset('darwin-x64', RELEASE)).toBe('')
|
||||
})
|
||||
|
||||
it('picks the highest semver when several difyctl versions are present', () => {
|
||||
const many = JSON.stringify({
|
||||
assets: [
|
||||
{ name: 'difyctl-v0.2.0-linux-x64' },
|
||||
{ name: 'difyctl-v0.10.0-linux-x64' },
|
||||
{ name: 'difyctl-v0.9.0-linux-x64' },
|
||||
],
|
||||
})
|
||||
expect(pickAsset('linux-x64', many)).toBe('difyctl-v0.10.0-linux-x64')
|
||||
})
|
||||
})
|
||||
|
||||
describe('install-cli asset_version', () => {
|
||||
it('extracts the version from a posix asset name', () => {
|
||||
expect(assetVersion('difyctl-v0.2.0-linux-x64', 'linux-x64')).toBe('0.2.0')
|
||||
})
|
||||
|
||||
it('extracts the version from a windows .exe asset name', () => {
|
||||
expect(assetVersion('difyctl-v0.2.0-windows-x64.exe', 'windows-x64')).toBe('0.2.0')
|
||||
})
|
||||
|
||||
it('extracts a prerelease version', () => {
|
||||
expect(assetVersion('difyctl-v0.1.0-rc.1-linux-x64', 'linux-x64')).toBe('0.1.0-rc.1')
|
||||
})
|
||||
})
|
||||
|
||||
describe('install-cli resolve_release', () => {
|
||||
it('DIFY_VERSION pins the release directly', () => {
|
||||
const r = runLib('resolve_release linux-x64; printf "%s" "$DIFY_TAG"', {
|
||||
DIFY_VERSION: '1.14.2',
|
||||
TAG_1_14_2: REL_1142,
|
||||
})
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('1.14.2')
|
||||
})
|
||||
|
||||
it('DIFY_VERSION that does not exist dies with a clear message', () => {
|
||||
const r = runLib('resolve_release linux-x64', { DIFY_VERSION: '9.9.9' })
|
||||
expect(r.code).not.toBe(0)
|
||||
expect(r.stderr).toContain('Dify release 9.9.9 not found')
|
||||
})
|
||||
|
||||
it('blank resolves to the latest stable release', () => {
|
||||
const r = runLib('resolve_release linux-x64; printf "%s" "$DIFY_TAG"', {
|
||||
LATEST_JSON: REL_1150,
|
||||
})
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('1.15.0')
|
||||
})
|
||||
|
||||
it('blank dies when the latest query fails (no silent fallback)', () => {
|
||||
const r = runLib('resolve_release linux-x64')
|
||||
expect(r.code).not.toBe(0)
|
||||
expect(r.stderr).toContain('failed to query latest Dify release')
|
||||
})
|
||||
|
||||
it('DIFYCTL_VERSION resolves to the release hosting that build', () => {
|
||||
const r = runLib('resolve_release linux-x64; printf "%s" "$DIFY_TAG"', {
|
||||
DIFYCTL_VERSION: '0.2.0',
|
||||
LIST_JSON: LIST_NEWEST_FIRST,
|
||||
TAG_1_15_0: REL_1150,
|
||||
TAG_1_14_2: REL_1142,
|
||||
})
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('1.14.2')
|
||||
})
|
||||
|
||||
it('DIFYCTL_VERSION not hosted anywhere dies', () => {
|
||||
const r = runLib('resolve_release linux-x64', {
|
||||
DIFYCTL_VERSION: '9.9.9',
|
||||
LIST_JSON: LIST_NEWEST_FIRST,
|
||||
TAG_1_15_0: REL_1150,
|
||||
TAG_1_14_2: REL_1142,
|
||||
})
|
||||
expect(r.code).not.toBe(0)
|
||||
expect(r.stderr).toContain('difyctl 9.9.9 not found on any Dify release')
|
||||
})
|
||||
})
|
||||
|
||||
describe('install-cli find_release_for_difyctl', () => {
|
||||
it('returns the newest release whose assets host the wanted build', () => {
|
||||
const r = runLib('find_release_for_difyctl 0.2.0 linux-x64', {
|
||||
LIST_JSON: LIST_NEWEST_FIRST,
|
||||
TAG_1_15_0: REL_1150,
|
||||
TAG_1_14_2: REL_1142,
|
||||
})
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('1.14.2')
|
||||
})
|
||||
|
||||
it('dies (not false-negative) when the releases listing fails', () => {
|
||||
const r = runLib('find_release_for_difyctl 0.2.0 linux-x64')
|
||||
expect(r.code).not.toBe(0)
|
||||
expect(r.stderr).toContain('failed to query')
|
||||
})
|
||||
|
||||
it('warns and skips a release whose fetch fails, then finds it later', () => {
|
||||
const r = runLib('find_release_for_difyctl 0.2.0 linux-x64', {
|
||||
LIST_JSON: LIST_NEWEST_FIRST,
|
||||
TAG_1_14_2: REL_1142,
|
||||
})
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('1.14.2')
|
||||
expect(r.stderr).toContain('fetch failed for 1.15.0')
|
||||
})
|
||||
})
|
||||
|
||||
describe('install-cli rate limit', () => {
|
||||
const futureReset = String(Math.floor(Date.now() / 1000) + 1800)
|
||||
|
||||
it('latest: reports the rate limit with reset ETA and remediation, not a generic error', () => {
|
||||
const r = runLibStub(failStub(`ratelimit:${futureReset}`), 'resolve_release linux-x64')
|
||||
expect(r.code).not.toBe(0)
|
||||
expect(r.stderr).toContain('rate limit exceeded')
|
||||
expect(r.stderr).toContain('resets in ~')
|
||||
expect(r.stderr).toContain('GITHUB_TOKEN')
|
||||
expect(r.stderr).not.toContain('failed to query latest')
|
||||
})
|
||||
|
||||
it('DIFY_VERSION: rate limit wins over the misleading "not found" message', () => {
|
||||
const r = runLibStub(failStub(`ratelimit:${futureReset}`), 'resolve_release linux-x64', {
|
||||
DIFY_VERSION: '1.15.0',
|
||||
})
|
||||
expect(r.code).not.toBe(0)
|
||||
expect(r.stderr).toContain('rate limit exceeded')
|
||||
expect(r.stderr).not.toContain('not found')
|
||||
})
|
||||
|
||||
it('DIFYCTL_VERSION: rate limit surfaces from the nested subshell, not "not found"', () => {
|
||||
const r = runLibStub(failStub(`ratelimit:${futureReset}`), 'resolve_release linux-x64', {
|
||||
DIFYCTL_VERSION: '0.2.0',
|
||||
})
|
||||
expect(r.code).not.toBe(0)
|
||||
expect(r.stderr).toContain('rate limit exceeded')
|
||||
expect(r.stderr).not.toContain('not found')
|
||||
})
|
||||
|
||||
it('omits the ETA line when the reset epoch is missing', () => {
|
||||
const r = runLibStub(failStub('ratelimit:'), 'resolve_release linux-x64')
|
||||
expect(r.code).not.toBe(0)
|
||||
expect(r.stderr).toContain('rate limit exceeded')
|
||||
expect(r.stderr).not.toContain('resets in ~')
|
||||
})
|
||||
|
||||
it('a non-rate-limit HTTP error falls back to the generic message (no false hint)', () => {
|
||||
const r = runLibStub(failStub('http:500'), 'resolve_release linux-x64')
|
||||
expect(r.code).not.toBe(0)
|
||||
expect(r.stderr).toContain('failed to query latest')
|
||||
expect(r.stderr).not.toContain('rate limit exceeded')
|
||||
})
|
||||
})
|
||||
|
||||
// A stand-in for curl that honours the flags fetch_json passes (-D/-o/-w/-H) and
|
||||
// fabricates a response per FAKE_MODE, so the tests exercise the REAL fetch_json
|
||||
// (its curl invocation, header parsing, classification and token handling) rather
|
||||
// than a stub. Header names it emits are lowercase, as HTTP/2 delivers them.
|
||||
const FAKE_CURL = `#!/bin/sh
|
||||
hdr=""; body=""
|
||||
while [ $# -gt 0 ]; do
|
||||
case "$1" in
|
||||
-D) hdr="$2"; shift 2 ;;
|
||||
-o) body="$2"; shift 2 ;;
|
||||
-H) printf '%s\\n' "$2" >> "\${FAKE_HDR_LOG:-/dev/null}"; shift 2 ;;
|
||||
-w) shift 2 ;;
|
||||
*) shift ;;
|
||||
esac
|
||||
done
|
||||
case "\${FAKE_MODE:-ok}" in
|
||||
ok)
|
||||
[ -n "$hdr" ] && printf 'HTTP/2 200\\r\\n\\r\\n' > "$hdr"
|
||||
[ -n "$body" ] && printf '%s' "\${FAKE_BODY:-}" > "$body"
|
||||
printf '200' ;;
|
||||
ratelimit)
|
||||
[ -n "$hdr" ] && printf 'HTTP/2 403\\r\\nx-ratelimit-remaining: 0\\r\\nx-ratelimit-reset: %s\\r\\n\\r\\n' "\${FAKE_RESET:-9999999999}" > "$hdr"
|
||||
printf '403' ;;
|
||||
perm403)
|
||||
[ -n "$hdr" ] && printf 'HTTP/2 403\\r\\nx-ratelimit-remaining: 59\\r\\n\\r\\n' > "$hdr"
|
||||
printf '403' ;;
|
||||
notfound)
|
||||
[ -n "$hdr" ] && printf 'HTTP/2 404\\r\\n\\r\\n' > "$hdr"
|
||||
printf '404' ;;
|
||||
net) exit 6 ;;
|
||||
esac
|
||||
`
|
||||
|
||||
// Drive the real fetch_json with FAKE_CURL first on PATH. Returns "OK|<body>" or
|
||||
// "FAIL|<FETCH_ERR_FILE contents>", plus any -H lines the fake curl received.
|
||||
function runRealFetch(
|
||||
mode: string,
|
||||
env: Record<string, string> = {},
|
||||
): { result: string; headers: string } {
|
||||
const dir = mkdtempSync(join(tmpdir(), 'difyctl-fakecurl-'))
|
||||
const hdrLog = join(dir, 'hdrlog')
|
||||
writeFileSync(join(dir, 'curl'), FAKE_CURL)
|
||||
chmodSync(join(dir, 'curl'), 0o755)
|
||||
const program =
|
||||
'if body=$(fetch_json "https://api.github.com/repos/x/releases/latest"); then printf \'OK|%s\' "$body"; else printf \'FAIL|%s\' "$(cat "$FETCH_ERR_FILE" 2>/dev/null)"; fi'
|
||||
const r = spawnSync('sh', ['-c', `. "${SCRIPT}"\n${program}`], {
|
||||
encoding: 'utf8',
|
||||
env: {
|
||||
...process.env,
|
||||
PATH: `${dir}:${process.env.PATH ?? ''}`,
|
||||
DIFYCTL_INSTALL_LIB: '1',
|
||||
DIFY_VERSION: '',
|
||||
DIFYCTL_VERSION: '',
|
||||
FAKE_MODE: mode,
|
||||
FAKE_HDR_LOG: hdrLog,
|
||||
...env,
|
||||
},
|
||||
})
|
||||
let headers = ''
|
||||
try {
|
||||
headers = readFileSync(hdrLog, 'utf8')
|
||||
} catch {
|
||||
/* no headers logged */
|
||||
}
|
||||
rmSync(dir, { recursive: true, force: true })
|
||||
return { result: (r.stdout ?? '').trim(), headers }
|
||||
}
|
||||
|
||||
describe('install-cli fetch_json (real, fake curl on PATH)', () => {
|
||||
it('returns the response body on 200', () => {
|
||||
expect(runRealFetch('ok', { FAKE_BODY: '{"tag_name":"1.15.0"}' }).result).toBe(
|
||||
'OK|{"tag_name":"1.15.0"}',
|
||||
)
|
||||
})
|
||||
|
||||
it('classifies a 403 with x-ratelimit-remaining:0 as a rate limit and captures the reset', () => {
|
||||
expect(runRealFetch('ratelimit', { FAKE_RESET: '1893456000' }).result).toBe(
|
||||
'FAIL|ratelimit:1893456000',
|
||||
)
|
||||
})
|
||||
|
||||
it('classifies a 403 with tokens left as a plain http error, not a rate limit', () => {
|
||||
expect(runRealFetch('perm403').result).toBe('FAIL|http:403')
|
||||
})
|
||||
|
||||
it('classifies a 404 as an http error', () => {
|
||||
expect(runRealFetch('notfound').result).toBe('FAIL|http:404')
|
||||
})
|
||||
|
||||
it('classifies a curl transport failure as a network error', () => {
|
||||
expect(runRealFetch('net').result).toBe('FAIL|network')
|
||||
})
|
||||
|
||||
it('sends an Authorization bearer header when GITHUB_TOKEN is set', () => {
|
||||
expect(runRealFetch('ok', { FAKE_BODY: '{}', GITHUB_TOKEN: 'ghp_secret' }).headers).toContain(
|
||||
'Authorization: Bearer ghp_secret',
|
||||
)
|
||||
})
|
||||
|
||||
it('falls back to GH_TOKEN when GITHUB_TOKEN is unset', () => {
|
||||
expect(
|
||||
runRealFetch('ok', { FAKE_BODY: '{}', GITHUB_TOKEN: '', GH_TOKEN: 'gho_fallback' }).headers,
|
||||
).toContain('Authorization: Bearer gho_fallback')
|
||||
})
|
||||
|
||||
it('sends no Authorization header when neither token is set', () => {
|
||||
expect(
|
||||
runRealFetch('ok', { FAKE_BODY: '{}', GITHUB_TOKEN: '', GH_TOKEN: '' }).headers,
|
||||
).not.toContain('Authorization')
|
||||
})
|
||||
})
|
||||
@@ -1,53 +0,0 @@
|
||||
import { spawnSync } from 'node:child_process'
|
||||
import { fileURLToPath } from 'node:url'
|
||||
import { describe, expect, it } from 'vitest'
|
||||
|
||||
const SCRIPT = fileURLToPath(new URL('./install-r2.ps1', import.meta.url))
|
||||
const hasPwsh = spawnSync('pwsh', ['-v'], { encoding: 'utf8' }).status === 0
|
||||
const d = hasPwsh ? describe : describe.skip
|
||||
|
||||
const MANIFEST = JSON.stringify({
|
||||
schema: 1,
|
||||
name: 'difyctl',
|
||||
channel: 'edge',
|
||||
version: '0.1.0-edge.2fd7b82',
|
||||
baseUrl: 'https://pub.example.r2.dev/difyctl/edge/0.1.0-edge.2fd7b82',
|
||||
targets: {
|
||||
'windows-x64': { asset: 'difyctl-v0.1.0-edge.2fd7b82-windows-x64.exe', sha256: 'deadbeef' },
|
||||
},
|
||||
})
|
||||
|
||||
function pwsh(program: string): { code: number; stdout: string; stderr: string } {
|
||||
const full = `. '${SCRIPT}'\n${program}`
|
||||
const r = spawnSync('pwsh', ['-NoProfile', '-Command', full], {
|
||||
encoding: 'utf8',
|
||||
env: { ...process.env, DIFYCTL_INSTALL_LIB: '1' },
|
||||
})
|
||||
return {
|
||||
code: r.status ?? 1,
|
||||
stdout: (r.stdout ?? '').replace(/\r\n/g, '\n').trim(),
|
||||
stderr: r.stderr ?? '',
|
||||
}
|
||||
}
|
||||
|
||||
d('install-r2.ps1', () => {
|
||||
it('parses a target asset + sha from the manifest', () => {
|
||||
const prog =
|
||||
`$m = ConvertFrom-Json @'\n${MANIFEST}\n'@\n` +
|
||||
`Write-Output (Get-TargetField $m 'windows-x64' 'asset')\n` +
|
||||
`Write-Output (Get-TargetField $m 'windows-x64' 'sha256')`
|
||||
const { stdout } = pwsh(prog)
|
||||
expect(stdout).toBe('difyctl-v0.1.0-edge.2fd7b82-windows-x64.exe\ndeadbeef')
|
||||
})
|
||||
|
||||
it('errors when DIFYCTL_R2_BASE is unset', () => {
|
||||
const r = spawnSync('pwsh', ['-NoProfile', '-File', SCRIPT], {
|
||||
encoding: 'utf8',
|
||||
env: { ...process.env, DIFYCTL_R2_BASE: '' },
|
||||
})
|
||||
if (hasPwsh) {
|
||||
expect(r.status).not.toBe(0)
|
||||
expect(r.stderr + r.stdout).toMatch(/DIFYCTL_R2_BASE/)
|
||||
}
|
||||
})
|
||||
})
|
||||
@@ -1,149 +0,0 @@
|
||||
import { spawnSync } from 'node:child_process'
|
||||
import { fileURLToPath } from 'node:url'
|
||||
import { describe, expect, it } from 'vitest'
|
||||
|
||||
const SCRIPT = fileURLToPath(new URL('./install-r2.sh', import.meta.url))
|
||||
|
||||
const MANIFEST = [
|
||||
'{',
|
||||
' "schema": 1,',
|
||||
' "name": "difyctl",',
|
||||
' "channel": "edge",',
|
||||
' "version": "0.1.0-edge.2fd7b82",',
|
||||
' "commit": "abc1234",',
|
||||
' "buildDate": "2026-06-14T12:00:00Z",',
|
||||
' "compat": {"minDify":"1.14.0","maxDify":"1.15.0"},',
|
||||
' "baseUrl": "https://pub.example.r2.dev/difyctl/edge/0.1.0-edge.2fd7b82",',
|
||||
' "targets": {',
|
||||
' "linux-x64": { "asset": "difyctl-v0.1.0-edge.2fd7b82-linux-x64", "sha256": "deadbeef" },',
|
||||
' "darwin-arm64": { "asset": "difyctl-v0.1.0-edge.2fd7b82-darwin-arm64", "sha256": "cafef00d" }',
|
||||
' }',
|
||||
'}',
|
||||
].join('\n')
|
||||
|
||||
const INDEX = [
|
||||
'{',
|
||||
' "schema": 1,',
|
||||
' "channel": "edge",',
|
||||
' "updated": "2026-06-15T00:00:00Z",',
|
||||
' "builds": [',
|
||||
' {',
|
||||
' "version": "0.1.0-edge.ce4af868",',
|
||||
' "commit": "ce4af868d653f405070fabb3be3303430cc030ad",',
|
||||
' "buildDate": "2026-06-15T00:00:00Z",',
|
||||
' "dir": "0.1.0-edge.ce4af868"',
|
||||
' },',
|
||||
' {',
|
||||
' "version": "0.1.0-edge.aaaa111",',
|
||||
' "commit": "aaaa111bbbbcccc000011112222333344445555",',
|
||||
' "buildDate": "2026-06-14T00:00:00Z",',
|
||||
' "dir": "0.1.0-edge.aaaa111"',
|
||||
' }',
|
||||
' ]',
|
||||
'}',
|
||||
].join('\n')
|
||||
|
||||
const CHECKSUMS = [
|
||||
'deadbeef difyctl-v0.1.0-edge.ce4af868-linux-x64',
|
||||
'cafef00d difyctl-v0.1.0-edge.ce4af868-darwin-arm64',
|
||||
'beadc0de difyctl-v0.1.0-edge.ce4af868-windows-x64.exe',
|
||||
].join('\n')
|
||||
|
||||
function lib(
|
||||
program: string,
|
||||
env: Record<string, string> = {},
|
||||
): { code: number; stdout: string; stderr: string } {
|
||||
const full = `. "${SCRIPT}"\n${program}`
|
||||
const r = spawnSync('sh', ['-c', full], {
|
||||
encoding: 'utf8',
|
||||
env: { ...process.env, DIFYCTL_INSTALL_LIB: '1', ...env },
|
||||
})
|
||||
return { code: r.status ?? 1, stdout: (r.stdout ?? '').trim(), stderr: r.stderr ?? '' }
|
||||
}
|
||||
|
||||
describe('install-r2 manifest parsing', () => {
|
||||
// install-r2.sh is POSIX-only; under git-bash on Windows `uname -s` is MINGW*,
|
||||
// so detect_target intentionally dies (Windows installs go through install-r2.ps1).
|
||||
it.skipIf(process.platform === 'win32')('detect_target maps to one of the 5 ids', () => {
|
||||
const { stdout } = lib('detect_target')
|
||||
expect(['linux-x64', 'linux-arm64', 'darwin-x64', 'darwin-arm64', 'windows-x64']).toContain(
|
||||
stdout,
|
||||
)
|
||||
})
|
||||
|
||||
it('manifest_str reads a top-level string field', () => {
|
||||
const { stdout } = lib(
|
||||
`printf '%s' '${MANIFEST}' > "$tmp_m"; manifest_str "$tmp_m" channel`,
|
||||
{},
|
||||
)
|
||||
expect(stdout).toBe('edge')
|
||||
})
|
||||
|
||||
it('manifest_target_field extracts per-target values from a single line', () => {
|
||||
const prog =
|
||||
`printf '%s' '${MANIFEST}' > "$tmp_m"\n` +
|
||||
'manifest_target_field "$tmp_m" darwin-arm64 asset\n' +
|
||||
'manifest_target_field "$tmp_m" darwin-arm64 sha256'
|
||||
const { stdout } = lib(prog)
|
||||
expect(stdout).toBe('difyctl-v0.1.0-edge.2fd7b82-darwin-arm64\ncafef00d')
|
||||
})
|
||||
|
||||
it('requires DIFYCTL_R2_BASE when run as the installer (not lib)', () => {
|
||||
const r = spawnSync('sh', [SCRIPT], {
|
||||
encoding: 'utf8',
|
||||
env: { ...process.env, DIFYCTL_R2_BASE: '' },
|
||||
})
|
||||
expect(r.status).not.toBe(0)
|
||||
expect(r.stderr).toMatch(/DIFYCTL_R2_BASE/)
|
||||
})
|
||||
|
||||
it('sha256_check aborts on a checksum mismatch', () => {
|
||||
const r = lib('f="$(mktemp)"; printf \'hello\' > "$f"; sha256_check "$f" deadbeef')
|
||||
expect(r.code).not.toBe(0)
|
||||
expect(r.stderr).toMatch(/checksum mismatch/)
|
||||
})
|
||||
|
||||
it('sha256_check passes on the correct hash', () => {
|
||||
// sha256('hello') = 2cf24dba5fb0a30e26e83b2ac5b9e29e1b161e5c1fa7425e73043362938b9824
|
||||
const r = lib(
|
||||
'f="$(mktemp)"; printf \'hello\' > "$f"; sha256_check "$f" 2cf24dba5fb0a30e26e83b2ac5b9e29e1b161e5c1fa7425e73043362938b9824 && echo OK',
|
||||
)
|
||||
expect(r.stdout).toBe('OK')
|
||||
})
|
||||
})
|
||||
|
||||
describe('install-r2 version/commit pin', () => {
|
||||
it('index_resolve matches a build by exact version', () => {
|
||||
const r = lib(
|
||||
`printf '%s' '${INDEX}' > "$tmp_m"; index_resolve "$tmp_m" version 0.1.0-edge.aaaa111`,
|
||||
)
|
||||
expect(r.stdout).toBe('0.1.0-edge.aaaa111\t0.1.0-edge.aaaa111')
|
||||
})
|
||||
|
||||
it('index_resolve matches a build by commit prefix', () => {
|
||||
const r = lib(`printf '%s' '${INDEX}' > "$tmp_m"; index_resolve "$tmp_m" commit ce4af868`)
|
||||
expect(r.stdout).toBe('0.1.0-edge.ce4af868\t0.1.0-edge.ce4af868')
|
||||
})
|
||||
|
||||
it('index_resolve matches the full 40-char commit too', () => {
|
||||
const r = lib(
|
||||
`printf '%s' '${INDEX}' > "$tmp_m"; index_resolve "$tmp_m" commit aaaa111bbbbcccc000011112222333344445555`,
|
||||
)
|
||||
expect(r.stdout).toBe('0.1.0-edge.aaaa111\t0.1.0-edge.aaaa111')
|
||||
})
|
||||
|
||||
it('index_resolve prints nothing when no build matches', () => {
|
||||
const r = lib(`printf '%s' '${INDEX}' > "$tmp_m"; index_resolve "$tmp_m" commit ffffff`)
|
||||
expect(r.stdout).toBe('')
|
||||
})
|
||||
|
||||
it('checksums_target extracts sha and asset for a posix target', () => {
|
||||
const r = lib(`printf '%s' '${CHECKSUMS}' > "$tmp_m"; checksums_target "$tmp_m" darwin-arm64`)
|
||||
expect(r.stdout).toBe('cafef00d\tdifyctl-v0.1.0-edge.ce4af868-darwin-arm64')
|
||||
})
|
||||
|
||||
it('checksums_target does not bleed x64 into arm64', () => {
|
||||
const r = lib(`printf '%s' '${CHECKSUMS}' > "$tmp_m"; checksums_target "$tmp_m" linux-x64`)
|
||||
expect(r.stdout).toBe('deadbeef\tdifyctl-v0.1.0-edge.ce4af868-linux-x64')
|
||||
})
|
||||
})
|
||||
@@ -1,297 +0,0 @@
|
||||
import { spawnSync } from 'node:child_process'
|
||||
import { fileURLToPath } from 'node:url'
|
||||
import { describe, expect, it } from 'vitest'
|
||||
|
||||
const SCRIPT = fileURLToPath(new URL('./install.ps1', import.meta.url))
|
||||
|
||||
function hasPwsh(): boolean {
|
||||
const r = spawnSync(
|
||||
'pwsh',
|
||||
['-NoProfile', '-NonInteractive', '-Command', '$PSVersionTable.PSVersion.Major'],
|
||||
{
|
||||
encoding: 'utf8',
|
||||
},
|
||||
)
|
||||
return r.status === 0
|
||||
}
|
||||
|
||||
const PWSH = hasPwsh()
|
||||
|
||||
const STUB = [
|
||||
'function Invoke-RestMethod {',
|
||||
' param([string]$Uri, $Headers)',
|
||||
" if ($Uri -like '*/releases/latest') {",
|
||||
" if (-not $env:HX_LATEST) { throw 'mock 404' }",
|
||||
' return ($env:HX_LATEST | ConvertFrom-Json)',
|
||||
' }',
|
||||
" elseif ($Uri -like '*/releases?per_page=100') {",
|
||||
" if (-not $env:HX_LIST) { throw 'mock 404' }",
|
||||
' return ($env:HX_LIST | ConvertFrom-Json)',
|
||||
' }',
|
||||
" elseif ($Uri -like '*/releases/tags/*') {",
|
||||
" $t = $Uri -replace '.*/releases/tags/', ''",
|
||||
" $k = 'HX_TAG_' + ($t -replace '[.\\-]', '_')",
|
||||
' $v = [Environment]::GetEnvironmentVariable($k)',
|
||||
" if (-not $v) { throw 'mock 404' }",
|
||||
' return ($v | ConvertFrom-Json)',
|
||||
' }',
|
||||
' throw "unexpected uri $Uri"',
|
||||
'}',
|
||||
].join('\n')
|
||||
|
||||
type Run = { code: number; stdout: string; stderr: string }
|
||||
|
||||
function runPwsh(body: string, env: Record<string, string> = {}): Run {
|
||||
const script = `$ErrorActionPreference='Stop'\n${STUB}\n. '${SCRIPT}'\n${body}`
|
||||
const r = spawnSync('pwsh', ['-NoProfile', '-NonInteractive', '-Command', script], {
|
||||
encoding: 'utf8',
|
||||
env: {
|
||||
...process.env,
|
||||
DIFYCTL_INSTALL_LIB: '1',
|
||||
DIFY_VERSION: '',
|
||||
DIFYCTL_VERSION: '',
|
||||
LOCALAPPDATA: process.env.LOCALAPPDATA || '/tmp',
|
||||
TEMP: process.env.TEMP || '/tmp',
|
||||
...env,
|
||||
},
|
||||
})
|
||||
return { code: r.status ?? 1, stdout: (r.stdout ?? '').trim(), stderr: r.stderr ?? '' }
|
||||
}
|
||||
|
||||
const REL_1142 = JSON.stringify({
|
||||
tag_name: '1.14.2',
|
||||
assets: [{ name: 'difyctl-v0.2.0-windows-x64.exe' }],
|
||||
})
|
||||
const REL_1150 = JSON.stringify({
|
||||
tag_name: '1.15.0',
|
||||
assets: [{ name: 'difyctl-v0.3.0-windows-x64.exe' }],
|
||||
})
|
||||
const LIST_NEWEST_FIRST = JSON.stringify([
|
||||
{ tag_name: '1.15.0', assets: [{ name: 'difyctl-v0.3.0-windows-x64.exe' }] },
|
||||
{ tag_name: '1.14.2', assets: [{ name: 'difyctl-v0.2.0-windows-x64.exe' }] },
|
||||
])
|
||||
|
||||
describe.skipIf(!PWSH)('install.ps1 Get-AssetSemver', () => {
|
||||
it('extracts the version from a windows .exe asset name', () => {
|
||||
const r = runPwsh("(Get-AssetSemver 'difyctl-v0.2.0-windows-x64.exe').Version")
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('0.2.0')
|
||||
})
|
||||
|
||||
it('extracts a prerelease version and its rc number', () => {
|
||||
const r = runPwsh(
|
||||
'$a = Get-AssetSemver \'difyctl-v0.1.0-rc.1-windows-x64.exe\'; "$($a.Version) $($a.Rc)"',
|
||||
)
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('0.1.0-rc.1 1')
|
||||
})
|
||||
|
||||
it('rejects a non-windows asset (returns null)', () => {
|
||||
const r = runPwsh(
|
||||
"if ($null -eq (Get-AssetSemver 'difyctl-v0.2.0-linux-x64')) { 'NULL' } else { 'OBJ' }",
|
||||
)
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('NULL')
|
||||
})
|
||||
|
||||
it('rejects a malformed core version (returns null)', () => {
|
||||
const r = runPwsh(
|
||||
"if ($null -eq (Get-AssetSemver 'difyctl-vx.y.z-windows-x64.exe')) { 'NULL' } else { 'OBJ' }",
|
||||
)
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('NULL')
|
||||
})
|
||||
})
|
||||
|
||||
describe.skipIf(!PWSH)('install.ps1 Select-Asset', () => {
|
||||
it('picks the highest semver among several windows builds', () => {
|
||||
const rel = JSON.stringify({
|
||||
assets: [
|
||||
{ name: 'difyctl-v0.2.0-windows-x64.exe' },
|
||||
{ name: 'difyctl-v0.10.0-windows-x64.exe' },
|
||||
{ name: 'difyctl-v0.9.0-windows-x64.exe' },
|
||||
],
|
||||
})
|
||||
const r = runPwsh(`(Select-Asset ('${rel}' | ConvertFrom-Json)).Version`)
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('0.10.0')
|
||||
})
|
||||
|
||||
it('prefers the stable release over an rc of the same core', () => {
|
||||
const rel = JSON.stringify({
|
||||
assets: [
|
||||
{ name: 'difyctl-v0.2.0-rc.1-windows-x64.exe' },
|
||||
{ name: 'difyctl-v0.2.0-windows-x64.exe' },
|
||||
],
|
||||
})
|
||||
const r = runPwsh(`(Select-Asset ('${rel}' | ConvertFrom-Json)).Version`)
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('0.2.0')
|
||||
})
|
||||
|
||||
it('ignores checksums and non-windows assets', () => {
|
||||
const rel = JSON.stringify({
|
||||
assets: [
|
||||
{ name: 'difyctl-v0.2.0-linux-x64' },
|
||||
{ name: 'difyctl-v0.2.0-checksums.txt' },
|
||||
{ name: 'difyctl-v0.2.0-windows-x64.exe' },
|
||||
{ name: 'some-other-asset.zip' },
|
||||
],
|
||||
})
|
||||
const r = runPwsh(`(Select-Asset ('${rel}' | ConvertFrom-Json)).Name`)
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('difyctl-v0.2.0-windows-x64.exe')
|
||||
})
|
||||
|
||||
it('yields null when no windows asset is present', () => {
|
||||
const rel = JSON.stringify({ assets: [{ name: 'difyctl-v0.2.0-linux-x64' }] })
|
||||
const r = runPwsh(
|
||||
`if ($null -eq (Select-Asset ('${rel}' | ConvertFrom-Json))) { 'NULL' } else { 'OBJ' }`,
|
||||
)
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('NULL')
|
||||
})
|
||||
})
|
||||
|
||||
describe.skipIf(!PWSH)('install.ps1 Resolve-Release', () => {
|
||||
it('DIFY_VERSION pins the release directly', () => {
|
||||
const r = runPwsh('(Resolve-Release).tag_name', {
|
||||
DIFY_VERSION: '1.14.2',
|
||||
HX_TAG_1_14_2: REL_1142,
|
||||
})
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('1.14.2')
|
||||
})
|
||||
|
||||
it('DIFY_VERSION that does not exist throws a clear message', () => {
|
||||
const r = runPwsh('(Resolve-Release).tag_name', { DIFY_VERSION: '9.9.9' })
|
||||
expect(r.code).not.toBe(0)
|
||||
expect(r.stderr).toContain('Dify release 9.9.9 not found')
|
||||
})
|
||||
|
||||
it('blank resolves to the latest release', () => {
|
||||
const r = runPwsh('(Resolve-Release).tag_name', { HX_LATEST: REL_1150 })
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('1.15.0')
|
||||
})
|
||||
|
||||
it('blank throws when the latest query fails (no silent fallback)', () => {
|
||||
const r = runPwsh('(Resolve-Release).tag_name')
|
||||
expect(r.code).not.toBe(0)
|
||||
expect(r.stderr).toContain('failed to query latest Dify release')
|
||||
})
|
||||
|
||||
it('DIFYCTL_VERSION resolves to the release hosting that build', () => {
|
||||
const r = runPwsh('(Resolve-Release).tag_name', {
|
||||
DIFYCTL_VERSION: '0.2.0',
|
||||
HX_LIST: LIST_NEWEST_FIRST,
|
||||
})
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('1.14.2')
|
||||
})
|
||||
|
||||
it('DIFYCTL_VERSION not hosted anywhere throws', () => {
|
||||
const r = runPwsh('(Resolve-Release).tag_name', {
|
||||
DIFYCTL_VERSION: '9.9.9',
|
||||
HX_LIST: LIST_NEWEST_FIRST,
|
||||
})
|
||||
expect(r.code).not.toBe(0)
|
||||
expect(r.stderr).toContain('difyctl 9.9.9 not found on any Dify release')
|
||||
})
|
||||
})
|
||||
|
||||
describe.skipIf(!PWSH)('install.ps1 Find-ReleaseForDifyctl', () => {
|
||||
it('returns the newest release whose assets host the wanted build', () => {
|
||||
const r = runPwsh("(Find-ReleaseForDifyctl '0.2.0').tag_name", { HX_LIST: LIST_NEWEST_FIRST })
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('1.14.2')
|
||||
})
|
||||
|
||||
it('returns nothing when no release hosts the wanted build', () => {
|
||||
const r = runPwsh(
|
||||
"$x = Find-ReleaseForDifyctl '9.9.9'; if ($null -eq $x) { 'NULL' } else { $x.tag_name }",
|
||||
{ HX_LIST: LIST_NEWEST_FIRST },
|
||||
)
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('NULL')
|
||||
})
|
||||
})
|
||||
|
||||
// Build a fake ErrorRecord whose Exception.Response is a real HttpResponseMessage
|
||||
// carrying the given status + headers, matching what Get-RateLimitInfo inspects.
|
||||
function fakeErr(status: number, headers: Record<string, string>): string {
|
||||
const adds = Object.entries(headers)
|
||||
.map(([k, v]) => `$resp.Headers.TryAddWithoutValidation('${k}','${v}') | Out-Null`)
|
||||
.join('\n')
|
||||
return [
|
||||
`$resp = [System.Net.Http.HttpResponseMessage]::new([System.Net.HttpStatusCode]${status})`,
|
||||
adds,
|
||||
'$err = [pscustomobject]@{ Exception = [pscustomobject]@{ Response = $resp } }',
|
||||
].join('\n')
|
||||
}
|
||||
|
||||
const futureReset = String(Math.floor(Date.now() / 1000) + 1800)
|
||||
|
||||
describe.skipIf(!PWSH)('install.ps1 rate limit', () => {
|
||||
it('classifies a 403 with x-ratelimit-remaining:0 as rate-limited, returning the reset', () => {
|
||||
const r =
|
||||
runPwsh(`${fakeErr(403, { 'x-ratelimit-remaining': '0', 'x-ratelimit-reset': futureReset })}
|
||||
$i = Get-RateLimitInfo $err; if ($null -eq $i) { 'NULL' } else { $i.Reset }`)
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe(futureReset)
|
||||
})
|
||||
|
||||
it('does not classify a 403 with remaining tokens as rate-limited', () => {
|
||||
const r = runPwsh(`${fakeErr(403, { 'x-ratelimit-remaining': '59' })}
|
||||
$i = Get-RateLimitInfo $err; if ($null -eq $i) { 'NULL' } else { 'LIMITED' }`)
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('NULL')
|
||||
})
|
||||
|
||||
it('always treats a 429 as rate-limited', () => {
|
||||
const r = runPwsh(`${fakeErr(429, {})}
|
||||
$i = Get-RateLimitInfo $err; if ($null -eq $i) { 'NULL' } else { 'LIMITED' }`)
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('LIMITED')
|
||||
})
|
||||
|
||||
it('returns null for an error without a response (e.g. a plain string throw)', () => {
|
||||
const r = runPwsh(
|
||||
"$err = [pscustomobject]@{ Exception = [pscustomobject]@{} }; if ($null -eq (Get-RateLimitInfo $err)) { 'NULL' } else { 'OBJ' }",
|
||||
)
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('NULL')
|
||||
})
|
||||
|
||||
it('Write-RateLimitHint prints cause, ETA, and remediation to stderr', () => {
|
||||
const r = runPwsh(`Write-RateLimitHint '${futureReset}'`)
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stderr).toContain('rate limit exceeded')
|
||||
expect(r.stderr).toContain('resets in ~')
|
||||
expect(r.stderr).toContain('GITHUB_TOKEN')
|
||||
})
|
||||
|
||||
it('Write-RateLimitHint omits the ETA line when the reset epoch is missing', () => {
|
||||
const r = runPwsh("Write-RateLimitHint ''")
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stderr).toContain('rate limit exceeded')
|
||||
expect(r.stderr).not.toContain('resets in ~')
|
||||
})
|
||||
|
||||
it('sends an Authorization header when GITHUB_TOKEN is set', () => {
|
||||
const r = runPwsh('$headers.Authorization', { GITHUB_TOKEN: 'ghp_secret123' })
|
||||
expect(r.code).toBe(0)
|
||||
expect(r.stdout).toBe('Bearer ghp_secret123')
|
||||
})
|
||||
|
||||
it('Resolve-Release surfaces the rate-limit hint and exits, not "not found"', () => {
|
||||
const stub = `function Invoke-RestMethod { throw [Microsoft.PowerShell.Commands.HttpResponseException]::new('rate limited', $script:rlResp) }`
|
||||
const setup = `$script:rlResp = [System.Net.Http.HttpResponseMessage]::new([System.Net.HttpStatusCode]403)
|
||||
$script:rlResp.Headers.TryAddWithoutValidation('x-ratelimit-remaining','0') | Out-Null
|
||||
$script:rlResp.Headers.TryAddWithoutValidation('x-ratelimit-reset','${futureReset}') | Out-Null`
|
||||
const r = runPwsh(`${setup}\n${stub}\nResolve-Release`, { DIFY_VERSION: '1.15.0' })
|
||||
expect(r.code).not.toBe(0)
|
||||
expect(r.stderr).toContain('rate limit exceeded')
|
||||
expect(r.stderr).not.toContain('not found')
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,78 @@
|
||||
import { assetNameFor } from './release-rules.mjs'
|
||||
|
||||
// checksums lines are "<sha256> <assetName>"
|
||||
export function parseChecksums(text) {
|
||||
const map = new Map()
|
||||
for (const line of text.split('\n')) {
|
||||
const m = line.match(/^([0-9a-f]{64})\s+(\S+)$/i)
|
||||
if (m) map.set(m[2], m[1])
|
||||
}
|
||||
return map
|
||||
}
|
||||
|
||||
// Newline-delimited dir names of binaries that still exist in R2.
|
||||
export function parseDirList(text) {
|
||||
const set = new Set()
|
||||
for (const line of text.split('\n')) {
|
||||
const d = line.trim()
|
||||
if (d) set.add(d)
|
||||
}
|
||||
return set
|
||||
}
|
||||
|
||||
export function resolveTargets(release, version, shas) {
|
||||
const targets = []
|
||||
const missing = []
|
||||
for (const t of release.targets) {
|
||||
const asset = assetNameFor(release, version, t.id)
|
||||
const sha = shas.get(asset)
|
||||
if (sha) targets.push({ id: t.id, asset, sha })
|
||||
else missing.push(asset)
|
||||
}
|
||||
return { targets, missing }
|
||||
}
|
||||
|
||||
export function renderManifest({
|
||||
binName,
|
||||
channel,
|
||||
version,
|
||||
commit,
|
||||
buildDate,
|
||||
compat,
|
||||
baseUrl,
|
||||
targets,
|
||||
}) {
|
||||
const head = {
|
||||
schema: 1,
|
||||
name: binName,
|
||||
channel,
|
||||
version,
|
||||
commit,
|
||||
buildDate,
|
||||
compat: { minDify: compat.minDify, maxDify: compat.maxDify },
|
||||
baseUrl,
|
||||
}
|
||||
const headLines = Object.entries(head)
|
||||
.map(([k, v]) => ` ${JSON.stringify(k)}: ${JSON.stringify(v)}`)
|
||||
.join(',\n')
|
||||
const targetLines = targets
|
||||
.map(
|
||||
(t) =>
|
||||
` ${JSON.stringify(t.id)}: { "asset": ${JSON.stringify(t.asset)}, "sha256": ${JSON.stringify(t.sha)} }`,
|
||||
)
|
||||
.join(',\n')
|
||||
return `{\n${headLines},\n "targets": {\n${targetLines}\n }\n}\n`
|
||||
}
|
||||
|
||||
// `current` is the parsed ledger or null (first publish). `existingDirs` is a
|
||||
// Set of build dirs still present in R2, or null when the caller could not list
|
||||
// them — lifecycle/TTL on the bin prefix is the only deletion mechanism, so the
|
||||
// ledger must never advertise a build whose binary is gone. The new build is
|
||||
// always kept (just uploaded). No count cap.
|
||||
export function buildIndex({ channel, version, commit, buildDate, current, existingDirs }) {
|
||||
const entry = { version, commit, buildDate, dir: version }
|
||||
const kept = (current?.builds ?? []).filter((b) => b.version !== entry.version)
|
||||
let builds = [entry, ...kept]
|
||||
if (existingDirs) builds = builds.filter((b) => b.dir === entry.dir || existingDirs.has(b.dir))
|
||||
return { schema: 1, channel, updated: buildDate, builds }
|
||||
}
|
||||
@@ -0,0 +1,155 @@
|
||||
import { describe, expect, it } from 'vitest'
|
||||
import {
|
||||
buildIndex,
|
||||
parseChecksums,
|
||||
parseDirList,
|
||||
renderManifest,
|
||||
resolveTargets,
|
||||
} from './edge-manifest.mjs'
|
||||
|
||||
const RELEASE = {
|
||||
tagPrefix: 'testctl-v',
|
||||
binName: 'testctl',
|
||||
checksumsSuffix: '.sums',
|
||||
targets: [
|
||||
{ id: 'linux-x64', bunTarget: 'bun-linux-x64', exe: false },
|
||||
{ id: 'windows-x64', bunTarget: 'bun-windows-x64', exe: true },
|
||||
],
|
||||
}
|
||||
|
||||
const VERSION = '7.7.7-edge.2fd7b82'
|
||||
const SHA_LINUX = 'a'.repeat(64)
|
||||
const SHA_WINDOWS = 'b'.repeat(64)
|
||||
|
||||
const CHECKSUMS = [
|
||||
`${SHA_LINUX} testctl-v${VERSION}-linux-x64`,
|
||||
`${SHA_WINDOWS} testctl-v${VERSION}-windows-x64.exe`,
|
||||
].join('\n')
|
||||
|
||||
describe('parseChecksums', () => {
|
||||
it('maps asset name to sha256', () => {
|
||||
const map = parseChecksums(CHECKSUMS)
|
||||
expect(map.get(`testctl-v${VERSION}-linux-x64`)).toBe(SHA_LINUX)
|
||||
expect(map.get(`testctl-v${VERSION}-windows-x64.exe`)).toBe(SHA_WINDOWS)
|
||||
})
|
||||
|
||||
it('ignores blank and malformed lines', () => {
|
||||
expect(parseChecksums('\nnot a checksum line\n\n').size).toBe(0)
|
||||
})
|
||||
})
|
||||
|
||||
describe('parseDirList', () => {
|
||||
it('collects non-blank trimmed lines', () => {
|
||||
expect([...parseDirList(' a \n\n b\n')]).toEqual(['a', 'b'])
|
||||
})
|
||||
|
||||
it('yields an empty set for empty input', () => {
|
||||
expect(parseDirList('').size).toBe(0)
|
||||
})
|
||||
})
|
||||
|
||||
describe('resolveTargets', () => {
|
||||
it('pairs every target with its sha', () => {
|
||||
const { targets, missing } = resolveTargets(RELEASE, VERSION, parseChecksums(CHECKSUMS))
|
||||
expect(missing).toEqual([])
|
||||
expect(targets).toEqual([
|
||||
{ id: 'linux-x64', asset: `testctl-v${VERSION}-linux-x64`, sha: SHA_LINUX },
|
||||
{ id: 'windows-x64', asset: `testctl-v${VERSION}-windows-x64.exe`, sha: SHA_WINDOWS },
|
||||
])
|
||||
})
|
||||
|
||||
it('reports assets with no checksum', () => {
|
||||
const partial = parseChecksums(`${SHA_LINUX} testctl-v${VERSION}-linux-x64`)
|
||||
const { targets, missing } = resolveTargets(RELEASE, VERSION, partial)
|
||||
expect(targets).toHaveLength(1)
|
||||
expect(missing).toEqual([`testctl-v${VERSION}-windows-x64.exe`])
|
||||
})
|
||||
})
|
||||
|
||||
describe('renderManifest', () => {
|
||||
const manifest = () =>
|
||||
renderManifest({
|
||||
binName: RELEASE.binName,
|
||||
channel: 'edge',
|
||||
version: VERSION,
|
||||
commit: 'abc1234',
|
||||
buildDate: '2026-06-14T12:00:00Z',
|
||||
compat: { minDify: '2.0.0', maxDify: '2.5.0' },
|
||||
baseUrl: 'https://example.r2.dev/testctl/edge',
|
||||
targets: resolveTargets(RELEASE, VERSION, parseChecksums(CHECKSUMS)).targets,
|
||||
})
|
||||
|
||||
it('emits the pointer fields', () => {
|
||||
const json = JSON.parse(manifest())
|
||||
expect(json).toMatchObject({
|
||||
schema: 1,
|
||||
name: 'testctl',
|
||||
channel: 'edge',
|
||||
version: VERSION,
|
||||
commit: 'abc1234',
|
||||
buildDate: '2026-06-14T12:00:00Z',
|
||||
baseUrl: 'https://example.r2.dev/testctl/edge',
|
||||
})
|
||||
})
|
||||
|
||||
it('carries both compat bounds through unswapped', () => {
|
||||
expect(JSON.parse(manifest()).compat).toEqual({ minDify: '2.0.0', maxDify: '2.5.0' })
|
||||
})
|
||||
|
||||
it('lists each target with asset name and sha256', () => {
|
||||
expect(JSON.parse(manifest()).targets).toEqual({
|
||||
'linux-x64': { asset: `testctl-v${VERSION}-linux-x64`, sha256: SHA_LINUX },
|
||||
'windows-x64': { asset: `testctl-v${VERSION}-windows-x64.exe`, sha256: SHA_WINDOWS },
|
||||
})
|
||||
})
|
||||
|
||||
it('renders each target on a single line (install-r2.sh greps it)', () => {
|
||||
expect(manifest()).toMatch(/^ {4}"linux-x64": \{ "asset": ".*", "sha256": ".*" \}/m)
|
||||
})
|
||||
})
|
||||
|
||||
describe('buildIndex', () => {
|
||||
const B1 = { version: '7.7.7-edge.aaaaaaa', commit: 'aaaaaaa', buildDate: '2026-06-14T09:00:00Z' }
|
||||
const B2 = { version: '7.7.7-edge.bbbbbbb', commit: 'bbbbbbb', buildDate: '2026-06-14T10:00:00Z' }
|
||||
const entryOf = (b: typeof B1) => ({ ...b, dir: b.version })
|
||||
const build = (over = {}) =>
|
||||
buildIndex({ channel: 'edge', ...B2, current: null, existingDirs: null, ...over })
|
||||
|
||||
it('creates a fresh ledger when there is no current index', () => {
|
||||
const index = build()
|
||||
expect(index).toMatchObject({ schema: 1, channel: 'edge', updated: B2.buildDate })
|
||||
expect(index.builds).toEqual([entryOf(B2)])
|
||||
})
|
||||
|
||||
it('prepends the new build, newest first', () => {
|
||||
const index = build({ current: { builds: [entryOf(B1)] } })
|
||||
expect(index.builds).toEqual([entryOf(B2), entryOf(B1)])
|
||||
})
|
||||
|
||||
it('replaces an existing entry for the same version rather than duplicating', () => {
|
||||
const stale = { ...entryOf(B2), commit: 'oldsha' }
|
||||
const index = build({ current: { builds: [stale, entryOf(B1)] } })
|
||||
expect(index.builds).toEqual([entryOf(B2), entryOf(B1)])
|
||||
})
|
||||
|
||||
it('drops builds whose binaries no longer exist in R2', () => {
|
||||
const index = build({
|
||||
current: { builds: [entryOf(B1)] },
|
||||
existingDirs: new Set<string>(),
|
||||
})
|
||||
expect(index.builds).toEqual([entryOf(B2)])
|
||||
})
|
||||
|
||||
it('keeps builds that still exist in R2', () => {
|
||||
const index = build({
|
||||
current: { builds: [entryOf(B1)] },
|
||||
existingDirs: new Set([B1.version]),
|
||||
})
|
||||
expect(index.builds).toEqual([entryOf(B2), entryOf(B1)])
|
||||
})
|
||||
|
||||
it('reconciles nothing when the caller could not list R2', () => {
|
||||
const index = build({ current: { builds: [entryOf(B1)] }, existingDirs: null })
|
||||
expect(index.builds).toEqual([entryOf(B2), entryOf(B1)])
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,115 @@
|
||||
// Pure release naming and version rules. Nothing here reads a file, prints, or
|
||||
// exits — the CLI shells own all of that.
|
||||
|
||||
const BUN_TARGET_RE = /^bun-(linux|darwin|windows)-(x64|arm64)$/
|
||||
const SEMVER_CORE_LEN = 3
|
||||
|
||||
// Add channels here: { name, prerelease, versionForm }.
|
||||
export const CHANNELS = [
|
||||
{ name: 'stable', prerelease: false, versionForm: /^\d+\.\d+\.\d+(\+[0-9A-Z.-]+)?$/i },
|
||||
{ name: 'alpha', prerelease: true, versionForm: /^\d+\.\d+\.\d+-alpha(\.\d+)?$/ },
|
||||
{ name: 'rc', prerelease: true, versionForm: /^\d+\.\d+\.\d+-rc\.\d+$/ },
|
||||
{ name: 'edge', prerelease: true, versionForm: /^\d+\.\d+\.\d+-edge\.[0-9a-f]{7,40}$/ },
|
||||
]
|
||||
|
||||
export const EDGE_SHA_RE = /^[0-9a-f]{7,40}$/
|
||||
|
||||
export const channelByName = (name) => CHANNELS.find((c) => c.name === name)
|
||||
export const channelNames = () => CHANNELS.map((c) => c.name).join(', ')
|
||||
|
||||
export function versionCore(v) {
|
||||
return String(v).replace(/^v/, '').replace(/\+.*$/, '').split('-')[0]
|
||||
}
|
||||
|
||||
// Compat ordering compares the numeric A.B.C core only: prerelease and build
|
||||
// suffixes are stripped first, so 1.16.0-rc.1 ranks equal to 1.16.0. This
|
||||
// matches the runtime check in src/version/compat.ts.
|
||||
function compareCore(a, b) {
|
||||
const A = versionCore(a).split('.').map(Number)
|
||||
const B = versionCore(b).split('.').map(Number)
|
||||
for (let i = 0; i < SEMVER_CORE_LEN; i++) {
|
||||
const x = A[i] ?? 0
|
||||
const y = B[i] ?? 0
|
||||
if (x !== y) return x < y ? -1 : 1
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
export function isWithinCompatWindow(version, { minDify, maxDify }) {
|
||||
return compareCore(version, minDify) >= 0 && compareCore(version, maxDify) <= 0
|
||||
}
|
||||
|
||||
// Returns a problem string if `version` cannot be resolved under `channel`,
|
||||
// else null.
|
||||
export function channelVersionProblem(version, channel) {
|
||||
if (typeof version !== 'string' || version.length === 0)
|
||||
return 'version must be a non-empty string'
|
||||
const ch = channelByName(channel)
|
||||
if (!ch) return `unknown channel: ${channel} (expected one of: ${channelNames()})`
|
||||
if (!ch.versionForm.test(version))
|
||||
return `version ${version} does not match the ${channel} channel form`
|
||||
return null
|
||||
}
|
||||
|
||||
// Returns a problem string if the edge base cannot be derived, else null.
|
||||
export function edgeVersionProblem(version, sha) {
|
||||
if (!EDGE_SHA_RE.test(sha ?? '')) return 'edge-version requires a git short sha (7-40 hex chars)'
|
||||
if (!/^\d+\.\d+\.\d+$/.test(versionCore(version)))
|
||||
return `cannot derive edge base from version: ${version}`
|
||||
return null
|
||||
}
|
||||
|
||||
export function edgeVersionFrom(version, sha) {
|
||||
return `${versionCore(version)}-edge.${sha}`
|
||||
}
|
||||
|
||||
export function tagName(release, version) {
|
||||
return `${release.tagPrefix}${version}`
|
||||
}
|
||||
|
||||
export function checksumsName(release, version) {
|
||||
return `${release.tagPrefix}${version}${release.checksumsSuffix}`
|
||||
}
|
||||
|
||||
export function assetNameFor(release, version, id) {
|
||||
const target = release.targets.find((t) => t.id === id)
|
||||
if (!target) return null
|
||||
return `${release.tagPrefix}${version}-${id}${target.exe ? '.exe' : ''}`
|
||||
}
|
||||
|
||||
export function compatProblems(compat) {
|
||||
const str = (v) => typeof v === 'string' && v.length > 0
|
||||
const problems = []
|
||||
if (!str(compat?.minDify)) problems.push('compat.minDify must be a non-empty string')
|
||||
if (!str(compat?.maxDify)) problems.push('compat.maxDify must be a non-empty string')
|
||||
if (problems.length === 0 && compareCore(compat.minDify, compat.maxDify) > 0)
|
||||
problems.push(
|
||||
`compat.minDify ${compat.minDify} must not exceed compat.maxDify ${compat.maxDify}`,
|
||||
)
|
||||
return problems
|
||||
}
|
||||
|
||||
export function releaseConfigProblems(release) {
|
||||
const problems = []
|
||||
const str = (v) => typeof v === 'string' && v.length > 0
|
||||
if (!str(release.tagPrefix)) problems.push('tagPrefix must be a non-empty string')
|
||||
if (!str(release.binName)) problems.push('binName must be a non-empty string')
|
||||
if (!str(release.checksumsSuffix)) problems.push('checksumsSuffix must be a non-empty string')
|
||||
if (!Array.isArray(release.targets) || release.targets.length === 0) {
|
||||
problems.push('targets must be a non-empty array')
|
||||
return problems
|
||||
}
|
||||
const seen = new Set()
|
||||
for (const t of release.targets) {
|
||||
const label = t?.id ?? JSON.stringify(t)
|
||||
if (!str(t?.id)) problems.push(`target ${label}: id must be a non-empty string`)
|
||||
else if (seen.has(t.id)) problems.push(`duplicate target id: ${t.id}`)
|
||||
else seen.add(t.id)
|
||||
if (!str(t?.bunTarget) || !BUN_TARGET_RE.test(t.bunTarget))
|
||||
problems.push(`target ${label}: bunTarget must match ${BUN_TARGET_RE}`)
|
||||
if (typeof t?.exe !== 'boolean') problems.push(`target ${label}: exe must be a boolean`)
|
||||
else if (str(t?.bunTarget) && t.exe !== t.bunTarget.startsWith('bun-windows-'))
|
||||
problems.push(`target ${label}: exe must be true iff bunTarget is bun-windows-*`)
|
||||
}
|
||||
return problems
|
||||
}
|
||||
@@ -0,0 +1,183 @@
|
||||
import { describe, expect, it } from 'vitest'
|
||||
import {
|
||||
assetNameFor,
|
||||
channelVersionProblem,
|
||||
checksumsName,
|
||||
compatProblems,
|
||||
edgeVersionFrom,
|
||||
edgeVersionProblem,
|
||||
isWithinCompatWindow,
|
||||
releaseConfigProblems,
|
||||
tagName,
|
||||
} from './release-rules.mjs'
|
||||
|
||||
const RELEASE = {
|
||||
tagPrefix: 'testctl-v',
|
||||
binName: 'testctl',
|
||||
checksumsSuffix: '.sums',
|
||||
targets: [
|
||||
{ id: 'linux-x64', bunTarget: 'bun-linux-x64', exe: false },
|
||||
{ id: 'windows-x64', bunTarget: 'bun-windows-x64', exe: true },
|
||||
],
|
||||
}
|
||||
|
||||
const WINDOW = { minDify: '2.0.0', maxDify: '2.5.0' }
|
||||
|
||||
describe('isWithinCompatWindow', () => {
|
||||
it.each([
|
||||
['2.3.0', true, 'inside the window'],
|
||||
['2.0.0', true, 'the inclusive lower bound'],
|
||||
['2.5.0', true, 'the inclusive upper bound'],
|
||||
['v2.3.0', true, 'a v-prefixed tag'],
|
||||
['1.9.9', false, 'below the lower bound'],
|
||||
['2.5.1', false, 'above the upper bound'],
|
||||
])('%s -> %s (%s)', (version, expected) => {
|
||||
expect(isWithinCompatWindow(version, WINDOW)).toBe(expected)
|
||||
})
|
||||
|
||||
it.each([
|
||||
['2.0.0-rc.1', true, 'a prerelease of the lower bound ranks equal to it'],
|
||||
['2.5.0-rc.1', true, 'a prerelease of the upper bound ranks equal to it'],
|
||||
['2.3.0-alpha.7+build9', true, 'both suffixes stripped before comparing'],
|
||||
['2.5.0+build123', true, 'build metadata on the upper bound'],
|
||||
['2.5.1-rc.1', false, 'a prerelease whose core is above the window'],
|
||||
['1.9.9-rc.1', false, 'a prerelease whose core is below the window'],
|
||||
['2.5.1+build123', false, 'build metadata out of range'],
|
||||
])('ignores suffixes: %s -> %s (%s)', (version, expected) => {
|
||||
expect(isWithinCompatWindow(version, WINDOW)).toBe(expected)
|
||||
})
|
||||
})
|
||||
|
||||
describe('channelVersionProblem', () => {
|
||||
it.each([
|
||||
['1.2.3', 'stable'],
|
||||
['1.2.3+build', 'stable'],
|
||||
['1.2.3-alpha', 'alpha'],
|
||||
['1.2.3-alpha.4', 'alpha'],
|
||||
['1.2.3-rc.4', 'rc'],
|
||||
['1.2.3-edge.2fd7b82', 'edge'],
|
||||
])('accepts %s on the %s channel', (version, channel) => {
|
||||
expect(channelVersionProblem(version, channel)).toBeNull()
|
||||
})
|
||||
|
||||
it.each([
|
||||
['1.2.3-rc.1', 'edge'],
|
||||
['1.2.3-alpha', 'stable'],
|
||||
['1.2.3', 'rc'],
|
||||
])('rejects %s on the %s channel', (version, channel) => {
|
||||
expect(channelVersionProblem(version, channel)).toMatch(/does not match the .* channel form/)
|
||||
})
|
||||
|
||||
it('rejects an unknown channel', () => {
|
||||
expect(channelVersionProblem('1.2.3', 'nightly')).toMatch(/unknown channel/)
|
||||
})
|
||||
|
||||
it('rejects an empty version', () => {
|
||||
expect(channelVersionProblem('', 'stable')).toMatch(/non-empty string/)
|
||||
})
|
||||
})
|
||||
|
||||
describe('edge version derivation', () => {
|
||||
it('derives <core>-edge.<sha>, dropping the prerelease', () => {
|
||||
expect(edgeVersionProblem('3.4.5-alpha', '2fd7b82')).toBeNull()
|
||||
expect(edgeVersionFrom('3.4.5-alpha', '2fd7b82')).toBe('3.4.5-edge.2fd7b82')
|
||||
})
|
||||
|
||||
it('accepts a 40-char sha', () => {
|
||||
const sha = '2fd7b829e1f0aaaabbbbccccddddeeeeffff0000'
|
||||
expect(edgeVersionProblem('3.4.5-alpha', sha)).toBeNull()
|
||||
expect(edgeVersionFrom('3.4.5-alpha', sha)).toBe(`3.4.5-edge.${sha}`)
|
||||
})
|
||||
|
||||
it.each([['nothex!'], [''], ['abc'], [undefined]])('rejects sha %s', (sha) => {
|
||||
expect(edgeVersionProblem('3.4.5-alpha', sha)).toMatch(/git short sha/)
|
||||
})
|
||||
|
||||
it('rejects a version with no derivable core', () => {
|
||||
expect(edgeVersionProblem('not-a-version', '2fd7b82')).toMatch(/cannot derive edge base/)
|
||||
})
|
||||
})
|
||||
|
||||
describe('artifact names', () => {
|
||||
it('builds the tag and checksums names from the prefix', () => {
|
||||
expect(tagName(RELEASE, '9.9.9')).toBe('testctl-v9.9.9')
|
||||
expect(checksumsName(RELEASE, '9.9.9')).toBe('testctl-v9.9.9.sums')
|
||||
})
|
||||
|
||||
it('appends .exe only for targets flagged exe', () => {
|
||||
expect(assetNameFor(RELEASE, '9.9.9', 'linux-x64')).toBe('testctl-v9.9.9-linux-x64')
|
||||
expect(assetNameFor(RELEASE, '9.9.9', 'windows-x64')).toBe('testctl-v9.9.9-windows-x64.exe')
|
||||
})
|
||||
|
||||
it('returns null for an unknown target id', () => {
|
||||
expect(assetNameFor(RELEASE, '9.9.9', 'solaris-sparc')).toBeNull()
|
||||
})
|
||||
})
|
||||
|
||||
describe('compatProblems', () => {
|
||||
it('accepts a well-formed window', () => {
|
||||
expect(compatProblems(WINDOW)).toEqual([])
|
||||
expect(compatProblems({ minDify: '2.0.0', maxDify: '2.0.0' })).toEqual([])
|
||||
})
|
||||
|
||||
it.each([
|
||||
[undefined, /minDify must be a non-empty string/],
|
||||
[{ maxDify: '2.5.0' }, /minDify must be a non-empty string/],
|
||||
[{ minDify: '2.0.0' }, /maxDify must be a non-empty string/],
|
||||
[{ minDify: '', maxDify: '' }, /minDify must be a non-empty string/],
|
||||
])('flags a missing bound', (compat, expected) => {
|
||||
expect(compatProblems(compat).join('\n')).toMatch(expected)
|
||||
})
|
||||
|
||||
it('compares bounds by core, so suffixes do not make a window inverted', () => {
|
||||
expect(compatProblems({ minDify: '2.0.0-rc.1', maxDify: '2.0.0' })).toEqual([])
|
||||
})
|
||||
|
||||
it('flags an inverted window', () => {
|
||||
expect(compatProblems({ minDify: '2.5.0', maxDify: '2.0.0' }).join('\n')).toMatch(
|
||||
/must not exceed/,
|
||||
)
|
||||
})
|
||||
})
|
||||
|
||||
describe('releaseConfigProblems', () => {
|
||||
it('accepts a well-formed config', () => {
|
||||
expect(releaseConfigProblems(RELEASE)).toEqual([])
|
||||
})
|
||||
|
||||
it.each([
|
||||
[{ ...RELEASE, tagPrefix: '' }, /tagPrefix must be a non-empty string/],
|
||||
[{ ...RELEASE, binName: '' }, /binName must be a non-empty string/],
|
||||
[{ ...RELEASE, checksumsSuffix: '' }, /checksumsSuffix must be a non-empty string/],
|
||||
[{ ...RELEASE, targets: [] }, /targets must be a non-empty array/],
|
||||
[{ ...RELEASE, targets: undefined }, /targets must be a non-empty array/],
|
||||
])('flags a malformed field', (release, expected) => {
|
||||
expect(releaseConfigProblems(release).join('\n')).toMatch(expected)
|
||||
})
|
||||
|
||||
it('flags a duplicate target id', () => {
|
||||
const release = { ...RELEASE, targets: [RELEASE.targets[0], { ...RELEASE.targets[0] }] }
|
||||
expect(releaseConfigProblems(release).join('\n')).toMatch(/duplicate target id: linux-x64/)
|
||||
})
|
||||
|
||||
it('flags a bunTarget that is not a bun triple', () => {
|
||||
const release = { ...RELEASE, targets: [{ id: 'a', bunTarget: 'linux-x64', exe: false }] }
|
||||
expect(releaseConfigProblems(release).join('\n')).toMatch(/bunTarget must match/)
|
||||
})
|
||||
|
||||
it('flags exe disagreeing with a windows bunTarget', () => {
|
||||
const release = {
|
||||
...RELEASE,
|
||||
targets: [{ id: 'windows-x64', bunTarget: 'bun-windows-x64', exe: false }],
|
||||
}
|
||||
expect(releaseConfigProblems(release).join('\n')).toMatch(/exe must be true iff/)
|
||||
})
|
||||
|
||||
it('flags a non-boolean exe', () => {
|
||||
const release = {
|
||||
...RELEASE,
|
||||
targets: [{ id: 'linux-x64', bunTarget: 'bun-linux-x64', exe: 'no' }],
|
||||
}
|
||||
expect(releaseConfigProblems(release).join('\n')).toMatch(/exe must be a boolean/)
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,72 @@
|
||||
import { readFileSync } from 'node:fs'
|
||||
import { fileURLToPath } from 'node:url'
|
||||
import { describe, expect, it } from 'vitest'
|
||||
import { main } from './release-naming.mjs'
|
||||
|
||||
const env = Object.fromEntries(
|
||||
main(['github-env'])
|
||||
.split('\n')
|
||||
.filter(Boolean)
|
||||
.map((line: string) => line.split(/=(.*)/s).slice(0, 2)),
|
||||
)
|
||||
|
||||
// Install scripts run standalone on an end user's machine — no repo, no node —
|
||||
// so they hardcode artifact names instead of asking release-naming.mjs. These
|
||||
// checks are what keeps the two copies from drifting: change the naming config
|
||||
// and they fail until the installers are updated to match.
|
||||
const INSTALLERS = ['install-cli.sh', 'install-r2.sh', 'install.ps1', 'install-r2.ps1']
|
||||
const installerText = (name: string) =>
|
||||
readFileSync(fileURLToPath(new URL(`./${name}`, import.meta.url)), 'utf8')
|
||||
|
||||
const CHECKSUMS_SUFFIX = main(['checksums', '1.2.3']).trim().replace(`${env.tagPrefix}1.2.3`, '')
|
||||
|
||||
const EXE_TARGET_ID = main(['targets'])
|
||||
.trim()
|
||||
.split('\n')
|
||||
.map((line: string) => line.split('\t'))
|
||||
.find(([, , exe]: string[]) => exe === '1')?.[1]
|
||||
|
||||
describe('install scripts stay aligned with the release naming config', () => {
|
||||
it.each(INSTALLERS)('%s uses the configured tag prefix', (name) => {
|
||||
expect(installerText(name)).toContain(env.tagPrefix)
|
||||
})
|
||||
|
||||
it.each(INSTALLERS)('%s uses the configured checksums suffix', (name) => {
|
||||
expect(installerText(name)).toContain(CHECKSUMS_SUFFIX)
|
||||
})
|
||||
|
||||
it.each(['install.ps1', 'install-r2.ps1'])('%s targets the declared windows build', (name) => {
|
||||
expect(EXE_TARGET_ID).toBeTruthy()
|
||||
expect(installerText(name)).toContain(EXE_TARGET_ID)
|
||||
})
|
||||
})
|
||||
|
||||
describe('the shipped difyctl release config', () => {
|
||||
it('passes the release gate (`release-naming.mjs validate`)', () => {
|
||||
expect(main(['validate'])).toMatch(/^difyctl release valid:/)
|
||||
})
|
||||
|
||||
it('declares a usable, correctly ordered compat window', () => {
|
||||
expect(main(['compat-check', env.minDify])).toContain('compatible')
|
||||
expect(main(['compat-check', env.maxDify])).toContain('compatible')
|
||||
})
|
||||
|
||||
it('gates a Dify version outside the declared window', () => {
|
||||
const aboveMax = `${Number(env.maxDify.split('.')[0]) + 1}.0.0`
|
||||
expect(() => main(['compat-check', aboveMax])).toThrow('outside difyctl compatibility window')
|
||||
expect(() => main(['compat-check', '0.0.1'])).toThrow('outside difyctl compatibility window')
|
||||
})
|
||||
|
||||
it('declares a version valid for the channel it ships on', () => {
|
||||
expect(main(['validate-version', env.version, env.channel])).toContain('valid')
|
||||
})
|
||||
|
||||
it('can name an artifact for every target it declares', () => {
|
||||
const ids = main(['targets'])
|
||||
.trim()
|
||||
.split('\n')
|
||||
.map((line: string) => line.split('\t')[1])
|
||||
expect(ids.length).toBeGreaterThan(0)
|
||||
for (const id of ids) expect(main(['asset', env.version, id])).toContain(id)
|
||||
})
|
||||
})
|
||||
+64
-167
@@ -1,9 +1,8 @@
|
||||
#!/usr/bin/env node
|
||||
// release-naming.mjs — single source of truth for difyctl release artifact
|
||||
// names and version/channel rules. Reads DATA from cli/package.json
|
||||
// `difyctl.release` (plus `version` and `difyctl.channel`) and owns the name
|
||||
// FORMAT and the per-channel version form. Producer scripts call this;
|
||||
// `validate` is the release gate.
|
||||
// `difyctl.release` (plus `version` and `difyctl.channel`); the name FORMAT and
|
||||
// the per-channel version form live in lib/release-rules.mjs. Producer scripts
|
||||
// call this; `validate` is the release gate.
|
||||
//
|
||||
// Subcommands:
|
||||
// tag <version> -> <tagPrefix><version>
|
||||
@@ -19,104 +18,32 @@
|
||||
|
||||
import { readFileSync, realpathSync } from 'node:fs'
|
||||
import { fileURLToPath } from 'node:url'
|
||||
import {
|
||||
assetNameFor,
|
||||
channelByName,
|
||||
channelNames,
|
||||
CHANNELS,
|
||||
channelVersionProblem,
|
||||
checksumsName,
|
||||
compatProblems,
|
||||
edgeVersionFrom,
|
||||
edgeVersionProblem,
|
||||
isWithinCompatWindow,
|
||||
releaseConfigProblems,
|
||||
tagName,
|
||||
} from './lib/release-rules.mjs'
|
||||
|
||||
const BUN_TARGET_RE = /^bun-(linux|darwin|windows)-(x64|arm64)$/
|
||||
const SEMVER_CORE_LEN = 3
|
||||
const PKG_URL = new URL('../package.json', import.meta.url)
|
||||
|
||||
// Add channels here: { name, prerelease, versionForm }.
|
||||
const CHANNELS = [
|
||||
{ name: 'stable', prerelease: false, versionForm: /^\d+\.\d+\.\d+(\+[0-9A-Z.-]+)?$/i },
|
||||
{ name: 'alpha', prerelease: true, versionForm: /^\d+\.\d+\.\d+-alpha(\.\d+)?$/ },
|
||||
{ name: 'rc', prerelease: true, versionForm: /^\d+\.\d+\.\d+-rc\.\d+$/ },
|
||||
{ name: 'edge', prerelease: true, versionForm: /^\d+\.\d+\.\d+-edge\.[0-9a-f]{7,40}$/ },
|
||||
]
|
||||
class UsageError extends Error {}
|
||||
|
||||
const channelByName = (name) => CHANNELS.find((c) => c.name === name)
|
||||
const channelNames = () => CHANNELS.map((c) => c.name).join(', ')
|
||||
|
||||
function parsePrecedence(v) {
|
||||
const s = String(v).replace(/^v/, '').replace(/\+.*$/, '')
|
||||
const i = s.indexOf('-')
|
||||
const core = i === -1 ? s : s.slice(0, i)
|
||||
const pre = i === -1 ? '' : s.slice(i + 1)
|
||||
return { nums: core.split('.').map(Number), pre }
|
||||
// Arrow so TS infers `never`, keeping expression positions like `x ?? die(...)` typed.
|
||||
const die = (msg) => {
|
||||
throw new UsageError(msg)
|
||||
}
|
||||
|
||||
function versionCore(v) {
|
||||
return String(v).replace(/^v/, '').replace(/\+.*$/, '').split('-')[0]
|
||||
}
|
||||
|
||||
function edgeVersion(sha) {
|
||||
if (!/^[0-9a-f]{7,40}$/.test(sha ?? ''))
|
||||
die('edge-version requires a git short sha (7-40 hex chars)')
|
||||
const { version } = loadPkg()
|
||||
const core = versionCore(version)
|
||||
if (!/^\d+\.\d+\.\d+$/.test(core)) die(`cannot derive edge base from version: ${version}`)
|
||||
return `${core}-edge.${sha}`
|
||||
}
|
||||
|
||||
// Returns a problem string if `version` cannot be resolved under `channel`, else
|
||||
// null. Shared by validateVersionForChannel (die-now) and validateVersionChannel
|
||||
// (collect for the `validate` gate).
|
||||
function channelVersionProblem(version, channel) {
|
||||
if (typeof version !== 'string' || version.length === 0)
|
||||
return 'version must be a non-empty string'
|
||||
const ch = channelByName(channel)
|
||||
if (!ch) return `unknown channel: ${channel} (expected one of: ${channelNames()})`
|
||||
if (!ch.versionForm.test(version))
|
||||
return `version ${version} does not match the ${channel} channel form`
|
||||
return null
|
||||
}
|
||||
|
||||
function validateVersionForChannel(version, channelName) {
|
||||
const problem = channelVersionProblem(version, channelName)
|
||||
if (problem) die(problem)
|
||||
return `valid: ${version} is a ${channelName} version`
|
||||
}
|
||||
|
||||
function comparePre(a, b) {
|
||||
const aparts = a.split('.')
|
||||
const bparts = b.split('.')
|
||||
const len = Math.max(aparts.length, bparts.length)
|
||||
for (let i = 0; i < len; i++) {
|
||||
if (aparts[i] === undefined) return -1
|
||||
if (bparts[i] === undefined) return 1
|
||||
const an = /^\d+$/.test(aparts[i])
|
||||
const bn = /^\d+$/.test(bparts[i])
|
||||
if (an && bn) {
|
||||
const d = Number(aparts[i]) - Number(bparts[i])
|
||||
if (d !== 0) return d < 0 ? -1 : 1
|
||||
} else if (an !== bn) {
|
||||
return an ? -1 : 1
|
||||
} else if (aparts[i] !== bparts[i]) {
|
||||
return aparts[i] < bparts[i] ? -1 : 1
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
function comparePrecedence(a, b) {
|
||||
const A = parsePrecedence(a)
|
||||
const B = parsePrecedence(b)
|
||||
for (let i = 0; i < SEMVER_CORE_LEN; i++) {
|
||||
const x = A.nums[i] ?? 0
|
||||
const y = B.nums[i] ?? 0
|
||||
if (x !== y) return x < y ? -1 : 1
|
||||
}
|
||||
if (A.pre === B.pre) return 0
|
||||
if (A.pre === '') return 1
|
||||
if (B.pre === '') return -1
|
||||
return comparePre(A.pre, B.pre)
|
||||
}
|
||||
|
||||
function die(msg) {
|
||||
process.stderr.write(`release-naming: ${msg}\n`)
|
||||
process.exit(1)
|
||||
}
|
||||
|
||||
function loadPkg() {
|
||||
const pkgUrl = new URL('../package.json', import.meta.url)
|
||||
const pkg = JSON.parse(readFileSync(pkgUrl, 'utf8'))
|
||||
function loadPkg(pkgPath = PKG_URL) {
|
||||
const pkg = JSON.parse(readFileSync(pkgPath, 'utf8'))
|
||||
if (!pkg.difyctl?.release) die('cli/package.json missing difyctl.release')
|
||||
return {
|
||||
version: pkg.version,
|
||||
@@ -136,7 +63,7 @@ function githubEnv() {
|
||||
minDify: compat.minDify,
|
||||
maxDify: compat.maxDify,
|
||||
tagPrefix: release.tagPrefix,
|
||||
difyctlTag: `${release.tagPrefix}${version}`,
|
||||
difyctlTag: tagName(release, version),
|
||||
}
|
||||
return Object.entries(fields)
|
||||
.map(([k, v]) => `${k}=${v}`)
|
||||
@@ -148,58 +75,19 @@ function requireVersion(version) {
|
||||
return version
|
||||
}
|
||||
|
||||
function assetName(release, version, id) {
|
||||
const target = release.targets.find((t) => t.id === id)
|
||||
if (!target) die(`unknown target id: ${id}`)
|
||||
const suffix = target.exe ? '.exe' : ''
|
||||
return `${release.tagPrefix}${version}-${id}${suffix}`
|
||||
}
|
||||
|
||||
function validateRelease(release) {
|
||||
const problems = []
|
||||
const str = (v) => typeof v === 'string' && v.length > 0
|
||||
if (!str(release.tagPrefix)) problems.push('tagPrefix must be a non-empty string')
|
||||
if (!str(release.binName)) problems.push('binName must be a non-empty string')
|
||||
if (!str(release.checksumsSuffix)) problems.push('checksumsSuffix must be a non-empty string')
|
||||
if (!Array.isArray(release.targets) || release.targets.length === 0) {
|
||||
problems.push('targets must be a non-empty array')
|
||||
return problems
|
||||
}
|
||||
const seen = new Set()
|
||||
for (const t of release.targets) {
|
||||
const label = t?.id ?? JSON.stringify(t)
|
||||
if (!str(t?.id)) problems.push(`target ${label}: id must be a non-empty string`)
|
||||
else if (seen.has(t.id)) problems.push(`duplicate target id: ${t.id}`)
|
||||
else seen.add(t.id)
|
||||
if (!str(t?.bunTarget) || !BUN_TARGET_RE.test(t.bunTarget))
|
||||
problems.push(`target ${label}: bunTarget must match ${BUN_TARGET_RE}`)
|
||||
if (typeof t?.exe !== 'boolean') problems.push(`target ${label}: exe must be a boolean`)
|
||||
else if (str(t?.bunTarget) && t.exe !== t.bunTarget.startsWith('bun-windows-'))
|
||||
problems.push(`target ${label}: exe must be true iff bunTarget is bun-windows-*`)
|
||||
}
|
||||
return problems
|
||||
}
|
||||
|
||||
function validateVersionChannel(version, channel) {
|
||||
const problem = channelVersionProblem(version, channel)
|
||||
return problem ? [problem] : []
|
||||
}
|
||||
|
||||
function main(argv) {
|
||||
const [cmd, ...rest] = argv
|
||||
switch (cmd) {
|
||||
case 'tag':
|
||||
return `${loadPkg().release.tagPrefix}${requireVersion(rest[0])}`
|
||||
case 'asset':
|
||||
return assetName(
|
||||
loadPkg().release,
|
||||
requireVersion(rest[0]),
|
||||
rest[1] ?? die('target id is required'),
|
||||
)
|
||||
case 'checksums': {
|
||||
return tagName(loadPkg().release, requireVersion(rest[0]))
|
||||
case 'asset': {
|
||||
const { release } = loadPkg()
|
||||
return `${release.tagPrefix}${requireVersion(rest[0])}${release.checksumsSuffix}`
|
||||
const version = requireVersion(rest[0])
|
||||
const id = rest[1] ?? die('target id is required')
|
||||
return assetNameFor(release, version, id) ?? die(`unknown target id: ${id}`)
|
||||
}
|
||||
case 'checksums':
|
||||
return checksumsName(loadPkg().release, requireVersion(rest[0]))
|
||||
case 'tag-prefix':
|
||||
return loadPkg().release.tagPrefix
|
||||
case 'targets':
|
||||
@@ -215,10 +103,7 @@ function main(argv) {
|
||||
const difyVersion = requireVersion(rest[0])
|
||||
if (!compat.minDify || !compat.maxDify)
|
||||
die('cli/package.json missing difyctl.compat.minDify/maxDify')
|
||||
if (
|
||||
comparePrecedence(difyVersion, compat.minDify) < 0 ||
|
||||
comparePrecedence(difyVersion, compat.maxDify) > 0
|
||||
)
|
||||
if (!isWithinCompatWindow(difyVersion, compat))
|
||||
die(
|
||||
`Dify ${difyVersion} is outside difyctl compatibility window ${compat.minDify}..${compat.maxDify}; bump difyctl.compat in cli/package.json`,
|
||||
)
|
||||
@@ -230,19 +115,31 @@ function main(argv) {
|
||||
return String(ch.prerelease)
|
||||
}
|
||||
case 'validate': {
|
||||
const { version, channel, release } = loadPkg()
|
||||
const problems = [...validateRelease(release), ...validateVersionChannel(version, channel)]
|
||||
const { version, channel, compat, release } = loadPkg()
|
||||
const versionProblem = channelVersionProblem(version, channel)
|
||||
const problems = [
|
||||
...releaseConfigProblems(release),
|
||||
...compatProblems(compat),
|
||||
...(versionProblem ? [versionProblem] : []),
|
||||
]
|
||||
if (problems.length > 0)
|
||||
die(`invalid difyctl release config:\n - ${problems.join('\n - ')}`)
|
||||
return `difyctl release valid: version=${version} channel=${channel} targets=${release.targets.length}`
|
||||
}
|
||||
case 'edge-version':
|
||||
return edgeVersion(rest[0])
|
||||
case 'validate-version':
|
||||
return validateVersionForChannel(
|
||||
requireVersion(rest[0]),
|
||||
rest[1] ?? die('channel argument is required'),
|
||||
)
|
||||
case 'edge-version': {
|
||||
const { version } = loadPkg()
|
||||
const sha = rest[0]
|
||||
const problem = edgeVersionProblem(version, sha)
|
||||
if (problem) die(problem)
|
||||
return edgeVersionFrom(version, sha)
|
||||
}
|
||||
case 'validate-version': {
|
||||
const version = requireVersion(rest[0])
|
||||
const channel = rest[1] ?? die('channel argument is required')
|
||||
const problem = channelVersionProblem(version, channel)
|
||||
if (problem) die(problem)
|
||||
return `valid: ${version} is a ${channel} version`
|
||||
}
|
||||
default:
|
||||
die(`unknown subcommand: ${cmd ?? '(none)'}`)
|
||||
}
|
||||
@@ -250,14 +147,14 @@ function main(argv) {
|
||||
|
||||
const invokedDirectly =
|
||||
process.argv[1] && realpathSync(process.argv[1]) === fileURLToPath(import.meta.url)
|
||||
if (invokedDirectly) process.stdout.write(`${main(process.argv.slice(2))}\n`)
|
||||
|
||||
export {
|
||||
assetName,
|
||||
channelByName,
|
||||
CHANNELS,
|
||||
edgeVersion,
|
||||
loadPkg,
|
||||
validateVersionForChannel,
|
||||
versionCore,
|
||||
if (invokedDirectly) {
|
||||
try {
|
||||
process.stdout.write(`${main(process.argv.slice(2))}\n`)
|
||||
} catch (e) {
|
||||
if (!(e instanceof UsageError)) throw e
|
||||
process.stderr.write(`release-naming: ${e.message}\n`)
|
||||
process.exit(1)
|
||||
}
|
||||
}
|
||||
|
||||
export { loadPkg, main }
|
||||
|
||||
@@ -1,102 +1,77 @@
|
||||
import { execFileSync } from 'node:child_process'
|
||||
import { fileURLToPath } from 'node:url'
|
||||
import { describe, expect, it } from 'vitest'
|
||||
import { main } from './release-naming.mjs'
|
||||
|
||||
const SCRIPT = fileURLToPath(new URL('./release-naming.mjs', import.meta.url))
|
||||
|
||||
function run(args: string[]): { code: number; stdout: string; stderr: string } {
|
||||
try {
|
||||
const stdout = execFileSync('node', [SCRIPT, ...args], { encoding: 'utf8' })
|
||||
return { code: 0, stdout, stderr: '' }
|
||||
} catch (e) {
|
||||
const err = e as { status?: number; stdout?: string; stderr?: string }
|
||||
return { code: err.status ?? 1, stdout: err.stdout ?? '', stderr: err.stderr ?? '' }
|
||||
}
|
||||
}
|
||||
|
||||
describe('release-naming compat-check (compat 1.16.0..1.16.0)', () => {
|
||||
it('accepts a version inside the window', () => {
|
||||
expect(run(['compat-check', '1.16.0']).code).toBe(0)
|
||||
describe('release-naming argument handling', () => {
|
||||
it.each([
|
||||
[['tag'], 'version argument is required'],
|
||||
[['asset'], 'version argument is required'],
|
||||
[['checksums'], 'version argument is required'],
|
||||
[['compat-check'], 'version argument is required'],
|
||||
[['validate-version'], 'version argument is required'],
|
||||
[['prerelease'], 'channel argument is required'],
|
||||
[['edge-version'], 'git short sha'],
|
||||
[['bogus'], 'unknown subcommand'],
|
||||
[[], 'unknown subcommand'],
|
||||
])('rejects %j', (args, message) => {
|
||||
expect(() => main(args)).toThrow(message)
|
||||
})
|
||||
|
||||
it('accepts the inclusive lower bound', () => {
|
||||
expect(run(['compat-check', '1.16.0']).code).toBe(0)
|
||||
})
|
||||
|
||||
it('accepts the inclusive upper bound', () => {
|
||||
expect(run(['compat-check', '1.16.0']).code).toBe(0)
|
||||
})
|
||||
|
||||
it('accepts a v-prefixed tag', () => {
|
||||
expect(run(['compat-check', 'v1.16.0']).code).toBe(0)
|
||||
})
|
||||
|
||||
it('rejects a version below the lower bound', () => {
|
||||
expect(run(['compat-check', '1.15.9']).code).not.toBe(0)
|
||||
})
|
||||
|
||||
it('rejects a version above the upper bound', () => {
|
||||
expect(run(['compat-check', '1.16.1']).code).not.toBe(0)
|
||||
})
|
||||
|
||||
it('treats a prerelease of the bound as below it (1.16.0-rc1 < 1.16.0)', () => {
|
||||
expect(run(['compat-check', '1.16.0-rc1']).code).not.toBe(0)
|
||||
})
|
||||
|
||||
it('ignores build metadata on the bound (1.16.0+build == 1.16.0)', () => {
|
||||
expect(run(['compat-check', '1.16.0+build123']).code).toBe(0)
|
||||
})
|
||||
|
||||
it('ignores build metadata when out of range (1.16.1+build still rejected)', () => {
|
||||
expect(run(['compat-check', '1.16.1+build123']).code).not.toBe(0)
|
||||
})
|
||||
|
||||
it('requires a version argument', () => {
|
||||
expect(run(['compat-check']).code).not.toBe(0)
|
||||
it('rejects an unknown target id', () => {
|
||||
expect(() => main(['asset', '9.9.9', 'solaris-sparc'])).toThrow('unknown target id')
|
||||
})
|
||||
})
|
||||
|
||||
describe('release-naming github-env', () => {
|
||||
it('emits difyctlTag = tagPrefix + version', () => {
|
||||
const { stdout } = run(['github-env'])
|
||||
expect(stdout).toMatch(/^difyctlTag=difyctl-v0\.2\.0-alpha$/m)
|
||||
describe('release-naming output shape', () => {
|
||||
it('emits every key CI reads from github-env', () => {
|
||||
const out = main(['github-env'])
|
||||
for (const key of [
|
||||
'version',
|
||||
'channel',
|
||||
'prerelease',
|
||||
'minDify',
|
||||
'maxDify',
|
||||
'tagPrefix',
|
||||
'difyctlTag',
|
||||
])
|
||||
expect(out).toMatch(new RegExp(`^${key}=.+$`, 'm'))
|
||||
expect(out).not.toMatch(/=undefined$/m)
|
||||
})
|
||||
|
||||
it('still emits the existing trace fields', () => {
|
||||
const { stdout } = run(['github-env'])
|
||||
for (const key of ['version', 'channel', 'prerelease', 'minDify', 'maxDify', 'tagPrefix'])
|
||||
expect(stdout).toMatch(new RegExp(`^${key}=`, 'm'))
|
||||
it('emits difyctlTag as tagPrefix immediately followed by version', () => {
|
||||
const env = Object.fromEntries(
|
||||
main(['github-env'])
|
||||
.split('\n')
|
||||
.filter(Boolean)
|
||||
.map((line: string) => line.split(/=(.*)/s).slice(0, 2)),
|
||||
)
|
||||
expect(env.difyctlTag).toBe(`${env.tagPrefix}${env.version}`)
|
||||
})
|
||||
})
|
||||
|
||||
describe('release-naming edge channel', () => {
|
||||
it('lists edge among channels', () => {
|
||||
expect(run(['channels']).stdout).toMatch(/^edge$/m)
|
||||
it('lists one channel per line, including edge', () => {
|
||||
const out = main(['channels'])
|
||||
expect(out.trim().split('\n').length).toBeGreaterThan(1)
|
||||
expect(out).toMatch(/^edge$/m)
|
||||
})
|
||||
|
||||
it('edge-version derives <pkgcore>-edge.<sha> from the package version', () => {
|
||||
// package.json version is 0.2.0-alpha -> core 0.2.0
|
||||
expect(run(['edge-version', '2fd7b82']).stdout.trim()).toBe('0.2.0-edge.2fd7b82')
|
||||
it('emits targets as bunTarget<TAB>id<TAB>0|1 (release-build.sh parses this)', () => {
|
||||
for (const line of main(['targets']).trim().split('\n'))
|
||||
expect(line).toMatch(/^\S+\t\S+\t[01]$/)
|
||||
})
|
||||
|
||||
it('edge-version accepts a 40-char sha', () => {
|
||||
const sha = '2fd7b829e1f0aaaabbbbccccddddeeeeffff0000'
|
||||
expect(run(['edge-version', sha]).stdout.trim()).toBe(`0.2.0-edge.${sha}`)
|
||||
it('reports prerelease as a boolean per channel', () => {
|
||||
expect(main(['prerelease', 'stable']).trim()).toBe('false')
|
||||
expect(main(['prerelease', 'alpha']).trim()).toBe('true')
|
||||
expect(() => main(['prerelease', 'nightly'])).toThrow('unknown channel')
|
||||
})
|
||||
|
||||
it('edge-version rejects a non-hex sha', () => {
|
||||
expect(run(['edge-version', 'nothex!']).code).not.toBe(0)
|
||||
it('derives an edge version from the packaged version and a sha', () => {
|
||||
expect(main(['edge-version', '2fd7b82']).trim()).toMatch(/-edge\.2fd7b82$/)
|
||||
})
|
||||
|
||||
it('edge-version requires a sha argument', () => {
|
||||
expect(run(['edge-version']).code).not.toBe(0)
|
||||
})
|
||||
|
||||
it('the edge version form matches a computed edge version', () => {
|
||||
expect(run(['validate-version', '0.1.0-edge.2fd7b82', 'edge']).code).toBe(0)
|
||||
})
|
||||
|
||||
it('validate-version rejects an rc string under the edge channel', () => {
|
||||
expect(run(['validate-version', '0.1.0-rc.1', 'edge']).code).not.toBe(0)
|
||||
it('validates a version against a channel form', () => {
|
||||
expect(main(['validate-version', '0.1.0-edge.2fd7b82', 'edge'])).toContain('valid')
|
||||
expect(() => main(['validate-version', '0.1.0-rc.1', 'edge'])).toThrow(
|
||||
'does not match the edge channel form',
|
||||
)
|
||||
})
|
||||
})
|
||||
|
||||
@@ -1,13 +1,23 @@
|
||||
#!/usr/bin/env node
|
||||
// release-r2-edge.mjs — edge/R2 release metadata generator. Two subcommands:
|
||||
// manifest -> the per-channel pointer manifest.json (the installer reads this)
|
||||
// index -> the per-channel build-history ledger index.json
|
||||
import { existsSync, readFileSync } from 'node:fs'
|
||||
import { assetName, loadPkg, validateVersionForChannel } from './release-naming.mjs'
|
||||
import { existsSync, readFileSync, realpathSync } from 'node:fs'
|
||||
import { fileURLToPath } from 'node:url'
|
||||
import {
|
||||
buildIndex,
|
||||
parseChecksums,
|
||||
parseDirList,
|
||||
renderManifest,
|
||||
resolveTargets,
|
||||
} from './lib/edge-manifest.mjs'
|
||||
import { channelVersionProblem } from './lib/release-rules.mjs'
|
||||
import { loadPkg } from './release-naming.mjs'
|
||||
|
||||
function die(msg) {
|
||||
process.stderr.write(`release-r2-edge: ${msg}\n`)
|
||||
process.exit(1)
|
||||
class UsageError extends Error {}
|
||||
|
||||
// Arrow so TS infers `never`, keeping expression positions like `return die(...)` typed.
|
||||
const die = (msg) => {
|
||||
throw new UsageError(msg)
|
||||
}
|
||||
|
||||
function parseArgs(argv) {
|
||||
@@ -28,104 +38,83 @@ function requireArgs(args, keys) {
|
||||
}
|
||||
}
|
||||
|
||||
// checksums lines are "<sha256> <assetName>"
|
||||
function shaMap(checksumsPath) {
|
||||
const map = new Map()
|
||||
for (const line of readFileSync(checksumsPath, 'utf8').split('\n')) {
|
||||
const m = line.match(/^([0-9a-f]{64})\s+(\S+)$/i)
|
||||
if (m) map.set(m[2], m[1])
|
||||
}
|
||||
return map
|
||||
}
|
||||
|
||||
function emitManifest(args) {
|
||||
requireArgs(args, ['channel', 'version', 'commit', 'build-date', 'base-url', 'checksums'])
|
||||
validateVersionForChannel(args.version, args.channel)
|
||||
const versionProblem = channelVersionProblem(args.version, args.channel)
|
||||
if (versionProblem) die(versionProblem)
|
||||
|
||||
const { release, compat } = loadPkg()
|
||||
const shas = shaMap(args.checksums)
|
||||
const shas = parseChecksums(readFileSync(args.checksums, 'utf8'))
|
||||
const { targets, missing } = resolveTargets(release, args.version, shas)
|
||||
if (missing.length > 0) die(`no sha256 for ${missing[0]} in ${args.checksums}`)
|
||||
|
||||
const targetLines = release.targets
|
||||
.map((t) => {
|
||||
const asset = assetName(release, args.version, t.id)
|
||||
const sha = shas.get(asset)
|
||||
if (!sha) die(`no sha256 for ${asset} in ${args.checksums}`)
|
||||
// one target per line: install-r2.sh grep/sed depends on this layout
|
||||
return ` ${JSON.stringify(t.id)}: { "asset": ${JSON.stringify(asset)}, "sha256": ${JSON.stringify(sha)} }`
|
||||
})
|
||||
.join(',\n')
|
||||
|
||||
const head = {
|
||||
schema: 1,
|
||||
name: release.binName,
|
||||
return renderManifest({
|
||||
binName: release.binName,
|
||||
channel: args.channel,
|
||||
version: args.version,
|
||||
commit: args.commit,
|
||||
buildDate: args['build-date'],
|
||||
compat: { minDify: compat.minDify, maxDify: compat.maxDify },
|
||||
compat,
|
||||
baseUrl: args['base-url'],
|
||||
}
|
||||
const headLines = Object.entries(head)
|
||||
.map(([k, v]) => ` ${JSON.stringify(k)}: ${JSON.stringify(v)}`)
|
||||
.join(',\n')
|
||||
process.stdout.write(`{\n${headLines},\n "targets": {\n${targetLines}\n }\n}\n`)
|
||||
targets,
|
||||
})
|
||||
}
|
||||
|
||||
// Newline-delimited dir names of binaries that still exist in R2. Absent file =
|
||||
// no reconciliation (caller could not list); empty file = no survivors.
|
||||
// empty / "-" / missing = no ledger yet (first publish)
|
||||
function loadCurrentIndex(path) {
|
||||
if (path === '-' || !existsSync(path)) return null
|
||||
const raw = readFileSync(path, 'utf8').trim()
|
||||
if (!raw || raw === '-') return null
|
||||
try {
|
||||
return JSON.parse(raw)
|
||||
} catch {
|
||||
return die(`current index at ${path} is not valid JSON`)
|
||||
}
|
||||
}
|
||||
|
||||
// Absent file = no reconciliation (caller could not list); empty file = no survivors.
|
||||
function loadExistingDirs(path) {
|
||||
if (!path || !existsSync(path)) return null
|
||||
const set = new Set()
|
||||
for (const line of readFileSync(path, 'utf8').split('\n')) {
|
||||
const d = line.trim()
|
||||
if (d) set.add(d)
|
||||
}
|
||||
return set
|
||||
return parseDirList(readFileSync(path, 'utf8'))
|
||||
}
|
||||
|
||||
function emitIndex(args) {
|
||||
requireArgs(args, ['current', 'channel', 'version', 'commit', 'build-date'])
|
||||
|
||||
// empty / "-" / missing = no ledger yet (first publish)
|
||||
let current = { schema: 1, channel: args.channel, builds: [] }
|
||||
if (args.current !== '-' && existsSync(args.current)) {
|
||||
const raw = readFileSync(args.current, 'utf8').trim()
|
||||
if (raw && raw !== '-') {
|
||||
try {
|
||||
current = JSON.parse(raw)
|
||||
} catch {
|
||||
die(`current index at ${args.current} is not valid JSON`)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
const entry = {
|
||||
const index = buildIndex({
|
||||
channel: args.channel,
|
||||
version: args.version,
|
||||
commit: args.commit,
|
||||
buildDate: args['build-date'],
|
||||
dir: args.version,
|
||||
current: loadCurrentIndex(args.current),
|
||||
existingDirs: loadExistingDirs(args['existing-dirs']),
|
||||
})
|
||||
return `${JSON.stringify(index, null, 2)}\n`
|
||||
}
|
||||
|
||||
// Returns the exact bytes to write to stdout.
|
||||
function main(argv) {
|
||||
const [cmd, ...rest] = argv
|
||||
const args = parseArgs(rest)
|
||||
switch (cmd) {
|
||||
case 'manifest':
|
||||
return emitManifest(args)
|
||||
case 'index':
|
||||
return emitIndex(args)
|
||||
default:
|
||||
return die(`unknown subcommand: ${cmd ?? '(none)'} (expected: manifest | index)`)
|
||||
}
|
||||
const kept = (current.builds ?? []).filter((b) => b.version !== entry.version)
|
||||
let builds = [entry, ...kept]
|
||||
|
||||
// Reconcile to binaries that still exist in R2: lifecycle/TTL on the bin prefix
|
||||
// is the only deletion mechanism, so the ledger never advertises a build whose
|
||||
// binary is gone. The new build is always kept (just uploaded). No count cap.
|
||||
const existing = loadExistingDirs(args['existing-dirs'])
|
||||
if (existing) builds = builds.filter((b) => b.dir === entry.dir || existing.has(b.dir))
|
||||
|
||||
const index = { schema: 1, channel: args.channel, updated: args['build-date'], builds }
|
||||
process.stdout.write(`${JSON.stringify(index, null, 2)}\n`)
|
||||
}
|
||||
|
||||
const [cmd, ...rest] = process.argv.slice(2)
|
||||
const args = parseArgs(rest)
|
||||
switch (cmd) {
|
||||
case 'manifest':
|
||||
emitManifest(args)
|
||||
break
|
||||
case 'index':
|
||||
emitIndex(args)
|
||||
break
|
||||
default:
|
||||
die(`unknown subcommand: ${cmd ?? '(none)'} (expected: manifest | index)`)
|
||||
const invokedDirectly =
|
||||
process.argv[1] && realpathSync(process.argv[1]) === fileURLToPath(import.meta.url)
|
||||
if (invokedDirectly) {
|
||||
try {
|
||||
process.stdout.write(main(process.argv.slice(2)))
|
||||
} catch (e) {
|
||||
if (!(e instanceof UsageError)) throw e
|
||||
process.stderr.write(`release-r2-edge: ${e.message}\n`)
|
||||
process.exit(1)
|
||||
}
|
||||
}
|
||||
|
||||
export { main }
|
||||
|
||||
+108
-235
@@ -1,43 +1,28 @@
|
||||
import { execFileSync, spawnSync } from 'node:child_process'
|
||||
import { mkdtempSync, writeFileSync } from 'node:fs'
|
||||
import { tmpdir } from 'node:os'
|
||||
import { join } from 'node:path'
|
||||
import { fileURLToPath } from 'node:url'
|
||||
import { describe, expect, it } from 'vitest'
|
||||
|
||||
const SCRIPT = fileURLToPath(new URL('./release-r2-edge.mjs', import.meta.url))
|
||||
|
||||
function run(args: string[]): { code: number; stdout: string; stderr: string } {
|
||||
try {
|
||||
return {
|
||||
code: 0,
|
||||
stdout: execFileSync('node', [SCRIPT, ...args], { encoding: 'utf8' }),
|
||||
stderr: '',
|
||||
}
|
||||
} catch (e) {
|
||||
const err = e as { status?: number; stdout?: string; stderr?: string }
|
||||
return { code: err.status ?? 1, stdout: err.stdout ?? '', stderr: err.stderr ?? '' }
|
||||
}
|
||||
}
|
||||
|
||||
// ---- manifest ----
|
||||
|
||||
function writeChecksums(version: string): string {
|
||||
const dir = mkdtempSync(join(tmpdir(), 'difyctl-manifest-'))
|
||||
const ids = ['linux-x64', 'linux-arm64', 'darwin-x64', 'darwin-arm64', 'windows-x64']
|
||||
const lines = ids.map((id, i) => {
|
||||
const exe = id === 'windows-x64' ? '.exe' : ''
|
||||
const sha = String(i).repeat(64)
|
||||
return `${sha} difyctl-v${version}-${id}${exe}`
|
||||
})
|
||||
const file = join(dir, `difyctl-v${version}-checksums.txt`)
|
||||
writeFileSync(file, `${lines.join('\n')}\n`)
|
||||
return file
|
||||
}
|
||||
import { main as naming } from './release-naming.mjs'
|
||||
import { main } from './release-r2-edge.mjs'
|
||||
|
||||
const VERSION = '0.1.0-edge.2fd7b82'
|
||||
const BASE_URL = 'https://example.r2.dev/difyctl/edge/0.1.0-edge.2fd7b82'
|
||||
|
||||
const TARGET_IDS = naming(['targets'])
|
||||
.trim()
|
||||
.split('\n')
|
||||
.map((line: string) => line.split('\t')[1])
|
||||
|
||||
const ASSETS = new Map(TARGET_IDS.map((id: string) => [id, naming(['asset', VERSION, id]).trim()]))
|
||||
const assetFor = (id: string) => ASSETS.get(id) ?? ''
|
||||
|
||||
const PKG_ENV = Object.fromEntries(
|
||||
naming(['github-env'])
|
||||
.split('\n')
|
||||
.filter(Boolean)
|
||||
.map((line: string) => line.split(/=(.*)/s).slice(0, 2)),
|
||||
)
|
||||
|
||||
type ManifestJson = {
|
||||
schema: number
|
||||
name: string
|
||||
@@ -50,28 +35,23 @@ type ManifestJson = {
|
||||
targets: Record<string, { asset: string; sha256: string }>
|
||||
}
|
||||
|
||||
type IndexBuild = {
|
||||
version: string
|
||||
commit: string
|
||||
buildDate: string
|
||||
dir: string
|
||||
}
|
||||
|
||||
type IndexJson = {
|
||||
schema: number
|
||||
channel: string
|
||||
updated: string
|
||||
builds: IndexBuild[]
|
||||
builds: { version: string; commit: string; buildDate: string; dir: string }[]
|
||||
}
|
||||
|
||||
function buildManifest(version = VERSION): {
|
||||
code: number
|
||||
json: ManifestJson
|
||||
stdout: string
|
||||
stderr: string
|
||||
} {
|
||||
const checksums = writeChecksums(version)
|
||||
const r = run([
|
||||
function writeChecksums(ids: string[]): string {
|
||||
const dir = mkdtempSync(join(tmpdir(), 'difyctl-manifest-'))
|
||||
const lines = ids.map((id, i) => `${String(i).repeat(64)} ${assetFor(id)}`)
|
||||
const file = join(dir, 'checksums.txt')
|
||||
writeFileSync(file, `${lines.join('\n')}\n`)
|
||||
return file
|
||||
}
|
||||
|
||||
function manifestArgs(version = VERSION, ids = TARGET_IDS): string[] {
|
||||
return [
|
||||
'manifest',
|
||||
'--channel',
|
||||
'edge',
|
||||
@@ -84,229 +64,122 @@ function buildManifest(version = VERSION): {
|
||||
'--base-url',
|
||||
BASE_URL,
|
||||
'--checksums',
|
||||
checksums,
|
||||
])
|
||||
return {
|
||||
code: r.code,
|
||||
json: (r.code === 0 ? JSON.parse(r.stdout) : null) as ManifestJson,
|
||||
stdout: r.stdout,
|
||||
stderr: r.stderr,
|
||||
}
|
||||
writeChecksums(ids),
|
||||
]
|
||||
}
|
||||
|
||||
const buildManifest = () => JSON.parse(main(manifestArgs())) as ManifestJson
|
||||
|
||||
describe('release-r2-edge manifest', () => {
|
||||
it('emits the core pointer fields', () => {
|
||||
const { json } = buildManifest()
|
||||
expect(json.schema).toBe(1)
|
||||
expect(json.name).toBe('difyctl')
|
||||
expect(json.channel).toBe('edge')
|
||||
expect(json.version).toBe(VERSION)
|
||||
expect(json.commit).toBe('abc1234')
|
||||
expect(json.buildDate).toBe('2026-06-14T12:00:00Z')
|
||||
expect(json.baseUrl).toBe(BASE_URL)
|
||||
it('passes the CLI arguments straight through to the manifest', () => {
|
||||
expect(buildManifest()).toMatchObject({
|
||||
schema: 1,
|
||||
channel: 'edge',
|
||||
version: VERSION,
|
||||
commit: 'abc1234',
|
||||
buildDate: '2026-06-14T12:00:00Z',
|
||||
baseUrl: BASE_URL,
|
||||
})
|
||||
})
|
||||
|
||||
it('carries the compat window from package.json', () => {
|
||||
const { json } = buildManifest()
|
||||
expect(json.compat).toEqual({ minDify: '1.16.0', maxDify: '1.16.0' })
|
||||
it('carries the compat window through from cli/package.json', () => {
|
||||
expect(buildManifest().compat).toEqual({ minDify: PKG_ENV.minDify, maxDify: PKG_ENV.maxDify })
|
||||
})
|
||||
|
||||
it('lists all 5 targets with asset name + sha256 from the checksums file', () => {
|
||||
const { json } = buildManifest()
|
||||
expect(Object.keys(json.targets).sort()).toEqual([
|
||||
'darwin-arm64',
|
||||
'darwin-x64',
|
||||
'linux-arm64',
|
||||
'linux-x64',
|
||||
'windows-x64',
|
||||
])
|
||||
expect(json.targets['linux-x64'].asset).toBe(`difyctl-v${VERSION}-linux-x64`)
|
||||
expect(json.targets['windows-x64'].asset).toBe(`difyctl-v${VERSION}-windows-x64.exe`)
|
||||
expect(json.targets['linux-x64'].sha256).toMatch(/^\d{64}$/)
|
||||
})
|
||||
|
||||
it('renders each target on a single line (installer greps it)', () => {
|
||||
const { stdout } = buildManifest()
|
||||
expect(stdout).toMatch(/^ {4}"linux-x64": \{ "asset": ".*", "sha256": ".*" \}/m)
|
||||
it('includes every target the release config declares', () => {
|
||||
expect(Object.keys(buildManifest().targets).sort()).toEqual([...TARGET_IDS].sort())
|
||||
})
|
||||
|
||||
it('rejects a version that does not match the channel form', () => {
|
||||
const { code } = buildManifest('0.1.0-rc.1')
|
||||
expect(code).not.toBe(0)
|
||||
expect(() => main(manifestArgs('0.1.0-rc.1'))).toThrow('does not match the edge channel form')
|
||||
})
|
||||
|
||||
it('dies when a target sha is missing from the checksums file', () => {
|
||||
const dir = mkdtempSync(join(tmpdir(), 'difyctl-manifest-'))
|
||||
const file = join(dir, `difyctl-v${VERSION}-checksums.txt`)
|
||||
writeFileSync(file, `${'0'.repeat(64)} difyctl-v${VERSION}-linux-x64\n`) // only 1 of 5
|
||||
const r = run([
|
||||
'manifest',
|
||||
'--channel',
|
||||
'edge',
|
||||
'--version',
|
||||
VERSION,
|
||||
'--commit',
|
||||
'abc1234',
|
||||
'--build-date',
|
||||
'2026-06-14T12:00:00Z',
|
||||
'--base-url',
|
||||
BASE_URL,
|
||||
'--checksums',
|
||||
file,
|
||||
])
|
||||
expect(r.code).not.toBe(0)
|
||||
expect(() => main(manifestArgs(VERSION, TARGET_IDS.slice(0, 1)))).toThrow('no sha256 for')
|
||||
})
|
||||
|
||||
it('rejects a malformed dropped-value argument (no silent misparse)', () => {
|
||||
// --version has no value; --commit must NOT be swallowed as the version
|
||||
const r = run([
|
||||
'manifest',
|
||||
'--channel',
|
||||
'edge',
|
||||
'--version',
|
||||
'--commit',
|
||||
'abc1234',
|
||||
'--build-date',
|
||||
'2026-06-14T12:00:00Z',
|
||||
'--base-url',
|
||||
'https://x',
|
||||
'--checksums',
|
||||
'/nonexistent',
|
||||
])
|
||||
expect(r.code).not.toBe(0)
|
||||
expect(() =>
|
||||
main([
|
||||
'manifest',
|
||||
'--channel',
|
||||
'edge',
|
||||
'--version',
|
||||
'--commit',
|
||||
'abc1234',
|
||||
'--build-date',
|
||||
'2026-06-14T12:00:00Z',
|
||||
'--base-url',
|
||||
'https://x',
|
||||
'--checksums',
|
||||
'/nonexistent',
|
||||
]),
|
||||
).toThrow('malformed argument')
|
||||
})
|
||||
|
||||
it('rejects an unknown subcommand', () => {
|
||||
expect(() => main(['bogus'])).toThrow('unknown subcommand')
|
||||
})
|
||||
})
|
||||
|
||||
// ---- index ----
|
||||
describe('release-r2-edge index', () => {
|
||||
const B1 = { version: '0.1.0-edge.aaaaaaa', commit: 'aaaaaaa', buildDate: '2026-06-14T09:00:00Z' }
|
||||
|
||||
function runIndex(
|
||||
currentContent: string | null,
|
||||
build: Record<string, string>,
|
||||
existingDirs?: string[],
|
||||
) {
|
||||
let currentArg = '-'
|
||||
if (currentContent !== null) {
|
||||
const dir = mkdtempSync(join(tmpdir(), 'difyctl-index-'))
|
||||
currentArg = join(dir, 'index.json')
|
||||
writeFileSync(currentArg, currentContent)
|
||||
}
|
||||
const extra: string[] = []
|
||||
if (existingDirs !== undefined) {
|
||||
const dir = mkdtempSync(join(tmpdir(), 'difyctl-existing-'))
|
||||
const f = join(dir, 'existing.txt')
|
||||
writeFileSync(f, `${existingDirs.join('\n')}\n`)
|
||||
extra.push('--existing-dirs', f)
|
||||
}
|
||||
const r = spawnSync(
|
||||
'node',
|
||||
[
|
||||
SCRIPT,
|
||||
function indexArgs(currentContent: string | null, existingDirs?: string[]): string[] {
|
||||
let currentArg = '-'
|
||||
if (currentContent !== null) {
|
||||
currentArg = join(mkdtempSync(join(tmpdir(), 'difyctl-index-')), 'index.json')
|
||||
writeFileSync(currentArg, currentContent)
|
||||
}
|
||||
const extra: string[] = []
|
||||
if (existingDirs !== undefined) {
|
||||
const f = join(mkdtempSync(join(tmpdir(), 'difyctl-existing-')), 'existing.txt')
|
||||
writeFileSync(f, `${existingDirs.join('\n')}\n`)
|
||||
extra.push('--existing-dirs', f)
|
||||
}
|
||||
return [
|
||||
'index',
|
||||
'--current',
|
||||
currentArg,
|
||||
'--channel',
|
||||
'edge',
|
||||
'--version',
|
||||
build.version,
|
||||
B1.version,
|
||||
'--commit',
|
||||
build.commit,
|
||||
B1.commit,
|
||||
'--build-date',
|
||||
build.buildDate,
|
||||
B1.buildDate,
|
||||
...extra,
|
||||
],
|
||||
{ encoding: 'utf8' },
|
||||
)
|
||||
return {
|
||||
code: r.status ?? 1,
|
||||
index: (r.status === 0 ? JSON.parse(r.stdout) : null) as IndexJson,
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
const B1 = { version: '0.1.0-edge.aaaaaaa', commit: 'aaaaaaa', buildDate: '2026-06-14T09:00:00Z' }
|
||||
const B2 = { version: '0.1.0-edge.bbbbbbb', commit: 'bbbbbbb', buildDate: '2026-06-14T10:00:00Z' }
|
||||
const buildIndex = (currentContent: string | null, existingDirs?: string[]) =>
|
||||
JSON.parse(main(indexArgs(currentContent, existingDirs))) as IndexJson
|
||||
|
||||
describe('release-r2-edge index', () => {
|
||||
it('creates a fresh index from a missing current (arg "-")', () => {
|
||||
const { index } = runIndex(null, B1)
|
||||
expect(index.schema).toBe(1)
|
||||
expect(index.channel).toBe('edge')
|
||||
expect(index.builds).toHaveLength(1)
|
||||
expect(index.builds[0]).toMatchObject({
|
||||
version: B1.version,
|
||||
commit: B1.commit,
|
||||
dir: B1.version,
|
||||
})
|
||||
it.each([
|
||||
['a missing current file (arg "-")', null],
|
||||
['an empty current file (first publish, curl wrote nothing)', ''],
|
||||
['a "-"-content current file (curl 404 fallback)', '-\n'],
|
||||
])('treats %s as a fresh ledger', (_label, content) => {
|
||||
const index = buildIndex(content)
|
||||
expect(index).toMatchObject({ schema: 1, channel: 'edge' })
|
||||
expect(index.builds).toEqual([{ ...B1, dir: B1.version }])
|
||||
})
|
||||
|
||||
it('treats an empty current file as fresh (first publish, curl wrote nothing)', () => {
|
||||
const { code, index } = runIndex('', B1)
|
||||
expect(code).toBe(0)
|
||||
expect(index.builds).toHaveLength(1)
|
||||
it('dies on a current index that is not valid JSON', () => {
|
||||
expect(() => main(indexArgs('{not json'))).toThrow('not valid JSON')
|
||||
})
|
||||
|
||||
it('treats a "-"-content current file as fresh (curl 404 fallback)', () => {
|
||||
const { code, index } = runIndex('-\n', B1)
|
||||
expect(code).toBe(0)
|
||||
expect(index.builds).toHaveLength(1)
|
||||
})
|
||||
|
||||
it('prepends the new build (publish order; newest at [0])', () => {
|
||||
const current = JSON.stringify({
|
||||
schema: 1,
|
||||
channel: 'edge',
|
||||
builds: [
|
||||
{ version: B1.version, commit: B1.commit, buildDate: B1.buildDate, dir: B1.version },
|
||||
],
|
||||
})
|
||||
const { index } = runIndex(current, B2)
|
||||
expect(index.builds.map((b) => b.version)).toEqual([B2.version, B1.version])
|
||||
})
|
||||
|
||||
it('dedups a re-cut of the same version (no duplicate, moves to top)', () => {
|
||||
const current = JSON.stringify({
|
||||
schema: 1,
|
||||
channel: 'edge',
|
||||
builds: [
|
||||
{ version: B2.version, commit: B2.commit, buildDate: B2.buildDate, dir: B2.version },
|
||||
{ version: B1.version, commit: B1.commit, buildDate: B1.buildDate, dir: B1.version },
|
||||
],
|
||||
})
|
||||
const { index } = runIndex(current, B1) // re-cut B1
|
||||
expect(index.builds.map((b) => b.version)).toEqual([B1.version, B2.version])
|
||||
})
|
||||
|
||||
it('reconciles to surviving binary dirs (drops a build whose binary expired)', () => {
|
||||
const current = JSON.stringify({
|
||||
schema: 1,
|
||||
channel: 'edge',
|
||||
builds: [
|
||||
{ version: B1.version, commit: B1.commit, buildDate: B1.buildDate, dir: B1.version },
|
||||
],
|
||||
})
|
||||
// B1's binary is gone (not in existing); the new B2 is always kept.
|
||||
const { index } = runIndex(current, B2, [B2.version])
|
||||
expect(index.builds.map((b) => b.version)).toEqual([B2.version])
|
||||
})
|
||||
|
||||
it('keeps the new build even when it is absent from the existing-dirs list', () => {
|
||||
const { index } = runIndex(null, B1, []) // empty survivors, fresh ledger
|
||||
expect(index.builds.map((b) => b.version)).toEqual([B1.version])
|
||||
})
|
||||
|
||||
it('does not reconcile when no --existing-dirs is given (list unavailable)', () => {
|
||||
const current = JSON.stringify({
|
||||
schema: 1,
|
||||
channel: 'edge',
|
||||
builds: [
|
||||
{ version: B1.version, commit: B1.commit, buildDate: B1.buildDate, dir: B1.version },
|
||||
],
|
||||
})
|
||||
const { index } = runIndex(current, B2) // no existing-dirs → keep all
|
||||
expect(index.builds.map((b) => b.version)).toEqual([B2.version, B1.version])
|
||||
})
|
||||
|
||||
it('dies on a non-empty current file that is not valid JSON', () => {
|
||||
const { code } = runIndex('{not json', B1)
|
||||
expect(code).not.toBe(0)
|
||||
it('reads the existing-dirs file to reconcile the ledger', () => {
|
||||
const older = {
|
||||
version: '0.1.0-edge.bbbbbbb',
|
||||
commit: 'bbbbbbb',
|
||||
buildDate: '2026-06-14T08:00:00Z',
|
||||
dir: '0.1.0-edge.bbbbbbb',
|
||||
}
|
||||
const current = JSON.stringify({ schema: 1, channel: 'edge', builds: [older] })
|
||||
expect(buildIndex(current, []).builds).toEqual([{ ...B1, dir: B1.version }])
|
||||
expect(buildIndex(current, [older.dir]).builds).toEqual([{ ...B1, dir: B1.version }, older])
|
||||
})
|
||||
})
|
||||
|
||||
@@ -1,91 +0,0 @@
|
||||
import { spawnSync } from 'node:child_process'
|
||||
import { fileURLToPath } from 'node:url'
|
||||
import { describe, expect, it } from 'vitest'
|
||||
|
||||
const SCRIPT = fileURLToPath(new URL('./release-r2-publish.sh', import.meta.url))
|
||||
|
||||
// Stub `aws` + `curl` + `node` as shell functions that just log action verbs to
|
||||
// $ORDER_LOG, then run the publish `main` and assert the order of operations.
|
||||
function runPublish(): { code: number; order: string[]; stderr: string } {
|
||||
const stub = [
|
||||
'ORDER_LOG="$(mktemp)"',
|
||||
'aws() {',
|
||||
' case "$*" in',
|
||||
' *"list-objects-v2"*) echo list-survivors >>"$ORDER_LOG" ;;',
|
||||
' *" cp "*"/index.json"*) echo put-index >>"$ORDER_LOG" ;;',
|
||||
' *" cp "*"/manifest.json"*) echo put-manifest >>"$ORDER_LOG" ;;',
|
||||
' *" sync "*) echo sync-binaries >>"$ORDER_LOG" ;;',
|
||||
' *"head-object"*) echo head-verify >>"$ORDER_LOG" ;;',
|
||||
' *) : ;;',
|
||||
' esac',
|
||||
'}',
|
||||
'curl() { echo "{}"; }',
|
||||
'node() {',
|
||||
' case "$*" in',
|
||||
' *release-naming.mjs*targets*)',
|
||||
" printf 'bun-linux-x64\\tlinux-x64\\t0\\nbun-linux-arm64\\tlinux-arm64\\t0\\nbun-darwin-x64\\tdarwin-x64\\t0\\nbun-darwin-arm64\\tdarwin-arm64\\t0\\nbun-windows-x64\\twindows-x64\\t1\\n' ;;",
|
||||
" *release-naming.mjs*' asset '*) printf 'difyctl-vX\\n' ;;",
|
||||
" *release-r2-edge.mjs*' index '*) echo '{}' ;;",
|
||||
" *release-r2-edge.mjs*' manifest '*) echo '{}' ;;",
|
||||
' *) : ;;',
|
||||
' esac',
|
||||
'}',
|
||||
].join('\n')
|
||||
const program = [
|
||||
stub,
|
||||
`. "${SCRIPT}"`,
|
||||
'publish_main edge 0.1.0-edge.2fd7b82',
|
||||
'cat "$ORDER_LOG"',
|
||||
].join('\n')
|
||||
// bash, NOT sh: the script uses BASH_SOURCE + process substitution.
|
||||
const r = spawnSync('bash', ['-c', program], {
|
||||
encoding: 'utf8',
|
||||
env: {
|
||||
...process.env,
|
||||
RELEASE_PUBLISH_LIB: '1',
|
||||
DIFYCTL_R2_S3_ENDPOINT: 'https://endpoint.example',
|
||||
DIFYCTL_R2_BUCKET: 'cli-dev',
|
||||
DIFYCTL_R2_PUBLIC_BASE: 'https://pub.example.r2.dev',
|
||||
DIST_DIR: '/tmp',
|
||||
},
|
||||
})
|
||||
return {
|
||||
code: r.status ?? 1,
|
||||
order: (r.stdout ?? '').trim().split('\n').filter(Boolean),
|
||||
stderr: r.stderr ?? '',
|
||||
}
|
||||
}
|
||||
|
||||
describe('release-r2-publish order', () => {
|
||||
it('uploads binaries, verifies, lists survivors, then index, then manifest', () => {
|
||||
const { code, order } = runPublish()
|
||||
expect(code).toBe(0)
|
||||
expect(order.indexOf('sync-binaries')).toBeLessThan(order.indexOf('head-verify'))
|
||||
expect(order.indexOf('head-verify')).toBeLessThan(order.indexOf('list-survivors'))
|
||||
expect(order.indexOf('list-survivors')).toBeLessThan(order.indexOf('put-index'))
|
||||
expect(order.indexOf('put-index')).toBeLessThan(order.indexOf('put-manifest'))
|
||||
// pointer is never pruned here — deletion is owned by the R2 lifecycle rule
|
||||
expect(order).not.toContain('prune')
|
||||
})
|
||||
|
||||
it('exits non-zero when no targets resolve (head-verify safety gate)', () => {
|
||||
const stub = [
|
||||
'aws() { :; }',
|
||||
'curl() { echo "{}"; }',
|
||||
'node() { case "$*" in *release-naming.mjs*targets*) : ;; *) echo "{}" ;; esac; }',
|
||||
].join('\n')
|
||||
const program = [stub, `. "${SCRIPT}"`, 'publish_main edge 0.1.0-edge.2fd7b82'].join('\n')
|
||||
const r = spawnSync('bash', ['-c', program], {
|
||||
encoding: 'utf8',
|
||||
env: {
|
||||
...process.env,
|
||||
RELEASE_PUBLISH_LIB: '1',
|
||||
DIFYCTL_R2_S3_ENDPOINT: 'https://endpoint.example',
|
||||
DIFYCTL_R2_BUCKET: 'cli-dev',
|
||||
DIFYCTL_R2_PUBLIC_BASE: 'https://pub.example.r2.dev',
|
||||
DIST_DIR: '/tmp',
|
||||
},
|
||||
})
|
||||
expect(r.status).not.toBe(0)
|
||||
})
|
||||
})
|
||||
@@ -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
|
||||
|
||||
@@ -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 () => {},
|
||||
}
|
||||
|
||||
@@ -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
@@ -56,5 +56,5 @@
|
||||
"engines": {
|
||||
"node": "^22.22.1"
|
||||
},
|
||||
"packageManager": "pnpm@11.15.0"
|
||||
"packageManager": "pnpm@11.17.0"
|
||||
}
|
||||
|
||||
@@ -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',
|
||||
|
||||
Generated
+1989
-1585
File diff suppressed because it is too large
Load Diff
+58
-58
@@ -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
|
||||
|
||||
+100
-144
@@ -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
|
||||
|
||||
+1
-1
@@ -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 = '/'
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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' ? (
|
||||
|
||||
+82
-2
@@ -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()
|
||||
})
|
||||
})
|
||||
|
||||
+5
-1
@@ -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>
|
||||
|
||||
+9
-10
@@ -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>
|
||||
|
||||
+1
-1
@@ -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}
|
||||
>
|
||||
|
||||
Reference in New Issue
Block a user