Compare commits

...
Author SHA1 Message Date
autofix-ci[bot]andGitHub 756ce8a10e [autofix.ci] apply automated fixes 2026-07-28 05:33:30 +00:00
-LAN- e07de297bf fix(workflow): migrate sys files to user input 2026-07-28 13:29:09 +08:00
yyhandGitHub 1e5e47b889 fix: validate plugin installation scope (#39669) 2026-07-28 04:13:20 +00:00
Asuka MinatoGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>Byron.wang
dd28b0d165 test: move OAuth server service coverage to unit tests (#38931)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: Byron.wang <byron@dify.ai>
2026-07-28 04:09:32 +00:00
JoelandGitHub 49e74e2f58 fix: prevent the model selector footer from covering options (#39668) 2026-07-28 04:02:28 +00:00
e6e5d761c2 fix: prevent Safari from clipping the settings close button (#39664)
Co-authored-by: yyh <92089059+lyzno1@users.noreply.github.com>
2026-07-28 04:01:12 +00:00
Asuka MinatoandGitHub e1ce808567 test: move data source controller coverage to unit tests (#38924) 2026-07-28 03:51:22 +00:00
Asuka MinatoandGitHub 3c7ad816d9 test: use SQLite sessions in commands (#39080) 2026-07-28 03:29:02 +00:00
Asuka MinatoandGitHub f44dd343da test: use sqlite3 session in test_plugin_service (#38727) 2026-07-28 03:27:58 +00:00
Asuka MinatoandGitHub 2d9b2d50f3 test: use sqlite3 session in test_wraps (#38770) 2026-07-28 03:27:22 +00:00
Asuka MinatoandGitHub 003e0f9614 test: move message cleanup coverage to unit tests (#38932) 2026-07-28 03:26:43 +00:00
Asuka MinatoGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>Byron.wang
da0979b373 test: use SQLite sessions in services core (#39090)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: Byron.wang <byron@dify.ai>
2026-07-28 03:20:27 +00:00
Asuka MinatoGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>Byron.wang
81fb57639e test: use SQLite sessions in controllers service api (#39098)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: Byron.wang <byron@dify.ai>
2026-07-28 03:19:17 +00:00
Asuka MinatoandGitHub 1e61078e93 test: use SQLite sessions in services core (#39088) 2026-07-28 03:18:32 +00:00
Escape0707andGitHub b597bb1b17 test: separate human input unit and database paths (#39655) 2026-07-28 03:18:01 +00:00
Asuka MinatoGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>Byron.wang
65ead05dfc test: use SQLite sessions in services plugin (#39084)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: Byron.wang <byron@dify.ai>
2026-07-28 03:16:19 +00:00
6f8ed69ee1 fix: fix mcp output_schema is optional (#39453)
Co-authored-by: yunlu.wen <yunlu.wen@dify.ai>
2026-07-28 02:31:18 +00:00
Yunlu WenGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
59b879d2df chore: bump version to 1.16.1 (#39653)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-07-28 01:09:25 +00:00
119 changed files with 3513 additions and 2074 deletions
+5 -6
View File
@@ -3,6 +3,7 @@ CLI command modules extracted from `commands.py`.
"""
from .account import create_tenant, reset_email, reset_password
from .app_maintenance import convert_to_agent_apps, fix_app_site_missing
from .data_migrate import data_migrate, legacy_model_types
from .data_migration import (
export_migration_data,
@@ -10,6 +11,7 @@ from .data_migration import (
import_migration_data,
migration_data_wizard,
)
from .database import upgrade_db
from .plugin import (
backfill_plugin_auto_upgrade,
extract_plugins,
@@ -36,12 +38,6 @@ from .retention import (
restore_workflow_runs,
)
from .storage import clear_orphaned_file_records, file_usage, migrate_oss, remove_orphaned_files_on_storage
from .system import (
convert_to_agent_apps,
fix_app_site_missing,
reset_encrypt_key_pair,
upgrade_db,
)
from .vector import (
add_qdrant_index,
migrate_annotation_vector_database,
@@ -49,6 +45,8 @@ from .vector import (
old_metadata_migration,
vdb_migrate,
)
from .workflow_migration import migrate_legacy_sys_files_workflows
from .workspace import reset_encrypt_key_pair
__all__ = [
"add_qdrant_index",
@@ -80,6 +78,7 @@ __all__ = [
"migrate_data_for_plugin",
"migrate_dataset_permissions_to_rbac",
"migrate_knowledge_vector_database",
"migrate_legacy_sys_files_workflows",
"migrate_member_roles_to_rbac",
"migrate_oss",
"migration_data_wizard",
@@ -1,86 +1,27 @@
"""App data maintenance CLI commands."""
import logging
import click
import sqlalchemy as sa
from sqlalchemy import delete, select, update
from sqlalchemy.orm import sessionmaker
from sqlalchemy import select, update
from configs import dify_config
from enums.deployment_edition import DeploymentEdition
from events.app_event import app_was_created
from extensions.ext_database import db
from extensions.ext_redis import redis_client
from libs.db_migration_lock import DbMigrationAutoRenewLock
from libs.rsa import generate_key_pair
from models import Tenant
from models.model import App, AppMode, Conversation
from models.provider import Provider, ProviderModel
from models.tools import ApiToolProvider, BuiltinToolProvider, MCPToolProvider
logger = logging.getLogger(__name__)
DB_UPGRADE_LOCK_TTL_SECONDS = 60
@click.command(
"reset-encrypt-key-pair",
help="Reset the asymmetric key pair of workspace for encrypt LLM credentials. "
"After the reset, all LLM credentials and tool provider credentials "
"(builtin / API / MCP) will be purged, requiring re-entry. "
"Only support SELF_HOSTED mode.",
)
@click.confirmation_option(
prompt=click.style(
"Are you sure you want to reset encrypt key pair? "
"This will also purge builtin / API / MCP tool provider records for every tenant. "
"This operation cannot be rolled back!",
fg="red",
)
)
def reset_encrypt_key_pair():
"""
Reset the encrypted key pair of workspace for encrypt LLM credentials.
After the reset, all LLM credentials will become invalid, requiring re-entry.
Only support SELF_HOSTED mode.
"""
if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD:
click.echo(click.style("This command is only for SELF_HOSTED installations.", fg="red"))
return
with sessionmaker(db.engine, expire_on_commit=False).begin() as session:
tenants = session.scalars(select(Tenant)).all()
for tenant in tenants:
if not tenant:
click.echo(click.style("No workspaces found. Run /install first.", fg="red"))
return
tenant.encrypt_public_key = generate_key_pair(tenant.id)
session.execute(delete(Provider).where(Provider.provider_type == "custom", Provider.tenant_id == tenant.id))
session.execute(delete(ProviderModel).where(ProviderModel.tenant_id == tenant.id))
# Purge tool provider records that hold credentials encrypted under the
# tenant key. Leaving them in place causes /console/api/workspaces/current/
# tool-providers to 500 because decryption fails on stale ciphertext (#35396).
session.execute(delete(BuiltinToolProvider).where(BuiltinToolProvider.tenant_id == tenant.id))
session.execute(delete(ApiToolProvider).where(ApiToolProvider.tenant_id == tenant.id))
session.execute(delete(MCPToolProvider).where(MCPToolProvider.tenant_id == tenant.id))
click.echo(
click.style(
f"Congratulations! The asymmetric key pair of workspace {tenant.id} has been reset.",
fg="green",
)
)
@click.command("convert-to-agent-apps", help="Convert Agent Assistant to Agent App.")
def convert_to_agent_apps():
def convert_to_agent_apps() -> None:
"""
Convert Agent Assistant to Agent App.
"""
click.echo(click.style("Starting convert to agent apps.", fg="green"))
proceeded_app_ids = []
proceeded_app_ids: list[str] = []
while True:
# fetch first 1000 apps
@@ -133,48 +74,14 @@ def convert_to_agent_apps():
click.echo(click.style(f"Conversion complete. Converted {len(proceeded_app_ids)} agent apps.", fg="green"))
@click.command("upgrade-db", help="Upgrade the database")
def upgrade_db():
click.echo("Preparing database migration...")
lock = DbMigrationAutoRenewLock(
redis_client=redis_client,
name="db_upgrade_lock",
ttl_seconds=DB_UPGRADE_LOCK_TTL_SECONDS,
logger=logger,
log_context="db_migration",
)
if lock.acquire(blocking=False):
migration_succeeded = False
try:
click.echo(click.style("Starting database migration.", fg="green"))
# run db migration
import flask_migrate
flask_migrate.upgrade()
migration_succeeded = True
click.echo(click.style("Database migration successful!", fg="green"))
except Exception as e:
logger.exception("Failed to execute database migration")
click.echo(click.style(f"Database migration failed: {e}", fg="red"))
raise SystemExit(1)
finally:
status = "successful" if migration_succeeded else "failed"
lock.release_safely(status=status)
else:
click.echo("Database migration skipped")
@click.command("fix-app-site-missing", help="Fix app related site missing issue.")
def fix_app_site_missing():
def fix_app_site_missing() -> None:
"""
Fix app related site missing issue.
"""
click.echo(click.style("Starting fix for missing app-related sites.", fg="green"))
failed_app_ids = []
failed_app_ids: list[str] = []
while True:
sql = """select apps.id as id from apps left join sites on sites.app_id=apps.id
where sites.id is null limit 1000"""
+45
View File
@@ -0,0 +1,45 @@
"""Database schema migration CLI commands."""
import logging
import click
from extensions.ext_redis import redis_client
from libs.db_migration_lock import DbMigrationAutoRenewLock
logger = logging.getLogger(__name__)
DB_UPGRADE_LOCK_TTL_SECONDS = 60
@click.command("upgrade-db", help="Upgrade the database")
def upgrade_db() -> None:
click.echo("Preparing database migration...")
lock = DbMigrationAutoRenewLock(
redis_client=redis_client,
name="db_upgrade_lock",
ttl_seconds=DB_UPGRADE_LOCK_TTL_SECONDS,
logger=logger,
log_context="db_migration",
)
if lock.acquire(blocking=False):
migration_succeeded = False
try:
click.echo(click.style("Starting database migration.", fg="green"))
import flask_migrate
flask_migrate.upgrade()
migration_succeeded = True
click.echo(click.style("Database migration successful!", fg="green"))
except Exception as e:
logger.exception("Failed to execute database migration")
click.echo(click.style(f"Database migration failed: {e}", fg="red"))
raise SystemExit(1)
finally:
status = "successful" if migration_succeeded else "failed"
lock.release_safely(status=status)
else:
click.echo("Database migration skipped")
+173
View File
@@ -0,0 +1,173 @@
"""Workflow data migration CLI commands.
TODO: Remove the legacy system file workflow migration command after the production migration is complete.
"""
import logging
from dataclasses import dataclass
import click
from sqlalchemy import select
from sqlalchemy.orm import Session, load_only, sessionmaker
from extensions.ext_database import db
from models.workflow import Workflow, WorkflowType
logger = logging.getLogger(__name__)
@dataclass
class LegacySysFilesWorkflowMigrationStats:
scanned: int = 0
migrated: int = 0
failed: int = 0
batches: int = 0
last_id: str | None = None
def _build_legacy_sys_files_workflow_query(
*,
start_after_id: str | None,
batch_size: int,
tenant_id: str | None,
app_id: str | None,
):
# Workflow IDs are UUID4, so this is not chronological pagination. The migration only needs a stable total
# order that matches the resume cursor; ordering by the same primary-key column used in the `id > cursor`
# predicate lets each batch continue deterministically without offset scans.
stmt = (
select(Workflow)
.options(load_only(Workflow.id, Workflow.type, Workflow.graph))
.where(Workflow.type.in_((WorkflowType.WORKFLOW, WorkflowType.CHAT)))
.order_by(Workflow.id)
.limit(batch_size)
)
if start_after_id:
stmt = stmt.where(Workflow.id > start_after_id)
if tenant_id:
stmt = stmt.where(Workflow.tenant_id == tenant_id)
if app_id:
stmt = stmt.where(Workflow.app_id == app_id)
return stmt
def _migrate_legacy_sys_files_workflow_batch(
*,
session: Session,
start_after_id: str | None,
batch_size: int,
tenant_id: str | None,
app_id: str | None,
dry_run: bool,
) -> LegacySysFilesWorkflowMigrationStats:
stats = LegacySysFilesWorkflowMigrationStats()
workflows = session.scalars(
_build_legacy_sys_files_workflow_query(
start_after_id=start_after_id,
batch_size=batch_size,
tenant_id=tenant_id,
app_id=app_id,
)
).all()
for workflow in workflows:
stats.scanned += 1
stats.last_id = workflow.id
try:
if workflow.migrate_legacy_sys_files_graph_in_place():
stats.migrated += 1
except Exception:
stats.failed += 1
logger.exception("Failed to migrate legacy sys.files workflow, workflow_id=%s", workflow.id)
if dry_run:
session.rollback()
else:
session.commit()
return stats
def run_legacy_sys_files_workflow_migration(
*,
batch_size: int,
limit: int | None,
start_after_id: str | None,
tenant_id: str | None,
app_id: str | None,
dry_run: bool,
) -> LegacySysFilesWorkflowMigrationStats:
"""Scan Workflow and Advanced Chat graphs in keyset-paginated batches."""
if batch_size <= 0:
raise click.UsageError("--batch-size must be greater than 0")
if limit is not None and limit <= 0:
raise click.UsageError("--limit must be greater than 0 when provided")
session_maker = sessionmaker(db.engine, expire_on_commit=False)
total = LegacySysFilesWorkflowMigrationStats(last_id=start_after_id)
next_start_after_id = start_after_id
while limit is None or total.scanned < limit:
remaining = None if limit is None else limit - total.scanned
current_batch_size = batch_size if remaining is None else min(batch_size, remaining)
if current_batch_size <= 0:
break
with session_maker() as session:
batch_stats = _migrate_legacy_sys_files_workflow_batch(
session=session,
start_after_id=next_start_after_id,
batch_size=current_batch_size,
tenant_id=tenant_id,
app_id=app_id,
dry_run=dry_run,
)
if batch_stats.scanned == 0:
break
total.scanned += batch_stats.scanned
total.migrated += batch_stats.migrated
total.failed += batch_stats.failed
total.batches += 1
total.last_id = batch_stats.last_id
next_start_after_id = batch_stats.last_id
if batch_stats.scanned < current_batch_size:
break
return total
@click.command(
"migrate-legacy-sys-files-workflows",
help="Migrate Workflow and Advanced Chat graphs that still reference deprecated sys.files.",
)
@click.option("--batch-size", default=1000, show_default=True, type=int, help="Number of workflows to scan per batch.")
@click.option("--limit", default=None, type=int, help="Maximum number of workflows to scan in this run.")
@click.option("--start-after-id", default=None, help="Resume scanning after this workflow ID.")
@click.option("--tenant-id", default=None, help="Limit migration to one tenant.")
@click.option("--app-id", default=None, help="Limit migration to one app.")
@click.option("--dry-run", is_flag=True, default=False, help="Scan and report without saving changes.")
def migrate_legacy_sys_files_workflows(
batch_size: int,
limit: int | None,
start_after_id: str | None,
tenant_id: str | None,
app_id: str | None,
dry_run: bool,
) -> None:
stats = run_legacy_sys_files_workflow_migration(
batch_size=batch_size,
limit=limit,
start_after_id=start_after_id,
tenant_id=tenant_id,
app_id=app_id,
dry_run=dry_run,
)
click.echo(
"Legacy sys.files workflow migration finished: "
f"scanned={stats.scanned} migrated={stats.migrated} failed={stats.failed} "
f"batches={stats.batches} last_id={stats.last_id or ''}"
)
if dry_run:
click.echo("Dry run only: no workflow graph changes were saved.")
+64
View File
@@ -0,0 +1,64 @@
"""Workspace maintenance CLI commands."""
import click
from sqlalchemy import delete, select
from sqlalchemy.orm import sessionmaker
from configs import dify_config
from enums.deployment_edition import DeploymentEdition
from extensions.ext_database import db
from libs.rsa import generate_key_pair
from models import Tenant
from models.provider import Provider, ProviderModel
from models.tools import ApiToolProvider, BuiltinToolProvider, MCPToolProvider
@click.command(
"reset-encrypt-key-pair",
help="Reset the asymmetric key pair of workspace for encrypt LLM credentials. "
"After the reset, all LLM credentials and tool provider credentials "
"(builtin / API / MCP) will be purged, requiring re-entry. "
"Only support SELF_HOSTED mode.",
)
@click.confirmation_option(
prompt=click.style(
"Are you sure you want to reset encrypt key pair? "
"This will also purge builtin / API / MCP tool provider records for every tenant. "
"This operation cannot be rolled back!",
fg="red",
)
)
def reset_encrypt_key_pair() -> None:
"""
Reset the encrypted key pair of workspace for encrypt LLM credentials.
After the reset, all LLM credentials will become invalid, requiring re-entry.
Only support SELF_HOSTED mode.
"""
if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD:
click.echo(click.style("This command is only for SELF_HOSTED installations.", fg="red"))
return
with sessionmaker(db.engine, expire_on_commit=False).begin() as session:
tenants = session.scalars(select(Tenant)).all()
for tenant in tenants:
if not tenant:
click.echo(click.style("No workspaces found. Run /install first.", fg="red"))
return
tenant.encrypt_public_key = generate_key_pair(tenant.id)
session.execute(delete(Provider).where(Provider.provider_type == "custom", Provider.tenant_id == tenant.id))
session.execute(delete(ProviderModel).where(ProviderModel.tenant_id == tenant.id))
# Purge tool provider records that hold credentials encrypted under the
# tenant key. Leaving them in place causes /console/api/workspaces/current/
# tool-providers to 500 because decryption fails on stale ciphertext (#35396).
session.execute(delete(BuiltinToolProvider).where(BuiltinToolProvider.tenant_id == tenant.id))
session.execute(delete(ApiToolProvider).where(ApiToolProvider.tenant_id == tenant.id))
session.execute(delete(MCPToolProvider).where(MCPToolProvider.tenant_id == tenant.id))
click.echo(
click.style(
f"Congratulations! The asymmetric key pair of workspace {tenant.id} has been reset.",
fg="green",
)
)
@@ -26,6 +26,10 @@ from controllers.service_api.app.error import (
ProviderQuotaExceededError,
WorkflowVersionExecutionNotAllowedError,
)
from controllers.service_api.app.legacy_system_files import (
attach_legacy_system_file_warning_for_service_api,
normalize_legacy_system_file_args_for_service_api,
)
from controllers.service_api.schema import (
InputFileList,
expect_user_json,
@@ -390,6 +394,7 @@ class ChatApi(Resource):
args["external_trace_id"] = external_trace_id
streaming = _resolve_agent_app_streaming(app_mode=app_mode, response_mode=payload.response_mode)
legacy_system_file_compat = None
try:
# Eagerly validate conversation to avoid hanging on invalid conversation_id
@@ -401,6 +406,14 @@ class ChatApi(Resource):
session=session,
)
if app_mode == AppMode.ADVANCED_CHAT:
args, legacy_system_file_compat = normalize_legacy_system_file_args_for_service_api(
session=session,
app_model=app_model,
args=args,
raw_payload=service_api_ns.payload,
workflow_id=args.get("workflow_id"),
)
response = AppGenerateService.generate(
session=session,
app_model=app_model,
@@ -409,6 +422,7 @@ class ChatApi(Resource):
invoke_from=InvokeFrom.SERVICE_API,
streaming=streaming,
)
response = attach_legacy_system_file_warning_for_service_api(response, legacy_system_file_compat)
# response-contract:ignore compact_generate_response
return helper.compact_generate_response(response)
@@ -0,0 +1,72 @@
"""Temporary Service API adapter for the deprecated workflow file input."""
from collections.abc import Generator, Mapping
from typing import Any
from sqlalchemy.orm import Session
from core.app.entities.app_invoke_entities import InvokeFrom
from core.app.features.rate_limiting.rate_limit import RateLimitGenerator
from core.workflow.legacy_system_files import (
LegacySysFilesCompatVariable,
attach_legacy_sys_files_warning,
normalize_legacy_sys_files_args,
)
from models.model import App
from services.app_generate_service import AppGenerateService
type ServiceAPIGenerateResponse = Mapping[str, Any] | Generator[str, None, None] | RateLimitGenerator
def normalize_legacy_system_file_args_for_service_api(
*,
session: Session,
app_model: App,
args: dict[str, Any],
raw_payload: Mapping[str, Any] | None,
workflow_id: str | None = None,
) -> tuple[dict[str, Any], LegacySysFilesCompatVariable | None]:
# TODO: Remove this hidden Service API compatibility path after all persisted workflows are migrated.
args_with_hidden_system = _copy_hidden_system_files_arg(args=args, raw_payload=raw_payload)
if not _has_legacy_file_arg(args_with_hidden_system):
return args, None
workflow = AppGenerateService.get_workflow(
app_model,
InvokeFrom.SERVICE_API,
workflow_id,
session=session,
)
return normalize_legacy_sys_files_args(graph=workflow.graph_dict, args=args_with_hidden_system)
def attach_legacy_system_file_warning_for_service_api(
response: ServiceAPIGenerateResponse,
compat_variable: LegacySysFilesCompatVariable | None,
) -> ServiceAPIGenerateResponse:
# TODO: Remove this warning once Service API clients no longer need the legacy migration notice.
if compat_variable is None:
return response
return attach_legacy_sys_files_warning(response, compat_variable)
def _copy_hidden_system_files_arg(
*,
args: dict[str, Any],
raw_payload: Mapping[str, Any] | None,
) -> dict[str, Any]:
system = raw_payload.get("system") if isinstance(raw_payload, Mapping) else None
if not isinstance(system, Mapping) or "files" not in system or system["files"] is None:
return args
copied_args = dict(args)
copied_args["system"] = {"files": system["files"]}
return copied_args
def _has_legacy_file_arg(args: Mapping[str, Any]) -> bool:
if args.get("files") is not None:
return True
system = args.get("system")
return isinstance(system, Mapping) and system.get("files") is not None
@@ -30,6 +30,10 @@ from controllers.service_api.app.error import (
ProviderQuotaExceededError,
WorkflowVersionExecutionNotAllowedError,
)
from controllers.service_api.app.legacy_system_files import (
attach_legacy_system_file_warning_for_service_api,
normalize_legacy_system_file_args_for_service_api,
)
from controllers.service_api.schema import (
expect_user_json,
expect_with_user,
@@ -344,6 +348,12 @@ class WorkflowRunApi(Resource):
streaming = payload.response_mode == "streaming"
try:
args, legacy_system_file_compat = normalize_legacy_system_file_args_for_service_api(
session=session,
app_model=app_model,
args=args,
raw_payload=service_api_ns.payload,
)
response = AppGenerateService.generate(
session=session,
app_model=app_model,
@@ -352,6 +362,7 @@ class WorkflowRunApi(Resource):
invoke_from=InvokeFrom.SERVICE_API,
streaming=streaming,
)
response = attach_legacy_system_file_warning_for_service_api(response, legacy_system_file_compat)
# response-contract:ignore compact_generate_response
return helper.compact_generate_response(response)
@@ -471,6 +482,13 @@ class WorkflowRunByIdApi(Resource):
streaming = payload.response_mode == "streaming"
try:
args, legacy_system_file_compat = normalize_legacy_system_file_args_for_service_api(
session=session,
app_model=app_model,
args=args,
raw_payload=service_api_ns.payload,
workflow_id=workflow_id,
)
response = AppGenerateService.generate(
session=session,
app_model=app_model,
@@ -479,6 +497,7 @@ class WorkflowRunByIdApi(Resource):
invoke_from=InvokeFrom.SERVICE_API,
streaming=streaming,
)
response = attach_legacy_system_file_warning_for_service_api(response, legacy_system_file_compat)
# response-contract:ignore compact_generate_response
return helper.compact_generate_response(response)
@@ -46,6 +46,7 @@ from core.ops.ops_trace_manager import TraceQueueManager
from core.prompt.utils.get_thread_messages_length import get_thread_messages_length
from core.repositories import DifyCoreRepositoryFactory
from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository
from core.workflow.legacy_system_files import normalize_legacy_sys_files_args
from extensions.ext_database import db
from factories import file_factory
from graphon.filters import ResponseStreamFilter
@@ -147,6 +148,8 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
if not args.get("query"):
raise ValueError("query is required")
# TODO: Remove this compatibility normalization after all persisted workflows are migrated.
args, _ = normalize_legacy_sys_files_args(graph=workflow.graph_dict, args=args)
query = args["query"]
if not isinstance(query, str):
raise ValueError("query must be a string")
@@ -155,8 +155,10 @@ class WorkflowResponseConverter:
# TODO(@future-refactor): store system variables separately from user inputs so we don't
# need to flatten `sys.*` entries into the input payload just for rerun/export tooling.
if field_name == SystemVariableKey.CONVERSATION_ID:
# Conversation IDs are session-scoped; omitting them keeps workflow inputs
# reusable without pinning new runs to a prior conversation.
# Conversation IDs are session-scoped; omitting them keeps workflow inputs reusable.
continue
if field_name == SystemVariableKey.FILES:
# When files are exposed as an input, application inputs use the canonical `userinput.files` key.
continue
inputs[f"sys.{field_name}"] = value
handled = WorkflowEntry.handle_special_values(inputs)
@@ -41,6 +41,7 @@ from core.helper.trace_id_helper import (
from core.ops.ops_trace_manager import TraceQueueManager
from core.repositories import DifyCoreRepositoryFactory
from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository
from core.workflow.legacy_system_files import normalize_legacy_sys_files_args
from extensions.ext_database import db
from factories import file_factory
from graphon.filters import ResponseStreamFilter
@@ -164,6 +165,8 @@ class WorkflowAppGenerator(BaseAppGenerator):
pause_state_config: PauseStateLayerConfig | None = None,
) -> Mapping[str, Any] | Generator[Mapping[str, Any] | str, None, None]:
with self._bind_file_access_scope(tenant_id=app_model.tenant_id, user=user, invoke_from=invoke_from):
# TODO: Remove this compatibility normalization after all persisted workflows are migrated.
args, _ = normalize_legacy_sys_files_args(graph=workflow.graph_dict, args=args)
files: Sequence[Mapping[str, Any]] = args.get("files") or []
# parse files
+27 -19
View File
@@ -66,7 +66,7 @@ from services.enterprise.plugin_manager_service import (
PreUninstallPluginRequest,
)
from services.errors.plugin import PluginInstallationForbiddenError
from services.feature_service import FeatureService, PluginInstallationScope
from services.feature_service import FeatureService, PluginInstallationPermissionModel, PluginInstallationScope
logger = logging.getLogger(__name__)
_provider_entities_adapter: TypeAdapter[list[ProviderEntity]] = TypeAdapter(list[ProviderEntity])
@@ -604,22 +604,30 @@ class PluginService:
return result
@staticmethod
def _check_marketplace_only_permission():
def _check_marketplace_only_permission() -> None:
"""
Check if the marketplace only permission is enabled
"""
features = FeatureService.get_system_features()
if features.plugin_installation_permission.restrict_to_marketplace_only:
permission = PluginService._get_plugin_installation_permission()
if permission.restrict_to_marketplace_only:
raise PluginInstallationForbiddenError("Plugin installation is restricted to marketplace only")
@staticmethod
def _check_plugin_installation_scope(plugin_verification: PluginVerification | None):
def _get_plugin_installation_permission() -> PluginInstallationPermissionModel:
"""Resolve the validated policy and reject deny-all before any installation side effect."""
permission = FeatureService.get_plugin_installation_permission()
if permission.plugin_installation_scope == PluginInstallationScope.NONE:
raise PluginInstallationForbiddenError("Installing plugins is not allowed")
return permission
@staticmethod
def _check_plugin_installation_scope(plugin_verification: PluginVerification | None) -> None:
"""
Check the plugin installation scope
"""
features = FeatureService.get_system_features()
permission = PluginService._get_plugin_installation_permission()
match features.plugin_installation_permission.plugin_installation_scope:
match permission.plugin_installation_scope:
case PluginInstallationScope.OFFICIAL_ONLY:
if (
plugin_verification is None
@@ -634,10 +642,10 @@ class PluginService:
raise PluginInstallationForbiddenError(
"Plugin installation is restricted to official and specific partners"
)
case PluginInstallationScope.NONE:
raise PluginInstallationForbiddenError("Installing plugins is not allowed")
case PluginInstallationScope.ALL:
pass
case _:
raise PluginInstallationForbiddenError("Plugin installation policy is invalid")
@staticmethod
def get_debugging_key(tenant_id: str) -> str:
@@ -907,7 +915,7 @@ class PluginService:
# check if plugin pkg is already downloaded
manager = PluginInstaller()
features = FeatureService.get_system_features()
permission = PluginService._get_plugin_installation_permission()
try:
manager.fetch_plugin_manifest(tenant_id, new_plugin_unique_identifier)
@@ -919,7 +927,7 @@ class PluginService:
response = manager.upload_pkg(
tenant_id,
pkg,
verify_signature=features.plugin_installation_permission.restrict_to_marketplace_only,
verify_signature=permission.restrict_to_marketplace_only,
)
# check if the plugin is available to install
@@ -974,11 +982,11 @@ class PluginService:
"""
PluginService._check_marketplace_only_permission()
manager = PluginInstaller()
features = FeatureService.get_system_features()
permission = PluginService._get_plugin_installation_permission()
response = manager.upload_pkg(
tenant_id,
pkg,
verify_signature=features.plugin_installation_permission.restrict_to_marketplace_only,
verify_signature=permission.restrict_to_marketplace_only,
)
PluginService._check_plugin_installation_scope(response.verification)
@@ -996,13 +1004,13 @@ class PluginService:
pkg = download_with_size_limit(
f"https://github.com/{repo}/releases/download/{version}/{package}", dify_config.PLUGIN_MAX_PACKAGE_SIZE
)
features = FeatureService.get_system_features()
permission = PluginService._get_plugin_installation_permission()
manager = PluginInstaller()
response = manager.upload_pkg(
tenant_id,
pkg,
verify_signature=features.plugin_installation_permission.restrict_to_marketplace_only,
verify_signature=permission.restrict_to_marketplace_only,
)
PluginService._check_plugin_installation_scope(response.verification)
@@ -1076,7 +1084,7 @@ class PluginService:
if not dify_config.MARKETPLACE_ENABLED:
raise ValueError("marketplace is not enabled")
features = FeatureService.get_system_features()
permission = PluginService._get_plugin_installation_permission()
manager = PluginInstaller()
try:
@@ -1086,7 +1094,7 @@ class PluginService:
response = manager.upload_pkg(
tenant_id,
pkg,
verify_signature=features.plugin_installation_permission.restrict_to_marketplace_only,
verify_signature=permission.restrict_to_marketplace_only,
)
# check if the plugin is available to install
PluginService._check_plugin_installation_scope(response.verification)
@@ -1108,7 +1116,7 @@ class PluginService:
# collect actual plugin_unique_identifiers
actual_plugin_unique_identifiers = []
metas = []
features = FeatureService.get_system_features()
permission = PluginService._get_plugin_installation_permission()
# check if already downloaded
for plugin_unique_identifier in plugin_unique_identifiers:
@@ -1126,7 +1134,7 @@ class PluginService:
response = manager.upload_pkg(
tenant_id,
pkg,
verify_signature=features.plugin_installation_permission.restrict_to_marketplace_only,
verify_signature=permission.restrict_to_marketplace_only,
)
# check if the plugin is available to install
PluginService._check_plugin_installation_scope(response.verification)
+2
View File
@@ -107,6 +107,8 @@ class MCPTool(Tool):
if self.entity.output_schema and result.structuredContent:
for k, v in result.structuredContent.items():
yield self.create_variable_message(k, v)
elif result.structuredContent:
yield self.create_json_message(result.structuredContent)
def _process_text_content(self, content: TextContent) -> Generator[ToolInvokeMessage, None, None]:
"""Process text content and yield appropriate messages."""
@@ -42,9 +42,9 @@ _NODE_SNIPPETS: dict[str, str] = {
["local_file", "remote_url"]. Only when you include "custom" must you
also set ``allowed_file_extensions`` to a non-empty list like
[".epub", ".rtf"]; otherwise leave it [].
In Advanced-Chat mode ``sys.query`` and ``sys.files`` are automatic
system variables — downstream nodes may reference them; do NOT add
them to ``variables``.""",
In Advanced-Chat mode ``sys.query`` is automatic. ``userinput.files`` is
the automatic file-upload variable in both app modes. Downstream nodes
may reference these variables; do NOT add them to ``variables``.""",
"end": """\
- end (Workflow mode only):
{"outputs": [
@@ -168,8 +168,8 @@ _NODE_SNIPPETS: dict[str, str] = {
Single output variable ``text``: a string when ``is_array_file`` is false,
an array of strings (one per file) when it is true. ``variable_selector``
MUST point at a ``start`` variable declared with type "file" / "file-list"
(or ``sys.files`` in Advanced-Chat mode). That start variable MUST set a
non-empty ``allowed_file_types`` (use ["document"] for document text).""",
(or the automatic ``userinput.files`` variable). A declared start variable
MUST set a non-empty ``allowed_file_types`` (use ["document"] for document text).""",
"variable-aggregator": """\
- variable-aggregator (merge mutually-exclusive branches into one output):
{"output_type": "string", # VarType of the merged value — one of
@@ -91,21 +91,20 @@ def format_parallel_plan(
def format_mode_section(mode: str) -> str:
"""Tell each builder which app mode it is configuring for.
Matters most in advanced-chat, where ``sys.query`` / ``sys.files`` are the
sanctioned way to reference the user's message — without this the model
invents start-node variables that postprocess then materializes as
spurious form inputs.
``sys.query`` is available in advanced-chat, while ``userinput.files`` is
available in both app modes. Without this guidance the model invents
start-node variables that postprocess then materializes as spurious inputs.
"""
if mode == "advanced-chat":
return (
"# App mode\n\n"
"advanced-chat: the user's chat message is available as sys.query and uploaded files "
'as sys.files placeholder {{#sys.query#}}, selector ["sys", "query"]. Reference them '
"directly; do NOT invent start-node variables for the chat message.\n\n"
"as userinput.files. Use placeholder {{#sys.query#}} or selector "
'["userinput", "files"] directly; do NOT invent start-node variables for them.\n\n'
)
return (
"# App mode\n\n"
"workflow: there are NO automatic system variables; reference user input only through "
"workflow: uploaded files are available as userinput.files; all other user input must use "
"the start node's declared variables.\n\n"
)
@@ -96,10 +96,10 @@ minimum set of Dify workflow nodes needed to fulfil it, in execution order.
- "text-input" for short single-line values (URLs, names),
- "paragraph" for free-form multi-line text (descriptions, queries),
- "number" / "select" / "file" / "file-list" for the obvious cases.
In Advanced-Chat mode the ``sys.query`` / ``sys.files`` system
variables are automatic — downstream nodes may reference them without
a ``start_inputs`` entry. In Workflow mode there is NO automatic
variable; everything the user supplies must be in ``start_inputs``.
In Advanced-Chat mode ``sys.query`` is automatic. ``userinput.files`` is
automatic in both app modes. Downstream nodes may reference these values
without a ``start_inputs`` entry; every other user-supplied Workflow value
must be declared in ``start_inputs``.
11. Give every node a unique runtime-safe ``id`` using only letters, digits,
and underscores. In create mode use ``node1``, ``node2``, ... in node-list
order. In refine mode preserve the existing id for every retained node.
+7 -4
View File
@@ -1382,8 +1382,8 @@ class WorkflowGenerator:
multiple outputs remain untouched so validation fails closed instead
of guessing which value the workflow should consume.
For Advanced-Chat mode, ``sys.query`` and ``sys.files`` are always
treated as resolved without any declaration. Tool nodes' parameter
``sys.query`` in Advanced-Chat mode and ``userinput.files`` in either mode
are treated as resolved without declarations. Tool nodes' parameter
references aren't validated here because we don't know each tool's
schema — the run time validates those.
"""
@@ -1398,9 +1398,12 @@ class WorkflowGenerator:
for node in nodes:
cls._collect_refs_in_data(node.get("data") or {}, refs)
automatic_refs = {("userinput", "files")}
if mode == "advanced-chat":
automatic_refs.add(("sys", "query"))
for node_id, var in refs:
# Advanced-Chat system variables are always resolved.
if mode == "advanced-chat" and node_id == "sys":
if (node_id, var) in automatic_refs:
continue
target = nodes_by_id.get(node_id)
if target is None:
+218
View File
@@ -0,0 +1,218 @@
"""Compatibility helpers for workflows that still reference deprecated `sys.files`.
TODO: Remove this module after all persisted Workflow and Advanced Chat graphs
have been migrated from the deprecated system file variable to `userinput.files`.
"""
from __future__ import annotations
import json
from collections.abc import Generator, Iterable, Mapping
from dataclasses import dataclass
from typing import Any
_LEGACY_SYSTEM_NODE_ID = "sys"
_USER_INPUT_NODE_ID = "userinput"
_LEGACY_FILES_VARIABLE = "files"
_USER_INPUT_FILE_SELECTOR = [_USER_INPUT_NODE_ID, _LEGACY_FILES_VARIABLE]
_USER_INPUT_FILE_INPUT_KEY = ".".join(_USER_INPUT_FILE_SELECTOR)
_LEGACY_FILES_TEMPLATE = "{{#sys.files#}}"
_USER_INPUT_FILES_TEMPLATE = "{{#userinput.files#}}"
@dataclass(frozen=True)
class LegacySysFilesCompatVariable:
node_id: str
variable_name: str
@dataclass(frozen=True)
class LegacySysFilesGraphMigrationResult:
graph: dict[str, Any]
changed: bool
def migrate_legacy_sys_files_graph_with_result(
graph: Mapping[str, Any],
) -> LegacySysFilesGraphMigrationResult:
"""Return the migrated graph and whether any legacy reference was rewritten."""
graph_copy = dict(graph)
nodes = graph_copy.get("nodes")
if not isinstance(nodes, list):
return LegacySysFilesGraphMigrationResult(graph=graph_copy, changed=False)
# Legacy references are stored in node data. Restricting both search and replacement to `nodes`
# avoids recursively scanning graph-level metadata and edges for every workflow load.
if not _contains_legacy_sys_files_reference(nodes):
return LegacySysFilesGraphMigrationResult(graph=graph_copy, changed=False)
graph_copy["nodes"] = _replace_legacy_sys_files_references(nodes)
return LegacySysFilesGraphMigrationResult(graph=graph_copy, changed=True)
def resolve_legacy_sys_files_compat_variable(graph: Mapping[str, Any]) -> LegacySysFilesCompatVariable | None:
"""Resolve the target variable used by the `sys.files` compatibility layer."""
nodes = graph.get("nodes")
if not isinstance(nodes, list):
return None
if not _contains_file_input_reference(nodes):
return None
return LegacySysFilesCompatVariable(node_id=_USER_INPUT_NODE_ID, variable_name=_LEGACY_FILES_VARIABLE)
def normalize_legacy_sys_files_args(
*,
graph: Mapping[str, Any],
args: Mapping[str, Any],
) -> tuple[dict[str, Any], LegacySysFilesCompatVariable | None]:
"""Map Service/Web API file arguments onto the `userinput.files` system alias.
The top-level `files` argument and hidden `system.files` payload both feed
the same runtime file collection. After graph references are migrated, the
file collection is exposed in the variable pool as `userinput.files`.
"""
normalized_args = dict(args)
files_from_input, input_files_used = _extract_userinput_files(args)
if input_files_used:
normalized_args["files"] = files_from_input
return normalized_args, None
compat_variable = resolve_legacy_sys_files_compat_variable(graph)
if compat_variable is None:
return normalized_args, None
files, legacy_files_used = _extract_legacy_files(args)
if not legacy_files_used:
return normalized_args, None
if normalized_args.get("files") is None:
normalized_args["files"] = files
raw_inputs = normalized_args.get("inputs")
inputs = dict(raw_inputs) if isinstance(raw_inputs, Mapping) else {}
inputs.setdefault(_USER_INPUT_FILE_INPUT_KEY, files)
normalized_args["inputs"] = inputs
return normalized_args, compat_variable
def attach_legacy_sys_files_warning(
response: Mapping[str, Any] | Iterable[Any],
compat_variable: LegacySysFilesCompatVariable,
) -> Mapping[str, Any] | Generator[str, None, None]:
warning = build_legacy_sys_files_warning(compat_variable)
if isinstance(response, Mapping):
response_with_warning = dict(response)
existing_warnings = response_with_warning.get("warnings")
warnings = list(existing_warnings) if isinstance(existing_warnings, list) else []
warnings.append(warning)
response_with_warning["warnings"] = warnings
return response_with_warning
def _with_warning() -> Generator[str, None, None]:
try:
yield f"data: {json.dumps({'event': 'warning', 'warning': warning})}\n\n"
yield from response
finally:
close = getattr(response, "close", None)
if callable(close):
close()
return _with_warning()
def build_legacy_sys_files_warning(compat_variable: LegacySysFilesCompatVariable) -> str:
variable_selector = ".".join((compat_variable.node_id, compat_variable.variable_name))
return (
"sys.files is deprecated. This workflow now reads files from "
f"`{variable_selector}`; update Service API calls to pass files in "
f"`inputs.{variable_selector}` instead of `system.files` or top-level `files`."
)
def _contains_legacy_sys_files_reference(value: Any) -> bool:
if _is_legacy_sys_files_selector(value):
return True
if isinstance(value, str):
return _LEGACY_FILES_TEMPLATE in value
if isinstance(value, Mapping):
return any(_contains_legacy_sys_files_reference(item) for item in value.values())
if isinstance(value, list):
return any(_contains_legacy_sys_files_reference(item) for item in value)
return False
def _contains_file_input_reference(value: Any) -> bool:
if _is_legacy_sys_files_selector(value) or _is_userinput_files_selector(value):
return True
if isinstance(value, str):
return _LEGACY_FILES_TEMPLATE in value or _USER_INPUT_FILES_TEMPLATE in value
if isinstance(value, Mapping):
return any(_contains_file_input_reference(item) for item in value.values())
if isinstance(value, list):
return any(_contains_file_input_reference(item) for item in value)
return False
def _replace_legacy_sys_files_references(value: Any) -> Any:
if _is_legacy_sys_files_selector(value):
return list(_USER_INPUT_FILE_SELECTOR)
if isinstance(value, str):
return value.replace(_LEGACY_FILES_TEMPLATE, _USER_INPUT_FILES_TEMPLATE)
if isinstance(value, Mapping):
return {key: _replace_legacy_sys_files_references(item) for key, item in value.items()}
if isinstance(value, list):
return [_replace_legacy_sys_files_references(item) for item in value]
return value
def _is_legacy_sys_files_selector(value: Any) -> bool:
return (
isinstance(value, list)
and len(value) == 2
and value[0] == _LEGACY_SYSTEM_NODE_ID
and value[1] == _LEGACY_FILES_VARIABLE
)
def _is_userinput_files_selector(value: Any) -> bool:
return isinstance(value, list) and value == _USER_INPUT_FILE_SELECTOR
def serialized_graph_may_contain_legacy_sys_files(serialized_graph: str) -> bool:
"""Cheaply reject stored graphs that cannot contain a legacy file reference."""
return _LEGACY_FILES_TEMPLATE in serialized_graph or ('"sys"' in serialized_graph and '"files"' in serialized_graph)
def _extract_legacy_files(args: Mapping[str, Any]) -> tuple[Any, bool]:
if "files" in args and args["files"] is not None:
return args["files"], True
system = args.get("system")
if isinstance(system, Mapping) and "files" in system and system["files"] is not None:
return system["files"], True
return None, False
def _extract_userinput_files(args: Mapping[str, Any]) -> tuple[Any, bool]:
inputs = args.get("inputs")
if isinstance(inputs, Mapping) and inputs.get(_USER_INPUT_FILE_INPUT_KEY) is not None:
return inputs[_USER_INPUT_FILE_INPUT_KEY], True
return None, False
@@ -368,7 +368,7 @@ class WorkflowAgentRuntimeRequestBuilder:
if uploaded_files is not None:
lines.append("- Uploaded workflow files:")
lines.append(f" - sys.files: {uploaded_files}")
lines.append(f" - userinput.files: {uploaded_files}")
if resolved_outputs:
lines.append("- Previous node outputs:")
+7
View File
@@ -16,6 +16,7 @@ from .variable_prefixes import (
ENVIRONMENT_VARIABLE_NODE_ID,
RAG_PIPELINE_VARIABLE_NODE_ID,
SYSTEM_VARIABLE_NODE_ID,
USER_INPUT_VARIABLE_NODE_ID,
)
@@ -118,6 +119,12 @@ def build_bootstrap_variables(
*(_with_selector(variable, ENVIRONMENT_VARIABLE_NODE_ID) for variable in environment_variables),
*(_with_selector(variable, CONVERSATION_VARIABLE_NODE_ID) for variable in conversation_variables),
]
# TODO: Stop emitting the legacy `sys.files` selector after stored graphs and Service API callers are migrated.
# `userinput.files` remains the canonical file-upload variable.
for variable in system_variables:
if variable.name == SystemVariableKey.FILES.value:
variables.append(_with_selector(variable, USER_INPUT_VARIABLE_NODE_ID))
break
rag_pipeline_variables_map: defaultdict[str, dict[str, Any]] = defaultdict(dict)
for rag_var in rag_pipeline_variables:
+1
View File
@@ -1,4 +1,5 @@
SYSTEM_VARIABLE_NODE_ID = "sys"
USER_INPUT_VARIABLE_NODE_ID = "userinput"
ENVIRONMENT_VARIABLE_NODE_ID = "env"
CONVERSATION_VARIABLE_NODE_ID = "conversation"
RAG_PIPELINE_VARIABLE_NODE_ID = "rag"
+2
View File
@@ -29,6 +29,7 @@ def init_app(app: DifyApp):
install_rag_pipeline_plugins,
migrate_data_for_plugin,
migrate_dataset_permissions_to_rbac,
migrate_legacy_sys_files_workflows,
migrate_member_roles_to_rbac,
migrate_oss,
migration_data_wizard,
@@ -57,6 +58,7 @@ def init_app(app: DifyApp):
data_migrate,
upgrade_db,
fix_app_site_missing,
migrate_legacy_sys_files_workflows,
migrate_data_for_plugin,
migrate_dataset_permissions_to_rbac,
migrate_member_roles_to_rbac,
+39 -3
View File
@@ -25,6 +25,10 @@ from typing_extensions import deprecated
from core.trigger.constants import TRIGGER_PLUGIN_NODE_TYPE
from core.workflow.human_input_adapter import adapt_node_config_for_graph
from core.workflow.legacy_system_files import (
migrate_legacy_sys_files_graph_with_result,
serialized_graph_may_contain_legacy_sys_files,
)
from core.workflow.nodes.human_input.pause_reason import (
HumanInputRequired,
)
@@ -325,7 +329,39 @@ class Workflow(Base): # bug
# Currently, the following functions / methods would mutate the returned dict:
#
# - `_get_graph_and_variable_pool_for_single_node_run`.
return json.loads(self.graph) if self.graph else {}
if not self.graph:
return {}
graph = json.loads(self.graph)
if not self._supports_legacy_sys_files_compatibility() or not serialized_graph_may_contain_legacy_sys_files(
self.graph
):
return graph
# TODO: Remove this load-time compatibility rewrite after all persisted workflows are migrated.
return migrate_legacy_sys_files_graph_with_result(graph).graph
def migrate_legacy_sys_files_graph_in_place(self) -> bool:
if (
not self.graph
or not self._supports_legacy_sys_files_compatibility()
or not serialized_graph_may_contain_legacy_sys_files(self.graph)
):
return False
# TODO: Remove this in-place compatibility rewrite after all persisted workflows are migrated.
migration_result = migrate_legacy_sys_files_graph_with_result(json.loads(self.graph))
if migration_result.changed:
self.graph = json.dumps(migration_result.graph)
return migration_result.changed
def _supports_legacy_sys_files_compatibility(self) -> bool:
return self.type in {
WorkflowType.WORKFLOW,
WorkflowType.CHAT,
WorkflowType.WORKFLOW.value,
WorkflowType.CHAT.value,
}
def get_node_config_by_id(self, node_id: str) -> NodeConfigDict:
"""Extract a node configuration from the workflow graph by node ID.
@@ -487,7 +523,7 @@ class Workflow(Base): # bug
"memory":
{
"window": { "enabled": false, "size": 10 },
"query_prompt_template": "{{#sys.query#}}\n\n{{#sys.files#}}",
"query_prompt_template": "{{#sys.query#}}\n\n{{#userinput.files#}}",
"role_prefix": { "user": "", "assistant": "" },
},
"selected": false,
@@ -1520,7 +1556,7 @@ class ConversationVariable(TypeBase):
return variable_factory.build_conversation_variable_from_mapping(mapping)
# Only `sys.query` and `sys.files` could be modified.
# TODO: Remove file-system-variable editability after all persisted workflows are migrated.
_EDITABLE_SYSTEM_VARIABLE = frozenset(("query", "files"))
+1 -1
View File
@@ -1,6 +1,6 @@
[project]
name = "dify-api"
version = "1.16.0"
version = "1.16.1"
requires-python = "~=3.12.0"
dependencies = [
+12 -1
View File
@@ -488,7 +488,7 @@ class AppGenerateService:
)
@classmethod
def _get_workflow(
def get_workflow(
cls,
app_model: App,
invoke_from: InvokeFrom,
@@ -533,6 +533,17 @@ class AppGenerateService:
return workflow
@classmethod
def _get_workflow(
cls,
app_model: App,
invoke_from: InvokeFrom,
workflow_id: str | None = None,
*,
session: Session,
) -> Workflow:
return cls.get_workflow(app_model, invoke_from, workflow_id, session=session)
@classmethod
def get_response_generator(
cls,
+48 -9
View File
@@ -1,6 +1,8 @@
import logging
from collections.abc import Mapping
from enum import StrEnum
from pydantic import BaseModel, ConfigDict, Field
from pydantic import BaseModel, ConfigDict, Field, ValidationError
from configs import dify_config
from constants.dsl_version import CURRENT_APP_DSL_VERSION
@@ -10,6 +12,8 @@ from enums.hosted_provider import HostedTrialProvider
from services.billing_service import BillingInfo, BillingService
from services.enterprise.enterprise_service import EnterpriseService
logger = logging.getLogger(__name__)
class FeatureResponseModel(BaseModel):
model_config = ConfigDict(json_schema_serialization_defaults_required=True, protected_namespaces=())
@@ -131,6 +135,13 @@ class PluginInstallationPermissionModel(FeatureResponseModel):
restrict_to_marketplace_only: bool = False
class _EnterprisePluginInstallationPermission(BaseModel):
model_config = ConfigDict(extra="ignore")
plugin_installation_scope: PluginInstallationScope = Field(alias="pluginInstallationScope")
restrict_to_marketplace_only: bool = Field(alias="restrictToMarketplaceOnly", strict=True)
class FeatureModel(FeatureResponseModel):
billing: BillingModel = BillingModel()
education: EducationModel = EducationModel()
@@ -285,6 +296,14 @@ class FeatureService:
"""Return whether Enterprise plugin credential policies must be enforced."""
return dify_config.ENTERPRISE_ENABLED
@classmethod
def get_plugin_installation_permission(cls) -> PluginInstallationPermissionModel:
"""Resolve the validated deployment-wide plugin installation policy."""
if not dify_config.ENTERPRISE_ENABLED:
return PluginInstallationPermissionModel()
return cls._resolve_plugin_installation_permission(EnterpriseService.get_info())
@classmethod
def get_license(cls) -> LicenseModel:
"""Return full license detail. Enterprise-only; requires an authenticated caller.
@@ -452,6 +471,33 @@ class FeatureService:
)
return license_model
@classmethod
def _resolve_plugin_installation_permission(
cls, enterprise_info: Mapping[str, object]
) -> PluginInstallationPermissionModel:
if "PluginInstallationPermission" not in enterprise_info:
return PluginInstallationPermissionModel()
try:
permission = _EnterprisePluginInstallationPermission.model_validate(
enterprise_info["PluginInstallationPermission"]
)
except ValidationError as exc:
# Do not attach the exception because it may contain raw Enterprise configuration values.
logger.error( # noqa: TRY400
"Invalid Enterprise plugin installation permission; denying all plugin installations: %s",
exc.errors(include_input=False),
)
return PluginInstallationPermissionModel(
plugin_installation_scope=PluginInstallationScope.NONE,
restrict_to_marketplace_only=True,
)
return PluginInstallationPermissionModel(
plugin_installation_scope=permission.plugin_installation_scope,
restrict_to_marketplace_only=permission.restrict_to_marketplace_only,
)
@classmethod
def _fulfill_params_from_enterprise(cls, features: SystemFeatureModel):
enterprise_info = EnterpriseService.get_info()
@@ -499,11 +545,4 @@ class FeatureService:
status=LicenseStatus(license_info.get("status", LicenseStatus.INACTIVE))
)
if "PluginInstallationPermission" in enterprise_info:
plugin_installation_info = enterprise_info["PluginInstallationPermission"]
features.plugin_installation_permission.plugin_installation_scope = plugin_installation_info[
"pluginInstallationScope"
]
features.plugin_installation_permission.restrict_to_marketplace_only = plugin_installation_info[
"restrictToMarketplaceOnly"
]
features.plugin_installation_permission = cls._resolve_plugin_installation_permission(enterprise_info)
+2 -1
View File
@@ -18,6 +18,7 @@ from core.app.apps.completion.app_config_manager import CompletionAppConfigManag
from core.helper import encrypter
from core.prompt.simple_prompt_transform import SimplePromptTransform
from core.prompt.utils.prompt_template_parser import PromptTemplateParser
from core.workflow.variable_prefixes import USER_INPUT_VARIABLE_NODE_ID
from events.app_event import app_was_created
from graphon.file import FileUploadConfig
from graphon.model_runtime.entities.llm_entities import LLMMode
@@ -597,7 +598,7 @@ class WorkflowConverter:
},
"vision": {
"enabled": file_upload is not None,
"variable_selector": ["sys", "files"] if file_upload is not None else None,
"variable_selector": [USER_INPUT_VARIABLE_NODE_ID, "files"] if file_upload is not None else None,
"configs": {"detail": file_upload.image_config.detail}
if file_upload is not None and file_upload.image_config is not None
else None,
+55 -2
View File
@@ -6,8 +6,10 @@ from collections.abc import Callable, Generator, Mapping, Sequence
from dataclasses import dataclass
from typing import Any, cast
from sqlalchemy import exists, select
from sqlalchemy import exists, inspect, select, update
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.orm import Session, sessionmaker
from sqlalchemy.orm.attributes import set_committed_value
from configs import dify_config
from core.app.apps.advanced_chat.app_config_manager import AdvancedChatAppConfigManager
@@ -143,6 +145,8 @@ from .human_input_delivery_test_service import (
from .workflow_draft_variable_service import DraftVariableSaver, DraftVarLoader, WorkflowDraftVariableService
from .workflow_restore import apply_published_workflow_snapshot_to_draft
logger = logging.getLogger(__name__)
_file_access_controller = DatabaseFileAccessController()
@@ -155,6 +159,7 @@ class WorkflowService:
"""Initialize WorkflowService with repository dependencies."""
if session_maker is None:
session_maker = sessionmaker(bind=db.engine, expire_on_commit=False)
self._session_maker = session_maker
self._node_execution_service_repo = DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository(
session_maker
)
@@ -211,7 +216,7 @@ class WorkflowService:
)
# return draft workflow
return workflow
return self._persist_legacy_sys_files_migration_on_load(workflow)
def get_published_workflow_by_id(self, app_model: App, workflow_id: str, *, session: Session) -> Workflow | None:
"""
@@ -236,6 +241,7 @@ class WorkflowService:
f"Cannot use draft workflow version. Workflow ID: {workflow_id}. "
f"Please use a published workflow version or leave workflow_id empty."
)
self._persist_legacy_sys_files_migration_on_load(workflow)
return workflow
def get_published_workflow(self, app_model: App, *, session: Session) -> Workflow | None:
@@ -259,6 +265,53 @@ class WorkflowService:
.limit(1)
)
return self._persist_legacy_sys_files_migration_on_load(workflow)
def _persist_legacy_sys_files_migration_on_load(self, workflow: Workflow | None) -> Workflow | None:
"""Persist a load-time graph rewrite without joining or dirtying the caller's transaction."""
if workflow is None:
return None
if inspect(workflow, raiseerr=False) is None:
return workflow
# TODO: Remove this load-time persistence path after the historical workflow migration is complete.
original_graph = workflow.graph
if not workflow.migrate_legacy_sys_files_graph_in_place():
return workflow
migrated_graph = workflow.graph
try:
with self._session_maker.begin() as session:
result = session.execute(
update(Workflow)
.where(
Workflow.id == workflow.id,
Workflow.tenant_id == workflow.tenant_id,
Workflow.graph == original_graph,
)
.values(graph=migrated_graph)
)
if getattr(result, "rowcount", None) == 0:
logger.warning(
"Skipped persisting legacy sys.files workflow migration because the workflow changed "
"concurrently, "
"workflow_id=%s tenant_id=%s",
workflow.id,
workflow.tenant_id,
)
except SQLAlchemyError:
logger.warning(
"Failed to persist legacy sys.files workflow migration, workflow_id=%s tenant_id=%s",
workflow.id,
workflow.tenant_id,
exc_info=True,
)
finally:
# The conditional update owns persistence. Mark the caller's instance clean so its later flush cannot
# overwrite a concurrent workflow edit with the compatibility rewrite.
set_committed_value(workflow, "graph", migrated_graph)
return workflow
def get_accessible_app_ids(self, app_ids: Sequence[str], tenant_id: str, *, session: Session) -> set[str]:
+1 -1
View File
@@ -61,7 +61,7 @@ workflow:
query_prompt_template: '{{#sys.query#}}
{{#sys.files#}}'
{{#userinput.files#}}'
window:
enabled: false
size: 10
@@ -162,7 +162,7 @@ workflow:
query_prompt_template: '{{#sys.query#}}
{{#sys.files#}}'
{{#userinput.files#}}'
role_prefix:
assistant: ''
user: ''
@@ -207,7 +207,7 @@ workflow:
query_prompt_template: '{{#sys.query#}}
{{#sys.files#}}'
{{#userinput.files#}}'
role_prefix:
assistant: ''
user: ''
+1 -1
View File
@@ -178,7 +178,7 @@ workflow:
query_prompt_template: '{{#sys.query#}}
{{#sys.files#}}'
{{#userinput.files#}}'
role_prefix:
assistant: ''
user: ''
@@ -1,494 +1,85 @@
"""Testcontainers integration tests for controllers.console.datasets.data_source endpoints."""
"""Integration coverage for Notion page bindings backed by persisted documents."""
from __future__ import annotations
import inspect
from collections.abc import Iterator
from datetime import UTC, datetime
from unittest.mock import MagicMock, PropertyMock, patch
from inspect import unwrap
from unittest.mock import MagicMock, patch
from uuid import uuid4
import pytest
from flask import Flask
from sqlalchemy.orm import Session
from werkzeug.exceptions import NotFound
from controllers.console.datasets import data_source
from controllers.console.datasets.data_source import (
DataSourceApi,
DataSourceNotionDatasetSyncApi,
DataSourceNotionDocumentSyncApi,
DataSourceNotionIndexingEstimateApi,
DataSourceNotionListApi,
DataSourceNotionPreviewApi,
)
from core.rag.index_processor.constant.index_type import IndexStructureType
from models import Account, DataSourceOauthBinding
from controllers.console.datasets.data_source import DataSourceNotionListApi
from models import Account
from models.dataset import Document
from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus
@pytest.fixture
def current_user() -> Account:
account = Account(name="Test User", email="u1@example.com")
account.id = "u1"
return account
def test_notion_page_is_marked_bound_from_persisted_document(
flask_app_with_containers: Flask,
db_session_with_containers: Session,
) -> None:
tenant_id = str(uuid4())
dataset_id = str(uuid4())
account = Account(name="Test User", email="user@example.com")
account.id = str(uuid4())
document = Document(
tenant_id=tenant_id,
dataset_id=dataset_id,
position=1,
data_source_type=DataSourceType.NOTION_IMPORT,
data_source_info='{"notion_page_id": "page-1"}',
batch=f"batch-{uuid4()}",
name="Notion Page",
created_from=DocumentCreatedFrom.WEB,
created_by=str(uuid4()),
indexing_status=IndexingStatus.COMPLETED,
enabled=True,
)
db_session_with_containers.add(document)
db_session_with_containers.commit()
runtime = MagicMock(
get_online_document_pages=lambda **_kwargs: iter(
[
MagicMock(
result=[
MagicMock(
workspace_id="workspace-1",
workspace_name="Workspace",
workspace_icon=None,
pages=[
MagicMock(
page_id="page-1",
page_name="Page",
type="page",
parent_id="parent",
page_icon=None,
)
],
)
]
)
]
),
datasource_provider_type=lambda: None,
)
@pytest.fixture
def mock_engine() -> Iterator[None]:
with patch.object(
type(data_source.db),
"engine",
new_callable=PropertyMock,
return_value=MagicMock(),
with (
flask_app_with_containers.test_request_context(f"/?credential_id=c1&dataset_id={dataset_id}"),
patch(
"controllers.console.datasets.data_source.DatasourceProviderService.get_datasource_credentials",
return_value={"token": "token"},
),
patch(
"controllers.console.datasets.data_source.DatasetService.get_dataset",
return_value=MagicMock(data_source_type="notion_import"),
),
patch(
"core.datasource.datasource_manager.DatasourceManager.get_datasource_runtime",
return_value=runtime,
),
):
yield
class TestDataSourceApi:
@pytest.fixture
def app(self, flask_app_with_containers: Flask) -> Flask:
return flask_app_with_containers
def test_get_success(self, app: Flask) -> None:
api = DataSourceApi()
method = inspect.unwrap(api.get)
binding = DataSourceOauthBinding(
tenant_id="tenant-1",
access_token="token",
provider="notion",
source_info={
"workspace_name": "Workspace",
"workspace_id": "workspace-1",
"workspace_icon": None,
"total": 1,
"pages": [
{
"page_id": "page-1",
"page_name": "Page",
"page_icon": {"type": "emoji", "emoji": "P", "url": None},
"parent_id": "parent-1",
"type": "page",
}
],
},
)
binding.id = "b1"
binding.created_at = datetime(2026, 5, 25, 1, 2, 3, tzinfo=UTC)
binding.disabled = False
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.data_source.db.session.scalars",
return_value=MagicMock(all=lambda: [binding]),
),
):
response, status = method(api, "tenant-1")
assert status == 200
assert response["data"][0] == {
"id": "b1",
"provider": "notion",
"created_at": 1779670923,
"is_bound": True,
"disabled": False,
"source_info": {
"workspace_name": "Workspace",
"workspace_id": "workspace-1",
"workspace_icon": None,
"pages": [
{
"page_name": "Page",
"page_id": "page-1",
"page_icon": {"type": "emoji", "url": None, "emoji": "P"},
"parent_id": "parent-1",
"type": "page",
}
],
"total": 1,
},
"link": "http://localhost/console/api/oauth/data-source/notion",
}
def test_get_no_bindings(self, app: Flask) -> None:
api = DataSourceApi()
method = inspect.unwrap(api.get)
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.data_source.db.session.scalars",
return_value=MagicMock(all=lambda: []),
),
):
response, status = method(api, "tenant-1")
assert status == 200
assert response["data"] == []
def test_patch_enable_binding(self, app: Flask) -> None:
api = DataSourceApi()
method = inspect.unwrap(api.patch)
binding = MagicMock(id="b1", disabled=True)
session = MagicMock()
session.scalar.return_value = binding
with app.test_request_context("/"):
response, status = method(api, session, "tenant-1", "b1", "enable")
assert status == 200
assert binding.disabled is False
def test_patch_disable_binding(self, app: Flask) -> None:
api = DataSourceApi()
method = inspect.unwrap(api.patch)
binding = MagicMock(id="b1", disabled=False)
session = MagicMock()
session.scalar.return_value = binding
with app.test_request_context("/"):
response, status = method(api, session, "tenant-1", "b1", "disable")
assert status == 200
assert binding.disabled is True
def test_patch_binding_not_found(self, app: Flask) -> None:
api = DataSourceApi()
method = inspect.unwrap(api.patch)
session = MagicMock()
session.scalar.return_value = None
with app.test_request_context("/"):
with pytest.raises(NotFound):
method(api, session, "tenant-1", "b1", "enable")
def test_patch_enable_already_enabled(self, app: Flask) -> None:
api = DataSourceApi()
method = inspect.unwrap(api.patch)
binding = MagicMock(id="b1", disabled=False)
session = MagicMock()
session.scalar.return_value = binding
with app.test_request_context("/"):
with pytest.raises(ValueError):
method(api, session, "tenant-1", "b1", "enable")
def test_patch_disable_already_disabled(self, app: Flask) -> None:
api = DataSourceApi()
method = inspect.unwrap(api.patch)
binding = MagicMock(id="b1", disabled=True)
session = MagicMock()
session.scalar.return_value = binding
with app.test_request_context("/"):
with pytest.raises(ValueError):
method(api, session, "tenant-1", "b1", "disable")
class TestDataSourceNotionListApi:
@pytest.fixture
def app(self, flask_app_with_containers: Flask) -> Flask:
return flask_app_with_containers
def test_get_credential_not_found(self, app: Flask, current_user: Account) -> None:
api = DataSourceNotionListApi()
method = inspect.unwrap(api.get)
with (
app.test_request_context("/?credential_id=c1"),
patch(
"controllers.console.datasets.data_source.DatasourceProviderService.get_datasource_credentials",
return_value=None,
),
):
with pytest.raises(NotFound):
method(api, MagicMock(), "tenant-1", current_user)
def test_get_success_no_dataset_id(self, app: Flask, current_user: Account, mock_engine: None) -> None:
api = DataSourceNotionListApi()
method = inspect.unwrap(api.get)
page = MagicMock(
page_id="p1",
page_name="Page 1",
type="page",
parent_id="parent",
page_icon=None,
response, status = unwrap(DataSourceNotionListApi().get)(
DataSourceNotionListApi(), db_session_with_containers, tenant_id, account
)
online_document_message = MagicMock(
result=[
MagicMock(
workspace_id="w1",
workspace_name="My Workspace",
workspace_icon="icon",
pages=[page],
)
]
)
with (
app.test_request_context("/?credential_id=c1"),
patch(
"controllers.console.datasets.data_source.DatasourceProviderService.get_datasource_credentials",
return_value={"token": "t"},
),
patch(
"core.datasource.datasource_manager.DatasourceManager.get_datasource_runtime",
return_value=MagicMock(
get_online_document_pages=lambda **kw: iter([online_document_message]),
datasource_provider_type=lambda: None,
),
),
):
response, status = method(api, MagicMock(), "tenant-1", current_user)
assert status == 200
def test_get_success_with_dataset_id(
self, app: Flask, current_user: Account, mock_engine: None, db_session_with_containers: Session
) -> None:
api = DataSourceNotionListApi()
method = inspect.unwrap(api.get)
tenant_id = str(uuid4())
dataset_id = str(uuid4())
page = MagicMock(
page_id="p1",
page_name="Page 1",
type="page",
parent_id="parent",
page_icon=None,
)
online_document_message = MagicMock(
result=[
MagicMock(
workspace_id="w1",
workspace_name="My Workspace",
workspace_icon="icon",
pages=[page],
)
]
)
dataset = MagicMock(data_source_type="notion_import")
document = Document(
tenant_id=tenant_id,
dataset_id=dataset_id,
position=1,
data_source_type=DataSourceType.NOTION_IMPORT,
data_source_info='{"notion_page_id": "p1"}',
batch=f"batch-{uuid4()}",
name="Notion Page",
created_from=DocumentCreatedFrom.WEB,
created_by=str(uuid4()),
indexing_status=IndexingStatus.COMPLETED,
enabled=True,
)
db_session_with_containers.add(document)
db_session_with_containers.commit()
with (
app.test_request_context(f"/?credential_id=c1&dataset_id={dataset_id}"),
patch(
"controllers.console.datasets.data_source.DatasourceProviderService.get_datasource_credentials",
return_value={"token": "t"},
),
patch(
"controllers.console.datasets.data_source.DatasetService.get_dataset",
return_value=dataset,
),
patch(
"core.datasource.datasource_manager.DatasourceManager.get_datasource_runtime",
return_value=MagicMock(
get_online_document_pages=lambda **kw: iter([online_document_message]),
datasource_provider_type=lambda: None,
),
),
):
response, status = method(api, db_session_with_containers, tenant_id, current_user)
assert status == 200
def test_get_invalid_dataset_type(self, app: Flask, current_user: Account, mock_engine: None) -> None:
api = DataSourceNotionListApi()
method = inspect.unwrap(api.get)
dataset = MagicMock(data_source_type="other_type")
with (
app.test_request_context("/?credential_id=c1&dataset_id=ds1"),
patch(
"controllers.console.datasets.data_source.DatasourceProviderService.get_datasource_credentials",
return_value={"token": "t"},
),
patch(
"controllers.console.datasets.data_source.DatasetService.get_dataset",
return_value=dataset,
),
):
with pytest.raises(ValueError):
method(api, MagicMock(), "tenant-1", current_user)
class TestDataSourceNotionPreviewApi:
@pytest.fixture
def app(self, flask_app_with_containers: Flask) -> Flask:
return flask_app_with_containers
def test_get_preview_success(self, app: Flask) -> None:
api = DataSourceNotionPreviewApi()
method = inspect.unwrap(api.get)
extractor = MagicMock(extract=lambda: [MagicMock(page_content="hello")])
with (
app.test_request_context("/?credential_id=c1"),
patch(
"controllers.console.datasets.data_source.DatasourceProviderService.get_datasource_credentials",
return_value={"integration_secret": "t"},
),
patch(
"controllers.console.datasets.data_source.NotionExtractor",
return_value=extractor,
),
):
response, status = method(api, "tenant-1", "p1", "page")
assert status == 200
class TestDataSourceNotionIndexingEstimateApi:
@pytest.fixture
def app(self, flask_app_with_containers: Flask) -> Flask:
return flask_app_with_containers
def test_post_indexing_estimate_success(self, app: Flask) -> None:
api = DataSourceNotionIndexingEstimateApi()
method = inspect.unwrap(api.post)
empty_rules: dict[str, object] = {}
payload: dict[str, object] = {
"notion_info_list": [
{
"workspace_id": "w1",
"credential_id": "c1",
"pages": [{"page_id": "p1", "type": "page"}],
}
],
"process_rule": {"rules": empty_rules},
"doc_form": IndexStructureType.PARAGRAPH_INDEX,
"doc_language": "English",
}
with (
app.test_request_context("/", method="POST", json=payload, headers={"Content-Type": "application/json"}),
patch(
"controllers.console.datasets.data_source.DocumentService.estimate_args_validate",
),
patch(
"controllers.console.datasets.data_source.IndexingRunner.indexing_estimate",
return_value=MagicMock(model_dump=lambda: {"total_pages": 1}),
),
):
response, status = method(api, MagicMock(), "tenant-1")
assert status == 200
class TestDataSourceNotionDatasetSyncApi:
@pytest.fixture
def app(self, flask_app_with_containers: Flask) -> Flask:
return flask_app_with_containers
def test_get_success(self, app: Flask) -> None:
api = DataSourceNotionDatasetSyncApi()
method = inspect.unwrap(api.get)
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.data_source.DatasetService.get_dataset",
return_value=MagicMock(),
),
patch(
"controllers.console.datasets.data_source.DocumentService.get_document_by_dataset_id",
return_value=[MagicMock(id="d1")],
),
patch(
"controllers.console.datasets.data_source.document_indexing_sync_task.delay",
return_value=None,
),
):
response, status = method(api, MagicMock(), "ds-1")
assert status == 200
def test_get_dataset_not_found(self, app: Flask) -> None:
api = DataSourceNotionDatasetSyncApi()
method = inspect.unwrap(api.get)
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.data_source.DatasetService.get_dataset",
return_value=None,
),
):
with pytest.raises(NotFound):
method(api, MagicMock(), "ds-1")
class TestDataSourceNotionDocumentSyncApi:
@pytest.fixture
def app(self, flask_app_with_containers: Flask) -> Flask:
return flask_app_with_containers
def test_get_success(self, app: Flask) -> None:
api = DataSourceNotionDocumentSyncApi()
method = inspect.unwrap(api.get)
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.data_source.DatasetService.get_dataset",
return_value=MagicMock(),
),
patch(
"controllers.console.datasets.data_source.DocumentService.get_document",
return_value=MagicMock(),
),
patch(
"controllers.console.datasets.data_source.document_indexing_sync_task.delay",
return_value=None,
),
):
response, status = method(api, MagicMock(), "ds-1", "doc-1")
assert status == 200
def test_get_document_not_found(self, app: Flask) -> None:
api = DataSourceNotionDocumentSyncApi()
method = inspect.unwrap(api.get)
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.data_source.DatasetService.get_dataset",
return_value=MagicMock(),
),
patch(
"controllers.console.datasets.data_source.DocumentService.get_document",
return_value=None,
),
):
with pytest.raises(NotFound):
method(api, MagicMock(), "ds-1", "doc-1")
assert status == 200
assert response["notion_info"][0]["pages"][0]["is_bound"] is True
@@ -4,7 +4,7 @@ import datetime
import json
import uuid
from decimal import Decimal
from unittest.mock import MagicMock, patch
from unittest.mock import patch
import pytest
from faker import Faker
@@ -1172,65 +1172,8 @@ class TestMessagesCleanServiceIntegration:
# Verify all messages were deleted
assert db_session_with_containers.query(Message).where(Message.id.in_(msg_ids)).count() == 0
def test_from_time_range_validation(self):
"""Test that from_time_range raises ValueError for invalid inputs."""
policy = MagicMock(spec=BillingDisabledPolicy)
now = datetime.datetime.now()
with pytest.raises(ValueError, match="start_from .* must be less than end_before"):
MessagesCleanService.from_time_range(policy, now, now)
with pytest.raises(ValueError, match="batch_size .* must be greater than 0"):
MessagesCleanService.from_time_range(policy, now - datetime.timedelta(days=1), now, batch_size=0)
def test_from_time_range_success(self):
"""Test that from_time_range creates a service with correct parameters."""
policy = MagicMock(spec=BillingDisabledPolicy)
start = datetime.datetime(2024, 1, 1)
end = datetime.datetime(2024, 2, 1)
service = MessagesCleanService.from_time_range(policy, start, end)
assert service._start_from == start
assert service._end_before == end
def test_from_days_validation(self):
"""Test that from_days raises ValueError for invalid inputs."""
policy = MagicMock(spec=BillingDisabledPolicy)
with pytest.raises(ValueError, match="days .* must be greater than or equal to 0"):
MessagesCleanService.from_days(policy, days=-1)
with pytest.raises(ValueError, match="batch_size .* must be greater than 0"):
MessagesCleanService.from_days(policy, days=30, batch_size=0)
def test_from_days_success(self):
"""Test that from_days creates a service with correct parameters."""
policy = MagicMock(spec=BillingDisabledPolicy)
with patch("services.retention.conversation.messages_clean_service.naive_utc_now") as mock_now:
fixed_now = datetime.datetime(2024, 6, 1)
mock_now.return_value = fixed_now
service = MessagesCleanService.from_days(policy, days=10)
assert service._start_from is None
assert service._end_before == fixed_now - datetime.timedelta(days=10)
def test_batch_delete_message_relations_empty(self, db_session_with_containers: Session):
"""Test that batch_delete_message_relations with empty list does nothing."""
# Get execute call count before
MessagesCleanService._batch_delete_message_relations(db_session_with_containers, [])
# No exception means success — empty list is a no-op
def test_run_calls_clean_messages(self):
"""Test that run() delegates to _clean_messages_by_time_range."""
policy = MagicMock(spec=BillingDisabledPolicy)
service = MessagesCleanService(
policy=policy,
end_before=datetime.datetime.now(),
batch_size=10,
)
with patch.object(service, "_clean_messages_by_time_range") as mock_clean:
mock_clean.return_value = {"total_deleted": 5}
result = service.run()
assert result == {"total_deleted": 5}
mock_clean.assert_called_once()
@@ -4,7 +4,7 @@ from unittest.mock import MagicMock
import pytest
from sqlalchemy.orm import Session
from commands import system as system_commands
from commands import app_maintenance as app_maintenance_commands
def test_fix_app_site_missing_passes_loaded_session_to_signal(monkeypatch: pytest.MonkeyPatch) -> None:
@@ -30,15 +30,15 @@ def test_fix_app_site_missing_passes_loaded_session_to_signal(monkeypatch: pytes
engine = MagicMock()
engine.begin.return_value.__enter__.return_value = connection
monkeypatch.setattr(system_commands, "db", SimpleNamespace(engine=engine, session=scoped_session))
monkeypatch.setattr(app_maintenance_commands, "db", SimpleNamespace(engine=engine, session=scoped_session))
send = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("signal"))
monkeypatch.setattr(system_commands.app_was_created, "send", send)
monkeypatch.setattr(app_maintenance_commands.app_was_created, "send", send)
system_commands.fix_app_site_missing.callback()
app_maintenance_commands.fix_app_site_missing.callback()
scoped_session.assert_called_once_with()
scalar.assert_called_once()
get.assert_called_once_with(system_commands.Tenant, app.tenant_id)
get.assert_called_once_with(app_maintenance_commands.Tenant, app.tenant_id)
tenant.get_accounts.assert_called_once_with(session=session)
send.assert_called_once_with(app, account=account, session=session)
commit.assert_called_once_with()
@@ -62,15 +62,19 @@ def test_fix_app_site_missing_rolls_back_when_signal_fails(monkeypatch: pytest.M
engine = MagicMock()
engine.begin.return_value.__enter__.return_value = connection
monkeypatch.setattr(system_commands, "db", SimpleNamespace(engine=engine, session=MagicMock(return_value=session)))
monkeypatch.setattr(
app_maintenance_commands,
"db",
SimpleNamespace(engine=engine, session=MagicMock(return_value=session)),
)
def fail_signal(*_args, **_kwargs) -> None:
phase_events.append("signal")
raise RuntimeError("failed")
monkeypatch.setattr(system_commands.app_was_created, "send", MagicMock(side_effect=fail_signal))
monkeypatch.setattr(app_maintenance_commands.app_was_created, "send", MagicMock(side_effect=fail_signal))
system_commands.fix_app_site_missing.callback()
app_maintenance_commands.fix_app_site_missing.callback()
session.rollback.assert_called_once_with()
session.commit.assert_not_called()
@@ -6,6 +6,7 @@ import json
import os
import threading
import time
from collections.abc import Iterator
from datetime import datetime, timedelta
from pathlib import Path
from types import SimpleNamespace
@@ -15,9 +16,12 @@ import pytest
import sqlalchemy as sa
from click.testing import CliRunner
from sqlalchemy.exc import OperationalError
from sqlalchemy.orm import Session, SessionTransaction, sessionmaker
from graphon.model_runtime.entities.model_entities import ModelType
from models import Dataset, DatasetPermission, DatasetPermissionEnum
from models.account import Tenant
from models.base import TypeBase
from models.enums import CredentialSourceType
from models.provider import ProviderModel
from tests.helpers.legacy_model_type_migration import (
@@ -59,6 +63,40 @@ def command_module():
)
@pytest.fixture
def rbac_session(sqlite_engine: sa.Engine, monkeypatch: pytest.MonkeyPatch) -> Iterator[Session]:
"""Bind RBAC command reads to persisted SQLite dataset rows."""
TypeBase.metadata.create_all(
sqlite_engine,
tables=[Dataset.__table__, DatasetPermission.__table__],
)
factory = sessionmaker(bind=sqlite_engine, expire_on_commit=False)
monkeypatch.setattr("commands.rbac.session_factory.create_session", factory)
with factory() as session:
yield session
def _persist_dataset(
session: Session,
*,
dataset_id: str = "dataset-1",
tenant_id: str = "tenant-1",
permission: DatasetPermissionEnum = DatasetPermissionEnum.ONLY_ME,
created_by: str = "creator-account-1",
) -> Dataset:
dataset = Dataset(
id=dataset_id,
tenant_id=tenant_id,
name=f"Dataset {dataset_id}",
permission=permission,
created_by=created_by,
)
session.add(dataset)
session.commit()
return dataset
def _parse_json_lines(output: io.StringIO) -> list[dict[str, object]]:
return [json.loads(line) for line in output.getvalue().splitlines() if line.strip()]
@@ -363,56 +401,35 @@ def test_dataset_permission_rbac_migration_maps_legacy_permissions_to_enum_scope
def test_dataset_permission_rbac_migration_uses_dataset_creator_as_operator(
command_module,
rbac_session: Session,
monkeypatch: pytest.MonkeyPatch,
) -> None:
rbac_module = importlib.import_module("commands.rbac")
dataset_row = SimpleNamespace(
id="dataset-1",
tenant_id="tenant-1",
permission="only_me",
created_by="creator-account-1",
)
execute_results = [[dataset_row], [], []]
_persist_dataset(rbac_session)
calls: list[dict[str, object]] = []
session_closed = False
class FakeExecuteResult:
def __init__(self, rows: list[object]) -> None:
self._rows = rows
def all(self) -> list[object]:
return self._rows
class FakeSession:
def __enter__(self):
return self
def __exit__(self, exc_type, exc, traceback) -> None:
nonlocal session_closed
session_closed = True
pass
def execute(self, stmt):
return FakeExecuteResult(execute_results.pop(0))
class FakeSessionFactory:
@staticmethod
def create_session() -> FakeSession:
return FakeSession()
read_transaction_ended = False
def fake_replace_whitelist(**kwargs):
assert session_closed is True
assert read_transaction_ended is True
calls.append(kwargs)
monkeypatch.setattr(rbac_module, "session_factory", FakeSessionFactory)
monkeypatch.setattr(rbac_module.RBACService.DatasetAccess, "replace_whitelist", fake_replace_whitelist)
def _record_transaction_end(session: Session, transaction: object) -> None:
nonlocal read_transaction_ended
del transaction
if session.get_bind() is rbac_session.get_bind():
read_transaction_ended = True
command_module.migrate_dataset_permissions_to_rbac.callback(
tenant_id=None,
dataset_id=None,
batch_size=500,
dry_run=False,
)
sa.event.listen(Session, "after_transaction_end", _record_transaction_end)
monkeypatch.setattr(rbac_module.RBACService.DatasetAccess, "replace_whitelist", fake_replace_whitelist)
try:
command_module.migrate_dataset_permissions_to_rbac.callback(
tenant_id=None,
dataset_id=None,
batch_size=500,
dry_run=False,
)
finally:
sa.event.remove(Session, "after_transaction_end", _record_transaction_end)
assert calls[0]["tenant_id"] == "tenant-1"
assert calls[0]["account_id"] == "creator-account-1"
@@ -422,41 +439,19 @@ def test_dataset_permission_rbac_migration_uses_dataset_creator_as_operator(
def test_dataset_permission_rbac_migration_dry_run_outputs_structured_proposed_changes(
command_module,
rbac_session: Session,
monkeypatch: pytest.MonkeyPatch,
) -> None:
rbac_module = importlib.import_module("commands.rbac")
dataset_row = SimpleNamespace(
id="dataset-1",
tenant_id="tenant-1",
permission="partial_members",
created_by="creator-account-1",
dataset = _persist_dataset(rbac_session, permission=DatasetPermissionEnum.PARTIAL_TEAM)
rbac_session.add(
DatasetPermission(
dataset_id=dataset.id,
account_id="member-account-1",
tenant_id=dataset.tenant_id,
)
)
permission_row = SimpleNamespace(dataset_id="dataset-1", account_id="member-account-1")
execute_results = [[dataset_row], [permission_row], []]
class FakeExecuteResult:
def __init__(self, rows: list[object]) -> None:
self._rows = rows
def all(self) -> list[object]:
return self._rows
class FakeSession:
def __enter__(self):
return self
def __exit__(self, exc_type, exc, traceback) -> None:
pass
def execute(self, stmt):
return FakeExecuteResult(execute_results.pop(0))
class FakeSessionFactory:
@staticmethod
def create_session() -> FakeSession:
return FakeSession()
monkeypatch.setattr(rbac_module, "session_factory", FakeSessionFactory)
rbac_session.commit()
monkeypatch.setattr(
rbac_module.RBACService.DatasetAccess,
"replace_whitelist",
@@ -1306,50 +1301,36 @@ def test_provider_models_processing_uses_same_plan_locking_and_transaction_entry
begin_calls: list[str] = []
configure_calls: list[str] = []
class _FakeBeginContext:
def __init__(self, phase: str) -> None:
self._phase = phase
def _record_begin(session: Session, transaction: SessionTransaction) -> None:
if session.get_bind() is sqlite_engine and transaction.parent is None:
begin_calls.append(current_phase["name"])
def __enter__(self) -> None:
begin_calls.append(self._phase)
def __exit__(self, exc_type, exc, tb) -> bool:
return False
class _FakeSession:
def __init__(self, phase: str) -> None:
self._phase = phase
def __enter__(self):
return self
def __exit__(self, exc_type, exc, tb) -> bool:
return False
def begin(self) -> _FakeBeginContext:
return _FakeBeginContext(self._phase)
def _fake_session_factory(engine: sa.Engine) -> _FakeSession:
return _FakeSession(current_phase["name"])
def _fake_build_plan(self, session, candidate, *, lock_rows: bool):
def _fake_build_plan(self, session: Session, candidate, *, lock_rows: bool):
assert session.get_bind() is sqlite_engine
lock_rows_seen.append((current_phase["name"], lock_rows))
return SimpleNamespace(group_row_ids=[str(candidate.row.id)], winner=None, loser_rows=[])
return migration_module._ProviderModelGroupPlan(
group_row_ids=[str(candidate.row.id)],
winner=None,
loser_rows=[],
)
def _fake_emit_plan(self, plan, *, session, tx_id: str, business_key: dict[str, object]) -> None:
return None
def _fake_configure(self, session) -> None:
def _fake_configure(self, session: Session) -> None:
assert session.get_bind() is sqlite_engine
configure_calls.append(current_phase["name"])
monkeypatch.setattr(migration_module, "_session_factory", _fake_session_factory)
monkeypatch.setattr(migration_module.Migration, "_build_provider_model_group_plan", _fake_build_plan)
monkeypatch.setattr(migration_module.Migration, "_emit_provider_model_group_plan", _fake_emit_plan)
monkeypatch.setattr(migration_module.Migration, "_configure_lock_timeout", _fake_configure)
dry_migration._process_provider_model_group(candidate, business_key)
current_phase["name"] = "apply"
apply_migration._process_provider_model_group(candidate, business_key)
sa.event.listen(Session, "after_transaction_create", _record_begin)
try:
dry_migration._process_provider_model_group(candidate, business_key)
current_phase["name"] = "apply"
apply_migration._process_provider_model_group(candidate, business_key)
finally:
sa.event.remove(Session, "after_transaction_create", _record_begin)
assert [phase for phase, _ in lock_rows_seen] == ["dry", "apply"]
assert lock_rows_seen[0][1] == lock_rows_seen[1][1]
@@ -1392,6 +1373,22 @@ def test_process_load_balancing_model_config_row_logs_stacktrace_for_lock_timeou
sqlite_engine: sa.Engine,
monkeypatch: pytest.MonkeyPatch,
) -> None:
create_minimal_legacy_model_type_schema(sqlite_engine)
created_at = datetime(2025, 1, 1, 12, 0, 0)
_insert_load_balancing_model_config(
sqlite_engine,
row_id="40000000-0000-0000-0000-000000000001",
tenant_id="tenant-1",
provider_name="openai",
model_name="gpt-4o-mini",
model_type="text-generation",
name="credential",
encrypted_config="{}",
credential_id="50000000-0000-0000-0000-000000000001",
enabled=True,
created_at=created_at,
updated_at=created_at,
)
output = io.StringIO()
migration = migration_module.Migration(
tenant_id="tenant-1",
@@ -1401,37 +1398,18 @@ def test_process_load_balancing_model_config_row_logs_stacktrace_for_lock_timeou
model_types=(ModelType.LLM,),
orm_models=(migration_module.LoadBalancingModelConfig,),
)
candidate = migration_module._RowWithRawModelType(
row=SimpleNamespace(id="lb-row-1"),
raw_model_type="text-generation",
canonical_model_type=ModelType.LLM,
)
candidate = migration._load_load_balancing_model_config_candidates(None)[0]
lock_timeout_exc = OperationalError("SELECT 1", {}, SimpleNamespace(pgcode="55P03"))
transaction_begins = 0
class _FakeBeginContext:
def __enter__(self) -> None:
return None
def __exit__(self, exc_type, exc, tb) -> bool:
return False
class _FakeSession:
def __enter__(self):
return self
def __exit__(self, exc_type, exc, tb) -> bool:
return False
def begin(self) -> _FakeBeginContext:
return _FakeBeginContext()
def _fake_session_factory(engine: sa.Engine) -> _FakeSession:
return _FakeSession()
def _record_begin(session: Session, transaction: SessionTransaction) -> None:
nonlocal transaction_begins
if session.get_bind() is sqlite_engine and transaction.parent is None:
transaction_begins += 1
def _fake_reload(self, session, original_candidate, *, lock_rows: bool):
raise lock_timeout_exc
monkeypatch.setattr(migration_module, "_session_factory", _fake_session_factory)
monkeypatch.setattr(migration_module.Migration, "_configure_lock_timeout", lambda self, session: None)
monkeypatch.setattr(
migration_module.Migration,
@@ -1439,17 +1417,22 @@ def test_process_load_balancing_model_config_row_logs_stacktrace_for_lock_timeou
_fake_reload,
)
migration._process_load_balancing_model_config_row(candidate)
sa.event.listen(Session, "after_transaction_create", _record_begin)
try:
migration._process_load_balancing_model_config_row(candidate)
finally:
sa.event.remove(Session, "after_transaction_create", _record_begin)
lines = _parse_json_lines(output)
assert len(lines) == 1
assert lines[0]["event"] == "lock_timeout_skipped"
attrs = cast(dict[str, object], lines[0]["attrs"])
assert attrs["table_name"] == "load_balancing_model_configs"
assert attrs["id"] == "lb-row-1"
assert attrs["id"] == str(candidate.row.id)
assert attrs["error"] == str(lock_timeout_exc)
assert isinstance(attrs["stacktrace"], str)
assert "OperationalError" in attrs["stacktrace"]
assert transaction_begins == 1
def test_process_load_balancing_model_config_row_logs_update_after_sql_execution(
@@ -1457,6 +1440,23 @@ def test_process_load_balancing_model_config_row_logs_update_after_sql_execution
sqlite_engine: sa.Engine,
monkeypatch: pytest.MonkeyPatch,
) -> None:
create_minimal_legacy_model_type_schema(sqlite_engine)
created_at = datetime(2025, 1, 1, 12, 0, 0)
row_id = "40000000-0000-0000-0000-000000000002"
_insert_load_balancing_model_config(
sqlite_engine,
row_id=row_id,
tenant_id="tenant-1",
provider_name="openai",
model_name="gpt-4o-mini",
model_type="text-generation",
name="credential",
encrypted_config="{}",
credential_id="50000000-0000-0000-0000-000000000002",
enabled=True,
created_at=created_at,
updated_at=created_at,
)
migration = migration_module.Migration(
tenant_id="tenant-1",
engine=sqlite_engine,
@@ -1465,42 +1465,33 @@ def test_process_load_balancing_model_config_row_logs_update_after_sql_execution
model_types=(ModelType.LLM,),
orm_models=(migration_module.LoadBalancingModelConfig,),
)
candidate = migration_module._RowWithRawModelType(
row=SimpleNamespace(id="lb-row-1"),
raw_model_type="text-generation",
canonical_model_type=ModelType.LLM,
)
candidate = migration._load_load_balancing_model_config_candidates(None)[0]
action_log: list[str] = []
class _FakeBeginContext:
def __enter__(self) -> None:
def _record_begin(session: Session, transaction: SessionTransaction) -> None:
if session.get_bind() is sqlite_engine and transaction.parent is None:
action_log.append("begin")
def __exit__(self, exc_type, exc, tb) -> bool:
return False
class _FakeSession:
def __enter__(self):
return self
def __exit__(self, exc_type, exc, tb) -> bool:
return False
def begin(self) -> _FakeBeginContext:
return _FakeBeginContext()
def execute(self, stmt) -> None:
def _record_sql(
connection: sa.Connection,
cursor: object,
statement: str,
parameters: object,
context: object,
executemany: bool,
) -> None:
del connection, cursor, parameters, context, executemany
if statement.lstrip().upper().startswith("UPDATE"):
action_log.append("sql_execute")
def _fake_session_factory(engine: sa.Engine) -> _FakeSession:
return _FakeSession()
def _fake_configure(self, session) -> None:
action_log.append("configure_lock_timeout")
def _fake_reload(self, session, original_candidate, *, lock_rows: bool):
original_reload = migration_module.Migration._reload_load_balancing_model_config_candidate
def _record_reload(self, session: Session, original_candidate, *, lock_rows: bool):
action_log.append(f"reload_candidate:{lock_rows}")
return candidate
return original_reload(self, session, original_candidate, lock_rows=lock_rows)
def _fake_log_row_updated(self, *args, **kwargs) -> None:
action_log.append("log_row_updated")
@@ -1508,12 +1499,11 @@ def test_process_load_balancing_model_config_row_logs_update_after_sql_execution
def _fake_cache_cleanup(self, *, row_id: str, tx_id: str) -> None:
action_log.append("cache_cleanup")
monkeypatch.setattr(migration_module, "_session_factory", _fake_session_factory)
monkeypatch.setattr(migration_module.Migration, "_configure_lock_timeout", _fake_configure)
monkeypatch.setattr(
migration_module.Migration,
"_reload_load_balancing_model_config_candidate",
_fake_reload,
_record_reload,
)
monkeypatch.setattr(migration_module.Migration, "_log_row_updated", _fake_log_row_updated)
monkeypatch.setattr(
@@ -1522,7 +1512,13 @@ def test_process_load_balancing_model_config_row_logs_update_after_sql_execution
_fake_cache_cleanup,
)
migration._process_load_balancing_model_config_row(candidate)
sa.event.listen(Session, "after_transaction_create", _record_begin)
sa.event.listen(sqlite_engine, "before_cursor_execute", _record_sql)
try:
migration._process_load_balancing_model_config_row(candidate)
finally:
sa.event.remove(sqlite_engine, "before_cursor_execute", _record_sql)
sa.event.remove(Session, "after_transaction_create", _record_begin)
assert action_log == [
"begin",
@@ -1532,6 +1528,10 @@ def test_process_load_balancing_model_config_row_logs_update_after_sql_execution
"log_row_updated",
"cache_cleanup",
]
with Session(sqlite_engine) as session:
persisted = session.get(migration_module.LoadBalancingModelConfig, row_id)
assert persisted is not None
assert persisted.model_type == ModelType.LLM
def test_load_balancing_model_config_cache_delete_failure_logs_stacktrace(
@@ -0,0 +1,236 @@
from types import SimpleNamespace
from unittest.mock import MagicMock
import click
import pytest
from commands import migrate_legacy_sys_files_workflows
from commands import workflow_migration as workflow_migration_commands
def test_migrate_legacy_sys_files_workflows_command_passes_batch_options(mocker, capsys):
runner = mocker.patch.object(
workflow_migration_commands,
"run_legacy_sys_files_workflow_migration",
return_value=workflow_migration_commands.LegacySysFilesWorkflowMigrationStats(
scanned=10,
migrated=2,
failed=0,
batches=1,
last_id="workflow-10",
),
)
migrate_legacy_sys_files_workflows.callback(
batch_size=200,
limit=500,
start_after_id="workflow-1",
tenant_id="tenant-1",
app_id="app-1",
dry_run=True,
)
runner.assert_called_once_with(
batch_size=200,
limit=500,
start_after_id="workflow-1",
tenant_id="tenant-1",
app_id="app-1",
dry_run=True,
)
captured = capsys.readouterr()
assert "scanned=10" in captured.out
assert "migrated=2" in captured.out
assert "last_id=workflow-10" in captured.out
def test_migrate_legacy_sys_files_workflows_rejects_non_positive_batch_size():
with pytest.raises(click.UsageError, match="batch-size"):
migrate_legacy_sys_files_workflows.callback(
batch_size=0,
limit=None,
start_after_id=None,
tenant_id=None,
app_id=None,
dry_run=False,
)
def test_migrate_legacy_sys_files_workflows_rejects_non_positive_limit():
with pytest.raises(click.UsageError, match="limit"):
migrate_legacy_sys_files_workflows.callback(
batch_size=100,
limit=0,
start_after_id=None,
tenant_id=None,
app_id=None,
dry_run=False,
)
def test_build_legacy_sys_files_workflow_query_uses_keyset_pagination():
stmt = workflow_migration_commands._build_legacy_sys_files_workflow_query(
start_after_id="workflow-1",
batch_size=200,
tenant_id="tenant-1",
app_id="app-1",
)
compiled = str(stmt.compile(compile_kwargs={"literal_binds": True}))
assert "workflows.id > 'workflow-1'" in compiled
assert "workflows.tenant_id = 'tenant-1'" in compiled
assert "workflows.app_id = 'app-1'" in compiled
assert "ORDER BY workflows.id" in compiled
assert "LIMIT 200" in compiled
assert "workflows.environment_variables" not in compiled
def test_migrate_legacy_sys_files_workflow_batch_dry_run_rolls_back():
migrated_workflow = MagicMock()
migrated_workflow.id = "workflow-1"
migrated_workflow.migrate_legacy_sys_files_graph_in_place.return_value = True
untouched_workflow = MagicMock()
untouched_workflow.id = "workflow-2"
untouched_workflow.migrate_legacy_sys_files_graph_in_place.return_value = False
session = MagicMock()
session.scalars.return_value.all.return_value = [migrated_workflow, untouched_workflow]
stats = workflow_migration_commands._migrate_legacy_sys_files_workflow_batch(
session=session,
start_after_id=None,
batch_size=200,
tenant_id=None,
app_id=None,
dry_run=True,
)
assert stats.scanned == 2
assert stats.migrated == 1
assert stats.failed == 0
assert stats.last_id == "workflow-2"
session.rollback.assert_called_once()
session.commit.assert_not_called()
def test_migrate_legacy_sys_files_workflow_batch_commits_and_counts_failures(caplog):
migrated_workflow = MagicMock()
migrated_workflow.id = "workflow-1"
migrated_workflow.migrate_legacy_sys_files_graph_in_place.return_value = True
failing_workflow = MagicMock()
failing_workflow.id = "workflow-2"
failing_workflow.migrate_legacy_sys_files_graph_in_place.side_effect = RuntimeError("boom")
session = MagicMock()
session.scalars.return_value.all.return_value = [migrated_workflow, failing_workflow]
stats = workflow_migration_commands._migrate_legacy_sys_files_workflow_batch(
session=session,
start_after_id=None,
batch_size=200,
tenant_id=None,
app_id=None,
dry_run=False,
)
assert stats.scanned == 2
assert stats.migrated == 1
assert stats.failed == 1
assert stats.last_id == "workflow-2"
assert "Failed to migrate legacy" in caplog.text
session.commit.assert_called_once()
session.rollback.assert_not_called()
def test_run_legacy_sys_files_workflow_migration_uses_keyset_batches(mocker):
session_maker = MagicMock()
sessions = [MagicMock(), MagicMock()]
session_maker.side_effect = sessions
mocker.patch.object(workflow_migration_commands, "sessionmaker", return_value=session_maker)
mocker.patch.object(workflow_migration_commands, "db", SimpleNamespace(engine=object()))
migrate_batch = mocker.patch.object(
workflow_migration_commands,
"_migrate_legacy_sys_files_workflow_batch",
side_effect=[
workflow_migration_commands.LegacySysFilesWorkflowMigrationStats(
scanned=2,
migrated=1,
failed=0,
last_id="workflow-2",
),
workflow_migration_commands.LegacySysFilesWorkflowMigrationStats(
scanned=1,
migrated=1,
failed=0,
last_id="workflow-3",
),
],
)
stats = workflow_migration_commands.run_legacy_sys_files_workflow_migration(
batch_size=2,
limit=3,
start_after_id="workflow-0",
tenant_id="tenant-1",
app_id="app-1",
dry_run=True,
)
assert stats.scanned == 3
assert stats.migrated == 2
assert stats.batches == 2
assert stats.last_id == "workflow-3"
assert migrate_batch.call_args_list[0].kwargs["start_after_id"] == "workflow-0"
assert migrate_batch.call_args_list[0].kwargs["batch_size"] == 2
assert migrate_batch.call_args_list[1].kwargs["start_after_id"] == "workflow-2"
assert migrate_batch.call_args_list[1].kwargs["batch_size"] == 1
def test_run_legacy_sys_files_workflow_migration_stops_on_empty_batch(mocker):
session_maker = MagicMock(return_value=MagicMock())
mocker.patch.object(workflow_migration_commands, "sessionmaker", return_value=session_maker)
mocker.patch.object(workflow_migration_commands, "db", SimpleNamespace(engine=object()))
mocker.patch.object(
workflow_migration_commands,
"_migrate_legacy_sys_files_workflow_batch",
return_value=workflow_migration_commands.LegacySysFilesWorkflowMigrationStats(scanned=0),
)
stats = workflow_migration_commands.run_legacy_sys_files_workflow_migration(
batch_size=2,
limit=None,
start_after_id=None,
tenant_id=None,
app_id=None,
dry_run=False,
)
assert stats.scanned == 0
assert stats.batches == 0
def test_run_legacy_sys_files_workflow_migration_stops_on_short_batch(mocker):
session_maker = MagicMock(return_value=MagicMock())
mocker.patch.object(workflow_migration_commands, "sessionmaker", return_value=session_maker)
mocker.patch.object(workflow_migration_commands, "db", SimpleNamespace(engine=object()))
migrate_batch = mocker.patch.object(
workflow_migration_commands,
"_migrate_legacy_sys_files_workflow_batch",
return_value=workflow_migration_commands.LegacySysFilesWorkflowMigrationStats(
scanned=1,
migrated=1,
failed=0,
last_id="workflow-1",
),
)
stats = workflow_migration_commands.run_legacy_sys_files_workflow_migration(
batch_size=2,
limit=None,
start_after_id=None,
tenant_id=None,
app_id=None,
dry_run=False,
)
assert stats.scanned == 1
assert stats.batches == 1
migrate_batch.assert_called_once()
@@ -14,7 +14,7 @@ from sqlalchemy import select
from sqlalchemy.orm import Session
import commands
from commands import system as system_commands
from commands import workspace as workspace_commands
from core.tools.entities.tool_entities import ApiProviderSchemaType
from graphon.model_runtime.entities.model_entities import ModelType
from models import Tenant
@@ -83,11 +83,11 @@ def _encrypted_rows(tenant_id: str, *, suffix: str = "1") -> tuple[object, ...]:
def _bind_command_to_sqlite(monkeypatch: pytest.MonkeyPatch, session: Session) -> None:
monkeypatch.setattr(system_commands, "db", SimpleNamespace(engine=session.get_bind()))
monkeypatch.setattr(workspace_commands, "db", SimpleNamespace(engine=session.get_bind()))
def test_reset_aborts_when_not_self_hosted(monkeypatch, capsys):
monkeypatch.setattr(system_commands.dify_config, "EDITION", "CLOUD")
monkeypatch.setattr(workspace_commands.dify_config, "EDITION", "CLOUD")
exit_code = _invoke_reset()
captured = capsys.readouterr()
@@ -106,8 +106,8 @@ def test_reset_purges_provider_and_tool_tables_for_each_tenant(
) -> None:
"""The command must purge LLM provider rows AND every tool provider table
that stores ciphertext encrypted under the tenant key (#35396)."""
monkeypatch.setattr(system_commands.dify_config, "EDITION", "SELF_HOSTED")
monkeypatch.setattr(system_commands, "generate_key_pair", lambda tenant_id: f"new-key-{tenant_id}")
monkeypatch.setattr(workspace_commands.dify_config, "EDITION", "SELF_HOSTED")
monkeypatch.setattr(workspace_commands, "generate_key_pair", lambda tenant_id: f"new-key-{tenant_id}")
_bind_command_to_sqlite(monkeypatch, sqlite_session)
tenant = _tenant(TENANT_ID)
@@ -146,8 +146,8 @@ def test_reset_purges_provider_and_tool_tables_for_each_tenant(
)
def test_reset_iterates_all_tenants(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
"""Multi-tenant deployments must purge every tenant, not just the first."""
monkeypatch.setattr(system_commands.dify_config, "EDITION", "SELF_HOSTED")
monkeypatch.setattr(system_commands, "generate_key_pair", lambda tenant_id: f"new-key-{tenant_id}")
monkeypatch.setattr(workspace_commands.dify_config, "EDITION", "SELF_HOSTED")
monkeypatch.setattr(workspace_commands, "generate_key_pair", lambda tenant_id: f"new-key-{tenant_id}")
_bind_command_to_sqlite(monkeypatch, sqlite_session)
tenant_ids = [f"11111111-1111-1111-1111-{index:012d}" for index in range(3)]
@@ -4,7 +4,7 @@ import types
from unittest.mock import MagicMock
import commands
from commands import system as system_commands
from commands import database as database_commands
from libs.db_migration_lock import LockNotOwnedError, RedisError
HEARTBEAT_WAIT_TIMEOUT_SECONDS = 5.0
@@ -25,11 +25,11 @@ def _invoke_upgrade_db() -> int:
def test_upgrade_db_skips_when_lock_not_acquired(monkeypatch, capsys):
monkeypatch.setattr(system_commands, "DB_UPGRADE_LOCK_TTL_SECONDS", 1234)
monkeypatch.setattr(database_commands, "DB_UPGRADE_LOCK_TTL_SECONDS", 1234)
lock = MagicMock()
lock.acquire.return_value = False
system_commands.redis_client.lock.return_value = lock
database_commands.redis_client.lock.return_value = lock
exit_code = _invoke_upgrade_db()
captured = capsys.readouterr()
@@ -37,18 +37,20 @@ def test_upgrade_db_skips_when_lock_not_acquired(monkeypatch, capsys):
assert exit_code == 0
assert "Database migration skipped" in captured.out
system_commands.redis_client.lock.assert_called_once_with(name="db_upgrade_lock", timeout=1234, thread_local=False)
database_commands.redis_client.lock.assert_called_once_with(
name="db_upgrade_lock", timeout=1234, thread_local=False
)
lock.acquire.assert_called_once_with(blocking=False)
lock.release.assert_not_called()
def test_upgrade_db_failure_not_masked_by_lock_release(monkeypatch, capsys):
monkeypatch.setattr(system_commands, "DB_UPGRADE_LOCK_TTL_SECONDS", 321)
monkeypatch.setattr(database_commands, "DB_UPGRADE_LOCK_TTL_SECONDS", 321)
lock = MagicMock()
lock.acquire.return_value = True
lock.release.side_effect = LockNotOwnedError("simulated")
system_commands.redis_client.lock.return_value = lock
database_commands.redis_client.lock.return_value = lock
def _upgrade():
raise RuntimeError("boom")
@@ -61,18 +63,18 @@ def test_upgrade_db_failure_not_masked_by_lock_release(monkeypatch, capsys):
assert exit_code == 1
assert "Database migration failed: boom" in captured.out
system_commands.redis_client.lock.assert_called_once_with(name="db_upgrade_lock", timeout=321, thread_local=False)
database_commands.redis_client.lock.assert_called_once_with(name="db_upgrade_lock", timeout=321, thread_local=False)
lock.acquire.assert_called_once_with(blocking=False)
lock.release.assert_called_once()
def test_upgrade_db_success_ignores_lock_not_owned_on_release(monkeypatch, capsys):
monkeypatch.setattr(system_commands, "DB_UPGRADE_LOCK_TTL_SECONDS", 999)
monkeypatch.setattr(database_commands, "DB_UPGRADE_LOCK_TTL_SECONDS", 999)
lock = MagicMock()
lock.acquire.return_value = True
lock.release.side_effect = LockNotOwnedError("simulated")
system_commands.redis_client.lock.return_value = lock
database_commands.redis_client.lock.return_value = lock
_install_fake_flask_migrate(monkeypatch, lambda: None)
@@ -82,7 +84,7 @@ def test_upgrade_db_success_ignores_lock_not_owned_on_release(monkeypatch, capsy
assert exit_code == 0
assert "Database migration successful!" in captured.out
system_commands.redis_client.lock.assert_called_once_with(name="db_upgrade_lock", timeout=999, thread_local=False)
database_commands.redis_client.lock.assert_called_once_with(name="db_upgrade_lock", timeout=999, thread_local=False)
lock.acquire.assert_called_once_with(blocking=False)
lock.release.assert_called_once()
@@ -93,11 +95,11 @@ def test_upgrade_db_renews_lock_during_migration(monkeypatch, capsys):
"""
# Use a small TTL so the heartbeat interval triggers quickly.
monkeypatch.setattr(system_commands, "DB_UPGRADE_LOCK_TTL_SECONDS", 0.3)
monkeypatch.setattr(database_commands, "DB_UPGRADE_LOCK_TTL_SECONDS", 0.3)
lock = MagicMock()
lock.acquire.return_value = True
system_commands.redis_client.lock.return_value = lock
database_commands.redis_client.lock.return_value = lock
renewed = threading.Event()
@@ -121,11 +123,11 @@ def test_upgrade_db_renews_lock_during_migration(monkeypatch, capsys):
def test_upgrade_db_ignores_reacquire_errors(monkeypatch, capsys):
# Use a small TTL so heartbeat runs during the upgrade call.
monkeypatch.setattr(system_commands, "DB_UPGRADE_LOCK_TTL_SECONDS", 0.3)
monkeypatch.setattr(database_commands, "DB_UPGRADE_LOCK_TTL_SECONDS", 0.3)
lock = MagicMock()
lock.acquire.return_value = True
system_commands.redis_client.lock.return_value = lock
database_commands.redis_client.lock.return_value = lock
attempted = threading.Event()
@@ -0,0 +1,49 @@
from unittest.mock import MagicMock
from commands import reset_encrypt_key_pair
from commands import workspace as workspace_commands
def test_reset_encrypt_key_pair_skips_non_self_hosted(monkeypatch, capsys):
monkeypatch.setattr(workspace_commands.dify_config, "EDITION", "CLOUD")
reset_encrypt_key_pair.callback()
captured = capsys.readouterr()
assert "only for SELF_HOSTED" in captured.out
def test_reset_encrypt_key_pair_rotates_keys_and_removes_custom_provider_data(monkeypatch, capsys):
monkeypatch.setattr(workspace_commands.dify_config, "EDITION", "SELF_HOSTED")
monkeypatch.setattr(workspace_commands, "generate_key_pair", lambda tenant_id: f"public-key-{tenant_id}")
tenant = MagicMock()
tenant.id = "tenant-1"
session = MagicMock()
session.scalars.return_value.all.return_value = [tenant]
session_manager = MagicMock()
session_manager.begin.return_value.__enter__.return_value = session
monkeypatch.setattr(workspace_commands, "sessionmaker", lambda *args, **kwargs: session_manager)
monkeypatch.setattr(workspace_commands, "db", MagicMock(engine=object()))
reset_encrypt_key_pair.callback()
assert tenant.encrypt_public_key == "public-key-tenant-1"
assert session.execute.call_count == 5
captured = capsys.readouterr()
assert "tenant-1 has been reset" in captured.out
def test_reset_encrypt_key_pair_stops_when_workspace_record_is_missing(monkeypatch, capsys):
monkeypatch.setattr(workspace_commands.dify_config, "EDITION", "SELF_HOSTED")
session = MagicMock()
session.scalars.return_value.all.return_value = [None]
session_manager = MagicMock()
session_manager.begin.return_value.__enter__.return_value = session
monkeypatch.setattr(workspace_commands, "sessionmaker", lambda *args, **kwargs: session_manager)
monkeypatch.setattr(workspace_commands, "db", MagicMock(engine=object()))
reset_encrypt_key_pair.callback()
session.execute.assert_not_called()
captured = capsys.readouterr()
assert "No workspaces found" in captured.out
@@ -1,18 +1,22 @@
from __future__ import annotations
import inspect
from collections.abc import Callable
from collections.abc import Callable, Iterator
from datetime import UTC, datetime
from typing import cast
from typing import Literal, cast
from unittest.mock import MagicMock, PropertyMock, patch
from uuid import uuid4
from uuid import UUID
import pytest
from flask import Flask
from sqlalchemy import select
from sqlalchemy.orm import Session
from werkzeug.exceptions import NotFound
from controllers.console.datasets import data_source as module
from controllers.console.datasets.data_source import DataSourceApi, DataSourceNotionListApi
from models import Account, DataSourceOauthBinding
from models.engine import db
ControllerMethod = Callable[..., tuple[dict[str, object], int]]
@@ -22,10 +26,15 @@ def unwrap(func: object) -> ControllerMethod:
@pytest.fixture
def flask_app() -> Flask:
def flask_app() -> Iterator[Flask]:
app = Flask(__name__)
app.config["TESTING"] = True
return app
app.config["SQLALCHEMY_DATABASE_URI"] = "sqlite:///:memory:"
db.init_app(app)
with app.app_context():
DataSourceOauthBinding.__table__.create(db.engine)
yield app
@pytest.fixture
@@ -35,9 +44,13 @@ def current_user() -> Account:
return account
def test_get_data_source_integrates_serializes_orm_binding(flask_app: Flask) -> None:
TENANT_ID = "11111111-1111-1111-1111-111111111111"
BINDING_ID = "22222222-2222-2222-2222-222222222222"
def _add_binding(session: Session, *, disabled: bool) -> DataSourceOauthBinding:
binding = DataSourceOauthBinding(
tenant_id="tenant-1",
tenant_id=TENANT_ID,
access_token="token",
provider="notion",
source_info={
@@ -55,24 +68,31 @@ def test_get_data_source_integrates_serializes_orm_binding(flask_app: Flask) ->
}
],
},
disabled=disabled,
)
binding.id = "binding-1"
binding.id = BINDING_ID
binding.created_at = datetime(2026, 5, 25, 1, 2, 3, tzinfo=UTC)
binding.disabled = False
session.add(binding)
session.commit()
return binding
with (
flask_app.test_request_context("/"),
patch.object(module.db.session, "scalars", return_value=MagicMock(all=lambda: [binding])),
):
response, status = unwrap(DataSourceApi().get)(DataSourceApi(), "tenant-1")
def test_get_data_source_integrates_serializes_orm_binding(
flask_app: Flask,
) -> None:
binding = _add_binding(db.session, disabled=False)
expected_created_at = int(binding.created_at.timestamp())
with flask_app.test_request_context("/"):
response, status = unwrap(DataSourceApi().get)(DataSourceApi(), TENANT_ID)
assert status == 200
assert response == {
"data": [
{
"id": "binding-1",
"id": BINDING_ID,
"provider": "notion",
"created_at": 1779670923,
"created_at": expected_created_at,
"is_bound": True,
"disabled": False,
"source_info": {
@@ -96,34 +116,75 @@ def test_get_data_source_integrates_serializes_orm_binding(flask_app: Flask) ->
}
def test_get_data_source_integrates_preserves_empty_list_when_no_binding(flask_app: Flask) -> None:
with (
flask_app.test_request_context("/"),
patch.object(module.db.session, "scalars", return_value=MagicMock(all=lambda: [])),
):
response, status = unwrap(DataSourceApi().get)(DataSourceApi(), "tenant-1")
def test_get_data_source_integrates_preserves_empty_list_when_no_binding(
flask_app: Flask,
) -> None:
with flask_app.test_request_context("/"):
response, status = unwrap(DataSourceApi().get)(DataSourceApi(), TENANT_ID)
assert status == 200
assert response == {"data": []}
def test_patch_data_source_binding_uses_injected_session(flask_app: Flask) -> None:
binding = MagicMock(disabled=True)
session = MagicMock()
session.scalar.return_value = binding
@pytest.mark.parametrize(
("disabled", "action", "expected_disabled"),
[(True, "enable", False), (False, "disable", True)],
)
@pytest.mark.parametrize("sqlite_session", [(DataSourceOauthBinding,)], indirect=True)
def test_patch_data_source_binding_updates_state(
flask_app: Flask,
sqlite_session: Session,
disabled: bool,
action: Literal["enable", "disable"],
expected_disabled: bool,
) -> None:
_add_binding(sqlite_session, disabled=disabled)
sqlite_session.expunge_all()
with flask_app.test_request_context("/"):
response, status = unwrap(DataSourceApi().patch)(DataSourceApi(), session, "tenant-1", uuid4(), "enable")
response, status = unwrap(DataSourceApi().patch)(
DataSourceApi(), sqlite_session, TENANT_ID, UUID(BINDING_ID), action
)
sqlite_session.flush()
sqlite_session.expire_all()
binding = sqlite_session.scalar(select(DataSourceOauthBinding).where(DataSourceOauthBinding.id == BINDING_ID))
assert status == 200
assert response == {"result": "success"}
assert binding.disabled is False
session.scalar.assert_called_once()
session.add.assert_not_called()
session.commit.assert_not_called()
assert binding is not None
assert binding.disabled is expected_disabled
def test_notion_pre_import_pages_serializes_frontend_list_shape(flask_app: Flask, current_user: Account) -> None:
@pytest.mark.parametrize("sqlite_session", [(DataSourceOauthBinding,)], indirect=True)
def test_patch_data_source_binding_rejects_unknown_binding(
flask_app: Flask,
sqlite_session: Session,
) -> None:
with flask_app.test_request_context("/"), pytest.raises(NotFound, match="Data source binding not found"):
unwrap(DataSourceApi().patch)(DataSourceApi(), sqlite_session, TENANT_ID, UUID(BINDING_ID), "enable")
@pytest.mark.parametrize(("disabled", "action"), [(False, "enable"), (True, "disable")])
@pytest.mark.parametrize("sqlite_session", [(DataSourceOauthBinding,)], indirect=True)
def test_patch_data_source_binding_rejects_current_state(
flask_app: Flask,
sqlite_session: Session,
disabled: bool,
action: Literal["enable", "disable"],
) -> None:
_add_binding(sqlite_session, disabled=disabled)
sqlite_session.expunge_all()
with flask_app.test_request_context("/"), pytest.raises(ValueError):
unwrap(DataSourceApi().patch)(DataSourceApi(), sqlite_session, TENANT_ID, UUID(BINDING_ID), action)
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_notion_pre_import_pages_serializes_frontend_list_shape(
flask_app: Flask,
current_user: Account,
sqlite_session: Session,
) -> None:
page = MagicMock(
page_id="page-1",
page_name="Page",
@@ -145,8 +206,6 @@ def test_notion_pre_import_pages_serializes_frontend_list_shape(flask_app: Flask
get_online_document_pages=MagicMock(return_value=iter([online_document_message])),
datasource_provider_type=MagicMock(return_value="online_document"),
)
session = MagicMock()
with (
flask_app.test_request_context("/?credential_id=credential-1"),
patch.object(
@@ -158,7 +217,7 @@ def test_notion_pre_import_pages_serializes_frontend_list_shape(flask_app: Flask
patch("core.datasource.datasource_manager.DatasourceManager.get_datasource_runtime", return_value=runtime),
):
response, status = unwrap(DataSourceNotionListApi().get)(
DataSourceNotionListApi(), session, "tenant-1", current_user
DataSourceNotionListApi(), sqlite_session, "tenant-1", current_user
)
assert status == 200
@@ -183,3 +242,38 @@ def test_notion_pre_import_pages_serializes_frontend_list_shape(flask_app: Flask
}
runtime.get_online_document_pages.assert_called_once()
assert runtime.get_online_document_pages.call_args.kwargs["datasource_parameters"] == {}
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_notion_pre_import_pages_rejects_missing_credential(
flask_app: Flask,
current_user: Account,
sqlite_session: Session,
) -> None:
with (
flask_app.test_request_context("/?credential_id=credential-1"),
patch.object(module.DatasourceProviderService, "get_datasource_credentials", return_value=None),
pytest.raises(NotFound, match="Credential not found"),
):
unwrap(DataSourceNotionListApi().get)(DataSourceNotionListApi(), sqlite_session, TENANT_ID, current_user)
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_notion_pre_import_pages_rejects_non_notion_dataset(
flask_app: Flask,
current_user: Account,
sqlite_session: Session,
) -> None:
dataset = MagicMock(data_source_type="other_type")
with (
flask_app.test_request_context("/?credential_id=credential-1&dataset_id=dataset-1"),
patch.object(
module.DatasourceProviderService,
"get_datasource_credentials",
return_value={"token": "token"},
),
patch.object(module.DatasetService, "get_dataset", return_value=dataset),
pytest.raises(ValueError, match="Dataset is not notion type"),
):
unwrap(DataSourceNotionListApi().get)(DataSourceNotionListApi(), sqlite_session, TENANT_ID, current_user)
@@ -0,0 +1,171 @@
"""Unit tests for controllers.console.datasets.data_source Notion endpoints."""
from __future__ import annotations
import inspect
from unittest.mock import MagicMock, patch
import pytest
from flask import Flask
from sqlalchemy.orm import Session
from werkzeug.exceptions import NotFound
from controllers.console.datasets.data_source import (
DataSourceNotionDatasetSyncApi,
DataSourceNotionDocumentSyncApi,
DataSourceNotionIndexingEstimateApi,
DataSourceNotionPreviewApi,
)
from core.rag.index_processor.constant.index_type import IndexStructureType
from models import Account
@pytest.fixture
def current_user() -> Account:
account = Account(name="Test User", email="u1@example.com")
account.id = "u1"
return account
class TestDataSourceNotionPreviewApi:
def test_get_preview_success(self, app: Flask) -> None:
api = DataSourceNotionPreviewApi()
method = inspect.unwrap(api.get)
extractor = MagicMock(extract=lambda: [MagicMock(page_content="hello")])
with (
app.test_request_context("/?credential_id=c1"),
patch(
"controllers.console.datasets.data_source.DatasourceProviderService.get_datasource_credentials",
return_value={"integration_secret": "t"},
),
patch(
"controllers.console.datasets.data_source.NotionExtractor",
return_value=extractor,
),
):
response, status = method(api, "tenant-1", "p1", "page")
assert status == 200
class TestDataSourceNotionIndexingEstimateApi:
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_post_indexing_estimate_success(self, app: Flask, sqlite_session: Session) -> None:
api = DataSourceNotionIndexingEstimateApi()
method = inspect.unwrap(api.post)
empty_rules: dict[str, object] = {}
payload: dict[str, object] = {
"notion_info_list": [
{
"workspace_id": "w1",
"credential_id": "c1",
"pages": [{"page_id": "p1", "type": "page"}],
}
],
"process_rule": {"rules": empty_rules},
"doc_form": IndexStructureType.PARAGRAPH_INDEX,
"doc_language": "English",
}
with (
app.test_request_context("/", method="POST", json=payload, headers={"Content-Type": "application/json"}),
patch(
"controllers.console.datasets.data_source.DocumentService.estimate_args_validate",
),
patch(
"controllers.console.datasets.data_source.IndexingRunner.indexing_estimate",
return_value=MagicMock(model_dump=lambda: {"total_pages": 1}),
),
):
response, status = method(api, sqlite_session, "tenant-1")
assert status == 200
class TestDataSourceNotionDatasetSyncApi:
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_get_success(self, app: Flask, sqlite_session: Session) -> None:
api = DataSourceNotionDatasetSyncApi()
method = inspect.unwrap(api.get)
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.data_source.DatasetService.get_dataset",
return_value=MagicMock(),
),
patch(
"controllers.console.datasets.data_source.DocumentService.get_document_by_dataset_id",
return_value=[MagicMock(id="d1")],
),
patch(
"controllers.console.datasets.data_source.document_indexing_sync_task.delay",
return_value=None,
),
):
response, status = method(api, sqlite_session, "ds-1")
assert status == 200
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_get_dataset_not_found(self, app: Flask, sqlite_session: Session) -> None:
api = DataSourceNotionDatasetSyncApi()
method = inspect.unwrap(api.get)
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.data_source.DatasetService.get_dataset",
return_value=None,
),
):
with pytest.raises(NotFound):
method(api, sqlite_session, "ds-1")
class TestDataSourceNotionDocumentSyncApi:
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_get_success(self, app: Flask, sqlite_session: Session) -> None:
api = DataSourceNotionDocumentSyncApi()
method = inspect.unwrap(api.get)
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.data_source.DatasetService.get_dataset",
return_value=MagicMock(),
),
patch(
"controllers.console.datasets.data_source.DocumentService.get_document",
return_value=MagicMock(),
),
patch(
"controllers.console.datasets.data_source.document_indexing_sync_task.delay",
return_value=None,
),
):
response, status = method(api, sqlite_session, "ds-1", "doc-1")
assert status == 200
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_get_document_not_found(self, app: Flask, sqlite_session: Session) -> None:
api = DataSourceNotionDocumentSyncApi()
method = inspect.unwrap(api.get)
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.data_source.DatasetService.get_dataset",
return_value=MagicMock(),
),
patch(
"controllers.console.datasets.data_source.DocumentService.get_document",
return_value=None,
),
):
with pytest.raises(NotFound):
method(api, sqlite_session, "ds-1", "doc-1")
@@ -0,0 +1,89 @@
from unittest.mock import MagicMock
from controllers.service_api.app.legacy_system_files import (
attach_legacy_system_file_warning_for_service_api,
normalize_legacy_system_file_args_for_service_api,
)
from core.app.entities.app_invoke_entities import InvokeFrom
from services.app_generate_service import AppGenerateService
_LEGACY_FILE_TEMPLATE = "{{#" + ".".join(("sys", "files")) + "#}}"
_USER_INPUT_FILE_INPUT_KEY = ".".join(("userinput", "files"))
def _legacy_file_graph() -> dict:
return {
"nodes": [
{"id": "start", "data": {"type": "start", "variables": []}},
{"id": "answer", "data": {"type": "answer", "answer": _LEGACY_FILE_TEMPLATE}},
],
"edges": [],
}
def test_hidden_service_api_file_payload_maps_to_userinput_files(mocker):
workflow = MagicMock()
workflow.graph_dict = _legacy_file_graph()
get_workflow = mocker.patch.object(AppGenerateService, "get_workflow", return_value=workflow)
app_model = MagicMock()
session = MagicMock()
files = [{"transfer_method": "remote_url", "url": "https://example.com/a.png"}]
args, compat_variable = normalize_legacy_system_file_args_for_service_api(
session=session,
app_model=app_model,
args={"inputs": {}, "files": None},
raw_payload={"system": {"files": files}},
)
get_workflow.assert_called_once_with(app_model, InvokeFrom.SERVICE_API, None, session=session)
assert compat_variable is not None
assert args["files"] == files
assert args["inputs"][_USER_INPUT_FILE_INPUT_KEY] == files
def test_service_api_file_payload_is_ignored_when_absent(mocker):
get_workflow = mocker.patch.object(AppGenerateService, "get_workflow")
app_model = MagicMock()
original_args = {"inputs": {}}
args, compat_variable = normalize_legacy_system_file_args_for_service_api(
session=MagicMock(),
app_model=app_model,
args=original_args,
raw_payload={},
)
assert args is original_args
assert compat_variable is None
get_workflow.assert_not_called()
def test_top_level_service_api_file_payload_still_checks_workflow_graph(mocker):
workflow = MagicMock()
workflow.graph_dict = {"nodes": []}
get_workflow = mocker.patch.object(AppGenerateService, "get_workflow", return_value=workflow)
app_model = MagicMock()
session = MagicMock()
files = [{"id": "file-1"}]
args, compat_variable = normalize_legacy_system_file_args_for_service_api(
session=session,
app_model=app_model,
args={"inputs": {}, "files": files},
raw_payload={},
)
get_workflow.assert_called_once_with(app_model, InvokeFrom.SERVICE_API, None, session=session)
assert args["files"] == files
assert compat_variable is None
def test_service_api_warning_is_attached_only_when_compatibility_was_used():
compat_variable = MagicMock(node_id="userinput", variable_name="files")
response = attach_legacy_system_file_warning_for_service_api({"answer": "ok"}, compat_variable)
response_without_warning = attach_legacy_system_file_warning_for_service_api({"answer": "ok"}, None)
assert response["warnings"]
assert response_without_warning == {"answer": "ok"}
@@ -15,12 +15,15 @@ Focus on:
"""
import uuid
from collections.abc import Iterator
from inspect import unwrap
from types import SimpleNamespace
from unittest.mock import Mock, patch
import pytest
from flask import Flask
from sqlalchemy import Engine
from sqlalchemy.orm import Session
from werkzeug.exceptions import BadRequest, InternalServerError, NotFound
from controllers.service_api.app.error import NotChatAppError
@@ -44,6 +47,14 @@ from services.errors.message import (
from services.message_service import MessageService
@pytest.fixture
def orm_session(sqlite_engine: Engine) -> Iterator[Session]:
"""Provide a real caller-owned session for MessageService interface tests."""
with Session(sqlite_engine, expire_on_commit=False) as session:
yield session
class TestMessageListQuery:
"""Test suite for MessageListQuery Pydantic model."""
@@ -253,7 +264,7 @@ class TestMessageService:
assert callable(MessageService.get_suggested_questions_after_answer)
@patch.object(MessageService, "pagination_by_first_id")
def test_pagination_by_first_id_returns_pagination_result(self, mock_pagination):
def test_pagination_by_first_id_returns_pagination_result(self, mock_pagination, orm_session: Session):
"""Test pagination_by_first_id returns expected format."""
mock_result = Mock()
mock_result.data = []
@@ -267,7 +278,7 @@ class TestMessageService:
conversation_id=str(uuid.uuid4()),
first_id=None,
limit=20,
session=Mock(),
session=orm_session,
)
assert hasattr(result, "data")
@@ -275,7 +286,7 @@ class TestMessageService:
assert hasattr(result, "has_more")
@patch.object(MessageService, "pagination_by_first_id")
def test_pagination_raises_conversation_not_exists_error(self, mock_pagination):
def test_pagination_raises_conversation_not_exists_error(self, mock_pagination, orm_session: Session):
"""Test pagination raises ConversationNotExistsError."""
import services.errors.conversation
@@ -288,11 +299,11 @@ class TestMessageService:
conversation_id="invalid_id",
first_id=None,
limit=20,
session=Mock(),
session=orm_session,
)
@patch.object(MessageService, "pagination_by_first_id")
def test_pagination_raises_first_message_not_exists_error(self, mock_pagination):
def test_pagination_raises_first_message_not_exists_error(self, mock_pagination, orm_session: Session):
"""Test pagination raises FirstMessageNotExistsError."""
mock_pagination.side_effect = FirstMessageNotExistsError()
@@ -303,11 +314,11 @@ class TestMessageService:
conversation_id=str(uuid.uuid4()),
first_id="invalid_first_id",
limit=20,
session=Mock(),
session=orm_session,
)
@patch.object(MessageService, "create_feedback")
def test_create_feedback_with_rating_and_content(self, mock_create_feedback):
def test_create_feedback_with_rating_and_content(self, mock_create_feedback, orm_session: Session):
"""Test create_feedback with rating and content."""
mock_create_feedback.return_value = None
@@ -317,13 +328,13 @@ class TestMessageService:
user=Mock(spec=EndUser),
rating=FeedbackRating.LIKE,
content="Great response!",
session=Mock(),
session=orm_session,
)
mock_create_feedback.assert_called_once()
@patch.object(MessageService, "create_feedback")
def test_create_feedback_raises_message_not_exists_error(self, mock_create_feedback):
def test_create_feedback_raises_message_not_exists_error(self, mock_create_feedback, orm_session: Session):
"""Test create_feedback raises MessageNotExistsError."""
mock_create_feedback.side_effect = MessageNotExistsError()
@@ -334,11 +345,11 @@ class TestMessageService:
user=Mock(spec=EndUser),
rating=FeedbackRating.LIKE,
content=None,
session=Mock(),
session=orm_session,
)
@patch.object(MessageService, "get_all_messages_feedbacks")
def test_get_all_messages_feedbacks_returns_list(self, mock_get_feedbacks):
def test_get_all_messages_feedbacks_returns_list(self, mock_get_feedbacks, orm_session: Session):
"""Test get_all_messages_feedbacks returns list of feedbacks."""
mock_feedbacks = [
{"message_id": str(uuid.uuid4()), "rating": "like"},
@@ -346,13 +357,15 @@ class TestMessageService:
]
mock_get_feedbacks.return_value = mock_feedbacks
result = MessageService.get_all_messages_feedbacks(app_model=Mock(spec=App), page=1, limit=20, session=Mock())
result = MessageService.get_all_messages_feedbacks(
app_model=Mock(spec=App), page=1, limit=20, session=orm_session
)
assert len(result) == 2
assert result[0]["rating"] == "like"
@patch.object(MessageService, "get_suggested_questions_after_answer")
def test_get_suggested_questions_returns_questions_list(self, mock_get_questions):
def test_get_suggested_questions_returns_questions_list(self, mock_get_questions, orm_session: Session):
"""Test get_suggested_questions_after_answer returns list of questions."""
mock_questions = ["What about this aspect?", "Can you elaborate on that?", "How does this relate to...?"]
mock_get_questions.return_value = mock_questions
@@ -362,14 +375,14 @@ class TestMessageService:
user=Mock(spec=EndUser),
message_id=str(uuid.uuid4()),
invoke_from=Mock(),
session=Mock(),
session=orm_session,
)
assert len(result) == 3
assert isinstance(result[0], str)
@patch.object(MessageService, "get_suggested_questions_after_answer")
def test_get_suggested_questions_raises_disabled_error(self, mock_get_questions):
def test_get_suggested_questions_raises_disabled_error(self, mock_get_questions, orm_session: Session):
"""Test get_suggested_questions_after_answer raises SuggestedQuestionsAfterAnswerDisabledError."""
mock_get_questions.side_effect = SuggestedQuestionsAfterAnswerDisabledError()
@@ -379,11 +392,11 @@ class TestMessageService:
user=Mock(spec=EndUser),
message_id=str(uuid.uuid4()),
invoke_from=Mock(),
session=Mock(),
session=orm_session,
)
@patch.object(MessageService, "get_suggested_questions_after_answer")
def test_get_suggested_questions_raises_message_not_exists_error(self, mock_get_questions):
def test_get_suggested_questions_raises_message_not_exists_error(self, mock_get_questions, orm_session: Session):
"""Test get_suggested_questions_after_answer raises MessageNotExistsError."""
mock_get_questions.side_effect = MessageNotExistsError()
@@ -393,7 +406,7 @@ class TestMessageService:
user=Mock(spec=EndUser),
message_id="invalid_message_id",
invoke_from=Mock(),
session=Mock(),
session=orm_session,
)
@@ -24,6 +24,7 @@ from unittest.mock import Mock, patch
import pytest
from flask import Flask
from sqlalchemy.orm import Session
from werkzeug.datastructures import FileStorage
from werkzeug.exceptions import Forbidden, NotFound
@@ -38,6 +39,7 @@ from controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow import (
)
from core.app.entities.app_invoke_entities import InvokeFrom
from models.account import Account
from models.dataset import Dataset
from services.errors.file import FileTooLargeError, UnsupportedFileTypeError
from services.rag_pipeline.entity.pipeline_service_api_entities import (
DatasourceNodeRunApiEntity,
@@ -46,6 +48,20 @@ from services.rag_pipeline.entity.pipeline_service_api_entities import (
from services.rag_pipeline.rag_pipeline import RagPipelineService
def _persist_dataset(session: Session, *, tenant_id: str, dataset_id: str) -> Dataset:
dataset = Dataset(
id=dataset_id,
tenant_id=tenant_id,
name="Pipeline dataset",
created_by="account-1",
data_source_type=None,
indexing_technique=None,
)
session.add(dataset)
session.commit()
return dataset
class TestDatasourceNodeRunPayload:
"""Test suite for DatasourceNodeRunPayload Pydantic model."""
@@ -550,13 +566,15 @@ class TestPipelineRunApiPost:
)
@patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.RagPipelineService")
@patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.service_api_ns")
def test_post_success_streaming(self, mock_ns, mock_svc_cls, mock_current_user, mock_gen_svc, mock_helper, app):
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
def test_post_success_streaming(
self, mock_ns, mock_svc_cls, mock_current_user, mock_gen_svc, mock_helper, app, sqlite_session: Session
):
"""Test successful pipeline run with streaming response."""
tenant_id = str(uuid.uuid4())
dataset_id = str(uuid.uuid4())
session = Mock()
session.scalar.return_value = Mock()
_persist_dataset(sqlite_session, tenant_id=tenant_id, dataset_id=dataset_id)
mock_ns.payload = {
"inputs": {"key": "val"},
@@ -577,33 +595,33 @@ class TestPipelineRunApiPost:
with app.test_request_context("/datasets/test/pipeline/run", method="POST"):
api = PipelineRunApi()
response = api.post.__wrapped__(api, session, tenant_id=tenant_id, dataset_id=dataset_id)
response = api.post.__wrapped__(api, sqlite_session, tenant_id=tenant_id, dataset_id=dataset_id)
assert response == {"result": "ok"}
mock_svc_cls.assert_called_once_with(session)
mock_svc_cls.assert_called_once_with(sqlite_session)
mock_gen_svc.generate.assert_called_once()
def test_post_not_found(self, app: Flask):
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
def test_post_not_found(self, app: Flask, sqlite_session: Session):
"""Test NotFound when dataset check fails."""
session = Mock()
session.scalar.return_value = None
with app.test_request_context("/datasets/test/pipeline/run", method="POST"):
api = PipelineRunApi()
with pytest.raises(NotFound):
api.post.__wrapped__(
api,
session,
sqlite_session,
tenant_id=str(uuid.uuid4()),
dataset_id=str(uuid.uuid4()),
)
@patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.current_user", new="not_account")
@patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.service_api_ns")
def test_post_forbidden_non_account_user(self, mock_ns, app: Flask):
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
def test_post_forbidden_non_account_user(self, mock_ns, app: Flask, sqlite_session: Session):
"""Test Forbidden when current_user is not an Account."""
session = Mock()
session.scalar.return_value = Mock()
tenant_id = str(uuid.uuid4())
dataset_id = str(uuid.uuid4())
_persist_dataset(sqlite_session, tenant_id=tenant_id, dataset_id=dataset_id)
mock_ns.payload = {
"inputs": {},
"datasource_type": "online_document",
@@ -618,9 +636,9 @@ class TestPipelineRunApiPost:
with pytest.raises(Forbidden):
api.post.__wrapped__(
api,
session,
tenant_id=str(uuid.uuid4()),
dataset_id=str(uuid.uuid4()),
sqlite_session,
tenant_id=tenant_id,
dataset_id=dataset_id,
)
@@ -3,10 +3,13 @@ Unit tests for Service API wraps (authentication decorators)
"""
import uuid
from types import SimpleNamespace
from unittest.mock import MagicMock, Mock, patch
import pytest
from flask import Flask
from sqlalchemy import select
from sqlalchemy.orm import Session
from werkzeug.exceptions import Forbidden, NotFound, Unauthorized
from controllers.service_api.wraps import (
@@ -21,12 +24,11 @@ from controllers.service_api.wraps import (
validate_dataset_token,
)
from enums.cloud_plan import CloudPlan
from models.account import TenantStatus
from models.model import ApiToken
from tests.unit_tests.conftest import (
setup_mock_dataset_owner_execute_result,
setup_mock_tenant_owner_execute_result,
)
from models import Account, Tenant, TenantAccountJoin
from models.account import TenantAccountRole
from models.dataset import Dataset, RateLimitLog
from models.enums import ApiTokenType
from models.model import ApiToken, App, AppMode, IconType
def _configure_current_app_mock(mock_current_app):
@@ -34,6 +36,51 @@ def _configure_current_app_mock(mock_current_app):
mock_current_app._get_current_object = Mock(return_value=Mock())
def _session_proxy(session: Session) -> MagicMock:
"""Emulate Flask-SQLAlchemy's callable scoped-session proxy around a test session."""
proxy = MagicMock(wraps=session)
proxy.return_value = session
return proxy
def _api_token(*, tenant_id: str, app_id: str | None = None, token_type: ApiTokenType) -> ApiToken:
return ApiToken(
id=str(uuid.uuid4()),
tenant_id=tenant_id,
app_id=app_id,
type=token_type,
token="test_token",
)
def _persist_workspace(session: Session) -> tuple[Tenant, Account, TenantAccountJoin]:
tenant = Tenant(name="Workspace")
account = Account(name="Owner", email=f"owner-{uuid.uuid4()}@example.com")
membership = TenantAccountJoin(
tenant_id=tenant.id,
account_id=account.id,
current=True,
role=TenantAccountRole.OWNER,
)
session.add_all([tenant, account, membership])
session.commit()
return tenant, account, membership
def _app_model(*, tenant_id: str, enable_api: bool = True) -> App:
return App(
id=str(uuid.uuid4()),
tenant_id=tenant_id,
name="Service API App",
mode=AppMode.CHAT,
icon_type=IconType.EMOJI,
icon="chat",
icon_background="#FFFFFF",
enable_site=False,
enable_api=enable_api,
)
class TestValidateAndGetApiToken:
"""Test suite for validate_and_get_api_token function"""
@@ -70,21 +117,24 @@ class TestValidateAndGetApiToken:
def test_valid_token_returns_api_token(self, mock_fetch_token, mock_cache_cls, mock_record_usage, app: Flask):
"""Test that valid token returns the ApiToken object."""
# Arrange
mock_api_token = Mock(spec=ApiToken)
mock_api_token.token = "valid_token_123"
mock_api_token.type = "app"
api_token = _api_token(
tenant_id=str(uuid.uuid4()),
app_id=str(uuid.uuid4()),
token_type=ApiTokenType.APP,
)
api_token.token = "valid_token_123"
mock_cache_instance = Mock()
mock_cache_instance.get.return_value = None # Cache miss
mock_cache_cls.get = mock_cache_instance.get
mock_fetch_token.return_value = mock_api_token
mock_fetch_token.return_value = api_token
# Act
with app.test_request_context("/", method="GET", headers={"Authorization": "Bearer valid_token_123"}):
result = validate_and_get_api_token("app")
# Assert
assert result == mock_api_token
assert result == api_token
@patch("controllers.service_api.wraps.record_token_usage")
@patch("controllers.service_api.wraps.ApiTokenCache")
@@ -117,116 +167,124 @@ class TestValidateAppToken:
return app
@patch("controllers.service_api.wraps.user_logged_in")
@patch("controllers.service_api.wraps.db")
@patch("controllers.service_api.wraps.validate_and_get_api_token")
@patch("controllers.service_api.wraps.current_app")
@pytest.mark.parametrize(
"sqlite_session",
[(App, ApiToken, Tenant, Account, TenantAccountJoin)],
indirect=True,
)
def test_valid_app_token_allows_access(
self, mock_current_app, mock_validate_token, mock_db, mock_user_logged_in, app
self,
mock_current_app,
mock_validate_token,
mock_user_logged_in,
app: Flask,
sqlite_session: Session,
):
"""Test that valid app token allows access to decorated view."""
# Arrange
_configure_current_app_mock(mock_current_app)
mock_api_token = Mock()
mock_api_token.app_id = str(uuid.uuid4())
mock_api_token.tenant_id = str(uuid.uuid4())
mock_validate_token.return_value = mock_api_token
mock_app = Mock()
mock_app.id = mock_api_token.app_id
mock_app.status = "normal"
mock_app.enable_api = True
mock_app.tenant_id = mock_api_token.tenant_id
mock_tenant = Mock()
mock_tenant.status = TenantStatus.NORMAL
mock_tenant.id = mock_api_token.tenant_id
mock_account = Mock()
mock_account.id = str(uuid.uuid4())
# Use side_effect to return app first, then tenant via session.get()
mock_db.session.get.side_effect = [mock_app, mock_tenant]
# Mock the tenant owner execute result (execute(select(...)).one_or_none())
setup_mock_tenant_owner_execute_result(mock_db, mock_tenant, mock_account)
tenant, account, _ = _persist_workspace(sqlite_session)
app_model = _app_model(tenant_id=tenant.id)
api_token = _api_token(tenant_id=tenant.id, app_id=app_model.id, token_type=ApiTokenType.APP)
sqlite_session.add_all([app_model, api_token])
sqlite_session.commit()
mock_validate_token.return_value = api_token
@validate_app_token
def protected_view(app_model):
return {"success": True, "app_id": app_model.id}
# Act
with app.test_request_context("/", method="GET", headers={"Authorization": "Bearer test_token"}):
with (
app.test_request_context("/", method="GET", headers={"Authorization": "Bearer test_token"}),
patch("controllers.service_api.wraps.db.session", _session_proxy(sqlite_session)),
):
result = protected_view()
# Assert
assert result["success"] is True
assert result["app_id"] == mock_app.id
assert result["app_id"] == app_model.id
assert account.current_tenant_id == tenant.id
@patch("controllers.service_api.wraps.db")
@patch("controllers.service_api.wraps.validate_and_get_api_token")
def test_app_not_found_raises_forbidden(self, mock_validate_token, mock_db, app: Flask):
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
def test_app_not_found_raises_forbidden(self, mock_validate_token, app: Flask, sqlite_session: Session):
"""Test that Forbidden is raised when app no longer exists."""
# Arrange
mock_api_token = Mock()
mock_api_token.app_id = str(uuid.uuid4())
mock_validate_token.return_value = mock_api_token
mock_db.session.get.return_value = None
api_token = _api_token(
tenant_id=str(uuid.uuid4()),
app_id=str(uuid.uuid4()),
token_type=ApiTokenType.APP,
)
mock_validate_token.return_value = api_token
@validate_app_token
def protected_view(**kwargs):
return {"success": True}
# Act & Assert
with app.test_request_context("/", method="GET"):
with (
app.test_request_context("/", method="GET"),
patch("controllers.service_api.wraps.db.session", sqlite_session),
):
with pytest.raises(Forbidden) as exc_info:
protected_view()
assert "no longer exists" in str(exc_info.value)
@patch("controllers.service_api.wraps.db")
@patch("controllers.service_api.wraps.validate_and_get_api_token")
def test_app_status_abnormal_raises_forbidden(self, mock_validate_token, mock_db, app: Flask):
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
def test_app_status_abnormal_raises_forbidden(self, mock_validate_token, app: Flask, sqlite_session: Session):
"""Test that Forbidden is raised when app status is abnormal."""
# Arrange
mock_api_token = Mock()
mock_api_token.app_id = str(uuid.uuid4())
mock_validate_token.return_value = mock_api_token
mock_app = Mock()
mock_app.status = "abnormal"
mock_db.session.get.return_value = mock_app
app_model = _app_model(tenant_id=str(uuid.uuid4()))
sqlite_session.add(app_model)
sqlite_session.commit()
app_model.status = "abnormal"
mock_validate_token.return_value = _api_token(
tenant_id=app_model.tenant_id,
app_id=app_model.id,
token_type=ApiTokenType.APP,
)
@validate_app_token
def protected_view(**kwargs):
return {"success": True}
# Act & Assert
with app.test_request_context("/", method="GET"):
with (
app.test_request_context("/", method="GET"),
patch("controllers.service_api.wraps.db.session", sqlite_session),
):
with pytest.raises(Forbidden) as exc_info:
protected_view()
assert "status is abnormal" in str(exc_info.value)
@patch("controllers.service_api.wraps.db")
@patch("controllers.service_api.wraps.validate_and_get_api_token")
def test_app_api_disabled_raises_forbidden(self, mock_validate_token, mock_db, app: Flask):
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
def test_app_api_disabled_raises_forbidden(self, mock_validate_token, app: Flask, sqlite_session: Session):
"""Test that Forbidden is raised when app API is disabled."""
# Arrange
mock_api_token = Mock()
mock_api_token.app_id = str(uuid.uuid4())
mock_validate_token.return_value = mock_api_token
mock_app = Mock()
mock_app.status = "normal"
mock_app.enable_api = False
mock_db.session.get.return_value = mock_app
app_model = _app_model(tenant_id=str(uuid.uuid4()), enable_api=False)
sqlite_session.add(app_model)
sqlite_session.commit()
mock_validate_token.return_value = _api_token(
tenant_id=app_model.tenant_id,
app_id=app_model.id,
token_type=ApiTokenType.APP,
)
@validate_app_token
def protected_view(**kwargs):
return {"success": True}
# Act & Assert
with app.test_request_context("/", method="GET"):
with (
app.test_request_context("/", method="GET"),
patch("controllers.service_api.wraps.db.session", sqlite_session),
):
with pytest.raises(Forbidden) as exc_info:
protected_view()
assert "API service has been disabled" in str(exc_info.value)
@@ -468,26 +526,35 @@ class TestCloudEditionBillingRateLimitCheck:
@patch("controllers.service_api.wraps.validate_and_get_api_token")
@patch("controllers.service_api.wraps.FeatureService.get_knowledge_rate_limit")
@patch("controllers.service_api.wraps.db")
@patch("controllers.service_api.wraps.sessionmaker")
@pytest.mark.parametrize("sqlite_session", [(RateLimitLog,)], indirect=True)
def test_rejects_over_rate_limit(
self, mock_sessionmaker, mock_db, mock_get_rate_limit, mock_validate_token, app: Flask
self,
mock_get_rate_limit,
mock_validate_token,
app: Flask,
sqlite_session: Session,
):
"""Test that Forbidden is raised when over rate limit."""
# Arrange
mock_validate_token.return_value = Mock(tenant_id="tenant123")
tenant_id = str(uuid.uuid4())
mock_validate_token.return_value = _api_token(
tenant_id=tenant_id,
token_type=ApiTokenType.DATASET,
)
mock_rate_limit = Mock()
mock_rate_limit.enabled = True
mock_rate_limit.limit = 10
mock_rate_limit.subscription_plan = "pro"
mock_get_rate_limit.return_value = mock_rate_limit
rate_limit_log_session = MagicMock()
session_factory = MagicMock()
session_factory.begin.return_value.__enter__.return_value = rate_limit_log_session
mock_sessionmaker.return_value = session_factory
with patch("controllers.service_api.wraps.redis_client") as mock_redis:
with (
patch("controllers.service_api.wraps.redis_client") as mock_redis,
patch(
"controllers.service_api.wraps.db",
SimpleNamespace(engine=sqlite_session.get_bind()),
),
):
mock_redis.zcard.return_value = 15 # Over limit
@cloud_edition_billing_rate_limit_check("knowledge", "dataset")
@@ -499,9 +566,12 @@ class TestCloudEditionBillingRateLimitCheck:
with pytest.raises(Forbidden) as exc_info:
knowledge_request()
assert "rate limit" in str(exc_info.value)
mock_sessionmaker.assert_called_once_with(bind=mock_db.engine, expire_on_commit=False)
rate_limit_log_session.add.assert_called_once()
mock_db.session.commit.assert_not_called()
persisted_logs = sqlite_session.scalars(select(RateLimitLog)).all()
assert len(persisted_logs) == 1
assert persisted_logs[0].tenant_id == tenant_id
assert persisted_logs[0].subscription_plan == "pro"
assert persisted_logs[0].operation == "knowledge"
class TestValidateDatasetToken:
@@ -515,65 +585,62 @@ class TestValidateDatasetToken:
return app
@patch("controllers.service_api.wraps.user_logged_in")
@patch("controllers.service_api.wraps.db")
@patch("controllers.service_api.wraps.validate_and_get_api_token")
@patch("controllers.service_api.wraps.current_app")
def test_valid_dataset_token(self, mock_current_app, mock_validate_token, mock_db, mock_user_logged_in, app: Flask):
@pytest.mark.parametrize(
"sqlite_session",
[(Tenant, Account, TenantAccountJoin)],
indirect=True,
)
def test_valid_dataset_token(
self,
mock_current_app,
mock_validate_token,
mock_user_logged_in,
app: Flask,
sqlite_session: Session,
):
"""Test that valid dataset token allows access."""
# Arrange
_configure_current_app_mock(mock_current_app)
tenant_id = str(uuid.uuid4())
mock_api_token = Mock()
mock_api_token.tenant_id = tenant_id
mock_validate_token.return_value = mock_api_token
mock_tenant = Mock()
mock_tenant.id = tenant_id
mock_tenant.status = TenantStatus.NORMAL
mock_ta = Mock()
mock_ta.account_id = str(uuid.uuid4())
mock_account = Mock()
mock_account.id = mock_ta.account_id
mock_account.current_tenant = mock_tenant
# Mock the tenant account join query (execute(select(...)).one_or_none())
setup_mock_dataset_owner_execute_result(mock_db, mock_tenant, mock_ta)
# Mock the account lookup via session.get()
mock_db.session.get.return_value = mock_account
tenant, account, _ = _persist_workspace(sqlite_session)
api_token = _api_token(tenant_id=tenant.id, token_type=ApiTokenType.DATASET)
mock_validate_token.return_value = api_token
@validate_dataset_token
def protected_view(tenant_id):
return {"success": True, "tenant_id": tenant_id}
# Act
with app.test_request_context("/", method="GET", headers={"Authorization": "Bearer test_token"}):
with (
app.test_request_context("/", method="GET", headers={"Authorization": "Bearer test_token"}),
patch("controllers.service_api.wraps.db.session", _session_proxy(sqlite_session)),
):
result = protected_view()
# Assert
assert result["success"] is True
assert result["tenant_id"] == tenant_id
assert result["tenant_id"] == tenant.id
assert account.current_tenant_id == tenant.id
@patch("controllers.service_api.wraps.db")
@patch("controllers.service_api.wraps.validate_and_get_api_token")
def test_dataset_not_found_raises_not_found(self, mock_validate_token, mock_db, app: Flask):
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
def test_dataset_not_found_raises_not_found(self, mock_validate_token, app: Flask, sqlite_session: Session):
"""Test that NotFound is raised when dataset doesn't exist."""
# Arrange
mock_api_token = Mock()
mock_api_token.tenant_id = str(uuid.uuid4())
mock_validate_token.return_value = mock_api_token
mock_db.session.scalar.return_value = None
api_token = _api_token(tenant_id=str(uuid.uuid4()), token_type=ApiTokenType.DATASET)
mock_validate_token.return_value = api_token
@validate_dataset_token
def protected_view(dataset_id=None, **kwargs):
return {"success": True}
# Act & Assert
with app.test_request_context("/", method="GET"):
with (
app.test_request_context("/", method="GET"),
patch("controllers.service_api.wraps.db.session", sqlite_session),
):
with pytest.raises(NotFound) as exc_info:
protected_view(dataset_id=str(uuid.uuid4()))
assert "Dataset not found" in str(exc_info.value)
@@ -46,7 +46,7 @@ class TestAdvancedChatAppGeneratorValidation:
with pytest.raises(ValueError, match="query must be a string"):
generator.generate(
app_model=SimpleNamespace(),
workflow=SimpleNamespace(),
workflow=SimpleNamespace(graph_dict={"nodes": []}),
user=SimpleNamespace(),
args={"inputs": {}, "query": 123},
invoke_from=InvokeFrom.WEB_APP,
@@ -186,7 +186,7 @@ class TestAdvancedChatAppGeneratorInternals:
result = generator.generate(
app_model=SimpleNamespace(id="app", tenant_id="tenant"),
workflow=SimpleNamespace(features_dict={}),
workflow=SimpleNamespace(features_dict={}, graph_dict={"nodes": []}),
user=user,
args={
"query": "hello",
@@ -1209,7 +1209,7 @@ class TestAdvancedChatAppGeneratorInternals:
monkeypatch.setattr(generator, "_generate", _fake_generate)
app_model = SimpleNamespace(id="app", tenant_id="tenant")
workflow = SimpleNamespace(features_dict={})
workflow = SimpleNamespace(features_dict={}, graph_dict={"nodes": []})
from models import Account
user = Account(name="Tester", email="tester@example.com")
@@ -1289,7 +1289,7 @@ class TestAdvancedChatAppGeneratorInternals:
monkeypatch.setattr(generator, "_generate", _fake_generate)
app_model = SimpleNamespace(id="app", tenant_id="tenant")
workflow = SimpleNamespace(features_dict={})
workflow = SimpleNamespace(features_dict={}, graph_dict={"nodes": []})
from models.model import EndUser
user = EndUser(tenant_id="tenant", type="session", name="tester", session_id="session")
@@ -22,7 +22,7 @@ def _build_converter() -> WorkflowResponseConverter:
app_config=SimpleNamespace(app_id="app-1", tenant_id="tenant-1"),
invoke_from=InvokeFrom.EXPLORE,
files=[],
inputs={},
inputs={"userinput.files": []},
workflow_execution_id="run-1",
call_depth=0,
)
@@ -54,3 +54,17 @@ def test_workflow_start_stream_response_carries_initial_reason():
reason=WorkflowStartReason.INITIAL,
)
assert resp.data.reason is WorkflowStartReason.INITIAL
def test_workflow_start_stream_response_exposes_only_canonical_file_input():
converter = _build_converter()
resp = converter.workflow_start_to_stream_response(
task_id="task-1",
workflow_run_id="run-1",
workflow_id="wf-1",
reason=WorkflowStartReason.INITIAL,
)
assert resp.data.inputs["userinput.files"] == []
assert "sys.files" not in resp.data.inputs
@@ -111,7 +111,7 @@ def test_generate_includes_parent_trace_context_in_extras(monkeypatch):
result = generator.generate(
app_model=SimpleNamespace(tenant_id="tenant-1", id="app-1"),
workflow=SimpleNamespace(features_dict={}),
workflow=SimpleNamespace(features_dict={}, graph_dict={"nodes": []}),
user=SimpleNamespace(id="user-1", session_id="session-1"),
args={
"inputs": {"query": "hello"},
@@ -377,7 +377,7 @@ class TestWorkflowAppGeneratorGenerate:
result = generator.generate(
app_model=SimpleNamespace(id="app", tenant_id="tenant"),
workflow=SimpleNamespace(features_dict={}),
workflow=SimpleNamespace(features_dict={}, graph_dict={"nodes": []}),
user=SimpleNamespace(id="user", session_id="session"),
args={"inputs": {}, SKIP_PREPARE_USER_INPUTS_KEY: True},
invoke_from=InvokeFrom.WEB_APP,
@@ -26,6 +26,7 @@ class TestPlannerSystemPrompt:
"""Auto-mode resolution rides on the planner echoing its mode choice."""
assert '"mode": "workflow | advanced-chat"' in PLANNER_SYSTEM_PROMPT
assert "When the ``# Mode`` section says auto, YOU decide" in PLANNER_SYSTEM_PROMPT
assert "userinput.files" in PLANNER_SYSTEM_PROMPT
class TestFormatIdealOutputSection:
@@ -75,6 +76,9 @@ class TestNodeBuilderPrompt:
assert '"viewport":' not in prompt
assert '"positionAbsolute":' not in prompt
def test_start_uses_user_input_files(self):
assert "userinput.files" in get_node_builder_system_prompt("start")
def test_supports_main_human_input_and_assigner_contracts(self):
human_input = get_node_builder_system_prompt("human-input")
assigner = get_node_builder_system_prompt("assigner")
@@ -130,17 +134,19 @@ class TestNodeBuilderUserSections:
class TestModeSection:
def test_advanced_chat_documents_system_variables(self):
def test_advanced_chat_documents_built_in_variables(self):
out = format_mode_section("advanced-chat")
assert "sys.query" in out
assert '["sys", "query"]' in out
assert "userinput.files" in out
assert '["userinput", "files"]' in out
assert "do NOT invent start-node variables" in out
def test_workflow_mode_forbids_system_variables(self):
def test_workflow_mode_documents_file_input(self):
out = format_mode_section("workflow")
assert "NO automatic system variables" in out
assert "userinput.files" in out
assert "start node's declared variables" in out
class TestExistingGraphSection:
@@ -228,12 +228,12 @@ def _previous_node_prompt_payload(result, selector: str) -> object:
def _uploaded_workflow_files_prompt_payload(result) -> object:
prefix = " - sys.files: "
prefix = " - userinput.files: "
user_prompt = _workflow_user_prompt(result)
for line in user_prompt.splitlines():
if line.startswith(prefix):
return json.loads(line.removeprefix(prefix))
raise AssertionError("missing prompt payload for sys.files")
raise AssertionError("missing prompt payload for userinput.files")
def test_builds_create_run_request_from_agent_soul_and_node_job():
@@ -0,0 +1,161 @@
from core.workflow.legacy_system_files import (
LegacySysFilesCompatVariable,
attach_legacy_sys_files_warning,
migrate_legacy_sys_files_graph_with_result,
normalize_legacy_sys_files_args,
resolve_legacy_sys_files_compat_variable,
)
_LEGACY_NODE_ID = "sys"
_LEGACY_ALIAS_NODE_ID = "userinput"
_LEGACY_VARIABLE_NAME = "files"
_LEGACY_SELECTOR = [_LEGACY_NODE_ID, _LEGACY_VARIABLE_NAME]
_LEGACY_TEMPLATE = "{{#" + ".".join((_LEGACY_NODE_ID, _LEGACY_VARIABLE_NAME)) + "#}}"
_LEGACY_ALIAS_SELECTOR = [_LEGACY_ALIAS_NODE_ID, _LEGACY_VARIABLE_NAME]
_LEGACY_ALIAS_TEMPLATE = "{{#" + ".".join((_LEGACY_ALIAS_NODE_ID, _LEGACY_VARIABLE_NAME)) + "#}}"
_LEGACY_ALIAS_INPUT_KEY = ".".join((_LEGACY_ALIAS_NODE_ID, _LEGACY_VARIABLE_NAME))
def test_migrate_legacy_sys_files_graph_ignores_invalid_or_unrelated_graphs():
assert not migrate_legacy_sys_files_graph_with_result({}).changed
assert not migrate_legacy_sys_files_graph_with_result({"nodes": [], "edges": [_LEGACY_SELECTOR]}).changed
assert not migrate_legacy_sys_files_graph_with_result({"nodes": [{"data": {"value": ["sys", "query"]}}]}).changed
def test_migrate_legacy_sys_files_graph_rewrites_sys_files_to_userinput_files_without_start_variable():
graph = {
"nodes": [
{"id": "start", "data": {"type": "start", "variables": [{"variable": "sys_files"}]}},
{
"id": "answer",
"data": {
"type": "answer",
"answer": _LEGACY_SELECTOR,
"template": _LEGACY_TEMPLATE,
},
},
],
}
result = migrate_legacy_sys_files_graph_with_result(graph)
assert result.changed
start_data = result.graph["nodes"][0]["data"]
assert start_data["variables"] == [{"variable": "sys_files"}]
assert result.graph["nodes"][1]["data"]["answer"] == _LEGACY_ALIAS_SELECTOR
assert result.graph["nodes"][1]["data"]["template"] == _LEGACY_ALIAS_TEMPLATE
assert graph["nodes"][1]["data"]["answer"] == _LEGACY_SELECTOR
assert graph["nodes"][1]["data"]["template"] == _LEGACY_TEMPLATE
def test_migrate_legacy_sys_files_graph_leaves_userinput_files_target_unchanged():
graph = {
"nodes": [
{"id": "start", "data": {"type": "start", "variables": []}},
{
"id": "answer",
"data": {
"type": "answer",
"answer": _LEGACY_ALIAS_SELECTOR,
"template": _LEGACY_ALIAS_TEMPLATE,
},
},
],
}
result = migrate_legacy_sys_files_graph_with_result(graph)
assert not result.changed
assert result.graph == graph
def test_resolve_legacy_sys_files_compat_variable_returns_userinput_files_target():
assert resolve_legacy_sys_files_compat_variable({}) is None
assert resolve_legacy_sys_files_compat_variable({"nodes": [{"data": {"value": ["sys", "query"]}}]}) is None
compat_variable = resolve_legacy_sys_files_compat_variable({"nodes": [{"data": {"value": _LEGACY_SELECTOR}}]})
assert compat_variable == LegacySysFilesCompatVariable(
node_id=_LEGACY_ALIAS_NODE_ID,
variable_name=_LEGACY_VARIABLE_NAME,
)
assert (
resolve_legacy_sys_files_compat_variable({"nodes": [{"data": {"value": _LEGACY_ALIAS_SELECTOR}}]})
== compat_variable
)
def test_normalize_legacy_sys_files_args_handles_no_compat_and_top_level_files():
args_without_legacy, compat_without_legacy = normalize_legacy_sys_files_args(
graph={"nodes": []},
args={"inputs": {}},
)
assert args_without_legacy == {"inputs": {}}
assert compat_without_legacy is None
files = [{"id": "file-1"}]
graph = {
"nodes": [
{"id": "start", "data": {"type": "start", "variables": []}},
{"id": "answer", "data": {"type": "answer", "answer": _LEGACY_TEMPLATE}},
],
}
normalized_args, compat_variable = normalize_legacy_sys_files_args(
graph=graph,
args={"inputs": {}, "files": files},
)
assert compat_variable is not None
assert normalized_args["files"] == files
assert normalized_args["inputs"][".".join((compat_variable.node_id, compat_variable.variable_name))] == files
def test_normalize_legacy_sys_files_args_maps_userinput_files_to_top_level_files_without_warning():
files = [{"id": "file-1"}]
normalized_args, compat_variable = normalize_legacy_sys_files_args(
graph={"nodes": []},
args={"inputs": {_LEGACY_ALIAS_INPUT_KEY: files}},
)
assert compat_variable is None
assert normalized_args["files"] == files
assert normalized_args["inputs"] == {_LEGACY_ALIAS_INPUT_KEY: files}
def test_normalize_legacy_sys_files_args_prefers_userinput_files_over_legacy_files():
legacy_files = [{"id": "legacy-file"}]
userinput_files = [{"id": "userinput-file"}]
normalized_args, compat_variable = normalize_legacy_sys_files_args(
graph={"nodes": [{"data": {"type": "answer", "answer": _LEGACY_ALIAS_TEMPLATE}}]},
args={
"inputs": {_LEGACY_ALIAS_INPUT_KEY: userinput_files},
"files": legacy_files,
},
)
assert compat_variable is None
assert normalized_args["files"] == userinput_files
def test_attach_legacy_sys_files_warning_wraps_stream_and_closes_source():
class CloseableStream:
closed = False
def __iter__(self):
yield "data: payload\n\n"
def close(self):
self.closed = True
stream = CloseableStream()
wrapped = attach_legacy_sys_files_warning(
stream,
LegacySysFilesCompatVariable(node_id=_LEGACY_ALIAS_NODE_ID, variable_name=_LEGACY_VARIABLE_NAME),
)
chunks = list(wrapped)
assert "warning" in chunks[0]
assert chunks[1] == "data: payload\n\n"
assert stream.closed
@@ -1,6 +1,7 @@
from types import SimpleNamespace
from core.workflow.system_variables import (
build_bootstrap_variables,
build_system_variables,
default_system_variables,
get_node_creation_preload_selectors,
@@ -56,6 +57,25 @@ def test_build_system_variables_preserves_file_values():
assert system_values["files"] == [file]
def test_build_bootstrap_variables_adds_userinput_files_alias():
file = File(
file_type=FileType.DOCUMENT,
transfer_method=FileTransferMethod.LOCAL_FILE,
related_id="file-id",
filename="test.txt",
extension=".txt",
mime_type="text/plain",
size=1,
storage_key="storage-key",
)
bootstrap_variables = build_bootstrap_variables(system_variables=build_system_variables(files=[file]))
file_variables_by_selector = {tuple(variable.selector): variable for variable in bootstrap_variables}
assert file_variables_by_selector[("sys", "files")].value == [file]
assert file_variables_by_selector[("userinput", "files")].value == [file]
def test_default_system_variables_generates_workflow_run_id():
system_variables = default_system_variables()
system_values = system_variables_to_mapping(system_variables)
@@ -18,6 +18,11 @@ from models.workflow import (
is_system_variable_editable,
)
_LEGACY_FILE_TEMPLATE = "{{#" + ".".join(("sys", "files")) + "#}}"
_LEGACY_FILE_SELECTOR = ["sys", "files"]
_USER_INPUT_FILE_TEMPLATE = "{{#" + ".".join(("userinput", "files")) + "#}}"
_USER_INPUT_FILE_SELECTOR = ["userinput", "files"]
def test_environment_variables():
# tenant_id context variable removed - using current_user.current_tenant_id directly
@@ -245,6 +250,144 @@ class TestIsSystemVariableEditable:
assert is_system_variable_editable("invalid_or_new_system_variable") == False
class TestWorkflowLegacySysFilesCompatibility:
def _make_workflow(self, graph: dict, *, features: dict | None = None) -> Workflow:
return Workflow(
tenant_id="tenant_id",
app_id="app_id",
type="workflow",
version="draft",
graph=json.dumps(graph),
features=json.dumps(features or {}),
created_by="account_id",
environment_variables=[],
conversation_variables=[],
)
def test_graph_dict_rewrites_legacy_sys_files_references_to_userinput_files(self):
workflow = self._make_workflow(
{
"nodes": [
{
"id": "start",
"data": {
"type": "start",
"title": "Start",
"variables": [],
},
},
{
"id": "llm",
"data": {
"type": "llm",
"prompt_template": [{"role": "user", "text": f"files: {_LEGACY_FILE_TEMPLATE}"}],
"context": {"variable_selector": _LEGACY_FILE_SELECTOR},
},
},
],
"edges": [],
}
)
stored_graph_before_read = workflow.graph
graph = workflow.graph_dict
start_node = next(node for node in graph["nodes"] if node["id"] == "start")
llm_node = next(node for node in graph["nodes"] if node["id"] == "llm")
assert start_node["data"]["variables"] == []
assert llm_node["data"]["prompt_template"][0]["text"] == f"files: {_USER_INPUT_FILE_TEMPLATE}"
assert llm_node["data"]["context"]["variable_selector"] == _USER_INPUT_FILE_SELECTOR
assert workflow.graph == stored_graph_before_read
def test_migrate_legacy_sys_files_graph_in_place_updates_stored_graph(self):
workflow = self._make_workflow(
{
"nodes": [
{"id": "answer", "data": {"type": "answer", "answer": _LEGACY_FILE_TEMPLATE}},
],
"edges": [],
}
)
assert workflow.migrate_legacy_sys_files_graph_in_place()
assert _LEGACY_FILE_TEMPLATE not in workflow.graph
assert _USER_INPUT_FILE_TEMPLATE in workflow.graph
def test_graph_dict_preserves_existing_start_variables_when_migrating_legacy_sys_files(self):
workflow = self._make_workflow(
{
"nodes": [
{
"id": "start",
"data": {
"type": "start",
"title": "Start",
"variables": [
{"variable": "sys_files", "label": "Existing", "type": "text-input"},
],
},
},
{
"id": "answer",
"data": {
"type": "answer",
"answer": _LEGACY_FILE_TEMPLATE,
},
},
],
"edges": [],
}
)
graph = workflow.graph_dict
start_node = next(node for node in graph["nodes"] if node["id"] == "start")
answer_node = next(node for node in graph["nodes"] if node["id"] == "answer")
assert [variable["variable"] for variable in start_node["data"]["variables"]] == ["sys_files"]
assert answer_node["data"]["answer"] == _USER_INPUT_FILE_TEMPLATE
def test_graph_dict_leaves_userinput_files_references_unchanged(self):
workflow = self._make_workflow(
{
"nodes": [
{
"id": "start",
"data": {
"type": "start",
"title": "Start",
"variables": [],
},
},
{
"id": "answer",
"data": {
"type": "answer",
"answer": _USER_INPUT_FILE_TEMPLATE,
},
},
],
"edges": [],
},
features={
"file_upload": {
"enabled": True,
"allowed_file_upload_methods": ["remote_url"],
"allowed_file_types": ["document", "custom"],
"allowed_file_extensions": [".pdf"],
"number_limits": 8,
}
},
)
graph = workflow.graph_dict
start_node = next(node for node in graph["nodes"] if node["id"] == "start")
assert start_node["data"]["variables"] == []
assert json.loads(workflow.graph) == graph
class TestWorkflowDraftVariableGetValue:
def test_get_value_by_case(self):
@dataclasses.dataclass
@@ -1,89 +1,105 @@
import logging
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
from uuid import uuid4
import pytest
from sqlalchemy import select
from sqlalchemy.orm import Session
from models.account import (
TenantPluginAutoUpgradeCategory,
TenantPluginAutoUpgradeMode,
TenantPluginAutoUpgradeStrategy,
TenantPluginAutoUpgradeStrategySetting,
)
from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService
MODULE = "services.plugin.plugin_auto_upgrade_service"
PLUGIN_CATEGORY = TenantPluginAutoUpgradeCategory.TOOL
STRATEGY_MODELS = (TenantPluginAutoUpgradeStrategy,)
def _patched_session():
"""Return a mock SQLAlchemy session for service calls."""
session = MagicMock()
return session
def _strategy(
tenant_id: str,
*,
category: TenantPluginAutoUpgradeCategory = PLUGIN_CATEGORY,
setting: TenantPluginAutoUpgradeStrategySetting = TenantPluginAutoUpgradeStrategySetting.FIX_ONLY,
mode: TenantPluginAutoUpgradeMode = TenantPluginAutoUpgradeMode.EXCLUDE,
exclude: list[str] | None = None,
include: list[str] | None = None,
upgrade_time: int = 0,
) -> TenantPluginAutoUpgradeStrategy:
return TenantPluginAutoUpgradeStrategy(
tenant_id=tenant_id,
category=category,
strategy_setting=setting,
upgrade_time_of_day=upgrade_time,
upgrade_mode=mode,
exclude_plugins=exclude or [],
include_plugins=include or [],
)
class TestGetStrategy:
def test_returns_strategy_when_found(self):
session = _patched_session()
strategy = MagicMock()
session.scalar.return_value = strategy
@pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True)
def test_returns_strategy_when_found(self, sqlite_session: Session) -> None:
tenant_id = str(uuid4())
strategy = _strategy(tenant_id)
sqlite_session.add(strategy)
sqlite_session.commit()
from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService
result = PluginAutoUpgradeService.get_strategy("t1", PLUGIN_CATEGORY, session=session)
result = PluginAutoUpgradeService.get_strategy(tenant_id, PLUGIN_CATEGORY, session=sqlite_session)
assert result is strategy
def test_returns_none_when_not_found(self):
session = _patched_session()
session.scalar.return_value = None
from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService
result = PluginAutoUpgradeService.get_strategy("t1", PLUGIN_CATEGORY, session=session)
assert result is None
@pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True)
def test_returns_none_when_not_found(self, sqlite_session: Session) -> None:
assert PluginAutoUpgradeService.get_strategy(str(uuid4()), PLUGIN_CATEGORY, session=sqlite_session) is None
class TestChangeStrategy:
def test_creates_new_strategy(self):
session = _patched_session()
session.scalar.return_value = None
with patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls:
strat_cls.return_value = MagicMock()
from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService
result = PluginAutoUpgradeService.change_strategy(
"t1",
TenantPluginAutoUpgradeStrategySetting.FIX_ONLY,
3,
TenantPluginAutoUpgradeMode.ALL,
[],
[],
category=PLUGIN_CATEGORY,
session=session,
)
assert result is True
session.add.assert_called_once()
def test_updates_existing_strategy(self):
session = _patched_session()
existing = MagicMock()
session.scalar.return_value = existing
from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService
@pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True)
def test_creates_new_strategy(self, sqlite_session: Session) -> None:
tenant_id = str(uuid4())
result = PluginAutoUpgradeService.change_strategy(
"t1",
tenant_id,
TenantPluginAutoUpgradeStrategySetting.FIX_ONLY,
3,
TenantPluginAutoUpgradeMode.ALL,
[],
[],
category=PLUGIN_CATEGORY,
session=sqlite_session,
)
strategy = sqlite_session.scalar(select(TenantPluginAutoUpgradeStrategy))
assert result is True
assert strategy is not None
assert strategy.tenant_id == tenant_id
assert strategy.upgrade_time_of_day == 3
assert strategy.upgrade_mode == TenantPluginAutoUpgradeMode.ALL
@pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True)
def test_updates_existing_strategy(self, sqlite_session: Session) -> None:
tenant_id = str(uuid4())
existing = _strategy(tenant_id)
sqlite_session.add(existing)
sqlite_session.commit()
result = PluginAutoUpgradeService.change_strategy(
tenant_id,
TenantPluginAutoUpgradeStrategySetting.LATEST,
5,
TenantPluginAutoUpgradeMode.PARTIAL,
["p1"],
["p2"],
category=PLUGIN_CATEGORY,
session=session,
session=sqlite_session,
)
sqlite_session.refresh(existing)
assert result is True
assert existing.strategy_setting == TenantPluginAutoUpgradeStrategySetting.LATEST
assert existing.upgrade_time_of_day == 5
@@ -93,157 +109,115 @@ class TestChangeStrategy:
class TestExcludePlugin:
def test_creates_default_strategy_when_none_exists(self):
session = _patched_session()
session.scalar.return_value = None
@pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True)
def test_creates_default_strategy_when_none_exists(self, sqlite_session: Session) -> None:
tenant_id = str(uuid4())
with (
patch(f"{MODULE}.select"),
patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy"),
):
from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService
result = PluginAutoUpgradeService.exclude_plugin(
"t1",
"plugin-1",
PLUGIN_CATEGORY,
session=session,
)
result = PluginAutoUpgradeService.exclude_plugin(tenant_id, "plugin-1", PLUGIN_CATEGORY, session=sqlite_session)
strategy = sqlite_session.scalar(select(TenantPluginAutoUpgradeStrategy))
assert result is True
session.add.assert_called_once()
assert strategy is not None
assert strategy.exclude_plugins == ["plugin-1"]
def test_appends_to_exclude_list_in_exclude_mode(self):
session = _patched_session()
existing = MagicMock()
existing.upgrade_mode = TenantPluginAutoUpgradeMode.EXCLUDE
existing.exclude_plugins = ["p-existing"]
session.scalar.return_value = existing
@pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True)
def test_appends_to_exclude_list_in_exclude_mode(self, sqlite_session: Session) -> None:
tenant_id = str(uuid4())
existing = _strategy(tenant_id, exclude=["p-existing"])
sqlite_session.add(existing)
sqlite_session.commit()
with patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls:
strat_cls.UpgradeMode.EXCLUDE = "exclude"
strat_cls.UpgradeMode.PARTIAL = "partial"
strat_cls.UpgradeMode.ALL = "all"
from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService
PluginAutoUpgradeService.exclude_plugin(tenant_id, "p-new", PLUGIN_CATEGORY, session=sqlite_session)
result = PluginAutoUpgradeService.exclude_plugin("t1", "p-new", PLUGIN_CATEGORY, session=session)
assert result is True
sqlite_session.refresh(existing)
assert existing.exclude_plugins == ["p-existing", "p-new"]
def test_removes_from_include_list_in_partial_mode(self):
session = _patched_session()
existing = MagicMock()
existing.upgrade_mode = TenantPluginAutoUpgradeMode.PARTIAL
existing.include_plugins = ["p1", "p2"]
session.scalar.return_value = existing
@pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True)
def test_removes_from_include_list_in_partial_mode(self, sqlite_session: Session) -> None:
tenant_id = str(uuid4())
existing = _strategy(tenant_id, mode=TenantPluginAutoUpgradeMode.PARTIAL, include=["p1", "p2"])
sqlite_session.add(existing)
sqlite_session.commit()
with patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls:
strat_cls.UpgradeMode.EXCLUDE = "exclude"
strat_cls.UpgradeMode.PARTIAL = "partial"
strat_cls.UpgradeMode.ALL = "all"
from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService
PluginAutoUpgradeService.exclude_plugin(tenant_id, "p1", PLUGIN_CATEGORY, session=sqlite_session)
result = PluginAutoUpgradeService.exclude_plugin("t1", "p1", PLUGIN_CATEGORY, session=session)
assert result is True
sqlite_session.refresh(existing)
assert existing.include_plugins == ["p2"]
def test_switches_to_exclude_mode_from_all(self):
session = _patched_session()
existing = MagicMock()
existing.upgrade_mode = TenantPluginAutoUpgradeMode.ALL
session.scalar.return_value = existing
@pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True)
def test_switches_to_exclude_mode_from_all(self, sqlite_session: Session) -> None:
tenant_id = str(uuid4())
existing = _strategy(tenant_id, mode=TenantPluginAutoUpgradeMode.ALL)
sqlite_session.add(existing)
sqlite_session.commit()
with patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls:
strat_cls.UpgradeMode.EXCLUDE = "exclude"
strat_cls.UpgradeMode.PARTIAL = "partial"
strat_cls.UpgradeMode.ALL = "all"
from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService
PluginAutoUpgradeService.exclude_plugin(tenant_id, "p1", PLUGIN_CATEGORY, session=sqlite_session)
result = PluginAutoUpgradeService.exclude_plugin("t1", "p1", PLUGIN_CATEGORY, session=session)
assert result is True
sqlite_session.refresh(existing)
assert existing.upgrade_mode == TenantPluginAutoUpgradeMode.EXCLUDE
assert existing.exclude_plugins == ["p1"]
def test_no_duplicate_in_exclude_list(self):
session = _patched_session()
existing = MagicMock()
existing.upgrade_mode = TenantPluginAutoUpgradeMode.EXCLUDE
existing.exclude_plugins = ["p1"]
session.scalar.return_value = existing
@pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True)
def test_no_duplicate_in_exclude_list(self, sqlite_session: Session) -> None:
tenant_id = str(uuid4())
existing = _strategy(tenant_id, exclude=["p1"])
sqlite_session.add(existing)
sqlite_session.commit()
with patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls:
strat_cls.UpgradeMode.EXCLUDE = "exclude"
strat_cls.UpgradeMode.PARTIAL = "partial"
strat_cls.UpgradeMode.ALL = "all"
from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService
PluginAutoUpgradeService.exclude_plugin("t1", "p1", PLUGIN_CATEGORY, session=session)
PluginAutoUpgradeService.exclude_plugin(tenant_id, "p1", PLUGIN_CATEGORY, session=sqlite_session)
sqlite_session.refresh(existing)
assert existing.exclude_plugins == ["p1"]
class TestBackfillStrategyCategories:
def test_creates_default_missing_categories_without_fetching_daemon(self):
session = _patched_session()
tool_strategy = SimpleNamespace(
category=TenantPluginAutoUpgradeCategory.TOOL,
strategy_setting=TenantPluginAutoUpgradeStrategySetting.FIX_ONLY,
upgrade_time_of_day=0,
upgrade_mode=TenantPluginAutoUpgradeMode.EXCLUDE,
exclude_plugins=[],
include_plugins=[],
)
session.scalars.return_value.all.return_value = [tool_strategy]
@pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True)
def test_creates_default_missing_categories_without_fetching_daemon(self, sqlite_session: Session) -> None:
tenant_id = str(uuid4())
tool_strategy = _strategy(tenant_id)
sqlite_session.add(tool_strategy)
sqlite_session.commit()
installer = MagicMock()
with patch(f"{MODULE}.PluginInstaller", return_value=installer):
from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService
result = PluginAutoUpgradeService.backfill_strategy_categories("t1", session=session)
expected_time = PluginAutoUpgradeService.default_upgrade_time_of_day("t1")
result = PluginAutoUpgradeService.backfill_strategy_categories(tenant_id, session=sqlite_session)
expected_time = PluginAutoUpgradeService.default_upgrade_time_of_day(tenant_id)
strategies = list(sqlite_session.scalars(select(TenantPluginAutoUpgradeStrategy)).all())
assert result.created_count == len(TenantPluginAutoUpgradeCategory) - 1
assert result.normalized is False
installer.list_plugins.assert_not_called()
assert len(strategies) == len(TenantPluginAutoUpgradeCategory)
assert tool_strategy.upgrade_time_of_day == expected_time
created_strategies = [call.args[0] for call in session.add.call_args_list]
model_strategy = next(
strategy for strategy in created_strategies if strategy.category == TenantPluginAutoUpgradeCategory.MODEL
strategy for strategy in strategies if strategy.category == TenantPluginAutoUpgradeCategory.MODEL
)
assert model_strategy.strategy_setting == TenantPluginAutoUpgradeStrategySetting.LATEST
assert model_strategy.upgrade_time_of_day == expected_time
def test_default_upgrade_time_is_aligned_to_fifteen_minutes(self):
from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService
default_time = PluginAutoUpgradeService.default_upgrade_time_of_day("t1")
def test_default_upgrade_time_is_aligned_to_fifteen_minutes(self) -> None:
default_time = PluginAutoUpgradeService.default_upgrade_time_of_day(str(uuid4()))
assert default_time % (15 * 60) == 0
assert 0 <= default_time < 24 * 60 * 60
def test_creates_missing_categories_and_splits_known_plugins(self, caplog: pytest.LogCaptureFixture):
session = _patched_session()
tool_strategy = SimpleNamespace(
category=TenantPluginAutoUpgradeCategory.TOOL,
strategy_setting=TenantPluginAutoUpgradeStrategySetting.FIX_ONLY,
upgrade_time_of_day=0,
upgrade_mode=TenantPluginAutoUpgradeMode.EXCLUDE,
exclude_plugins=["tool-plugin", "model-plugin", "unknown-plugin"],
include_plugins=["model-plugin", "tool-plugin"],
@pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True)
def test_creates_missing_categories_and_splits_known_plugins(
self, sqlite_session: Session, caplog: pytest.LogCaptureFixture
) -> None:
tenant_id = str(uuid4())
tool_strategy = _strategy(
tenant_id,
exclude=["tool-plugin", "model-plugin", "unknown-plugin"],
include=["model-plugin", "tool-plugin"],
)
model_strategy = SimpleNamespace(
model_strategy = _strategy(
tenant_id,
category=TenantPluginAutoUpgradeCategory.MODEL,
strategy_setting=TenantPluginAutoUpgradeStrategySetting.FIX_ONLY,
upgrade_time_of_day=0,
upgrade_mode=TenantPluginAutoUpgradeMode.EXCLUDE,
exclude_plugins=["tool-plugin", "model-plugin", "unknown-plugin"],
include_plugins=["model-plugin", "tool-plugin"],
exclude=["tool-plugin", "model-plugin", "unknown-plugin"],
include=["model-plugin", "tool-plugin"],
)
session.scalars.return_value.all.return_value = [tool_strategy, model_strategy]
sqlite_session.add_all([tool_strategy, model_strategy])
sqlite_session.commit()
installed_plugins = [
SimpleNamespace(
plugin_id="tool-plugin",
@@ -261,18 +235,17 @@ class TestBackfillStrategyCategories:
patch(f"{MODULE}.PluginInstaller", return_value=installer),
caplog.at_level(logging.WARNING, logger=MODULE),
):
from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService
result = PluginAutoUpgradeService.backfill_strategy_categories("t1", session=session)
result = PluginAutoUpgradeService.backfill_strategy_categories(tenant_id, session=sqlite_session)
strategies = list(sqlite_session.scalars(select(TenantPluginAutoUpgradeStrategy)).all())
assert result.created_count == len(TenantPluginAutoUpgradeCategory) - 2
assert result.normalized is True
assert session.add.call_count == len(TenantPluginAutoUpgradeCategory) - 2
assert len(strategies) == len(TenantPluginAutoUpgradeCategory)
assert tool_strategy.exclude_plugins == ["tool-plugin"]
assert tool_strategy.include_plugins == ["tool-plugin"]
assert model_strategy.exclude_plugins == ["model-plugin"]
assert model_strategy.include_plugins == ["model-plugin"]
assert (
"Skipped unknown plugin IDs while backfilling plugin auto-upgrade strategies: "
"tenant_id=t1, field=exclude_plugins, plugin_ids=['unknown-plugin']" in caplog.messages
f"tenant_id={tenant_id}, field=exclude_plugins, plugin_ids=['unknown-plugin']" in caplog.messages
)
@@ -7,28 +7,19 @@ import pytest
import zstandard
from pydantic import TypeAdapter
from redis import RedisError
from sqlalchemy.orm import Session
from core.helper.model_provider_cache import ProviderCredentialsCacheType
from core.plugin.entities.plugin import PluginCategory, PluginInstallationSource
from core.plugin.entities.plugin_daemon import PluginInstallTask, PluginInstallTaskStatus, PluginModelProviderEntity
from graphon.model_runtime.entities.common_entities import I18nObject
from graphon.model_runtime.entities.provider_entities import ConfigurateMethod, ProviderEntity
from models.provider import Provider, ProviderCredential, ProviderType, TenantPreferredModelProvider
MODULE = "core.plugin.plugin_service"
class _FakeSession:
def __init__(self) -> None:
self.execute = Mock()
self.scalars = Mock(return_value=SimpleNamespace(all=Mock(return_value=[])))
def __enter__(self) -> "_FakeSession":
return self
def __exit__(self, exc_type, exc, traceback) -> None:
return None
def begin(self) -> "_FakeSession":
return self
TENANT_ID = "11111111-1111-1111-1111-111111111111"
OTHER_TENANT_ID = "22222222-2222-2222-2222-222222222222"
USER_ID = "33333333-3333-3333-3333-333333333333"
def _build_provider_entity(provider: str = "openai") -> ProviderEntity:
@@ -1166,19 +1157,72 @@ class TestPluginModelProviderCacheInvalidation:
assert result is True
invalidate_cache.assert_called_once_with("tenant-1")
def test_uninstall_existing_plugin_invalidates_cache_after_credential_cleanup(self) -> None:
@pytest.mark.parametrize(
"sqlite_session", [(Provider, ProviderCredential, TenantPreferredModelProvider)], indirect=True
)
def test_uninstall_existing_plugin_invalidates_cache_after_credential_cleanup(
self, sqlite_session: Session
) -> None:
"""Successful uninstall with plugin metadata also invalidates the mutated tenant provider cache."""
plugin_id = "langgenius/openai"
provider_name = f"{plugin_id}/openai"
plugin = SimpleNamespace(
installation_id="installation-1",
plugin_id="langgenius/openai",
plugin_id=plugin_id,
plugin_unique_identifier="langgenius/openai:1.0.0",
)
session = _FakeSession()
credential = ProviderCredential(
tenant_id=TENANT_ID,
provider_name=provider_name,
credential_name="Target credential",
encrypted_config="{}",
user_id=USER_ID,
)
other_credential = ProviderCredential(
tenant_id=OTHER_TENANT_ID,
provider_name=provider_name,
credential_name="Other credential",
encrypted_config="{}",
user_id=USER_ID,
)
sqlite_session.add_all([credential, other_credential])
sqlite_session.flush()
provider = Provider(
tenant_id=TENANT_ID,
provider_name=provider_name,
provider_type=ProviderType.CUSTOM,
credential_id=credential.id,
)
other_provider = Provider(
tenant_id=OTHER_TENANT_ID,
provider_name=provider_name,
provider_type=ProviderType.CUSTOM,
credential_id=other_credential.id,
)
preferred_provider = TenantPreferredModelProvider(
tenant_id=TENANT_ID,
provider_name=provider_name,
preferred_provider_type=ProviderType.CUSTOM,
)
other_preferred_provider = TenantPreferredModelProvider(
tenant_id=OTHER_TENANT_ID,
provider_name=provider_name,
preferred_provider_type=ProviderType.CUSTOM,
)
sqlite_session.add_all([provider, other_provider, preferred_provider, other_preferred_provider])
sqlite_session.commit()
credential_id = credential.id
other_credential_id = other_credential.id
provider_id = provider.id
other_provider_id = other_provider.id
preferred_provider_id = preferred_provider.id
other_preferred_provider_id = other_preferred_provider.id
with (
patch(f"{MODULE}.db", SimpleNamespace(engine=object())),
patch(f"{MODULE}.db", SimpleNamespace(engine=sqlite_session.get_bind())),
patch(f"{MODULE}.dify_config") as mock_config,
patch(f"{MODULE}.PluginInstaller") as installer_cls,
patch(f"{MODULE}.Session", return_value=session),
patch(f"{MODULE}.ProviderCredentialsCache") as credentials_cache,
patch(f"{MODULE}.PluginService.invalidate_plugin_model_providers_cache") as invalidate_cache,
):
mock_config.ENTERPRISE_ENABLED = False
@@ -1188,8 +1232,26 @@ class TestPluginModelProviderCacheInvalidation:
from core.plugin.plugin_service import PluginService
result = PluginService.uninstall("tenant-1", "installation-1")
result = PluginService.uninstall(TENANT_ID, "installation-1")
assert result is True
installer.uninstall.assert_called_once_with("tenant-1", "installation-1")
invalidate_cache.assert_called_once_with("tenant-1")
installer.uninstall.assert_called_once_with(TENANT_ID, "installation-1")
invalidate_cache.assert_called_once_with(TENANT_ID)
credentials_cache.assert_called_once_with(
tenant_id=TENANT_ID,
identity_id=provider_id,
cache_type=ProviderCredentialsCacheType.PROVIDER,
)
credentials_cache.return_value.delete.assert_called_once_with()
sqlite_session.expunge_all()
assert sqlite_session.get(ProviderCredential, credential_id) is None
persisted_provider = sqlite_session.get(Provider, provider_id)
assert persisted_provider is not None
assert persisted_provider.credential_id is None
assert sqlite_session.get(TenantPreferredModelProvider, preferred_provider_id) is None
assert sqlite_session.get(ProviderCredential, other_credential_id) is not None
persisted_other_provider = sqlite_session.get(Provider, other_provider_id)
assert persisted_other_provider is not None
assert persisted_other_provider.credential_id == other_credential_id
assert sqlite_session.get(TenantPreferredModelProvider, other_preferred_provider_id) is not None
@@ -8,6 +8,7 @@ verification, marketplace upgrade flows, and uninstall with credential cleanup.
from __future__ import annotations
from collections.abc import Iterator
from typing import cast
from unittest.mock import MagicMock, patch
from uuid import uuid4
@@ -19,7 +20,6 @@ from sqlalchemy.orm import Session
from core.plugin.entities.plugin import PluginInstallationSource
from core.plugin.entities.plugin_daemon import PluginVerification
from core.plugin.plugin_service import PluginService
from enums.deployment_edition import DeploymentEdition
from models import ProviderType
from models.engine import db
from models.provider import Provider, ProviderCredential, TenantPreferredModelProvider
@@ -27,20 +27,16 @@ from services.errors.plugin import PluginInstallationForbiddenError
from services.feature_service import (
PluginInstallationPermissionModel,
PluginInstallationScope,
SystemFeatureModel,
)
def _make_features(
def _make_permission(
restrict_to_marketplace: bool = False,
scope: PluginInstallationScope = PluginInstallationScope.ALL,
) -> SystemFeatureModel:
return SystemFeatureModel(
deployment_edition=DeploymentEdition.COMMUNITY,
plugin_installation_permission=PluginInstallationPermissionModel(
restrict_to_marketplace_only=restrict_to_marketplace,
plugin_installation_scope=scope,
),
) -> PluginInstallationPermissionModel:
return PluginInstallationPermissionModel(
restrict_to_marketplace_only=restrict_to_marketplace,
plugin_installation_scope=scope,
)
@@ -119,22 +115,31 @@ class TestFetchLatestPluginVersion:
class TestCheckMarketplaceOnlyPermission:
@patch("core.plugin.plugin_service.FeatureService")
def test_raises_when_restricted(self, mock_fs):
mock_fs.get_system_features.return_value = _make_features(restrict_to_marketplace=True)
mock_fs.get_plugin_installation_permission.return_value = _make_permission(restrict_to_marketplace=True)
with pytest.raises(PluginInstallationForbiddenError):
PluginService._check_marketplace_only_permission()
@patch("core.plugin.plugin_service.FeatureService")
def test_passes_when_not_restricted(self, mock_fs):
mock_fs.get_system_features.return_value = _make_features(restrict_to_marketplace=False)
mock_fs.get_plugin_installation_permission.return_value = _make_permission(restrict_to_marketplace=False)
PluginService._check_marketplace_only_permission() # should not raise
@patch("core.plugin.plugin_service.FeatureService")
def test_raises_when_scope_denies_all(self, mock_fs):
mock_fs.get_plugin_installation_permission.return_value = _make_permission(scope=PluginInstallationScope.NONE)
with pytest.raises(PluginInstallationForbiddenError, match="not allowed"):
PluginService._check_marketplace_only_permission()
class TestCheckPluginInstallationScope:
@patch("core.plugin.plugin_service.FeatureService")
def test_official_only_allows_langgenius(self, mock_fs):
mock_fs.get_system_features.return_value = _make_features(scope=PluginInstallationScope.OFFICIAL_ONLY)
mock_fs.get_plugin_installation_permission.return_value = _make_permission(
scope=PluginInstallationScope.OFFICIAL_ONLY
)
verification = MagicMock()
verification.authorized_category = PluginVerification.AuthorizedCategory.Langgenius
@@ -142,14 +147,16 @@ class TestCheckPluginInstallationScope:
@patch("core.plugin.plugin_service.FeatureService")
def test_official_only_rejects_third_party(self, mock_fs):
mock_fs.get_system_features.return_value = _make_features(scope=PluginInstallationScope.OFFICIAL_ONLY)
mock_fs.get_plugin_installation_permission.return_value = _make_permission(
scope=PluginInstallationScope.OFFICIAL_ONLY
)
with pytest.raises(PluginInstallationForbiddenError):
PluginService._check_plugin_installation_scope(None)
@patch("core.plugin.plugin_service.FeatureService")
def test_official_and_partners_allows_partner(self, mock_fs):
mock_fs.get_system_features.return_value = _make_features(
mock_fs.get_plugin_installation_permission.return_value = _make_permission(
scope=PluginInstallationScope.OFFICIAL_AND_SPECIFIC_PARTNERS
)
verification = MagicMock()
@@ -159,7 +166,7 @@ class TestCheckPluginInstallationScope:
@patch("core.plugin.plugin_service.FeatureService")
def test_official_and_partners_rejects_none(self, mock_fs):
mock_fs.get_system_features.return_value = _make_features(
mock_fs.get_plugin_installation_permission.return_value = _make_permission(
scope=PluginInstallationScope.OFFICIAL_AND_SPECIFIC_PARTNERS
)
@@ -168,7 +175,7 @@ class TestCheckPluginInstallationScope:
@patch("core.plugin.plugin_service.FeatureService")
def test_none_scope_always_raises(self, mock_fs):
mock_fs.get_system_features.return_value = _make_features(scope=PluginInstallationScope.NONE)
mock_fs.get_plugin_installation_permission.return_value = _make_permission(scope=PluginInstallationScope.NONE)
verification = MagicMock()
verification.authorized_category = PluginVerification.AuthorizedCategory.Langgenius
@@ -177,10 +184,19 @@ class TestCheckPluginInstallationScope:
@patch("core.plugin.plugin_service.FeatureService")
def test_all_scope_passes_any(self, mock_fs):
mock_fs.get_system_features.return_value = _make_features(scope=PluginInstallationScope.ALL)
mock_fs.get_plugin_installation_permission.return_value = _make_permission(scope=PluginInstallationScope.ALL)
PluginService._check_plugin_installation_scope(None) # should not raise
@patch("core.plugin.plugin_service.FeatureService")
def test_unknown_scope_always_raises(self, mock_fs):
permission = _make_permission()
permission.plugin_installation_scope = cast(PluginInstallationScope, "unknown-scope")
mock_fs.get_plugin_installation_permission.return_value = permission
with pytest.raises(PluginInstallationForbiddenError, match="policy is invalid"):
PluginService._check_plugin_installation_scope(None)
class TestGetPluginIconUrl:
@patch("core.plugin.plugin_service.dify_config")
@@ -248,7 +264,7 @@ class TestUpgradePluginWithMarketplace:
@patch("core.plugin.plugin_service.dify_config")
def test_skips_download_when_already_installed(self, mock_config, mock_installer_cls, mock_fs, mock_marketplace):
mock_config.MARKETPLACE_ENABLED = True
mock_fs.get_system_features.return_value = _make_features()
mock_fs.get_plugin_installation_permission.return_value = _make_permission()
installer = mock_installer_cls.return_value
installer.fetch_plugin_manifest.return_value = MagicMock()
installer.upgrade_plugin.return_value = MagicMock()
@@ -264,7 +280,7 @@ class TestUpgradePluginWithMarketplace:
@patch("core.plugin.plugin_service.dify_config")
def test_downloads_when_not_installed(self, mock_config, mock_installer_cls, mock_fs, mock_download):
mock_config.MARKETPLACE_ENABLED = True
mock_fs.get_system_features.return_value = _make_features()
mock_fs.get_plugin_installation_permission.return_value = _make_permission()
installer = mock_installer_cls.return_value
installer.fetch_plugin_manifest.side_effect = RuntimeError("not found")
mock_download.return_value = b"pkg-bytes"
@@ -283,7 +299,7 @@ class TestUpgradePluginWithGithub:
@patch("core.plugin.plugin_service.FeatureService")
@patch("core.plugin.plugin_service.PluginInstaller")
def test_checks_marketplace_permission_and_delegates(self, mock_installer_cls: MagicMock, mock_fs: MagicMock):
mock_fs.get_system_features.return_value = _make_features()
mock_fs.get_plugin_installation_permission.return_value = _make_permission()
installer = mock_installer_cls.return_value
installer.upgrade_plugin.return_value = MagicMock()
@@ -298,7 +314,7 @@ class TestUploadPkg:
@patch("core.plugin.plugin_service.FeatureService")
@patch("core.plugin.plugin_service.PluginInstaller")
def test_runs_permission_and_scope_checks(self, mock_installer_cls: MagicMock, mock_fs: MagicMock):
mock_fs.get_system_features.return_value = _make_features()
mock_fs.get_plugin_installation_permission.return_value = _make_permission()
upload_resp = MagicMock()
upload_resp.verification = None
mock_installer_cls.return_value.upload_pkg.return_value = upload_resp
@@ -322,7 +338,7 @@ class TestInstallFromMarketplacePkg:
@patch("core.plugin.plugin_service.dify_config")
def test_downloads_when_not_cached(self, mock_config, mock_installer_cls, mock_fs, mock_download):
mock_config.MARKETPLACE_ENABLED = True
mock_fs.get_system_features.return_value = _make_features()
mock_fs.get_plugin_installation_permission.return_value = _make_permission()
installer = mock_installer_cls.return_value
installer.fetch_plugin_manifest.side_effect = RuntimeError("not found")
mock_download.return_value = b"pkg"
@@ -344,7 +360,7 @@ class TestInstallFromMarketplacePkg:
@patch("core.plugin.plugin_service.dify_config")
def test_uses_cached_when_already_downloaded(self, mock_config, mock_installer_cls: MagicMock, mock_fs: MagicMock):
mock_config.MARKETPLACE_ENABLED = True
mock_fs.get_system_features.return_value = _make_features()
mock_fs.get_plugin_installation_permission.return_value = _make_permission()
installer = mock_installer_cls.return_value
installer.fetch_plugin_manifest.return_value = MagicMock()
decode_resp = MagicMock()
@@ -0,0 +1,92 @@
import logging
import pytest
from enums.deployment_edition import DeploymentEdition
from services import feature_service as feature_service_module
from services.feature_service import FeatureService, PluginInstallationScope, SystemFeatureModel
def test_get_plugin_installation_permission_defaults_to_all_for_non_enterprise(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(feature_service_module.dify_config, "ENTERPRISE_ENABLED", False)
permission = FeatureService.get_plugin_installation_permission()
assert permission.plugin_installation_scope is PluginInstallationScope.ALL
assert permission.restrict_to_marketplace_only is False
def test_get_plugin_installation_permission_parses_enterprise_policy(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(feature_service_module.dify_config, "ENTERPRISE_ENABLED", True)
monkeypatch.setattr(
feature_service_module.EnterpriseService,
"get_info",
staticmethod(
lambda: {
"PluginInstallationPermission": {
"pluginInstallationScope": "official_only",
"restrictToMarketplaceOnly": True,
}
}
),
)
permission = FeatureService.get_plugin_installation_permission()
assert permission.plugin_installation_scope is PluginInstallationScope.OFFICIAL_ONLY
assert permission.restrict_to_marketplace_only is True
@pytest.mark.parametrize(
"invalid_permission",
[
{
"pluginInstallationScope": "unknown-scope",
"restrictToMarketplaceOnly": False,
},
{
"pluginInstallationScope": "all",
"restrictToMarketplaceOnly": "false",
},
],
ids=["unknown_scope", "non_boolean_marketplace_restriction"],
)
def test_invalid_enterprise_policy_denies_all_plugin_installations(
caplog: pytest.LogCaptureFixture,
invalid_permission: dict[str, object],
) -> None:
with caplog.at_level(logging.ERROR, logger="services.feature_service"):
permission = FeatureService._resolve_plugin_installation_permission(
{"PluginInstallationPermission": invalid_permission}
)
assert permission.plugin_installation_scope is PluginInstallationScope.NONE
assert permission.restrict_to_marketplace_only is True
assert "denying all plugin installations" in caplog.text
def test_system_features_exposes_only_validated_plugin_installation_policy(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(
feature_service_module.EnterpriseService,
"get_info",
staticmethod(
lambda: {
"PluginInstallationPermission": {
"pluginInstallationScope": "unknown-scope",
"restrictToMarketplaceOnly": False,
}
}
),
)
features = SystemFeatureModel(deployment_edition=DeploymentEdition.ENTERPRISE)
FeatureService._fulfill_params_from_enterprise(features)
assert features.plugin_installation_permission.plugin_installation_scope is PluginInstallationScope.NONE
assert features.plugin_installation_permission.restrict_to_marketplace_only is True
+111 -124
View File
@@ -1,6 +1,8 @@
import base64
import hashlib
import os
from collections.abc import Iterator
from datetime import UTC, datetime
from unittest.mock import MagicMock, patch
import pytest
@@ -9,6 +11,8 @@ from sqlalchemy.orm import Session, sessionmaker
from werkzeug.exceptions import NotFound
from configs import dify_config
from extensions.storage.storage_type import StorageType
from models.base import TypeBase
from models.enums import CreatorUserRole
from models.model import Account, EndUser, UploadFile
from services.errors.file import BlockedFileExtensionError, FileTooLargeError, UnsupportedFileTypeError
@@ -17,31 +21,54 @@ from services.file_service import FileService
class TestFileService:
@pytest.fixture
def mock_db_session(self):
session = MagicMock(spec=Session)
# Mock context manager behavior
session.__enter__.return_value = session
return session
def sqlite_session_maker(self, sqlite_engine: Engine) -> sessionmaker[Session]:
TypeBase.metadata.create_all(sqlite_engine, tables=[TypeBase.metadata.tables[UploadFile.__tablename__]])
return sessionmaker(bind=sqlite_engine, expire_on_commit=False)
@pytest.fixture
def mock_session_maker(self, mock_db_session):
maker = MagicMock(spec=sessionmaker)
maker.return_value = mock_db_session
return maker
def db_session(self, sqlite_session_maker: sessionmaker[Session]) -> Iterator[Session]:
with sqlite_session_maker() as session:
yield session
@pytest.fixture
def file_service(self, mock_session_maker):
return FileService(session_factory=mock_session_maker)
def file_service(self, sqlite_session_maker: sessionmaker[Session]) -> FileService:
return FileService(session_factory=sqlite_session_maker)
def test_init_with_engine(self):
engine = MagicMock(spec=Engine)
service = FileService(session_factory=engine)
@staticmethod
def _persist_upload_file(
session: Session,
*,
file_id: str = "file_id",
tenant_id: str = "tenant_id",
extension: str = "txt",
mime_type: str = "text/plain",
key: str = "key",
) -> UploadFile:
upload_file = UploadFile(
tenant_id=tenant_id,
storage_type=StorageType.LOCAL,
key=key,
name=f"test.{extension}",
size=10,
extension=extension,
mime_type=mime_type,
created_by_role=CreatorUserRole.ACCOUNT,
created_by="user_id",
created_at=datetime(2024, 1, 1, tzinfo=UTC),
used=False,
)
upload_file.id = file_id
session.add(upload_file)
session.commit()
return upload_file
def test_init_with_engine(self, sqlite_engine: Engine):
service = FileService(session_factory=sqlite_engine)
assert isinstance(service._session_maker, sessionmaker)
def test_init_with_sessionmaker(self):
maker = MagicMock(spec=sessionmaker)
service = FileService(session_factory=maker)
assert service._session_maker == maker
def test_init_with_sessionmaker(self, sqlite_session_maker: sessionmaker[Session]):
service = FileService(session_factory=sqlite_session_maker)
assert service._session_maker == sqlite_session_maker
def test_init_invalid_factory(self):
with pytest.raises(AssertionError, match="must be a sessionmaker or an Engine."):
@@ -52,11 +79,11 @@ class TestFileService:
@patch("services.file_service.extract_tenant_id")
@patch("services.file_service.file_helpers.get_signed_file_url")
def test_upload_file_success(
self, mock_get_url, mock_tenant_id, mock_now, mock_storage, file_service: FileService, mock_db_session
self, mock_get_url, mock_tenant_id, mock_now, mock_storage, file_service: FileService, db_session: Session
):
# Setup
mock_tenant_id.return_value = "tenant_id"
mock_now.return_value = "2024-01-01"
mock_now.return_value = datetime(2024, 1, 1, tzinfo=UTC)
mock_get_url.return_value = "http://signed-url"
user = MagicMock(spec=Account)
@@ -81,8 +108,9 @@ class TestFileService:
assert result.source_url == "http://signed-url"
mock_storage.save.assert_called_once()
mock_db_session.add.assert_called_once_with(result)
mock_db_session.commit.assert_called_once()
persisted = db_session.get(UploadFile, result.id)
assert persisted is not None
assert persisted.hash == result.hash
def test_upload_file_uses_explicit_resource_tenant(self, file_service: FileService):
user = MagicMock(spec=Account)
@@ -109,7 +137,7 @@ class TestFileService:
with pytest.raises(ValueError, match="Filename contains invalid characters"):
file_service.upload_file(filename="invalid/file.txt", content=b"", mimetype="text/plain", user=MagicMock())
def test_upload_file_long_filename(self, file_service: FileService, mock_db_session):
def test_upload_file_long_filename(self, file_service: FileService, db_session: Session):
# Setup
long_name = "a" * 210 + ".txt"
user = MagicMock(spec=Account)
@@ -124,6 +152,7 @@ class TestFileService:
result = file_service.upload_file(filename=long_name, content=b"test", mimetype="text/plain", user=user)
assert len(result.name) <= 205 # 200 + . + extension
assert result.name.endswith(".txt")
assert db_session.get(UploadFile, result.id) is not None
def test_upload_file_blocked_extension(self, file_service):
with patch.object(dify_config, "inner_UPLOAD_FILE_EXTENSION_BLACKLIST", "exe"):
@@ -145,7 +174,7 @@ class TestFileService:
with pytest.raises(FileTooLargeError):
file_service.upload_file(filename="test.jpg", content=content, mimetype="image/jpeg", user=MagicMock())
def test_upload_file_end_user(self, file_service: FileService, mock_db_session):
def test_upload_file_end_user(self, file_service: FileService, db_session: Session):
user = MagicMock(spec=EndUser)
user.id = "end_user_id"
@@ -157,6 +186,7 @@ class TestFileService:
mock_tenant.return_value = "tenant"
result = file_service.upload_file(filename="test.txt", content=b"test", mimetype="text/plain", user=user)
assert result.created_by_role == CreatorUserRole.END_USER
assert db_session.get(UploadFile, result.id) is not None
def test_is_file_size_within_limit(self):
with (
@@ -181,12 +211,8 @@ class TestFileService:
assert FileService.is_file_size_within_limit(extension="txt", file_size=5 * 1024 * 1024) is True
assert FileService.is_file_size_within_limit(extension="pdf", file_size=6 * 1024 * 1024) is False
def test_get_file_base64_success(self, file_service: FileService, mock_db_session):
# Setup
upload_file = MagicMock(spec=UploadFile)
upload_file.id = "file_id"
upload_file.key = "test_key"
mock_db_session.scalar.return_value = upload_file
def test_get_file_base64_success(self, file_service: FileService, db_session: Session):
self._persist_upload_file(db_session, key="test_key")
with patch("services.file_service.storage") as mock_storage:
mock_storage.load_once.return_value = b"test content"
@@ -198,16 +224,17 @@ class TestFileService:
assert result == base64.b64encode(b"test content").decode()
mock_storage.load_once.assert_called_once_with("test_key")
def test_get_file_base64_not_found(self, file_service: FileService, mock_db_session):
mock_db_session.scalar.return_value = None
def test_get_file_base64_not_found(self, file_service: FileService):
with pytest.raises(NotFound, match="File not found"):
file_service.get_file_base64("non_existent")
def test_get_file_presigned_url_success(self, file_service: FileService, mock_db_session):
upload_file = MagicMock(spec=UploadFile)
upload_file.key = "upload_files/tenant_id/icon.png"
upload_file.mime_type = "image/png"
mock_db_session.scalar.return_value = upload_file
def test_get_file_presigned_url_success(self, file_service: FileService, db_session: Session):
self._persist_upload_file(
db_session,
extension="png",
mime_type="image/png",
key="upload_files/tenant_id/icon.png",
)
with (
patch.object(dify_config, "FILES_ACCESS_TIMEOUT", 300),
@@ -224,13 +251,11 @@ class TestFileService:
content_type="image/png",
)
def test_get_file_presigned_url_not_found(self, file_service: FileService, mock_db_session):
mock_db_session.scalar.return_value = None
def test_get_file_presigned_url_not_found(self, file_service: FileService):
with pytest.raises(NotFound, match="File not found"):
file_service.get_file_presigned_url(file_id="file_id", tenant_id="tenant_id")
def test_upload_text_success(self, file_service: FileService, mock_db_session):
def test_upload_text_success(self, file_service: FileService, db_session: Session):
# Setup
text = "sample text"
text_name = "test.txt"
@@ -249,21 +274,17 @@ class TestFileService:
assert result.used is True
assert result.extension == "txt"
mock_storage.save.assert_called_once()
mock_db_session.add.assert_called_once()
mock_db_session.commit.assert_called_once()
assert db_session.get(UploadFile, result.id) is not None
def test_upload_text_long_name(self, file_service: FileService, mock_db_session):
def test_upload_text_long_name(self, file_service: FileService, db_session: Session):
long_name = "a" * 210
with patch("services.file_service.storage"):
result = file_service.upload_text("text", long_name, "user", "tenant")
assert len(result.name) == 200
assert db_session.get(UploadFile, result.id) is not None
def test_get_file_preview_success(self, file_service: FileService, mock_db_session):
# Setup
upload_file = MagicMock(spec=UploadFile)
upload_file.id = "file_id"
upload_file.extension = "pdf"
mock_db_session.scalar.return_value = upload_file
def test_get_file_preview_success(self, file_service: FileService, db_session: Session):
self._persist_upload_file(db_session, extension="pdf", mime_type="application/pdf")
with patch("services.file_service.ExtractProcessor.load_from_upload_file") as mock_extract:
mock_extract.return_value = "Extracted text content"
@@ -274,27 +295,17 @@ class TestFileService:
# Assert
assert result == "Extracted text content"
def test_get_file_preview_not_found(self, file_service: FileService, mock_db_session):
mock_db_session.scalar.return_value = None
def test_get_file_preview_not_found(self, file_service: FileService):
with pytest.raises(NotFound, match="File not found"):
file_service.get_file_preview("non_existent", "tenant_id")
def test_get_file_preview_unsupported_type(self, file_service: FileService, mock_db_session):
upload_file = MagicMock(spec=UploadFile)
upload_file.id = "file_id"
upload_file.extension = "exe"
mock_db_session.scalar.return_value = upload_file
def test_get_file_preview_unsupported_type(self, file_service: FileService, db_session: Session):
self._persist_upload_file(db_session, extension="exe", mime_type="application/octet-stream")
with pytest.raises(UnsupportedFileTypeError):
file_service.get_file_preview("file_id", "tenant_id")
def test_get_image_preview_success(self, file_service: FileService, mock_db_session):
# Setup
upload_file = MagicMock(spec=UploadFile)
upload_file.id = "file_id"
upload_file.extension = "jpg"
upload_file.mime_type = "image/jpeg"
upload_file.key = "key"
mock_db_session.scalar.return_value = upload_file
def test_get_image_preview_success(self, file_service: FileService, db_session: Session):
self._persist_upload_file(db_session, extension="jpg", mime_type="image/jpeg")
with (
patch("services.file_service.file_helpers.verify_image_signature") as mock_verify,
@@ -316,28 +327,21 @@ class TestFileService:
with pytest.raises(NotFound, match="File not found or signature is invalid"):
file_service.get_image_preview("file_id", "ts", "nonce", "sign")
def test_get_image_preview_not_found(self, file_service: FileService, mock_db_session):
mock_db_session.scalar.return_value = None
def test_get_image_preview_not_found(self, file_service: FileService):
with patch("services.file_service.file_helpers.verify_image_signature") as mock_verify:
mock_verify.return_value = True
with pytest.raises(NotFound, match="File not found or signature is invalid"):
file_service.get_image_preview("file_id", "ts", "nonce", "sign")
def test_get_image_preview_unsupported_type(self, file_service: FileService, mock_db_session):
upload_file = MagicMock(spec=UploadFile)
upload_file.id = "file_id"
upload_file.extension = "txt"
mock_db_session.scalar.return_value = upload_file
def test_get_image_preview_unsupported_type(self, file_service: FileService, db_session: Session):
self._persist_upload_file(db_session)
with patch("services.file_service.file_helpers.verify_image_signature") as mock_verify:
mock_verify.return_value = True
with pytest.raises(UnsupportedFileTypeError):
file_service.get_image_preview("file_id", "ts", "nonce", "sign")
def test_get_file_generator_by_file_id_success(self, file_service: FileService, mock_db_session):
upload_file = MagicMock(spec=UploadFile)
upload_file.id = "file_id"
upload_file.key = "key"
mock_db_session.scalar.return_value = upload_file
def test_get_file_generator_by_file_id_success(self, file_service: FileService, db_session: Session):
upload_file = self._persist_upload_file(db_session)
with (
patch("services.file_service.file_helpers.verify_file_signature") as mock_verify,
@@ -348,7 +352,8 @@ class TestFileService:
gen, file = file_service.get_file_generator_by_file_id("file_id", "ts", "nonce", "sign")
assert list(gen) == [b"chunk"]
assert file == upload_file
assert file.id == upload_file.id
assert file.key == upload_file.key
def test_get_file_generator_by_file_id_invalid_sig(self, file_service):
with patch("services.file_service.file_helpers.verify_file_signature") as mock_verify:
@@ -356,20 +361,14 @@ class TestFileService:
with pytest.raises(NotFound, match="File not found or signature is invalid"):
file_service.get_file_generator_by_file_id("file_id", "ts", "nonce", "sign")
def test_get_file_generator_by_file_id_not_found(self, file_service: FileService, mock_db_session):
mock_db_session.scalar.return_value = None
def test_get_file_generator_by_file_id_not_found(self, file_service: FileService):
with patch("services.file_service.file_helpers.verify_file_signature") as mock_verify:
mock_verify.return_value = True
with pytest.raises(NotFound, match="File not found or signature is invalid"):
file_service.get_file_generator_by_file_id("file_id", "ts", "nonce", "sign")
def test_get_public_image_preview_success(self, file_service: FileService, mock_db_session):
upload_file = MagicMock(spec=UploadFile)
upload_file.id = "file_id"
upload_file.extension = "png"
upload_file.mime_type = "image/png"
upload_file.key = "key"
mock_db_session.scalar.return_value = upload_file
def test_get_public_image_preview_success(self, file_service: FileService, db_session: Session):
self._persist_upload_file(db_session, extension="png", mime_type="image/png")
with patch("services.file_service.storage") as mock_storage:
mock_storage.load.return_value = b"image content"
@@ -377,66 +376,56 @@ class TestFileService:
assert gen == b"image content"
assert mime == "image/png"
def test_get_public_image_preview_not_found(self, file_service: FileService, mock_db_session):
mock_db_session.scalar.return_value = None
def test_get_public_image_preview_not_found(self, file_service: FileService):
with pytest.raises(NotFound, match="File not found or signature is invalid"):
file_service.get_public_image_preview("file_id")
def test_get_public_image_preview_unsupported_type(self, file_service: FileService, mock_db_session):
upload_file = MagicMock(spec=UploadFile)
upload_file.id = "file_id"
upload_file.extension = "txt"
mock_db_session.scalar.return_value = upload_file
def test_get_public_image_preview_unsupported_type(self, file_service: FileService, db_session: Session):
self._persist_upload_file(db_session)
with pytest.raises(UnsupportedFileTypeError):
file_service.get_public_image_preview("file_id")
def test_get_file_content_success(self, file_service: FileService, mock_db_session):
upload_file = MagicMock(spec=UploadFile)
upload_file.id = "file_id"
upload_file.key = "key"
mock_db_session.scalar.return_value = upload_file
def test_get_file_content_success(self, file_service: FileService, db_session: Session):
self._persist_upload_file(db_session)
with patch("services.file_service.storage") as mock_storage:
mock_storage.load.return_value = b"hello world"
result = file_service.get_file_content("file_id")
assert result == "hello world"
def test_get_file_content_not_found(self, file_service: FileService, mock_db_session):
mock_db_session.scalar.return_value = None
def test_get_file_content_not_found(self, file_service: FileService):
with pytest.raises(NotFound, match="File not found"):
file_service.get_file_content("file_id")
def test_delete_file_success(self, file_service: FileService, mock_db_session):
upload_file = MagicMock(spec=UploadFile)
upload_file.id = "file_id"
upload_file.key = "key"
# For session.scalar(select(...))
mock_db_session.scalar.return_value = upload_file
def test_delete_file_success(self, file_service: FileService, db_session: Session):
self._persist_upload_file(db_session)
with patch("services.file_service.storage") as mock_storage:
file_service.delete_file("file_id")
mock_storage.delete.assert_called_once_with("key")
mock_db_session.delete.assert_called_once_with(upload_file)
db_session.expire_all()
assert db_session.get(UploadFile, "file_id") is None
def test_delete_file_not_found(self, file_service: FileService, mock_db_session):
mock_db_session.scalar.return_value = None
def test_delete_file_not_found(self, file_service: FileService):
file_service.delete_file("file_id")
# Should return without doing anything
def test_get_upload_files_by_ids_empty(self):
session = MagicMock()
result = FileService.get_upload_files_by_ids("tenant_id", [], session=session)
def test_get_upload_files_by_ids_empty(self, db_session: Session):
result = FileService.get_upload_files_by_ids("tenant_id", [], session=db_session)
assert result == {}
def test_get_upload_files_by_ids(self):
upload_file = MagicMock(spec=UploadFile)
upload_file.id = "550e8400-e29b-41d4-a716-446655440000"
upload_file.tenant_id = "tenant_id"
session = MagicMock()
session.scalars().all.return_value = [upload_file]
def test_get_upload_files_by_ids(self, db_session: Session):
upload_file = self._persist_upload_file(db_session, file_id="550e8400-e29b-41d4-a716-446655440000")
self._persist_upload_file(
db_session,
file_id="550e8400-e29b-41d4-a716-446655440001",
tenant_id="other-tenant",
)
result = FileService.get_upload_files_by_ids(
"tenant_id", ["550e8400-e29b-41d4-a716-446655440000"], session=session
"tenant_id",
["550e8400-e29b-41d4-a716-446655440000", "550e8400-e29b-41d4-a716-446655440001"],
session=db_session,
)
assert result["550e8400-e29b-41d4-a716-446655440000"] == upload_file
@@ -453,10 +442,8 @@ class TestFileService:
used.add("a (1).txt")
assert FileService._dedupe_zip_entry_name("a.txt", used) == "a (2).txt"
def test_build_upload_files_zip_tempfile(self):
upload_file = MagicMock(spec=UploadFile)
upload_file.name = "test.txt"
upload_file.key = "key"
def test_build_upload_files_zip_tempfile(self, db_session: Session):
upload_file = self._persist_upload_file(db_session)
with (
patch("services.file_service.storage") as mock_storage,
@@ -1,6 +1,5 @@
import dataclasses
import logging
from collections.abc import Iterator
from datetime import datetime, timedelta
from unittest.mock import MagicMock
@@ -42,14 +41,13 @@ from services.human_input_service import (
@pytest.fixture
def sqlite_session_factory(sqlite_engine: Engine) -> Iterator[tuple[sessionmaker[Session], Session]]:
factory = sessionmaker(bind=sqlite_engine, expire_on_commit=False)
with factory() as session:
yield factory, session
def unbound_session_factory() -> sessionmaker[Session]:
"""Supply the required constructor dependency without enabling database access."""
return sessionmaker()
def _persist_app(sqlite_session: Session, mode: AppMode) -> App:
app = App(
def _make_app(mode: AppMode) -> App:
return App(
id="app-id",
tenant_id="tenant-id",
name="Test App",
@@ -60,9 +58,6 @@ def _persist_app(sqlite_session: Session, mode: AppMode) -> App:
enable_api=True,
max_active_requests=0,
)
sqlite_session.add(app)
sqlite_session.commit()
return app
@pytest.fixture
@@ -97,14 +92,11 @@ def sample_form_record():
)
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
def test_enqueue_resume_dispatches_task_for_workflow(
mocker: MockerFixture,
sqlite_session_factory,
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
):
session_factory, _ = sqlite_session_factory
service = HumanInputService(session_factory)
service = HumanInputService(sqlite_session_factory)
workflow_run = MagicMock()
workflow_run.app_id = "app-id"
@@ -116,7 +108,8 @@ def test_enqueue_resume_dispatches_task_for_workflow(
return_value=workflow_run_repo,
)
_persist_app(sqlite_session, AppMode.WORKFLOW)
with sqlite_session_factory.begin() as arrange_session:
arrange_session.add(_make_app(AppMode.WORKFLOW))
resume_task = mocker.patch("services.human_input_service.resume_app_execution")
@@ -128,10 +121,9 @@ def test_enqueue_resume_dispatches_task_for_workflow(
def test_ensure_form_active_respects_global_timeout(
monkeypatch, sample_form_record: HumanInputFormRecord, sqlite_session_factory
monkeypatch, sample_form_record: HumanInputFormRecord, unbound_session_factory
):
session_factory, _ = sqlite_session_factory
service = HumanInputService(session_factory)
service = HumanInputService(unbound_session_factory)
expired_record = dataclasses.replace(
sample_form_record,
created_at=naive_utc_now() - timedelta(hours=2),
@@ -143,14 +135,11 @@ def test_ensure_form_active_respects_global_timeout(
service.ensure_form_active(Form(expired_record))
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
def test_enqueue_resume_dispatches_task_for_advanced_chat(
mocker: MockerFixture,
sqlite_session_factory,
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
):
session_factory, _ = sqlite_session_factory
service = HumanInputService(session_factory)
service = HumanInputService(sqlite_session_factory)
workflow_run = MagicMock()
workflow_run.app_id = "app-id"
@@ -162,7 +151,8 @@ def test_enqueue_resume_dispatches_task_for_advanced_chat(
return_value=workflow_run_repo,
)
_persist_app(sqlite_session, AppMode.ADVANCED_CHAT)
with sqlite_session_factory.begin() as arrange_session:
arrange_session.add(_make_app(AppMode.ADVANCED_CHAT))
resume_task = mocker.patch("services.human_input_service.resume_app_execution")
@@ -173,14 +163,11 @@ def test_enqueue_resume_dispatches_task_for_advanced_chat(
assert call_kwargs["kwargs"]["payload"]["workflow_run_id"] == "workflow-run-id"
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
def test_enqueue_resume_skips_unsupported_app_mode(
mocker: MockerFixture,
sqlite_session_factory,
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
):
session_factory, _ = sqlite_session_factory
service = HumanInputService(session_factory)
service = HumanInputService(sqlite_session_factory)
workflow_run = MagicMock()
workflow_run.app_id = "app-id"
@@ -192,7 +179,8 @@ def test_enqueue_resume_skips_unsupported_app_mode(
return_value=workflow_run_repo,
)
_persist_app(sqlite_session, AppMode.COMPLETION)
with sqlite_session_factory.begin() as arrange_session:
arrange_session.add(_make_app(AppMode.COMPLETION))
resume_task = mocker.patch("services.human_input_service.resume_app_execution")
@@ -202,14 +190,13 @@ def test_enqueue_resume_skips_unsupported_app_mode(
def test_get_form_definition_by_token_for_console_uses_repository(
sample_form_record: HumanInputFormRecord, sqlite_session_factory
sample_form_record: HumanInputFormRecord, unbound_session_factory
):
session_factory, _ = sqlite_session_factory
repo = MagicMock(spec=HumanInputFormSubmissionRepository)
console_record = dataclasses.replace(sample_form_record, recipient_type=RecipientType.CONSOLE)
repo.get_by_token.return_value = console_record
service = HumanInputService(session_factory, form_repository=repo)
service = HumanInputService(unbound_session_factory, form_repository=repo)
form = service.get_form_definition_by_token_for_console("token")
repo.get_by_token.assert_called_once_with("token")
@@ -245,9 +232,8 @@ def _build_resumption_context_state(*, options: list[str], workflow_run_id: str)
def test_resolve_form_inputs_uses_runtime_select_options(
sample_form_record: HumanInputFormRecord, sqlite_session_factory, mocker: MockerFixture
sample_form_record: HumanInputFormRecord, unbound_session_factory, mocker: MockerFixture
):
session_factory, _ = sqlite_session_factory
configured_input = SelectInputConfig(
output_variable_name="decision",
option_source=StringListSource(
@@ -272,7 +258,7 @@ def test_resolve_form_inputs_uses_runtime_select_options(
"services.human_input_service.DifyAPIRepositoryFactory.create_api_workflow_run_repository",
return_value=workflow_run_repo,
)
service = HumanInputService(session_factory)
service = HumanInputService(unbound_session_factory)
resolved_inputs = service.resolve_form_inputs(Form(record))
@@ -284,13 +270,12 @@ def test_resolve_form_inputs_uses_runtime_select_options(
def test_submit_form_by_token_calls_repository_and_enqueue(
sample_form_record: HumanInputFormRecord, sqlite_session_factory, mocker: MockerFixture
sample_form_record: HumanInputFormRecord, unbound_session_factory, mocker: MockerFixture
):
session_factory, _ = sqlite_session_factory
repo = MagicMock(spec=HumanInputFormSubmissionRepository)
repo.get_by_token.return_value = sample_form_record
repo.mark_submitted.return_value = sample_form_record
service = HumanInputService(session_factory, form_repository=repo)
service = HumanInputService(unbound_session_factory, form_repository=repo)
enqueue_spy = mocker.patch.object(service, "enqueue_resume")
service.submit_form_by_token(
@@ -313,11 +298,10 @@ def test_submit_form_by_token_calls_repository_and_enqueue(
def test_submit_form_by_token_enqueues_agent_app_resume_for_conversation_form(
sample_form_record, sqlite_session_factory, mocker: MockerFixture
sample_form_record, unbound_session_factory, mocker: MockerFixture
):
# ENG-635: a conversation-owned (Agent v2 chat) form routes to the chat
# resume, not the workflow resume.
session_factory, _ = sqlite_session_factory
repo = MagicMock(spec=HumanInputFormSubmissionRepository)
conversation_record = dataclasses.replace(
sample_form_record,
@@ -326,7 +310,7 @@ def test_submit_form_by_token_enqueues_agent_app_resume_for_conversation_form(
)
repo.get_by_token.return_value = conversation_record
repo.mark_submitted.return_value = conversation_record
service = HumanInputService(session_factory, form_repository=repo)
service = HumanInputService(unbound_session_factory, form_repository=repo)
workflow_enqueue_spy = mocker.patch.object(service, "enqueue_resume")
chat_enqueue_spy = mocker.patch.object(service, "enqueue_agent_app_resume")
@@ -343,9 +327,8 @@ def test_submit_form_by_token_enqueues_agent_app_resume_for_conversation_form(
def test_submit_form_by_token_skips_enqueue_for_delivery_test(
sample_form_record: HumanInputFormRecord, sqlite_session_factory, mocker: MockerFixture
sample_form_record: HumanInputFormRecord, unbound_session_factory, mocker: MockerFixture
):
session_factory, _ = sqlite_session_factory
repo = MagicMock(spec=HumanInputFormSubmissionRepository)
test_record = dataclasses.replace(
sample_form_record,
@@ -354,7 +337,7 @@ def test_submit_form_by_token_skips_enqueue_for_delivery_test(
)
repo.get_by_token.return_value = test_record
repo.mark_submitted.return_value = test_record
service = HumanInputService(session_factory, form_repository=repo)
service = HumanInputService(unbound_session_factory, form_repository=repo)
enqueue_spy = mocker.patch.object(service, "enqueue_resume")
service.submit_form_by_token(
@@ -368,13 +351,12 @@ def test_submit_form_by_token_skips_enqueue_for_delivery_test(
def test_submit_form_by_token_passes_submission_user_id(
sample_form_record: HumanInputFormRecord, sqlite_session_factory, mocker: MockerFixture
sample_form_record: HumanInputFormRecord, unbound_session_factory, mocker: MockerFixture
):
session_factory, _ = sqlite_session_factory
repo = MagicMock(spec=HumanInputFormSubmissionRepository)
repo.get_by_token.return_value = sample_form_record
repo.mark_submitted.return_value = sample_form_record
service = HumanInputService(session_factory, form_repository=repo)
service = HumanInputService(unbound_session_factory, form_repository=repo)
enqueue_spy = mocker.patch.object(service, "enqueue_resume")
service.submit_form_by_token(
@@ -391,11 +373,10 @@ def test_submit_form_by_token_passes_submission_user_id(
enqueue_spy.assert_called_once_with(sample_form_record.workflow_run_id)
def test_submit_form_by_token_invalid_action(sample_form_record: HumanInputFormRecord, sqlite_session_factory):
session_factory, _ = sqlite_session_factory
def test_submit_form_by_token_invalid_action(sample_form_record: HumanInputFormRecord, unbound_session_factory):
repo = MagicMock(spec=HumanInputFormSubmissionRepository)
repo.get_by_token.return_value = dataclasses.replace(sample_form_record)
service = HumanInputService(session_factory, form_repository=repo)
service = HumanInputService(unbound_session_factory, form_repository=repo)
with pytest.raises(InvalidFormDataError) as exc_info:
service.submit_form_by_token(
@@ -409,8 +390,7 @@ def test_submit_form_by_token_invalid_action(sample_form_record: HumanInputFormR
repo.mark_submitted.assert_not_called()
def test_submit_form_by_token_missing_inputs(sample_form_record: HumanInputFormRecord, sqlite_session_factory):
session_factory, _ = sqlite_session_factory
def test_submit_form_by_token_missing_inputs(sample_form_record: HumanInputFormRecord, unbound_session_factory):
repo = MagicMock(spec=HumanInputFormSubmissionRepository)
definition_with_input = FormDefinition(
@@ -422,7 +402,7 @@ def test_submit_form_by_token_missing_inputs(sample_form_record: HumanInputFormR
)
form_with_input = dataclasses.replace(sample_form_record, definition=definition_with_input)
repo.get_by_token.return_value = form_with_input
service = HumanInputService(session_factory, form_repository=repo)
service = HumanInputService(unbound_session_factory, form_repository=repo)
with pytest.raises(InvalidFormDataError) as exc_info:
service.submit_form_by_token(
@@ -436,42 +416,6 @@ def test_submit_form_by_token_missing_inputs(sample_form_record: HumanInputFormR
repo.mark_submitted.assert_not_called()
def test_validate_human_input_submission_accepts_select_file_and_file_list(sqlite_session_factory):
session_factory, _ = sqlite_session_factory
service = HumanInputService(session_factory)
definition = FormDefinition.model_validate(
{
"form_content": "Pick one and upload files",
"inputs": [
{
"type": "select",
"output_variable_name": "decision",
"option_source": {
"type": "constant",
"value": ["approve", "reject"],
},
},
{
"type": "file",
"output_variable_name": "attachment",
"allowed_file_types": ["document"],
"allowed_file_upload_methods": ["remote_url"],
},
{
"type": "file-list",
"output_variable_name": "attachments",
"allowed_file_types": ["document"],
"allowed_file_upload_methods": ["remote_url"],
"number_limits": 3,
},
],
"user_actions": [{"id": "submit", "title": "Submit"}],
"rendered_content": "<p>Pick one and upload files</p>",
"expiration_time": naive_utc_now() + timedelta(hours=1),
}
)
@pytest.mark.parametrize(
("input_definition", "submitted_value", "expected_message"),
[
@@ -522,12 +466,11 @@ def test_validate_human_input_submission_accepts_select_file_and_file_list(sqlit
)
def test_validate_human_input_submission_rejects_invalid_select_and_file_payloads(
sample_form_record,
sqlite_session_factory,
unbound_session_factory,
input_definition,
submitted_value,
expected_message,
):
session_factory, _ = sqlite_session_factory
repo = MagicMock(spec=HumanInputFormSubmissionRepository)
definition = FormDefinition.model_validate(
{
@@ -539,7 +482,7 @@ def test_validate_human_input_submission_rejects_invalid_select_and_file_payload
}
)
repo.get_by_token.return_value = dataclasses.replace(sample_form_record, definition=definition)
service = HumanInputService(session_factory, form_repository=repo)
service = HumanInputService(unbound_session_factory, form_repository=repo)
with pytest.raises(InvalidFormDataError) as exc_info:
service.submit_form_by_token(
@@ -569,7 +512,7 @@ def test_form_properties(sample_form_record: HumanInputFormRecord):
def test_form_submitted_error_init():
error = FormSubmittedError(form_id="test-form")
assert "form_id=test-form" in error.description
assert error.description == "This form has already been submitted by another user, form_id=test-form"
assert error.code == 412
@@ -580,61 +523,55 @@ def test_human_input_service_init_with_engine(sqlite_engine: Engine):
assert service._session_factory.kw["bind"] is sqlite_engine
def test_get_form_by_token_none(sqlite_session_factory):
session_factory, _ = sqlite_session_factory
def test_get_form_by_token_none(unbound_session_factory):
repo = MagicMock(spec=HumanInputFormSubmissionRepository)
repo.get_by_token.return_value = None
service = HumanInputService(session_factory, form_repository=repo)
service = HumanInputService(unbound_session_factory, form_repository=repo)
assert service.get_form_by_token("invalid") is None
def test_get_form_definition_by_token_mismatch(sample_form_record: HumanInputFormRecord, sqlite_session_factory):
session_factory, _ = sqlite_session_factory
def test_get_form_definition_by_token_mismatch(sample_form_record: HumanInputFormRecord, unbound_session_factory):
repo = MagicMock(spec=HumanInputFormSubmissionRepository)
repo.get_by_token.return_value = sample_form_record
service = HumanInputService(session_factory, form_repository=repo)
service = HumanInputService(unbound_session_factory, form_repository=repo)
# RecipientType mismatch
assert service.get_form_definition_by_token(RecipientType.CONSOLE, "token") is None
def test_get_form_definition_by_token_success(sample_form_record: HumanInputFormRecord, sqlite_session_factory):
session_factory, _ = sqlite_session_factory
def test_get_form_definition_by_token_success(sample_form_record: HumanInputFormRecord, unbound_session_factory):
repo = MagicMock(spec=HumanInputFormSubmissionRepository)
repo.get_by_token.return_value = sample_form_record
service = HumanInputService(session_factory, form_repository=repo)
service = HumanInputService(unbound_session_factory, form_repository=repo)
form = service.get_form_definition_by_token(RecipientType.STANDALONE_WEB_APP, "token")
assert form is not None
assert form.id == sample_form_record.form_id
def test_get_form_definition_by_token_for_console_mismatch(
sample_form_record: HumanInputFormRecord, sqlite_session_factory
sample_form_record: HumanInputFormRecord, unbound_session_factory
):
session_factory, _ = sqlite_session_factory
repo = MagicMock(spec=HumanInputFormSubmissionRepository)
repo.get_by_token.return_value = sample_form_record # is STANDALONE_WEB_APP
service = HumanInputService(session_factory, form_repository=repo)
service = HumanInputService(unbound_session_factory, form_repository=repo)
assert service.get_form_definition_by_token_for_console("token") is None
def test_submit_form_by_token_delivery_not_enabled(sqlite_session_factory):
session_factory, _ = sqlite_session_factory
def test_submit_form_by_token_delivery_not_enabled(unbound_session_factory):
repo = MagicMock(spec=HumanInputFormSubmissionRepository)
repo.get_by_token.return_value = None
service = HumanInputService(session_factory, form_repository=repo)
service = HumanInputService(unbound_session_factory, form_repository=repo)
with pytest.raises(human_input_service_module.WebAppDeliveryNotEnabledError):
service.submit_form_by_token(RecipientType.STANDALONE_WEB_APP, "token", "action", {})
def test_submit_form_by_token_no_workflow_run_id(
sample_form_record: HumanInputFormRecord, sqlite_session_factory, mocker: MockerFixture
sample_form_record: HumanInputFormRecord, unbound_session_factory, mocker: MockerFixture
):
session_factory, _ = sqlite_session_factory
repo = MagicMock(spec=HumanInputFormSubmissionRepository)
repo.get_by_token.return_value = sample_form_record
@@ -642,16 +579,15 @@ def test_submit_form_by_token_no_workflow_run_id(
result_record = dataclasses.replace(sample_form_record, workflow_run_id=None)
repo.mark_submitted.return_value = result_record
service = HumanInputService(session_factory, form_repository=repo)
service = HumanInputService(unbound_session_factory, form_repository=repo)
enqueue_spy = mocker.patch.object(service, "enqueue_resume")
service.submit_form_by_token(RecipientType.STANDALONE_WEB_APP, "token", "submit", {})
enqueue_spy.assert_not_called()
def test_ensure_form_active_errors(sample_form_record: HumanInputFormRecord, sqlite_session_factory):
session_factory, _ = sqlite_session_factory
service = HumanInputService(session_factory)
def test_ensure_form_active_errors(sample_form_record: HumanInputFormRecord, unbound_session_factory):
service = HumanInputService(unbound_session_factory)
# Submitted
submitted_record = dataclasses.replace(sample_form_record, submitted_at=naive_utc_now())
@@ -671,18 +607,16 @@ def test_ensure_form_active_errors(sample_form_record: HumanInputFormRecord, sql
service.ensure_form_active(Form(expired_time_record))
def test_ensure_not_submitted_raises(sample_form_record: HumanInputFormRecord, sqlite_session_factory):
session_factory, _ = sqlite_session_factory
service = HumanInputService(session_factory)
def test_ensure_not_submitted_raises(sample_form_record: HumanInputFormRecord, unbound_session_factory):
service = HumanInputService(unbound_session_factory)
submitted_record = dataclasses.replace(sample_form_record, submitted_at=naive_utc_now())
with pytest.raises(human_input_service_module.FormSubmittedError):
service._ensure_not_submitted(Form(submitted_record))
def test_enqueue_resume_workflow_not_found(mocker: MockerFixture, sqlite_session_factory):
session_factory, _ = sqlite_session_factory
service = HumanInputService(session_factory)
def test_enqueue_resume_workflow_not_found(mocker: MockerFixture, unbound_session_factory):
service = HumanInputService(unbound_session_factory)
workflow_run_repo = MagicMock()
workflow_run_repo.get_workflow_run_by_id_without_tenant.return_value = None
@@ -696,15 +630,12 @@ def test_enqueue_resume_workflow_not_found(mocker: MockerFixture, sqlite_session
assert "WorkflowRun not found" in str(excinfo.value)
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
def test_enqueue_resume_app_not_found(
mocker,
sqlite_session_factory,
sqlite_session: Session,
sqlite_session_factory: sessionmaker[Session],
caplog: pytest.LogCaptureFixture,
):
session_factory, _ = sqlite_session_factory
service = HumanInputService(session_factory)
service = HumanInputService(sqlite_session_factory)
workflow_run = MagicMock()
workflow_run.app_id = "app-id"
@@ -715,26 +646,31 @@ def test_enqueue_resume_app_not_found(
"services.human_input_service.DifyAPIRepositoryFactory.create_api_workflow_run_repository",
return_value=workflow_run_repo,
)
resume_task = mocker.patch("services.human_input_service.resume_app_execution")
with caplog.at_level(logging.ERROR, logger="services.human_input_service"):
service.enqueue_resume("workflow-run-id")
assert any(r.levelno >= logging.ERROR for r in caplog.records)
assert (
"services.human_input_service",
logging.ERROR,
"App not found for WorkflowRun, workflow_run_id=workflow-run-id, app_id=app-id",
) in caplog.record_tuples
resume_task.apply_async.assert_not_called()
def test_is_globally_expired_zero_timeout(
monkeypatch: pytest.MonkeyPatch, sample_form_record: HumanInputFormRecord, sqlite_session_factory
monkeypatch: pytest.MonkeyPatch, sample_form_record: HumanInputFormRecord, unbound_session_factory
):
session_factory, _ = sqlite_session_factory
service = HumanInputService(session_factory)
service = HumanInputService(unbound_session_factory)
monkeypatch.setattr(human_input_service_module.dify_config, "HUMAN_INPUT_GLOBAL_TIMEOUT_SECONDS", 0)
assert service._is_globally_expired(Form(sample_form_record)) is False
def test_submit_form_by_token_normalizes_select_and_files(
sample_form_record: HumanInputFormRecord, sqlite_session_factory, mocker: MockerFixture
sample_form_record: HumanInputFormRecord, unbound_session_factory, mocker: MockerFixture
) -> None:
session_factory, _ = sqlite_session_factory
repo = MagicMock(spec=HumanInputFormSubmissionRepository)
definition = FormDefinition(
form_content="hello",
@@ -753,7 +689,7 @@ def test_submit_form_by_token_normalizes_select_and_files(
form_with_inputs = dataclasses.replace(sample_form_record, definition=definition)
repo.get_by_token.return_value = form_with_inputs
repo.mark_submitted.return_value = form_with_inputs
service = HumanInputService(session_factory, form_repository=repo)
service = HumanInputService(unbound_session_factory, form_repository=repo)
single_file = File(
file_id="file-1",
@@ -815,9 +751,8 @@ def test_submit_form_by_token_normalizes_select_and_files(
def test_submit_form_by_token_invalid_select_value(
sample_form_record: HumanInputFormRecord, sqlite_session_factory
sample_form_record: HumanInputFormRecord, unbound_session_factory
) -> None:
session_factory, _ = sqlite_session_factory
repo = MagicMock(spec=HumanInputFormSubmissionRepository)
definition = FormDefinition(
form_content="hello",
@@ -832,7 +767,7 @@ def test_submit_form_by_token_invalid_select_value(
expiration_time=sample_form_record.expiration_time,
)
repo.get_by_token.return_value = dataclasses.replace(sample_form_record, definition=definition)
service = HumanInputService(session_factory, form_repository=repo)
service = HumanInputService(unbound_session_factory, form_repository=repo)
with pytest.raises(InvalidFormDataError, match="Invalid value for select input 'decision'"):
service.submit_form_by_token(
@@ -844,9 +779,8 @@ def test_submit_form_by_token_invalid_select_value(
def test_submit_form_by_token_invalid_file_list_item(
sample_form_record: HumanInputFormRecord, sqlite_session_factory
sample_form_record: HumanInputFormRecord, unbound_session_factory
) -> None:
session_factory, _ = sqlite_session_factory
repo = MagicMock(spec=HumanInputFormSubmissionRepository)
definition = FormDefinition(
form_content="hello",
@@ -856,7 +790,7 @@ def test_submit_form_by_token_invalid_file_list_item(
expiration_time=sample_form_record.expiration_time,
)
repo.get_by_token.return_value = dataclasses.replace(sample_form_record, definition=definition)
service = HumanInputService(session_factory, form_repository=repo)
service = HumanInputService(unbound_session_factory, form_repository=repo)
with pytest.raises(
InvalidFormDataError,
@@ -871,9 +805,8 @@ def test_submit_form_by_token_invalid_file_list_item(
def test_submit_form_by_token_rejects_cross_tenant_file(
sample_form_record: HumanInputFormRecord, sqlite_session_factory, mocker: MockerFixture
sample_form_record: HumanInputFormRecord, unbound_session_factory, mocker: MockerFixture
) -> None:
session_factory, _ = sqlite_session_factory
repo = MagicMock(spec=HumanInputFormSubmissionRepository)
definition = FormDefinition(
form_content="hello",
@@ -883,7 +816,7 @@ def test_submit_form_by_token_rejects_cross_tenant_file(
expiration_time=sample_form_record.expiration_time,
)
repo.get_by_token.return_value = dataclasses.replace(sample_form_record, definition=definition)
service = HumanInputService(session_factory, form_repository=repo)
service = HumanInputService(unbound_session_factory, form_repository=repo)
mocker.patch("services.human_input_service.build_from_mapping", side_effect=ValueError("Invalid upload file"))
with pytest.raises(InvalidFormDataError, match="Invalid value for file input 'attachment'"):
@@ -904,9 +837,8 @@ def test_submit_form_by_token_rejects_cross_tenant_file(
def test_submit_form_by_token_rejects_cross_tenant_file_list(
sample_form_record: HumanInputFormRecord, sqlite_session_factory, mocker: MockerFixture
sample_form_record: HumanInputFormRecord, unbound_session_factory, mocker: MockerFixture
) -> None:
session_factory, _ = sqlite_session_factory
repo = MagicMock(spec=HumanInputFormSubmissionRepository)
definition = FormDefinition(
form_content="hello",
@@ -916,7 +848,7 @@ def test_submit_form_by_token_rejects_cross_tenant_file_list(
expiration_time=sample_form_record.expiration_time,
)
repo.get_by_token.return_value = dataclasses.replace(sample_form_record, definition=definition)
service = HumanInputService(session_factory, form_repository=repo)
service = HumanInputService(unbound_session_factory, form_repository=repo)
mocker.patch("services.human_input_service.build_from_mappings", side_effect=ValueError("Invalid upload file"))
with pytest.raises(
@@ -1,16 +1,20 @@
"""Testcontainers integration tests for OAuthServerService."""
"""Unit tests for OAuthServerService with SQLite-backed database access."""
from __future__ import annotations
import uuid
from collections.abc import Iterator
from typing import cast
from unittest.mock import MagicMock, patch
from uuid import uuid4
import pytest
from flask import Flask
from sqlalchemy import Engine
from sqlalchemy.orm import Session
from werkzeug.exceptions import BadRequest
from models.engine import db
from models.model import OAuthProviderApp
from services.oauth_server import (
OAUTH_ACCESS_TOKEN_EXPIRES_IN,
@@ -23,10 +27,24 @@ from services.oauth_server import (
)
class TestOAuthServerServiceGetProviderApp:
"""DB-backed tests for get_oauth_provider_app."""
@pytest.fixture
def oauth_db() -> Iterator[Session]:
"""Provide the production database extension with an isolated SQLite provider table."""
app = Flask(__name__)
app.config["SQLALCHEMY_DATABASE_URI"] = "sqlite:///:memory:"
db.init_app(app)
def _create_oauth_provider_app(self, db_session_with_containers: Session, *, client_id: str) -> OAuthProviderApp:
with app.app_context():
OAuthProviderApp.__table__.create(db.engine)
with Session(db.engine, expire_on_commit=False) as session:
yield session
class TestOAuthServerServiceGetProviderApp:
"""Verify provider lookup against a real SQLAlchemy database."""
def test_get_oauth_provider_app_returns_app_when_exists(self, oauth_db: Session) -> None:
client_id = f"client-{uuid4()}"
app = OAuthProviderApp(
app_icon="icon.png",
client_id=client_id,
@@ -35,35 +53,30 @@ class TestOAuthServerServiceGetProviderApp:
redirect_uris=["https://example.com/callback"],
scope="read",
)
db_session_with_containers.add(app)
db_session_with_containers.commit()
return app
def test_get_oauth_provider_app_returns_app_when_exists(self, db_session_with_containers: Session):
client_id = f"client-{uuid4()}"
created = self._create_oauth_provider_app(db_session_with_containers, client_id=client_id)
oauth_db.add(app)
oauth_db.commit()
result = OAuthServerService.get_oauth_provider_app(client_id)
assert result is not None
assert result.client_id == client_id
assert result.id == created.id
assert result.id == app.id
def test_get_oauth_provider_app_returns_none_when_not_exists(self, db_session_with_containers: Session):
def test_get_oauth_provider_app_returns_none_when_not_exists(self, oauth_db: Session) -> None:
result = OAuthServerService.get_oauth_provider_app(f"nonexistent-{uuid4()}")
assert result is None
class TestOAuthServerServiceTokenOperations:
"""Redis-backed tests for token sign/validate operations."""
"""Verify Redis-backed token signing and validation branches."""
@pytest.fixture
def mock_redis(self):
with patch("services.oauth_server.redis_client") as mock:
yield mock
def test_sign_authorization_code_stores_and_returns_code(self, mock_redis):
def test_sign_authorization_code_stores_and_returns_code(self, mock_redis) -> None:
deterministic_uuid = uuid.UUID("00000000-0000-0000-0000-000000000111")
with patch("services.oauth_server.uuid.uuid4", return_value=deterministic_uuid):
code = OAuthServerService.sign_oauth_authorization_code("client-1", "user-1")
@@ -75,7 +88,7 @@ class TestOAuthServerServiceTokenOperations:
ex=600,
)
def test_sign_access_token_raises_bad_request_for_invalid_code(self, mock_redis):
def test_sign_access_token_raises_bad_request_for_invalid_code(self, mock_redis) -> None:
mock_redis.get.return_value = None
with pytest.raises(BadRequest, match="invalid code"):
@@ -85,14 +98,13 @@ class TestOAuthServerServiceTokenOperations:
client_id="client-1",
)
def test_sign_access_token_issues_tokens_for_valid_code(self, mock_redis):
def test_sign_access_token_issues_tokens_for_valid_code(self, mock_redis) -> None:
token_uuids = [
uuid.UUID("00000000-0000-0000-0000-000000000201"),
uuid.UUID("00000000-0000-0000-0000-000000000202"),
]
with patch("services.oauth_server.uuid.uuid4", side_effect=token_uuids):
mock_redis.get.return_value = b"user-1"
access_token, refresh_token = OAuthServerService.sign_oauth_access_token(
grant_type=OAuthGrantType.AUTHORIZATION_CODE,
code="code-1",
@@ -114,7 +126,7 @@ class TestOAuthServerServiceTokenOperations:
ex=OAUTH_REFRESH_TOKEN_EXPIRES_IN,
)
def test_sign_access_token_raises_bad_request_for_invalid_refresh_token(self, mock_redis):
def test_sign_access_token_raises_bad_request_for_invalid_refresh_token(self, mock_redis) -> None:
mock_redis.get.return_value = None
with pytest.raises(BadRequest, match="invalid refresh token"):
@@ -124,11 +136,10 @@ class TestOAuthServerServiceTokenOperations:
client_id="client-1",
)
def test_sign_access_token_issues_new_token_for_valid_refresh(self, mock_redis):
def test_sign_access_token_issues_new_token_for_valid_refresh(self, mock_redis) -> None:
deterministic_uuid = uuid.UUID("00000000-0000-0000-0000-000000000301")
with patch("services.oauth_server.uuid.uuid4", return_value=deterministic_uuid):
mock_redis.get.return_value = b"user-1"
access_token, returned_refresh = OAuthServerService.sign_oauth_access_token(
grant_type=OAuthGrantType.REFRESH_TOKEN,
refresh_token="refresh-1",
@@ -138,14 +149,14 @@ class TestOAuthServerServiceTokenOperations:
assert access_token == str(deterministic_uuid)
assert returned_refresh == "refresh-1"
def test_sign_access_token_returns_none_for_unknown_grant_type(self, mock_redis):
def test_sign_access_token_returns_none_for_unknown_grant_type(self, mock_redis) -> None:
grant_type = cast(OAuthGrantType, "invalid-grant-type")
result = OAuthServerService.sign_oauth_access_token(grant_type=grant_type, client_id="client-1")
assert result is None
def test_sign_refresh_token_stores_with_expected_expiry(self, mock_redis):
def test_sign_refresh_token_stores_with_expected_expiry(self, mock_redis) -> None:
deterministic_uuid = uuid.UUID("00000000-0000-0000-0000-000000000401")
with patch("services.oauth_server.uuid.uuid4", return_value=deterministic_uuid):
refresh_token = OAuthServerService._sign_oauth_refresh_token("client-2", "user-2")
@@ -157,22 +168,21 @@ class TestOAuthServerServiceTokenOperations:
ex=OAUTH_REFRESH_TOKEN_EXPIRES_IN,
)
def test_validate_access_token_returns_none_when_not_found(self, mock_redis, db_session_with_containers: Session):
def test_validate_access_token_returns_none_when_not_found(self, mock_redis, sqlite_engine: Engine) -> None:
mock_redis.get.return_value = None
session = MagicMock()
result = OAuthServerService.validate_oauth_access_token("client-1", "missing-token", db_session_with_containers)
with Session(sqlite_engine) as session:
result = OAuthServerService.validate_oauth_access_token("client-1", "missing-token", session)
assert result is None
def test_validate_access_token_loads_user_when_exists(self, mock_redis, db_session_with_containers: Session):
def test_validate_access_token_loads_user_when_exists(self, mock_redis, sqlite_engine: Engine) -> None:
mock_redis.get.return_value = b"user-88"
expected_user = MagicMock()
with patch("services.oauth_server.AccountService.load_user", return_value=expected_user) as mock_load:
result = OAuthServerService.validate_oauth_access_token(
"client-1", "access-token", db_session_with_containers
)
with Session(sqlite_engine) as session:
with patch("services.oauth_server.AccountService.load_user", return_value=expected_user) as mock_load:
result = OAuthServerService.validate_oauth_access_token("client-1", "access-token", session)
mock_load.assert_called_once_with("user-88", session)
assert result is expected_user
mock_load.assert_called_once_with("user-88", db_session_with_containers)
@@ -1,178 +1,57 @@
"""Comprehensive unit tests for WorkflowRunService class.
"""Tests for the session lifecycle owned by ``WorkflowRunService``."""
This test suite covers all pause state management operations including:
- Retrieving pause state for workflow runs
- Saving pause state with file uploads
- Marking paused workflows as resumed
- Error handling and edge cases
- Database transaction management
- Repository-based approach testing
"""
from datetime import datetime
from unittest.mock import MagicMock, create_autospec, patch
from unittest.mock import create_autospec, patch
import pytest
from sqlalchemy import Engine
from sqlalchemy import Engine, text
from sqlalchemy.orm import Session, sessionmaker
from graphon.enums import WorkflowExecutionStatus
from models.workflow import WorkflowPause
from repositories.api_workflow_run_repository import APIWorkflowRunRepository
from repositories.sqlalchemy_api_workflow_run_repository import _PrivateWorkflowPauseEntity
from services.workflow_run_service import (
WorkflowRunService,
)
from services.workflow_run_service import WorkflowRunService
class TestDataFactory:
"""Factory class for creating test data objects."""
@staticmethod
def create_workflow_run_mock(
id: str = "workflow-run-123",
tenant_id: str = "tenant-456",
app_id: str = "app-789",
workflow_id: str = "workflow-101",
status: str | WorkflowExecutionStatus = "paused",
**kwargs,
) -> MagicMock:
"""Create a mock WorkflowRun object."""
mock_run = MagicMock()
mock_run.id = id
mock_run.tenant_id = tenant_id
mock_run.app_id = app_id
mock_run.workflow_id = workflow_id
mock_run.status = status
for key, value in kwargs.items():
setattr(mock_run, key, value)
return mock_run
@staticmethod
def create_workflow_pause_mock(
id: str = "pause-123",
tenant_id: str = "tenant-456",
app_id: str = "app-789",
workflow_id: str = "workflow-101",
workflow_execution_id: str = "workflow-execution-123",
state_file_id: str = "file-456",
resumed_at: datetime | None = None,
**kwargs,
) -> MagicMock:
"""Create a mock WorkflowPauseModel object."""
mock_pause = MagicMock(spec=WorkflowPause)
mock_pause.id = id
mock_pause.tenant_id = tenant_id
mock_pause.app_id = app_id
mock_pause.workflow_id = workflow_id
mock_pause.workflow_execution_id = workflow_execution_id
mock_pause.state_file_id = state_file_id
mock_pause.resumed_at = resumed_at
for key, value in kwargs.items():
setattr(mock_pause, key, value)
return mock_pause
@staticmethod
def create_pause_entity_mock(
pause_model: MagicMock | None = None,
) -> _PrivateWorkflowPauseEntity:
"""Create a mock _PrivateWorkflowPauseEntity object."""
if pause_model is None:
pause_model = TestDataFactory.create_workflow_pause_mock()
return _PrivateWorkflowPauseEntity(pause_model=pause_model, reason_models=[], human_input_form=[])
@pytest.fixture
def sqlite_session_factory(sqlite_engine: Engine) -> sessionmaker[Session]:
"""Return a real factory whose sessions are bound to the isolated SQLite engine."""
return sessionmaker(bind=sqlite_engine, expire_on_commit=False)
class TestWorkflowRunService:
"""Comprehensive unit tests for WorkflowRunService class."""
@pytest.fixture
def workflow_run_repository():
"""Keep the repository boundary mocked while exercising real session construction."""
return create_autospec(APIWorkflowRunRepository)
@pytest.fixture
def mock_session_factory(self):
"""Create a mock session factory with proper session management."""
mock_session = create_autospec(Session)
# Create a mock context manager for the session
mock_session_cm = MagicMock()
mock_session_cm.__enter__ = MagicMock(return_value=mock_session)
mock_session_cm.__exit__ = MagicMock(return_value=None)
def test_init_with_session_factory(
sqlite_session_factory: sessionmaker[Session], workflow_run_repository: APIWorkflowRunRepository
) -> None:
with patch("services.workflow_run_service.DifyAPIRepositoryFactory", autospec=True) as repository_factory:
repository_factory.create_api_workflow_run_repository.return_value = workflow_run_repository
# Create a mock context manager for the transaction
mock_transaction_cm = MagicMock()
mock_transaction_cm.__enter__ = MagicMock(return_value=mock_session)
mock_transaction_cm.__exit__ = MagicMock(return_value=None)
service = WorkflowRunService(sqlite_session_factory)
mock_session.begin = MagicMock(return_value=mock_transaction_cm)
assert service._session_factory is sqlite_session_factory
repository_factory.create_api_workflow_run_repository.assert_called_once_with(sqlite_session_factory)
with service._session_factory() as session:
assert session.scalar(text("SELECT 1")) == 1
# Create mock factory that returns the context manager
mock_factory = MagicMock(spec=sessionmaker)
mock_factory.return_value = mock_session_cm
return mock_factory, mock_session
def test_init_with_engine_creates_bound_session_factory(
sqlite_engine: Engine, workflow_run_repository: APIWorkflowRunRepository
) -> None:
with patch("services.workflow_run_service.DifyAPIRepositoryFactory", autospec=True) as repository_factory:
repository_factory.create_api_workflow_run_repository.return_value = workflow_run_repository
@pytest.fixture
def mock_workflow_run_repository(self):
"""Create a mock APIWorkflowRunRepository."""
mock_repo = create_autospec(APIWorkflowRunRepository)
return mock_repo
service = WorkflowRunService(sqlite_engine)
@pytest.fixture
def workflow_run_service(self, mock_session_factory, mock_workflow_run_repository):
"""Create WorkflowRunService instance with mocked dependencies."""
session_factory, _ = mock_session_factory
assert service._session_factory.kw["bind"] is sqlite_engine
assert service._session_factory.kw["expire_on_commit"] is False
repository_factory.create_api_workflow_run_repository.assert_called_once_with(service._session_factory)
with service._session_factory() as session:
assert session.scalar(text("SELECT 1")) == 1
with patch("services.workflow_run_service.DifyAPIRepositoryFactory", autospec=True) as mock_factory:
mock_factory.create_api_workflow_run_repository.return_value = mock_workflow_run_repository
service = WorkflowRunService(session_factory)
return service
@pytest.fixture
def workflow_run_service_with_engine(self, mock_session_factory, mock_workflow_run_repository):
"""Create WorkflowRunService instance with Engine input."""
mock_engine = create_autospec(Engine)
session_factory, _ = mock_session_factory
def test_init_with_default_repository_dependencies(sqlite_session_factory: sessionmaker[Session]) -> None:
service = WorkflowRunService(sqlite_session_factory)
with patch("services.workflow_run_service.DifyAPIRepositoryFactory", autospec=True) as mock_factory:
mock_factory.create_api_workflow_run_repository.return_value = mock_workflow_run_repository
service = WorkflowRunService(mock_engine)
return service
# ==================== Initialization Tests ====================
def test_init_with_session_factory(self, mock_session_factory, mock_workflow_run_repository):
"""Test WorkflowRunService initialization with session_factory."""
session_factory, _ = mock_session_factory
with patch("services.workflow_run_service.DifyAPIRepositoryFactory", autospec=True) as mock_factory:
mock_factory.create_api_workflow_run_repository.return_value = mock_workflow_run_repository
service = WorkflowRunService(session_factory)
assert service._session_factory == session_factory
mock_factory.create_api_workflow_run_repository.assert_called_once_with(session_factory)
def test_init_with_engine(self, mock_session_factory, mock_workflow_run_repository):
"""Test WorkflowRunService initialization with Engine (should convert to sessionmaker)."""
mock_engine = create_autospec(Engine)
session_factory, _ = mock_session_factory
with patch("services.workflow_run_service.DifyAPIRepositoryFactory", autospec=True) as mock_factory:
mock_factory.create_api_workflow_run_repository.return_value = mock_workflow_run_repository
with patch(
"services.workflow_run_service.sessionmaker", return_value=session_factory, autospec=True
) as mock_sessionmaker:
service = WorkflowRunService(mock_engine)
mock_sessionmaker.assert_called_once_with(bind=mock_engine, expire_on_commit=False)
assert service._session_factory == session_factory
mock_factory.create_api_workflow_run_repository.assert_called_once_with(session_factory)
def test_init_with_default_dependencies(self, mock_session_factory):
"""Test WorkflowRunService initialization with default dependencies."""
session_factory, _ = mock_session_factory
service = WorkflowRunService(session_factory)
assert service._session_factory == session_factory
assert service._session_factory is sqlite_session_factory
@@ -50,6 +50,9 @@ from services.workflow_service import (
_setup_variable_pool,
)
_LEGACY_FILE_TEMPLATE = "{{#" + ".".join(("sys", "files")) + "#}}"
_USER_INPUT_FILE_TEMPLATE = "{{#" + ".".join(("userinput", "files")) + "#}}"
class TestWorkflowAssociatedDataFactory:
"""
@@ -277,6 +280,31 @@ class TestWorkflowService:
assert result is workflow
def test_get_draft_workflow_persists_legacy_system_files_migration(
self, workflow_service: WorkflowService, sqlite_session: Session
):
# TODO: Remove this compatibility test after the historical workflow migration is complete.
app = TestWorkflowAssociatedDataFactory.create_app()
workflow = TestWorkflowAssociatedDataFactory.create_workflow(
graph={
"nodes": [
{"id": "start", "data": {"type": "start", "variables": []}},
{"id": "answer", "data": {"type": "answer", "answer": _LEGACY_FILE_TEMPLATE}},
],
"edges": [],
}
)
sqlite_session.add(workflow)
sqlite_session.commit()
result = workflow_service.get_draft_workflow(app, session=sqlite_session)
assert result is workflow
assert _LEGACY_FILE_TEMPLATE not in workflow.graph
assert _USER_INPUT_FILE_TEMPLATE in workflow.graph
sqlite_session.expire_all()
assert _USER_INPUT_FILE_TEMPLATE in sqlite_session.get(Workflow, workflow.id).graph
def test_get_draft_workflow_returns_none(self, workflow_service: WorkflowService, sqlite_session: Session):
"""Test get_draft_workflow returns None when no draft exists."""
app = TestWorkflowAssociatedDataFactory.create_app()
@@ -239,6 +239,29 @@ def test__convert_to_llm_node_for_chatbot_simple_chat_model(default_variables: l
assert "{{#start.text_input#}}" in node["data"]["prompt_template"][0]["text"]
def test__convert_to_llm_node_uses_userinput_files_for_vision(default_variables: list[VariableEntity]) -> None:
workflow_converter = WorkflowConverter()
graph = {"nodes": [workflow_converter._convert_to_start_node(default_variables)], "edges": []}
model_config = ModelConfigEntity(provider="openai", model="gpt-4", mode=LLMMode.CHAT.value, parameters={}, stop=[])
prompt_template = PromptTemplateEntity(
prompt_type=PromptTemplateEntity.PromptType.SIMPLE,
simple_prompt_template="Describe the upload",
)
file_upload = MagicMock()
file_upload.image_config = None
node = workflow_converter._convert_to_llm_node(
original_app_mode=AppMode.CHAT,
new_app_mode=AppMode.ADVANCED_CHAT,
model_config=model_config,
graph=graph,
prompt_template=prompt_template,
file_upload=file_upload,
)
assert node["data"]["vision"]["variable_selector"] == ["userinput", "files"]
def test__convert_to_llm_node_for_chatbot_simple_chat_model_with_empty_template(
default_variables: list[VariableEntity],
monkeypatch: pytest.MonkeyPatch,
@@ -34,6 +34,9 @@ from services.workflow_draft_variable_service import (
_model_to_insertion_dict,
)
# TODO: Remove this compatibility key after the historical workflow migration is complete.
_SYSTEM_FILE_OUTPUT_KEY = ".".join((SYSTEM_VARIABLE_NODE_ID, "files"))
SQLITE_MODELS = (Workflow, WorkflowDraftVariable, WorkflowDraftVariableFile, WorkflowNodeExecutionModel)
pytestmark = [
pytest.mark.usefixtures("sqlite_session"),
@@ -127,7 +130,7 @@ class TestDraftVariableSaver:
assert node_id == c.expected_node_id, fail_msg
assert name == c.expected_name, fail_msg
def test_build_variables_from_start_mapping_rebuilds_system_files(self, sqlite_session: Session):
def test_build_variables_from_start_mapping_rebuilds_system_file_variable(self, sqlite_session: Session):
mock_user = MagicMock(spec=Account)
mock_user.id = str(uuid.uuid4())
saver = DraftVariableSaver(
@@ -159,7 +162,7 @@ class TestDraftVariableSaver:
"services.workflow_draft_variable_service.build_file_from_stored_mapping",
return_value=rebuilt_file,
) as rebuild_file:
draft_vars = saver._build_variables_from_start_mapping({"sys.files": [raw_file]})
draft_vars = saver._build_variables_from_start_mapping({_SYSTEM_FILE_OUTPUT_KEY: [raw_file]})
sys_var = draft_vars[0]
assert sys_var.get_value().value[0] == rebuilt_file
@@ -267,7 +270,7 @@ class TestDraftVariableSaver:
def test_start_node_save_persists_sys_timestamp_and_workflow_run_id(
self, mock_batch_upsert, sqlite_session: Session
):
"""Start node should persist common `sys.*` variables, not only `sys.files`."""
"""Start node should persist common system variables."""
mock_user = MagicMock(spec=Account)
mock_user.id = "test-user-id"
mock_user.tenant_id = "test-tenant-id"
@@ -551,7 +554,7 @@ class TestWorkflowDraftVariableService:
# Create mock execution record
mock_execution = Mock(spec=WorkflowNodeExecutionModel)
mock_execution.load_full_outputs.return_value = {"sys.files": "[]"}
mock_execution.load_full_outputs.return_value = {_SYSTEM_FILE_OUTPUT_KEY: "[]"}
# Mock the repository to return the execution record
service._api_node_execution_repo = Mock()
@@ -135,6 +135,19 @@ class TestMCPToolInvoke:
values = {m.message.variable_name: m.message.variable_value for m in var_msgs}
assert values == {"a": 1, "b": "x"}
def test_invoke_yields_json_when_structured_content_has_no_output_schema(self, orm_session: Session) -> None:
tool = _make_mcp_tool()
result = CallToolResult(content=[], structuredContent={"a": 1, "b": "x"})
with patch.object(tool, "invoke_remote_mcp_tool", return_value=result):
messages = list(tool._invoke(session=orm_session, user_id="test_user", tool_parameters={}))
assert len(messages) == 1
msg = messages[0]
assert msg.type == ToolInvokeMessage.MessageType.JSON
assert isinstance(msg.message, ToolInvokeMessage.JsonMessage)
assert msg.message.json_object == {"a": 1, "b": "x"}
class TestMCPToolUsageExtraction:
"""Test usage metadata extraction from MCP tool results."""
Generated
+2 -2
View File
@@ -1281,7 +1281,7 @@ wheels = [
[[package]]
name = "dify-agent"
version = "1.16.0"
version = "1.16.1"
source = { editable = "../dify-agent" }
dependencies = [
{ name = "httpx" },
@@ -1331,7 +1331,7 @@ docs = [
[[package]]
name = "dify-api"
version = "1.16.0"
version = "1.16.1"
source = { virtual = "." }
dependencies = [
{ name = "aliyun-log-python-sdk" },
+1 -1
View File
@@ -71,7 +71,7 @@
"channel": "alpha",
"compat": {
"minDify": "1.16.0",
"maxDify": "1.16.0"
"maxDify": "1.16.1"
},
"release": {
"tagPrefix": "difyctl-v",
+1 -1
View File
@@ -1,6 +1,6 @@
[project]
name = "dify-agent"
version = "1.16.0"
version = "1.16.1"
description = "Add your description here"
readme = "README.md"
requires-python = ">=3.12,<4.0"
+1 -1
View File
@@ -581,7 +581,7 @@ wheels = [
[[package]]
name = "dify-agent"
version = "1.16.0"
version = "1.16.1"
source = { editable = "." }
dependencies = [
{ name = "httpx" },
+7 -7
View File
@@ -220,7 +220,7 @@ services:
# API service
api:
<<: *shared-api-worker-config
image: langgenius/dify-api:1.16.0
image: langgenius/dify-api:1.16.1
environment:
MODE: api
SENTRY_DSN: ${API_SENTRY_DSN:-}
@@ -271,7 +271,7 @@ services:
# WebSocket service for workflow collaboration.
api_websocket:
<<: *shared-api-worker-config
image: langgenius/dify-api:1.16.0
image: langgenius/dify-api:1.16.1
profiles:
- collaboration
environment:
@@ -297,7 +297,7 @@ services:
# The Celery worker for processing all queues (dataset, workflow, mail, etc.)
worker:
<<: *shared-worker-config
image: langgenius/dify-api:1.16.0
image: langgenius/dify-api:1.16.1
environment:
MODE: worker
SENTRY_DSN: ${API_SENTRY_DSN:-}
@@ -347,7 +347,7 @@ services:
# Celery beat for scheduling periodic tasks.
worker_beat:
<<: *shared-worker-beat-config
image: langgenius/dify-api:1.16.0
image: langgenius/dify-api:1.16.1
environment:
MODE: beat
depends_on:
@@ -380,7 +380,7 @@ services:
# Frontend web application.
web:
image: langgenius/dify-web:1.16.0
image: langgenius/dify-web:1.16.1
restart: always
env_file:
- path: ./envs/core-services/web.env
@@ -542,7 +542,7 @@ services:
# on port 3128, which only allows agent_backend /agent-stub/ and the Dify API
# /files/* endpoints (see ssrf_proxy/squid-agent.conf.template).
local_sandbox:
image: langgenius/dify-agent-local-sandbox:1.16.0
image: langgenius/dify-agent-local-sandbox:1.16.1
restart: always
env_file:
- path: ./envs/core-services/local-sandbox.env
@@ -651,7 +651,7 @@ services:
# Dify Agent backend service.
agent_backend:
image: langgenius/dify-agent-backend:1.16.0
image: langgenius/dify-agent-backend:1.16.1
restart: always
env_file:
- path: ./envs/core-services/dify-agent.env
+7 -7
View File
@@ -226,7 +226,7 @@ services:
# API service
api:
<<: *shared-api-worker-config
image: langgenius/dify-api:1.16.0
image: langgenius/dify-api:1.16.1
environment:
MODE: api
SENTRY_DSN: ${API_SENTRY_DSN:-}
@@ -277,7 +277,7 @@ services:
# WebSocket service for workflow collaboration.
api_websocket:
<<: *shared-api-worker-config
image: langgenius/dify-api:1.16.0
image: langgenius/dify-api:1.16.1
profiles:
- collaboration
environment:
@@ -303,7 +303,7 @@ services:
# The Celery worker for processing all queues (dataset, workflow, mail, etc.)
worker:
<<: *shared-worker-config
image: langgenius/dify-api:1.16.0
image: langgenius/dify-api:1.16.1
environment:
MODE: worker
SENTRY_DSN: ${API_SENTRY_DSN:-}
@@ -353,7 +353,7 @@ services:
# Celery beat for scheduling periodic tasks.
worker_beat:
<<: *shared-worker-beat-config
image: langgenius/dify-api:1.16.0
image: langgenius/dify-api:1.16.1
environment:
MODE: beat
depends_on:
@@ -386,7 +386,7 @@ services:
# Frontend web application.
web:
image: langgenius/dify-web:1.16.0
image: langgenius/dify-web:1.16.1
restart: always
env_file:
- path: ./envs/core-services/web.env
@@ -548,7 +548,7 @@ services:
# on port 3128, which only allows agent_backend /agent-stub/ and the Dify API
# /files/* endpoints (see ssrf_proxy/squid-agent.conf.template).
local_sandbox:
image: langgenius/dify-agent-local-sandbox:1.16.0
image: langgenius/dify-agent-local-sandbox:1.16.1
restart: always
env_file:
- path: ./envs/core-services/local-sandbox.env
@@ -657,7 +657,7 @@ services:
# Dify Agent backend service.
agent_backend:
image: langgenius/dify-agent-backend:1.16.0
image: langgenius/dify-agent-backend:1.16.1
restart: always
env_file:
- path: ./envs/core-services/dify-agent.env
-18
View File
@@ -3883,11 +3883,6 @@
"count": 2
}
},
"web/app/components/workflow-app/hooks/use-workflow-template.ts": {
"typescript/no-explicit-any": {
"count": 2
}
},
"web/app/components/workflow-app/store/workflow/workflow-slice.ts": {
"typescript/no-explicit-any": {
"count": 2
@@ -4217,9 +4212,6 @@
"web/app/components/workflow/nodes/_base/components/memory-config.tsx": {
"no-restricted-imports": {
"count": 1
},
"unicorn/prefer-number-properties": {
"count": 1
}
},
"web/app/components/workflow/nodes/_base/components/next-step/operator.tsx": {
@@ -5071,21 +5063,11 @@
"count": 8
}
},
"web/app/components/workflow/nodes/start/panel.tsx": {
"typescript/no-explicit-any": {
"count": 2
}
},
"web/app/components/workflow/nodes/start/use-config.ts": {
"typescript/no-explicit-any": {
"count": 1
}
},
"web/app/components/workflow/nodes/start/use-single-run-form-params.ts": {
"typescript/no-explicit-any": {
"count": 3
}
},
"web/app/components/workflow/nodes/template-transform/use-config.ts": {
"typescript/no-explicit-any": {
"count": 4
@@ -15,7 +15,7 @@ export const mockedWorkflowProcess = {
predecessor_node_id: null,
inputs: {
'sys.query': 'hi',
'sys.files': [],
'userinput.files': [],
'sys.conversation_id': '92ce0a3e-8f15-43d1-b31d-32716c4b10a7',
'sys.user_id': 'fbff43f9-d5a4-4e85-b63b-d3a91d806c6f',
'sys.dialogue_count': 1,
@@ -26,7 +26,7 @@ export const mockedWorkflowProcess = {
process_data: null,
outputs: {
'sys.query': 'hi',
'sys.files': [],
'userinput.files': [],
'sys.conversation_id': '92ce0a3e-8f15-43d1-b31d-32716c4b10a7',
'sys.user_id': 'fbff43f9-d5a4-4e85-b63b-d3a91d806c6f',
'sys.dialogue_count': 1,
@@ -494,7 +494,7 @@ describe('ComponentPicker (component-picker-block/index.tsx)', () => {
expect(dispatchSpy).not.toHaveBeenCalled()
})
it('handles workflow variable selection for nested fields: sys.query, sys.files, and normal paths', async () => {
it('handles workflow variable selection for nested fields: built-in inputs and normal paths', async () => {
const captures: Captures = { editor: null, eventEmitter: null }
const user = userEvent.setup()
@@ -503,7 +503,7 @@ describe('ComponentPicker (component-picker-block/index.tsx)', () => {
makeWorkflowNodeVar('sys.query', VarType.object, [
makeWorkflowNodeVar('q', VarType.string),
]),
makeWorkflowNodeVar('sys.files', VarType.object, [
makeWorkflowNodeVar('userinput.files', VarType.object, [
makeWorkflowNodeVar('f', VarType.string),
]),
makeWorkflowNodeVar('output', VarType.object, [makeWorkflowNodeVar('x', VarType.string)]),
@@ -545,8 +545,10 @@ describe('ComponentPicker (component-picker-block/index.tsx)', () => {
expect(dispatchSpy).toHaveBeenCalledWith(INSERT_WORKFLOW_VARIABLE_BLOCK_COMMAND, ['sys.query'])
await waitFor(() => expect(readEditorText(editor)).not.toContain('{'))
await openPickerAndSelectField('sys.files', 'f')
expect(dispatchSpy).toHaveBeenCalledWith(INSERT_WORKFLOW_VARIABLE_BLOCK_COMMAND, ['sys.files'])
await openPickerAndSelectField('userinput.files', 'f')
expect(dispatchSpy).toHaveBeenCalledWith(INSERT_WORKFLOW_VARIABLE_BLOCK_COMMAND, [
'userinput.files',
])
await waitFor(() => expect(readEditorText(editor)).not.toContain('{'))
await openPickerAndSelectField('output', 'x')
@@ -191,6 +191,7 @@ const ComponentPicker = ({
if (needRemove) needRemove.remove()
})
const isFlat = variables.length === 1
const builtInVariable = variables[1]
if (isFlat) {
const varName = variables[0]
if (varName === 'current')
@@ -198,8 +199,8 @@ const ComponentPicker = ({
else if (varName === 'error_message')
editor.dispatchCommand(INSERT_ERROR_MESSAGE_BLOCK_COMMAND, null)
else if (varName === 'last_run') editor.dispatchCommand(INSERT_LAST_RUN_BLOCK_COMMAND, null)
} else if (variables[1] === 'sys.query' || variables[1] === 'sys.files') {
editor.dispatchCommand(INSERT_WORKFLOW_VARIABLE_BLOCK_COMMAND, [variables[1]])
} else if (builtInVariable && ['sys.query', 'userinput.files'].includes(builtInVariable)) {
editor.dispatchCommand(INSERT_WORKFLOW_VARIABLE_BLOCK_COMMAND, [builtInVariable])
} else {
editor.dispatchCommand(INSERT_WORKFLOW_VARIABLE_BLOCK_COMMAND, variables)
}
@@ -437,7 +437,7 @@ Chat applications support session persistence, allowing previous chat history to
"id": "a4959eb4-c852-4e0c-ac7a-348233f7f345",
"workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
"inputs": {
"sys.files": [],
"userinput.files": [],
"sys.user_id": "abc-123",
"sys.app_id": "d1074979-f67e-4114-8691-e35878df9a89",
"sys.workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
@@ -481,7 +481,7 @@ Chat applications support session persistence, allowing previous chat history to
"index": 1,
"predecessor_node_id": null,
"inputs": {
"sys.files": [],
"userinput.files": [],
"sys.user_id": "abc-123",
"sys.app_id": "d1074979-f67e-4114-8691-e35878df9a89",
"sys.workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
@@ -492,7 +492,7 @@ Chat applications support session persistence, allowing previous chat history to
"process_data": {},
"process_data_truncated": false,
"outputs": {
"sys.files": [],
"userinput.files": [],
"sys.user_id": "abc-123",
"sys.app_id": "d1074979-f67e-4114-8691-e35878df9a89",
"sys.workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
@@ -1055,7 +1055,7 @@ Chat applications support session persistence, allowing previous chat history to
```streaming {{ title: 'Response' }}
event: ping
data: {"event":"workflow_started","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"sys.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","sys.timestamp":1776087863},"created_at":1776087863,"reason":"initial"}}
data: {"event":"workflow_started","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"userinput.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","sys.timestamp":1776087863},"created_at":1776087863,"reason":"initial"}}
data: {"event":"node_started","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"id":"b552d685-1119-4e6a-9a81-e91a23e5324b","node_id":"1775717266623","node_type":"start","title":"User Input","index":1,"predecessor_node_id":null,"inputs":null,"created_at":1776087863,"extras":{},"iteration_id":null,"loop_id":null}}
@@ -1067,7 +1067,7 @@ Chat applications support session persistence, allowing previous chat history to
data: {"event":"workflow_paused","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","paused_nodes":["1775717346519"],"outputs":{},"reasons":[{"form_id":"019d8716-0fde-75da-8207-1458ccde76e5","form_content":"this is form 1:\n{{#$output.some_field#}}\n","inputs":[{"type":"paragraph","output_variable_name":"some_field","default":{"type":"variable","selector":["sys","workflow_run_id"],"value":""}}],"actions":[{"id":"approve","title":"YES","button_style":"default"},{"id":"reject","title":"NO","button_style":"default"}],"display_in_ui":true,"node_id":"1775717346519","node_title":"Human Input","resolved_default_values":{"some_field":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c"},"form_token":"n7hFG4ZDYdGcgZ5VDc7EGM","type":"human_input_required"}],"status":"paused","created_at":1776087863,"elapsed_time":0.0,"total_tokens":0,"total_steps":2}}
data: {"event":"workflow_started","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"sys.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c"},"created_at":1776087877,"reason":"resumption"}}
data: {"event":"workflow_started","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"userinput.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c"},"created_at":1776087877,"reason":"resumption"}}
data: {"event":"node_started","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","data":{"id":"8d7e8e01-5159-4089-a4b6-3aa394992cc2","node_id":"1775717346519","node_type":"human-input","title":"Human Input","index":1,"predecessor_node_id":null,"inputs":null,"inputs_truncated":false,"created_at":1776087877,"extras":{},"iteration_id":null,"loop_id":null,"agent_strategy":null}}
@@ -437,7 +437,7 @@ import { WorkflowVersionApiUpgradeNotice } from '../workflow-version-api-upgrade
"id": "a4959eb4-c852-4e0c-ac7a-348233f7f345",
"workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
"inputs": {
"sys.files": [],
"userinput.files": [],
"sys.user_id": "abc-123",
"sys.app_id": "d1074979-f67e-4114-8691-e35878df9a89",
"sys.workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
@@ -481,7 +481,7 @@ import { WorkflowVersionApiUpgradeNotice } from '../workflow-version-api-upgrade
"index": 1,
"predecessor_node_id": null,
"inputs": {
"sys.files": [],
"userinput.files": [],
"sys.user_id": "abc-123",
"sys.app_id": "d1074979-f67e-4114-8691-e35878df9a89",
"sys.workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
@@ -492,7 +492,7 @@ import { WorkflowVersionApiUpgradeNotice } from '../workflow-version-api-upgrade
"process_data": {},
"process_data_truncated": false,
"outputs": {
"sys.files": [],
"userinput.files": [],
"sys.user_id": "abc-123",
"sys.app_id": "d1074979-f67e-4114-8691-e35878df9a89",
"sys.workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
@@ -1056,7 +1056,7 @@ import { WorkflowVersionApiUpgradeNotice } from '../workflow-version-api-upgrade
```streaming {{ title: '応答' }}
event: ping
data: {"event":"workflow_started","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"sys.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","sys.timestamp":1776087863},"created_at":1776087863,"reason":"initial"}}
data: {"event":"workflow_started","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"userinput.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","sys.timestamp":1776087863},"created_at":1776087863,"reason":"initial"}}
data: {"event":"node_started","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"id":"b552d685-1119-4e6a-9a81-e91a23e5324b","node_id":"1775717266623","node_type":"start","title":"User Input","index":1,"predecessor_node_id":null,"inputs":null,"created_at":1776087863,"extras":{},"iteration_id":null,"loop_id":null}}
@@ -1068,7 +1068,7 @@ import { WorkflowVersionApiUpgradeNotice } from '../workflow-version-api-upgrade
data: {"event":"workflow_paused","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","paused_nodes":["1775717346519"],"outputs":{},"reasons":[{"form_id":"019d8716-0fde-75da-8207-1458ccde76e5","form_content":"this is form 1:\n{{#$output.some_field#}}\n","inputs":[{"type":"paragraph","output_variable_name":"some_field","default":{"type":"variable","selector":["sys","workflow_run_id"],"value":""}}],"actions":[{"id":"approve","title":"YES","button_style":"default"},{"id":"reject","title":"NO","button_style":"default"}],"display_in_ui":true,"node_id":"1775717346519","node_title":"Human Input","resolved_default_values":{"some_field":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c"},"form_token":"n7hFG4ZDYdGcgZ5VDc7EGM","type":"human_input_required"}],"status":"paused","created_at":1776087863,"elapsed_time":0.0,"total_tokens":0,"total_steps":2}}
data: {"event":"workflow_started","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"sys.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c"},"created_at":1776087877,"reason":"resumption"}}
data: {"event":"workflow_started","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"userinput.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c"},"created_at":1776087877,"reason":"resumption"}}
data: {"event":"node_started","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","data":{"id":"8d7e8e01-5159-4089-a4b6-3aa394992cc2","node_id":"1775717346519","node_type":"human-input","title":"Human Input","index":1,"predecessor_node_id":null,"inputs":null,"inputs_truncated":false,"created_at":1776087877,"extras":{},"iteration_id":null,"loop_id":null,"agent_strategy":null}}
@@ -436,7 +436,7 @@ import { WorkflowVersionApiUpgradeNotice } from '../workflow-version-api-upgrade
"id": "a4959eb4-c852-4e0c-ac7a-348233f7f345",
"workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
"inputs": {
"sys.files": [],
"userinput.files": [],
"sys.user_id": "abc-123",
"sys.app_id": "d1074979-f67e-4114-8691-e35878df9a89",
"sys.workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
@@ -480,7 +480,7 @@ import { WorkflowVersionApiUpgradeNotice } from '../workflow-version-api-upgrade
"index": 1,
"predecessor_node_id": null,
"inputs": {
"sys.files": [],
"userinput.files": [],
"sys.user_id": "abc-123",
"sys.app_id": "d1074979-f67e-4114-8691-e35878df9a89",
"sys.workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
@@ -491,7 +491,7 @@ import { WorkflowVersionApiUpgradeNotice } from '../workflow-version-api-upgrade
"process_data": {},
"process_data_truncated": false,
"outputs": {
"sys.files": [],
"userinput.files": [],
"sys.user_id": "abc-123",
"sys.app_id": "d1074979-f67e-4114-8691-e35878df9a89",
"sys.workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
@@ -1049,7 +1049,7 @@ import { WorkflowVersionApiUpgradeNotice } from '../workflow-version-api-upgrade
```streaming {{ title: 'Response' }}
event: ping
data: {"event":"workflow_started","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"sys.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","sys.timestamp":1776087863},"created_at":1776087863,"reason":"initial"}}
data: {"event":"workflow_started","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"userinput.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","sys.timestamp":1776087863},"created_at":1776087863,"reason":"initial"}}
data: {"event":"node_started","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"id":"b552d685-1119-4e6a-9a81-e91a23e5324b","node_id":"1775717266623","node_type":"start","title":"User Input","index":1,"predecessor_node_id":null,"inputs":null,"created_at":1776087863,"extras":{},"iteration_id":null,"loop_id":null}}
@@ -1061,7 +1061,7 @@ import { WorkflowVersionApiUpgradeNotice } from '../workflow-version-api-upgrade
data: {"event":"workflow_paused","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","paused_nodes":["1775717346519"],"outputs":{},"reasons":[{"form_id":"019d8716-0fde-75da-8207-1458ccde76e5","form_content":"this is form 1:\n{{#$output.some_field#}}\n","inputs":[{"type":"paragraph","output_variable_name":"some_field","default":{"type":"variable","selector":["sys","workflow_run_id"],"value":""}}],"actions":[{"id":"approve","title":"YES","button_style":"default"},{"id":"reject","title":"NO","button_style":"default"}],"display_in_ui":true,"node_id":"1775717346519","node_title":"Human Input","resolved_default_values":{"some_field":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c"},"form_token":"n7hFG4ZDYdGcgZ5VDc7EGM","type":"human_input_required"}],"status":"paused","created_at":1776087863,"elapsed_time":0.0,"total_tokens":0,"total_steps":2}}
data: {"event":"workflow_started","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"sys.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c"},"created_at":1776087877,"reason":"resumption"}}
data: {"event":"workflow_started","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"userinput.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c"},"created_at":1776087877,"reason":"resumption"}}
data: {"event":"node_started","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","data":{"id":"8d7e8e01-5159-4089-a4b6-3aa394992cc2","node_id":"1775717346519","node_type":"human-input","title":"Human Input","index":1,"predecessor_node_id":null,"inputs":null,"inputs_truncated":false,"created_at":1776087877,"extras":{},"iteration_id":null,"loop_id":null,"agent_strategy":null}}
@@ -354,7 +354,7 @@ Workflow applications offers non-session support and is ideal for translation, a
"id": "a4959eb4-c852-4e0c-ac7a-348233f7f345",
"workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
"inputs": {
"sys.files": [],
"userinput.files": [],
"sys.user_id": "abc-123",
"sys.app_id": "d1074979-f67e-4114-8691-e35878df9a89",
"sys.workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
@@ -398,7 +398,7 @@ Workflow applications offers non-session support and is ideal for translation, a
"index": 1,
"predecessor_node_id": null,
"inputs": {
"sys.files": [],
"userinput.files": [],
"sys.user_id": "abc-123",
"sys.app_id": "d1074979-f67e-4114-8691-e35878df9a89",
"sys.workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
@@ -409,7 +409,7 @@ Workflow applications offers non-session support and is ideal for translation, a
"process_data": {},
"process_data_truncated": false,
"outputs": {
"sys.files": [],
"userinput.files": [],
"sys.user_id": "abc-123",
"sys.app_id": "d1074979-f67e-4114-8691-e35878df9a89",
"sys.workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
@@ -942,7 +942,7 @@ Workflow applications offers non-session support and is ideal for translation, a
"id": "b1ad3277-089e-42c6-9dff-6820d94fbc76",
"workflow_id": "19eff89f-ec03-4f75-b0fc-897e7effea02",
"status": "succeeded",
"inputs": "{\"sys.files\": [], \"sys.user_id\": \"abc-123\"}",
"inputs": "{\"userinput.files\": [], \"sys.user_id\": \"abc-123\"}",
"outputs": null,
"error": null,
"total_steps": 3,
@@ -1154,7 +1154,7 @@ Workflow applications offers non-session support and is ideal for translation, a
```streaming {{ title: 'Response' }}
event: ping
data: {"event":"workflow_started","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"sys.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","sys.timestamp":1776087863},"created_at":1776087863,"reason":"initial"}}
data: {"event":"workflow_started","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"userinput.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","sys.timestamp":1776087863},"created_at":1776087863,"reason":"initial"}}
data: {"event":"node_started","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"id":"b552d685-1119-4e6a-9a81-e91a23e5324b","node_id":"1775717266623","node_type":"start","title":"User Input","index":1,"predecessor_node_id":null,"inputs":null,"created_at":1776087863,"extras":{},"iteration_id":null,"loop_id":null}}
@@ -1166,7 +1166,7 @@ Workflow applications offers non-session support and is ideal for translation, a
data: {"event":"workflow_paused","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","paused_nodes":["1775717346519"],"outputs":{},"reasons":[{"form_id":"019d8716-0fde-75da-8207-1458ccde76e5","form_content":"this is form 1:\n{{#$output.some_field#}}\n","inputs":[{"type":"paragraph","output_variable_name":"some_field","default":{"type":"variable","selector":["sys","workflow_run_id"],"value":""}}],"actions":[{"id":"approve","title":"YES","button_style":"default"},{"id":"reject","title":"NO","button_style":"default"}],"display_in_ui":true,"node_id":"1775717346519","node_title":"Human Input","resolved_default_values":{"some_field":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c"},"form_token":"n7hFG4ZDYdGcgZ5VDc7EGM","type":"human_input_required"}],"status":"paused","created_at":1776087863,"elapsed_time":0.0,"total_tokens":0,"total_steps":2}}
data: {"event":"workflow_started","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"sys.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c"},"created_at":1776087877,"reason":"resumption"}}
data: {"event":"workflow_started","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"userinput.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c"},"created_at":1776087877,"reason":"resumption"}}
data: {"event":"node_started","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","data":{"id":"8d7e8e01-5159-4089-a4b6-3aa394992cc2","node_id":"1775717346519","node_type":"human-input","title":"Human Input","index":1,"predecessor_node_id":null,"inputs":null,"inputs_truncated":false,"created_at":1776087877,"extras":{},"iteration_id":null,"loop_id":null,"agent_strategy":null}}
@@ -354,7 +354,7 @@ import { WorkflowVersionApiContent, WorkflowVersionApiUpgradeNotice } from '../w
"id": "a4959eb4-c852-4e0c-ac7a-348233f7f345",
"workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
"inputs": {
"sys.files": [],
"userinput.files": [],
"sys.user_id": "abc-123",
"sys.app_id": "d1074979-f67e-4114-8691-e35878df9a89",
"sys.workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
@@ -398,7 +398,7 @@ import { WorkflowVersionApiContent, WorkflowVersionApiUpgradeNotice } from '../w
"index": 1,
"predecessor_node_id": null,
"inputs": {
"sys.files": [],
"userinput.files": [],
"sys.user_id": "abc-123",
"sys.app_id": "d1074979-f67e-4114-8691-e35878df9a89",
"sys.workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
@@ -409,7 +409,7 @@ import { WorkflowVersionApiContent, WorkflowVersionApiUpgradeNotice } from '../w
"process_data": {},
"process_data_truncated": false,
"outputs": {
"sys.files": [],
"userinput.files": [],
"sys.user_id": "abc-123",
"sys.app_id": "d1074979-f67e-4114-8691-e35878df9a89",
"sys.workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
@@ -937,7 +937,7 @@ import { WorkflowVersionApiContent, WorkflowVersionApiUpgradeNotice } from '../w
"id": "b1ad3277-089e-42c6-9dff-6820d94fbc76",
"workflow_id": "19eff89f-ec03-4f75-b0fc-897e7effea02",
"status": "succeeded",
"inputs": "{\"sys.files\": [], \"sys.user_id\": \"abc-123\"}",
"inputs": "{\"userinput.files\": [], \"sys.user_id\": \"abc-123\"}",
"outputs": null,
"error": null,
"total_steps": 3,
@@ -1149,7 +1149,7 @@ import { WorkflowVersionApiContent, WorkflowVersionApiUpgradeNotice } from '../w
```streaming {{ title: '応答' }}
event: ping
data: {"event":"workflow_started","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"sys.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","sys.timestamp":1776087863},"created_at":1776087863,"reason":"initial"}}
data: {"event":"workflow_started","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"userinput.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","sys.timestamp":1776087863},"created_at":1776087863,"reason":"initial"}}
data: {"event":"node_started","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"id":"b552d685-1119-4e6a-9a81-e91a23e5324b","node_id":"1775717266623","node_type":"start","title":"User Input","index":1,"predecessor_node_id":null,"inputs":null,"created_at":1776087863,"extras":{},"iteration_id":null,"loop_id":null}}
@@ -1161,7 +1161,7 @@ import { WorkflowVersionApiContent, WorkflowVersionApiUpgradeNotice } from '../w
data: {"event":"workflow_paused","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","paused_nodes":["1775717346519"],"outputs":{},"reasons":[{"form_id":"019d8716-0fde-75da-8207-1458ccde76e5","form_content":"this is form 1:\n{{#$output.some_field#}}\n","inputs":[{"type":"paragraph","output_variable_name":"some_field","default":{"type":"variable","selector":["sys","workflow_run_id"],"value":""}}],"actions":[{"id":"approve","title":"YES","button_style":"default"},{"id":"reject","title":"NO","button_style":"default"}],"display_in_ui":true,"node_id":"1775717346519","node_title":"Human Input","resolved_default_values":{"some_field":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c"},"form_token":"n7hFG4ZDYdGcgZ5VDc7EGM","type":"human_input_required"}],"status":"paused","created_at":1776087863,"elapsed_time":0.0,"total_tokens":0,"total_steps":2}}
data: {"event":"workflow_started","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"sys.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c"},"created_at":1776087877,"reason":"resumption"}}
data: {"event":"workflow_started","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"userinput.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c"},"created_at":1776087877,"reason":"resumption"}}
data: {"event":"node_started","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","data":{"id":"8d7e8e01-5159-4089-a4b6-3aa394992cc2","node_id":"1775717346519","node_type":"human-input","title":"Human Input","index":1,"predecessor_node_id":null,"inputs":null,"inputs_truncated":false,"created_at":1776087877,"extras":{},"iteration_id":null,"loop_id":null,"agent_strategy":null}}
@@ -344,7 +344,7 @@ Workflow 应用无会话支持,适合用于翻译/文章写作/总结 AI 等
"id": "a4959eb4-c852-4e0c-ac7a-348233f7f345",
"workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
"inputs": {
"sys.files": [],
"userinput.files": [],
"sys.user_id": "abc-123",
"sys.app_id": "d1074979-f67e-4114-8691-e35878df9a89",
"sys.workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
@@ -388,7 +388,7 @@ Workflow 应用无会话支持,适合用于翻译/文章写作/总结 AI 等
"index": 1,
"predecessor_node_id": null,
"inputs": {
"sys.files": [],
"userinput.files": [],
"sys.user_id": "abc-123",
"sys.app_id": "d1074979-f67e-4114-8691-e35878df9a89",
"sys.workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
@@ -399,7 +399,7 @@ Workflow 应用无会话支持,适合用于翻译/文章写作/总结 AI 等
"process_data": {},
"process_data_truncated": false,
"outputs": {
"sys.files": [],
"userinput.files": [],
"sys.user_id": "abc-123",
"sys.app_id": "d1074979-f67e-4114-8691-e35878df9a89",
"sys.workflow_id": "e46514f1-c008-41ff-94b0-4f33d4b97d36",
@@ -930,7 +930,7 @@ Workflow 应用无会话支持,适合用于翻译/文章写作/总结 AI 等
"id": "b1ad3277-089e-42c6-9dff-6820d94fbc76",
"workflow_id": "19eff89f-ec03-4f75-b0fc-897e7effea02",
"status": "succeeded",
"inputs": "{\"sys.files\": [], \"sys.user_id\": \"abc-123\"}",
"inputs": "{\"userinput.files\": [], \"sys.user_id\": \"abc-123\"}",
"outputs": null,
"error": null,
"total_steps": 3,
@@ -1142,7 +1142,7 @@ Workflow 应用无会话支持,适合用于翻译/文章写作/总结 AI 等
```streaming {{ title: 'Response' }}
event: ping
data: {"event":"workflow_started","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"sys.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","sys.timestamp":1776087863},"created_at":1776087863,"reason":"initial"}}
data: {"event":"workflow_started","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"userinput.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","sys.timestamp":1776087863},"created_at":1776087863,"reason":"initial"}}
data: {"event":"node_started","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"id":"b552d685-1119-4e6a-9a81-e91a23e5324b","node_id":"1775717266623","node_type":"start","title":"User Input","index":1,"predecessor_node_id":null,"inputs":null,"created_at":1776087863,"extras":{},"iteration_id":null,"loop_id":null}}
@@ -1154,7 +1154,7 @@ Workflow 应用无会话支持,适合用于翻译/文章写作/总结 AI 等
data: {"event":"workflow_paused","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","data":{"workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","paused_nodes":["1775717346519"],"outputs":{},"reasons":[{"form_id":"019d8716-0fde-75da-8207-1458ccde76e5","form_content":"this is form 1:\n{{#$output.some_field#}}\n","inputs":[{"type":"paragraph","output_variable_name":"some_field","default":{"type":"variable","selector":["sys","workflow_run_id"],"value":""}}],"actions":[{"id":"approve","title":"YES","button_style":"default"},{"id":"reject","title":"NO","button_style":"default"}],"display_in_ui":true,"node_id":"1775717346519","node_title":"Human Input","resolved_default_values":{"some_field":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c"},"form_token":"n7hFG4ZDYdGcgZ5VDc7EGM","type":"human_input_required"}],"status":"paused","created_at":1776087863,"elapsed_time":0.0,"total_tokens":0,"total_steps":2}}
data: {"event":"workflow_started","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"sys.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c"},"created_at":1776087877,"reason":"resumption"}}
data: {"event":"workflow_started","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","data":{"id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","inputs":{"userinput.files":[],"sys.user_id":"abc-123","sys.app_id":"d1074979-f67e-4114-8691-e35878df9a89","sys.workflow_id":"e46514f1-c008-41ff-94b0-4f33d4b97d36","sys.workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c"},"created_at":1776087877,"reason":"resumption"}}
data: {"event":"node_started","workflow_run_id":"5d7ef348-e1c1-4f6d-bb9b-62cc2fb2ef3c","task_id":"1784c3dd-20eb-4919-bd5d-a8d800b74ada","data":{"id":"8d7e8e01-5159-4089-a4b6-3aa394992cc2","node_id":"1775717346519","node_type":"human-input","title":"Human Input","index":1,"predecessor_node_id":null,"inputs":null,"inputs_truncated":false,"created_at":1776087877,"extras":{},"iteration_id":null,"loop_id":null,"agent_strategy":null}}
@@ -211,6 +211,18 @@ export default function AccountSetting({
return (
<MenuDialog show onClose={handleClose}>
<div className="fixed top-6 right-6 z-20 flex shrink-0 flex-col items-center">
<Button
variant="tertiary"
size="large"
className="px-2"
aria-label={t(($) => $['operation.close'], { ns: 'common' })}
onClick={handleClose}
>
<span className="i-ri-close-line size-5" />
</Button>
<div className="mt-1 system-2xs-medium-uppercase text-text-tertiary">ESC</div>
</div>
<div className="flex h-screen w-full max-w-full pl-0 sm:pl-[232px]">
<div className="flex w-[44px] shrink-0 flex-col pr-6 pl-4 sm:w-[224px]">
<div className="mt-6 mb-8 flex h-[38px] items-center px-3 title-2xl-semi-bold whitespace-nowrap text-text-primary">
@@ -275,18 +287,6 @@ export default function AccountSetting({
</div>
)}
</div>
<div className="fixed top-6 right-6 flex shrink-0 flex-col items-center">
<Button
variant="tertiary"
size="large"
className="px-2"
aria-label={t(($) => $['operation.close'], { ns: 'common' })}
onClick={handleClose}
>
<span className="i-ri-close-line size-5" />
</Button>
<div className="mt-1 system-2xs-medium-uppercase text-text-tertiary">ESC</div>
</div>
</div>
<div className="max-w-full min-w-0 px-4 pt-6 sm:px-8">
{activeMenu === ACCOUNT_SETTING_TAB.PROVIDER && (
@@ -272,8 +272,10 @@ function Popup({
return (
<ModelSelectorPopupFrame>
<ModelSelectorSearchHeader inputValue={inputValue} onInputValueChange={onInputValueChange} />
{showCreditsExhaustedAlert && <CreditsExhaustedAlert hasApiKeyFallback={hasApiKeyFallback} />}
<ModelSelectorScrollBody label={t(($) => $['modelProvider.models'], { ns: 'common' })}>
{showCreditsExhaustedAlert && (
<CreditsExhaustedAlert hasApiKeyFallback={hasApiKeyFallback} />
)}
<ComboboxList className="max-h-none overflow-visible p-0">
<div className="pb-1">
{filteredModelList.map((model) => (
@@ -4,125 +4,101 @@ import { InstallationScope } from '@/features/system-features/constants'
import { renderHookWithConsoleQuery as renderHook } from '@/test/console/query-data'
import { pluginInstallLimit } from '../use-install-plugin-limit'
type PluginInstallCandidate = Parameters<typeof pluginInstallLimit>[0]
type SystemFeatures = Parameters<typeof pluginInstallLimit>[1]
const basePlugin = {
from: 'marketplace' as const,
verification: { authorized_category: 'langgenius' },
} satisfies PluginInstallCandidate
function makeSystemFeatures(
scope: PluginInstallationScope,
restrictToMarketplaceOnly = false,
): SystemFeatures {
return {
plugin_installation_permission: {
restrict_to_marketplace_only: restrictToMarketplaceOnly,
plugin_installation_scope: scope,
},
}
}
describe('pluginInstallLimit', () => {
it('should allow all plugins when scope is ALL', () => {
const features = {
plugin_installation_permission: {
restrict_to_marketplace_only: false,
plugin_installation_scope: InstallationScope.ALL,
},
}
const features = makeSystemFeatures(InstallationScope.ALL)
expect(pluginInstallLimit(basePlugin as never, features as never).canInstall).toBe(true)
expect(pluginInstallLimit(basePlugin, features).canInstall).toBe(true)
})
it('should deny all plugins when scope is NONE', () => {
const features = {
plugin_installation_permission: {
restrict_to_marketplace_only: false,
plugin_installation_scope: InstallationScope.NONE,
},
}
const features = makeSystemFeatures(InstallationScope.NONE)
expect(pluginInstallLimit(basePlugin as never, features as never).canInstall).toBe(false)
expect(pluginInstallLimit(basePlugin, features).canInstall).toBe(false)
})
it('should allow langgenius plugins when scope is OFFICIAL_ONLY', () => {
const features = {
plugin_installation_permission: {
restrict_to_marketplace_only: false,
plugin_installation_scope: InstallationScope.OFFICIAL_ONLY,
},
}
const features = makeSystemFeatures(InstallationScope.OFFICIAL_ONLY)
expect(pluginInstallLimit(basePlugin as never, features as never).canInstall).toBe(true)
expect(pluginInstallLimit(basePlugin, features).canInstall).toBe(true)
})
it('should deny non-official plugins when scope is OFFICIAL_ONLY', () => {
const features = {
plugin_installation_permission: {
restrict_to_marketplace_only: false,
plugin_installation_scope: InstallationScope.OFFICIAL_ONLY,
},
}
const plugin = { ...basePlugin, verification: { authorized_category: 'community' } }
const features = makeSystemFeatures(InstallationScope.OFFICIAL_ONLY)
const plugin = {
...basePlugin,
verification: { authorized_category: 'community' as const },
} satisfies PluginInstallCandidate
expect(pluginInstallLimit(plugin as never, features as never).canInstall).toBe(false)
expect(pluginInstallLimit(plugin, features).canInstall).toBe(false)
})
it('should allow partner plugins when scope is OFFICIAL_AND_PARTNER', () => {
const features = {
plugin_installation_permission: {
restrict_to_marketplace_only: false,
plugin_installation_scope: InstallationScope.OFFICIAL_AND_PARTNER,
},
}
const plugin = { ...basePlugin, verification: { authorized_category: 'partner' } }
const features = makeSystemFeatures(InstallationScope.OFFICIAL_AND_PARTNER)
const plugin = {
...basePlugin,
verification: { authorized_category: 'partner' as const },
} satisfies PluginInstallCandidate
expect(pluginInstallLimit(plugin as never, features as never).canInstall).toBe(true)
expect(pluginInstallLimit(plugin, features).canInstall).toBe(true)
})
it('should deny github plugins when restrict_to_marketplace_only is true', () => {
const features = {
plugin_installation_permission: {
restrict_to_marketplace_only: true,
plugin_installation_scope: InstallationScope.ALL,
},
}
const plugin = { ...basePlugin, from: 'github' as const }
const features = makeSystemFeatures(InstallationScope.ALL, true)
const plugin = { ...basePlugin, from: 'github' as const } satisfies PluginInstallCandidate
expect(pluginInstallLimit(plugin as never, features as never).canInstall).toBe(false)
expect(pluginInstallLimit(plugin, features).canInstall).toBe(false)
})
it('should deny package plugins when restrict_to_marketplace_only is true', () => {
const features = {
plugin_installation_permission: {
restrict_to_marketplace_only: true,
plugin_installation_scope: InstallationScope.ALL,
},
}
const plugin = { ...basePlugin, from: 'package' as const }
const features = makeSystemFeatures(InstallationScope.ALL, true)
const plugin = { ...basePlugin, from: 'package' as const } satisfies PluginInstallCandidate
expect(pluginInstallLimit(plugin as never, features as never).canInstall).toBe(false)
expect(pluginInstallLimit(plugin, features).canInstall).toBe(false)
})
it('should allow marketplace plugins even when restrict_to_marketplace_only is true', () => {
const features = {
plugin_installation_permission: {
restrict_to_marketplace_only: true,
plugin_installation_scope: InstallationScope.ALL,
},
}
const features = makeSystemFeatures(InstallationScope.ALL, true)
expect(pluginInstallLimit(basePlugin as never, features as never).canInstall).toBe(true)
expect(pluginInstallLimit(basePlugin, features).canInstall).toBe(true)
})
it('should default to langgenius when no verification info', () => {
const features = {
plugin_installation_permission: {
restrict_to_marketplace_only: false,
plugin_installation_scope: InstallationScope.OFFICIAL_ONLY,
},
}
const plugin = { from: 'marketplace' as const }
const features = makeSystemFeatures(InstallationScope.OFFICIAL_ONLY)
const plugin = { from: 'marketplace' as const } satisfies PluginInstallCandidate
expect(pluginInstallLimit(plugin as never, features as never).canInstall).toBe(true)
expect(pluginInstallLimit(plugin, features).canInstall).toBe(true)
})
it('should fallback to canInstall true for unrecognized scope', () => {
it('should deny installation for an unrecognized runtime scope', () => {
const features = {
plugin_installation_permission: {
restrict_to_marketplace_only: false,
plugin_installation_scope: 'unknown-scope' as unknown as PluginInstallationScope,
plugin_installation_scope: 'unknown-scope',
},
}
} as unknown as SystemFeatures
expect(pluginInstallLimit(basePlugin as never, features as never).canInstall).toBe(true)
expect(pluginInstallLimit(basePlugin, features).canInstall).toBe(false)
})
})
@@ -132,9 +108,9 @@ describe('usePluginInstallLimit', () => {
const plugin = {
from: 'marketplace' as const,
verification: { authorized_category: 'langgenius' },
}
} satisfies PluginInstallCandidate
const { result } = renderHook(() => usePluginInstallLimit(plugin as never))
const { result } = renderHook(() => usePluginInstallLimit(plugin))
expect(result.current.canInstall).toBe(true)
})
@@ -1,68 +1,55 @@
import type { GetSystemFeaturesResponse } from '@dify/contracts/api/console/system-features/types.gen'
import type { Plugin, PluginManifestInMarket } from '../../types'
import type {
PluginBundleDependencyType,
PluginVerification,
} from '@dify/contracts/api/console/workspaces/types.gen'
import { useSuspenseQuery } from '@tanstack/react-query'
import { systemFeaturesQueryOptions } from '@/features/system-features/client'
import { InstallationScope } from '@/features/system-features/constants'
type PluginProps = (Plugin | PluginManifestInMarket) & {
from: 'github' | 'marketplace' | 'package'
type PluginInstallCandidate = {
from: PluginBundleDependencyType
verification?: PluginVerification | null
}
type PluginInstallLimitResult = {
canInstall: boolean
}
function denyUnsupportedInstallationScope(_scope: never): PluginInstallLimitResult {
return { canInstall: false }
}
export function pluginInstallLimit(
plugin: PluginProps,
plugin: PluginInstallCandidate,
systemFeatures: Pick<GetSystemFeaturesResponse, 'plugin_installation_permission'>,
) {
if (systemFeatures.plugin_installation_permission.restrict_to_marketplace_only) {
const permission = systemFeatures.plugin_installation_permission
if (permission.restrict_to_marketplace_only) {
if (plugin.from === 'github' || plugin.from === 'package') return { canInstall: false }
}
if (
systemFeatures.plugin_installation_permission.plugin_installation_scope ===
InstallationScope.ALL
) {
return {
canInstall: true,
}
}
if (
systemFeatures.plugin_installation_permission.plugin_installation_scope ===
InstallationScope.NONE
) {
return {
canInstall: false,
}
}
const verification = plugin.verification || {}
if (!plugin.verification || !plugin.verification.authorized_category)
verification.authorized_category = 'langgenius'
const authorizedCategory = plugin.verification?.authorized_category ?? 'langgenius'
const scope = permission.plugin_installation_scope
if (
systemFeatures.plugin_installation_permission.plugin_installation_scope ===
InstallationScope.OFFICIAL_ONLY
) {
return {
canInstall: verification.authorized_category === 'langgenius',
}
}
if (
systemFeatures.plugin_installation_permission.plugin_installation_scope ===
InstallationScope.OFFICIAL_AND_PARTNER
) {
return {
canInstall:
verification.authorized_category === 'langgenius' ||
verification.authorized_category === 'partner',
}
}
return {
canInstall: true,
switch (scope) {
case InstallationScope.ALL:
return { canInstall: true }
case InstallationScope.NONE:
return { canInstall: false }
case InstallationScope.OFFICIAL_ONLY:
return { canInstall: authorizedCategory === 'langgenius' }
case InstallationScope.OFFICIAL_AND_PARTNER:
return {
canInstall: authorizedCategory === 'langgenius' || authorizedCategory === 'partner',
}
default:
return denyUnsupportedInstallationScope(scope)
}
}
export default function usePluginInstallLimit(plugin: PluginProps): PluginInstallLimitResult {
export default function usePluginInstallLimit(
plugin: PluginInstallCandidate,
): PluginInstallLimitResult {
const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions())
return pluginInstallLimit(plugin, systemFeatures)
@@ -149,7 +149,7 @@ describe('useCreateSnippetFromSelection', () => {
createNode('variable-aggregator', {
type: BlockEnum.VariableAggregator,
variables: [
['sys', 'files'],
['userinput', 'files'],
['llm', 'text'],
],
advanced_settings: {
@@ -115,6 +115,9 @@ describe('useWorkflowTemplate', () => {
expect(generateNewNodeCalls[1]!.data).toMatchObject({
type: 'llm',
title: 'workflow.blocks.llm',
memory: {
query_prompt_template: '{{#sys.query#}}\n\n{{#userinput.files#}}',
},
})
expect(generateNewNodeCalls[2]!.data).toMatchObject({
type: 'answer',
@@ -1,3 +1,5 @@
import type { AnswerNodeType } from '@/app/components/workflow/nodes/answer/types'
import type { LLMNodeType } from '@/app/components/workflow/nodes/llm/types'
import type { StartNodeType } from '@/app/components/workflow/nodes/start/types'
import { useTranslation } from 'react-i18next'
import { useStore as useAppStore } from '@/app/components/app/store'
@@ -31,37 +33,41 @@ export const useWorkflowTemplate = () => {
if (isChatMode) {
const startNode = createStartNode()
const llmData: LLMNodeType = {
...(llmDefault.defaultValue as LLMNodeType),
desc: '',
memory: {
window: { enabled: false, size: 10 },
query_prompt_template: '{{#sys.query#}}\n\n{{#userinput.files#}}',
},
selected: true,
type: llmDefault.metaData.type,
title: t(($) => $[`blocks.${llmDefault.metaData.type}`], { ns: 'workflow' }),
}
const { newNode: llmNode } = generateNewNode({
id: 'llm',
data: {
...llmDefault.defaultValue,
memory: {
window: { enabled: false, size: 10 },
query_prompt_template: '{{#sys.query#}}\n\n{{#sys.files#}}',
},
selected: true,
type: llmDefault.metaData.type,
title: t(($) => $[`blocks.${llmDefault.metaData.type}`], { ns: 'workflow' }),
},
data: llmData,
position: {
x: START_INITIAL_POSITION.x + NODE_WIDTH_X_OFFSET,
y: START_INITIAL_POSITION.y,
},
} as any)
})
const answerData: AnswerNodeType = {
...(answerDefault.defaultValue as AnswerNodeType),
answer: `{{#${llmNode.id}.text#}}`,
desc: '',
type: answerDefault.metaData.type,
title: t(($) => $[`blocks.${answerDefault.metaData.type}`], { ns: 'workflow' }),
}
const { newNode: answerNode } = generateNewNode({
id: 'answer',
data: {
...answerDefault.defaultValue,
answer: `{{#${llmNode.id}.text#}}`,
type: answerDefault.metaData.type,
title: t(($) => $[`blocks.${answerDefault.metaData.type}`], { ns: 'workflow' }),
},
data: answerData,
position: {
x: START_INITIAL_POSITION.x + NODE_WIDTH_X_OFFSET * 2,
y: START_INITIAL_POSITION.y,
},
} as any)
})
const startToLlmEdge = {
id: `${startNode.id}-${llmNode.id}`,
+1 -1
View File
@@ -77,7 +77,7 @@ export const getGlobalVars = (isChatMode: boolean): Var[] => {
export const VAR_SHOW_NAME_MAP: Record<string, string> = {
'sys.query': 'query',
'sys.files': 'files',
'userinput.files': 'files',
}
export const RETRIEVAL_OUTPUT_STRUCT = `{
@@ -42,7 +42,7 @@ describe('useConfigVision', () => {
})
})
it('should expose vision capability and enable default chat configs for vision models', () => {
it('should use userinput.files by default for vision models in chat mode', () => {
const onChange = vi.fn()
mockUseIsChatMode.mockReturnValue(true)
mockUseTextGenerationCurrentProviderAndModelAndModelList.mockReturnValue({
@@ -68,7 +68,7 @@ describe('useConfigVision', () => {
enabled: true,
configs: {
detail: Resolution.high,
variable_selector: ['sys', 'files'],
variable_selector: ['userinput', 'files'],
},
})
})
@@ -131,7 +131,7 @@ describe('useConfigVision', () => {
enabled: true,
configs: {
detail: Resolution.high,
variable_selector: ['sys', 'files'],
variable_selector: ['userinput', 'files'],
},
}),
onChange,
@@ -149,6 +149,7 @@ describe('useConfigVision', () => {
it('should reset enabled vision configs when the model changes but still supports vision', () => {
const onChange = vi.fn()
mockUseIsChatMode.mockReturnValue(true)
mockUseTextGenerationCurrentProviderAndModelAndModelList.mockReturnValue({
currentModel: {
features: [ModelFeatureEnum.vision],
@@ -172,6 +173,34 @@ describe('useConfigVision', () => {
result.current.handleModelChanged()
})
expect(onChange).toHaveBeenCalledWith({
enabled: true,
configs: {
detail: Resolution.high,
variable_selector: ['userinput', 'files'],
},
})
})
it('should require an explicit file variable in workflow mode', () => {
const onChange = vi.fn()
mockUseTextGenerationCurrentProviderAndModelAndModelList.mockReturnValue({
currentModel: {
features: [ModelFeatureEnum.vision],
},
})
const { result } = renderHook(() =>
useConfigVision(createModel(), {
payload: createVisionPayload(),
onChange,
}),
)
act(() => {
result.current.handleVisionResolutionEnabledChange(true)
})
expect(onChange).toHaveBeenCalledWith({
enabled: true,
configs: {
@@ -52,6 +52,7 @@ describe('useInspectVarsCrud', () => {
createInspectVar({
id: 'files-var',
name: 'files',
// TODO: Remove this legacy API fixture when the system-variable endpoint returns userinput.files.
selector: ['sys', 'files'],
}),
createInspectVar({
@@ -134,6 +135,10 @@ describe('useInspectVarsCrud', () => {
'query',
'files',
])
expect(result.current.nodesWithInspectVars[0]?.vars.at(-1)?.selector).toEqual([
'userinput',
'files',
])
expect(result.current.hasNodeInspectVars).toBe(hasNodeInspectVars)
expect(result.current.fetchInspectVarValue).toBe(fetchInspectVarValue)
expect(result.current.deleteAllInspectorVars).toBe(deleteAllInspectorVars)
@@ -84,7 +84,7 @@ const outputVarsWithSystemVars: NodeOutPutVar[] = [
type: VarType.string,
},
{
variable: 'sys.files',
variable: 'userinput.files',
type: VarType.arrayFile,
},
] satisfies Var[],
@@ -41,10 +41,10 @@ const useConfigVision = (
(enabled: boolean) => {
const newPayload = produce(payload, (draft) => {
draft.enabled = enabled
if (enabled && isChatMode) {
if (enabled) {
draft.configs = {
detail: Resolution.high,
variable_selector: ['sys', 'files'],
variable_selector: isChatMode ? ['userinput', 'files'] : [],
}
} else if (!enabled) {
delete draft.configs
@@ -76,11 +76,11 @@ const useConfigVision = (
enabled: true,
configs: {
detail: Resolution.high,
variable_selector: [],
variable_selector: isChatMode ? ['userinput', 'files'] : [],
},
})
}
}, [getIsVisionModel, handleVisionResolutionEnabledChange, onChange, payload.enabled])
}, [getIsVisionModel, handleVisionResolutionEnabledChange, isChatMode, onChange, payload.enabled])
return {
isVisionModel,
@@ -17,7 +17,14 @@ const useInspectVarsCrud = () => {
const { varsAppendStartNode, systemVars } = (() => {
if (allSystemVars?.length === 0) return { varsAppendStartNode: [], systemVars: [] }
const varsAppendStartNode =
allSystemVars?.filter(({ name }) => varsAppendStartNodeKeys.includes(name)) || []
allSystemVars
?.filter(({ name }) => varsAppendStartNodeKeys.includes(name))
.map((variable) => {
if (variable.name !== 'files') return variable
// TODO: Remove this normalization after the system-variable API stops returning the legacy file selector.
return { ...variable, selector: ['userinput', 'files'] }
}) || []
const systemVars =
allSystemVars?.filter(({ name }) => !varsAppendStartNodeKeys.includes(name)) || []
return { varsAppendStartNode, systemVars }
@@ -55,7 +55,7 @@ type Props = Readonly<{
const MEMORY_DEFAULT: Memory = {
window: { enabled: false, size: WINDOW_SIZE_DEFAULT },
query_prompt_template: '{{#sys.query#}}\n\n{{#sys.files#}}',
query_prompt_template: '{{#sys.query#}}\n\n{{#userinput.files#}}',
}
const MemoryConfig: FC<Props> = ({
@@ -97,7 +97,7 @@ const MemoryConfig: FC<Props> = ({
limitedSize = null
} else {
limitedSize = Number.parseInt(limitedSize as string, 10)
if (isNaN(limitedSize)) limitedSize = WINDOW_SIZE_DEFAULT
if (Number.isNaN(limitedSize)) limitedSize = WINDOW_SIZE_DEFAULT
if (limitedSize < WINDOW_SIZE_MIN) limitedSize = WINDOW_SIZE_MIN
@@ -4,6 +4,7 @@ import type { HumanInputNodeType } from '@/app/components/workflow/nodes/human-i
import type { LLMNodeType } from '@/app/components/workflow/nodes/llm/types'
import type { Node, PromptItem } from '@/app/components/workflow/types'
import { describe, expect, it } from 'vitest'
import { createStartNode } from '@/app/components/workflow/__tests__/fixtures'
import { DeliveryMethodType } from '@/app/components/workflow/nodes/human-input/types'
import {
BlockEnum,
@@ -13,7 +14,14 @@ import {
VarType,
} from '@/app/components/workflow/types'
import { AppModeEnum } from '@/types/app'
import { getNodeUsedVars, toNodeAvailableVars, updateNodeVars } from '../utils'
import {
getNodeOutputVars,
getNodeUsedVars,
isGlobalVar,
isSystemVar,
toNodeAvailableVars,
updateNodeVars,
} from '../utils'
const createNode = <T>(data: Node<T>['data']): Node<T> => ({
id: 'node-1',
@@ -49,6 +57,15 @@ const createLLMNodeData = (promptTemplate: PromptItem[]): LLMNodeType => ({
})
describe('variable utils', () => {
describe('virtual input variables', () => {
it('recognizes userinput.files without treating it as a global variable', () => {
expect(isSystemVar(['userinput', 'files'])).toBe(true)
expect(isSystemVar(['start', 'userinput', 'files'])).toBe(true)
expect(isGlobalVar(['start', 'userinput', 'files'])).toBe(false)
expect(isGlobalVar(['sys', 'workflow_id'])).toBe(true)
})
})
describe('toNodeAvailableVars', () => {
it('uses Agent v2 default declared outputs for agent nodes', () => {
const node = createNode<AgentV2NodeType>({
@@ -213,6 +230,51 @@ describe('variable utils', () => {
})
})
describe('node output variables', () => {
it('should expose sys.query and userinput.files for start nodes in chat mode', () => {
const startNode = createStartNode({
id: 'start',
data: {
type: BlockEnum.Start,
variables: [
{
label: 'Files',
variable: 'files',
type: InputVarType.multiFiles,
required: false,
},
],
},
})
expect(getNodeOutputVars(startNode, true)).toEqual([
['start', 'files'],
['start', 'sys', 'query'],
['start', 'userinput', 'files'],
])
const availableVars = toNodeAvailableVars({
beforeNodes: [startNode],
isChatMode: true,
filterVar: () => true,
allPluginInfoList: {},
})
expect(availableVars).toEqual(
expect.arrayContaining([
expect.objectContaining({
nodeId: 'start',
vars: expect.arrayContaining([
expect.objectContaining({ variable: 'files', type: VarType.arrayFile }),
expect.objectContaining({ variable: 'sys.query', type: VarType.string }),
expect.objectContaining({ variable: 'userinput.files', type: VarType.arrayFile }),
]),
}),
]),
)
})
})
describe('updateNodeVars', () => {
it('should replace answer prompt references', () => {
const node = createNode<AnswerNodeType>({
@@ -80,9 +80,9 @@ describe('var-reference-picker.helpers', () => {
isLoopVar: false,
iterationNode: null,
loopNode: null,
outputVarNodeId: 'sys',
outputVarNodeId: 'userinput',
startNode,
value: ['sys', 'files'],
value: ['userinput', 'files'],
}),
).toEqual(startNode.data)
@@ -287,7 +287,7 @@ describe('var-reference-picker.helpers', () => {
})
it('should keep mapped variable names for known workflow aliases', () => {
expect(getVarDisplayName(true, ['sys', 'files'])).toBe('files')
expect(getVarDisplayName(true, ['userinput', 'files'])).toBe('files')
expect(
getVariableMeta({ type: VarType.string }, ['conversation', 'name'], 'name'),
).toMatchObject({
@@ -9,7 +9,7 @@ import {
describe('var-reference-vars helpers', () => {
it('should derive display names for flat and mapped variables', () => {
expect(getVariableDisplayName('sys.files', false)).toBe('files')
expect(getVariableDisplayName('userinput.files', false)).toBe('files')
expect(getVariableDisplayName('current', true, true)).toBe('current_code')
expect(getVariableDisplayName('foo', true, false)).toBe('foo')
})
@@ -106,6 +106,32 @@ describe('VarReferenceVars', () => {
)
})
it('should resolve userinput.files as a virtual input selector', () => {
const onChange = vi.fn()
render(
<VarReferenceVars
vars={createVars([
{
title: 'Start',
nodeId: 'start-node',
vars: [{ variable: 'userinput.files', type: VarType.arrayFile }],
},
])}
onChange={onChange}
/>,
)
fireEvent.click(screen.getByText('files'))
expect(onChange).toHaveBeenCalledWith(
['userinput', 'files'],
expect.objectContaining({
variable: 'userinput.files',
}),
)
})
it('should render empty state and manage input action', () => {
const onManageInputField = vi.fn()

Some files were not shown because too many files have changed in this diff Show More