Compare commits
62
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
1bcc1092dd | ||
|
|
9666cdff03 | ||
|
|
26905dfb44 | ||
|
|
3183980d71 | ||
|
|
df784e7a7d | ||
|
|
5c0a0759f0 | ||
|
|
3443054d96 | ||
|
|
be737a3504 | ||
|
|
e0b910d43d | ||
|
|
44048320bd | ||
|
|
3dc85bd398 | ||
|
|
56d87d6318 | ||
|
|
73db4d3932 | ||
|
|
13c6556e6f | ||
|
|
7570427f51 | ||
|
|
48717e5981 | ||
|
|
5805b94e23 | ||
|
|
c64c772b5b | ||
|
|
b4537455cf | ||
|
|
5fe010b653 | ||
|
|
66cd69b07e | ||
|
|
09eca00370 | ||
|
|
12cf69e5c8 | ||
|
|
4332023806 | ||
|
|
dd9f1e041e | ||
|
|
2271257c2f | ||
|
|
8adbda988b | ||
|
|
876ac24496 | ||
|
|
8b0892fd79 | ||
|
|
50d2ad8479 | ||
|
|
1ef8d3a2a9 | ||
|
|
b8adb0cac2 | ||
|
|
df88c45e3c | ||
|
|
d3df1b2847 | ||
|
|
709aac9337 | ||
|
|
865affdd9b | ||
|
|
71840c3f71 | ||
|
|
dc09aea7b3 | ||
|
|
99507a5ae5 | ||
|
|
55ae0aadf8 | ||
|
|
6c4ba71a23 | ||
|
|
e5f53053e3 | ||
|
|
3e392527e9 | ||
|
|
36d3a9c7fb | ||
|
|
914598d7e3 | ||
|
|
602c24cf8c | ||
|
|
e0b5ecc50f | ||
|
|
a83c7e89e0 | ||
|
|
7a4252b3de | ||
|
|
5c3dd25b32 | ||
|
|
30c8e9b08d | ||
|
|
3e3a16feaf | ||
|
|
0e1fad88c0 | ||
|
|
f442559543 | ||
|
|
2a63f26796 | ||
|
|
7e325b2ff2 | ||
|
|
26487d19b4 | ||
|
|
71a936e89a | ||
|
|
f2c19b456b | ||
|
|
2d3da2a49f | ||
|
|
11b393db9f | ||
|
|
9a01b16d6c |
@@ -22,7 +22,7 @@ from .plugin import (
|
||||
setup_system_trigger_oauth_client,
|
||||
transform_datasource_credentials,
|
||||
)
|
||||
from .rbac import migrate_member_roles_to_rbac
|
||||
from .rbac import migrate_dataset_permissions_to_rbac, migrate_member_roles_to_rbac
|
||||
from .retention import (
|
||||
archive_workflow_runs,
|
||||
archive_workflow_runs_plan,
|
||||
@@ -76,6 +76,7 @@ __all__ = [
|
||||
"legacy_model_types",
|
||||
"migrate_annotation_vector_database",
|
||||
"migrate_data_for_plugin",
|
||||
"migrate_dataset_permissions_to_rbac",
|
||||
"migrate_knowledge_vector_database",
|
||||
"migrate_member_roles_to_rbac",
|
||||
"migrate_oss",
|
||||
|
||||
@@ -7,6 +7,7 @@ from typing import cast
|
||||
|
||||
import click
|
||||
|
||||
from commands.rbac import migrate_dataset_permissions_to_rbac
|
||||
from extensions.ext_database import db
|
||||
from graphon.model_runtime.entities.model_entities import ModelType
|
||||
from services.legacy_model_type_migration import (
|
||||
@@ -177,3 +178,4 @@ def legacy_model_types(
|
||||
|
||||
|
||||
data_migrate.add_command(legacy_model_types)
|
||||
data_migrate.add_command(migrate_dataset_permissions_to_rbac)
|
||||
|
||||
+438
-66
@@ -1,11 +1,55 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Iterator
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
|
||||
import click
|
||||
from sqlalchemy import select
|
||||
|
||||
from configs import dify_config
|
||||
from core.db.session_factory import session_factory
|
||||
from models import TenantAccountJoin, TenantAccountRole
|
||||
from services.enterprise.rbac_service import ListOption, RBACService
|
||||
from core.rbac import RBACResourceWhitelistScope
|
||||
from models import Dataset, DatasetPermission, DatasetPermissionEnum, TenantAccountJoin, TenantAccountRole
|
||||
from services.enterprise.rbac_service import ListOption, RBACService, ReplaceMemberBindings, ReplaceUserAccessPolicies
|
||||
|
||||
_RBAC_DEFAULT_ACCESS_POLICY_ID = "default"
|
||||
|
||||
_LEGACY_ROLE_TO_BUILTIN_TAG = {
|
||||
TenantAccountRole.OWNER.value: "owner",
|
||||
TenantAccountRole.ADMIN.value: "admin",
|
||||
TenantAccountRole.EDITOR.value: "editor",
|
||||
TenantAccountRole.NORMAL.value: "normal",
|
||||
TenantAccountRole.DATASET_OPERATOR.value: "dataset_operator",
|
||||
}
|
||||
|
||||
|
||||
def _resolve_builtin_role_ids(tenant_id: str, operator_account_id: str) -> dict[str, str]:
|
||||
"""Resolve every legacy workspace role to the current tenant's builtin RBAC role id.
|
||||
|
||||
The migration replays the old `TenantAccountJoin.role` values onto the
|
||||
RBAC member-role binding API. Builtin RBAC roles are tenant-scoped and
|
||||
identified by runtime ids, so the command must look them up per tenant.
|
||||
"""
|
||||
roles = RBACService.Roles.list(
|
||||
tenant_id=tenant_id,
|
||||
account_id=operator_account_id,
|
||||
options=ListOption(page_number=1, results_per_page=100),
|
||||
).data
|
||||
role_id_by_tag = {
|
||||
role.role_tag: role.id
|
||||
for role in roles
|
||||
if role.is_builtin and role.category == "global_system_default" and role.role_tag
|
||||
}
|
||||
resolved: dict[str, str] = {}
|
||||
for legacy_role, expected_builtin_tag in _LEGACY_ROLE_TO_BUILTIN_TAG.items():
|
||||
role_id = role_id_by_tag.get(expected_builtin_tag)
|
||||
if expected_builtin_tag == "dataset_operator" and not dify_config.DATASET_OPERATOR_ENABLED:
|
||||
continue
|
||||
if not role_id:
|
||||
raise ValueError(f"Builtin RBAC role not found for tenant={tenant_id}, legacy_role={legacy_role}")
|
||||
resolved[legacy_role] = role_id
|
||||
return resolved
|
||||
|
||||
|
||||
def _resolve_builtin_role_id(tenant_id: str, operator_account_id: str, legacy_role: str) -> str:
|
||||
@@ -15,26 +59,86 @@ def _resolve_builtin_role_id(tenant_id: str, operator_account_id: str, legacy_ro
|
||||
RBAC member-role binding API. Builtin RBAC roles are tenant-scoped and
|
||||
identified by runtime ids, so the command must look them up per tenant.
|
||||
"""
|
||||
expected_builtin_tag = {
|
||||
TenantAccountRole.OWNER.value: "owner",
|
||||
TenantAccountRole.ADMIN.value: "admin",
|
||||
TenantAccountRole.EDITOR.value: "editor",
|
||||
TenantAccountRole.NORMAL.value: "normal",
|
||||
TenantAccountRole.DATASET_OPERATOR.value: "dataset_operator",
|
||||
}.get(legacy_role)
|
||||
if not expected_builtin_tag:
|
||||
if legacy_role not in _LEGACY_ROLE_TO_BUILTIN_TAG:
|
||||
raise ValueError(f"Unsupported legacy workspace role: {legacy_role}")
|
||||
|
||||
roles = RBACService.Roles.list(
|
||||
return _resolve_builtin_role_ids(tenant_id, operator_account_id)[legacy_role]
|
||||
|
||||
|
||||
def _iter_tenant_member_batches(
|
||||
tenant_id: str | None,
|
||||
*,
|
||||
db_batch_size: int,
|
||||
api_batch_size: int,
|
||||
) -> Iterator[tuple[str, str, list[tuple[str, str]]]]:
|
||||
"""Yield legacy member roles in tenant-scoped API-sized batches.
|
||||
|
||||
Rows are projected to primitive values and streamed from the database, so
|
||||
the command never materializes every TenantAccountJoin ORM object. The
|
||||
iterator only keeps one tenant's API-sized batches in memory while it
|
||||
finds that tenant's owner account.
|
||||
"""
|
||||
with session_factory.create_session() as session:
|
||||
stmt = (
|
||||
select(TenantAccountJoin.tenant_id, TenantAccountJoin.account_id, TenantAccountJoin.role)
|
||||
.order_by(TenantAccountJoin.tenant_id.asc(), TenantAccountJoin.id.asc())
|
||||
.execution_options(yield_per=db_batch_size)
|
||||
)
|
||||
if tenant_id:
|
||||
stmt = stmt.where(TenantAccountJoin.tenant_id == tenant_id)
|
||||
|
||||
current_tenant_id: str | None = None
|
||||
owner_account_id: str | None = None
|
||||
batches: list[list[tuple[str, str]]] = []
|
||||
batch: list[tuple[str, str]] = []
|
||||
|
||||
def flush_current_tenant() -> Iterator[tuple[str, str, list[tuple[str, str]]]]:
|
||||
if current_tenant_id is None:
|
||||
return
|
||||
if batch:
|
||||
batches.append(batch.copy())
|
||||
if not owner_account_id:
|
||||
raise ValueError(f"Workspace owner not found for tenant={current_tenant_id}")
|
||||
for item in batches:
|
||||
yield current_tenant_id, owner_account_id, item
|
||||
|
||||
for row in session.execute(stmt):
|
||||
workspace_id = str(row.tenant_id)
|
||||
if current_tenant_id is not None and workspace_id != current_tenant_id:
|
||||
yield from flush_current_tenant()
|
||||
owner_account_id = None
|
||||
batches = []
|
||||
batch = []
|
||||
current_tenant_id = workspace_id
|
||||
account_id = str(row.account_id)
|
||||
role = str(row.role)
|
||||
if role == TenantAccountRole.OWNER.value:
|
||||
owner_account_id = account_id
|
||||
batch.append((account_id, role))
|
||||
if len(batch) >= api_batch_size:
|
||||
batches.append(batch)
|
||||
batch = []
|
||||
|
||||
yield from flush_current_tenant()
|
||||
|
||||
|
||||
def _member_already_has_role(current_roles_by_account_id: dict[str, set[str]], account_id: str, role_id: str) -> bool:
|
||||
return current_roles_by_account_id.get(account_id) == {role_id}
|
||||
|
||||
|
||||
def _replace_member_role(
|
||||
tenant_id: str,
|
||||
operator_account_id: str,
|
||||
member_account_id: str,
|
||||
role_id: str,
|
||||
) -> str:
|
||||
RBACService.MemberRoles.replace(
|
||||
tenant_id=tenant_id,
|
||||
account_id=operator_account_id,
|
||||
options=ListOption(page_number=1, results_per_page=100),
|
||||
).data
|
||||
for role in roles:
|
||||
if role.is_builtin and role.category == "global_system_default" and role.role_tag == expected_builtin_tag:
|
||||
return str(role.id)
|
||||
|
||||
raise ValueError(f"Builtin RBAC role not found for tenant={tenant_id}, legacy_role={legacy_role}")
|
||||
member_account_id=member_account_id,
|
||||
role_ids=[role_id],
|
||||
)
|
||||
return member_account_id
|
||||
|
||||
|
||||
@click.command(
|
||||
@@ -42,7 +146,16 @@ def _resolve_builtin_role_id(tenant_id: str, operator_account_id: str, legacy_ro
|
||||
)
|
||||
@click.option("--tenant-id", help="Only migrate a single workspace.")
|
||||
@click.option("--dry-run", is_flag=True, default=False, help="Preview the migration without writing RBAC bindings.")
|
||||
def migrate_member_roles_to_rbac(tenant_id: str | None, dry_run: bool) -> None:
|
||||
@click.option("--db-batch-size", default=5000, show_default=True, help="Rows fetched per database batch.")
|
||||
@click.option("--api-batch-size", default=200, show_default=True, help="Members checked per RBAC batch_get call.")
|
||||
@click.option("--workers", default=1, show_default=True, help="Concurrent member role replace calls per tenant batch.")
|
||||
def migrate_member_roles_to_rbac(
|
||||
tenant_id: str | None,
|
||||
dry_run: bool,
|
||||
db_batch_size: int,
|
||||
api_batch_size: int,
|
||||
workers: int,
|
||||
) -> None:
|
||||
"""Backfill RBAC member-role bindings from legacy `TenantAccountJoin.role` data.
|
||||
|
||||
This is an offline migration command for workspaces that already have
|
||||
@@ -50,63 +163,322 @@ def migrate_member_roles_to_rbac(tenant_id: str | None, dry_run: bool) -> None:
|
||||
member-role binding store.
|
||||
"""
|
||||
click.echo(click.style("Starting RBAC member-role migration.", fg="green"))
|
||||
if workers < 1:
|
||||
raise click.BadParameter("workers must be >= 1", param_hint="--workers")
|
||||
|
||||
with session_factory.create_session() as session:
|
||||
stmt = select(TenantAccountJoin).order_by(TenantAccountJoin.tenant_id.asc(), TenantAccountJoin.id.asc())
|
||||
if tenant_id:
|
||||
stmt = stmt.where(TenantAccountJoin.tenant_id == tenant_id)
|
||||
tenant_count = 0
|
||||
scanned_count = 0
|
||||
skipped_count = 0
|
||||
migrated_count = 0
|
||||
current_tenant_id: str | None = None
|
||||
role_ids_by_legacy_role: dict[str, str] = {}
|
||||
|
||||
joins = list(session.scalars(stmt).all())
|
||||
for workspace_id, owner_account_id, batch in _iter_tenant_member_batches(
|
||||
tenant_id,
|
||||
db_batch_size=db_batch_size,
|
||||
api_batch_size=api_batch_size,
|
||||
):
|
||||
scanned_count += len(batch)
|
||||
if workspace_id != current_tenant_id:
|
||||
tenant_count += 1
|
||||
current_tenant_id = workspace_id
|
||||
role_ids_by_legacy_role = _resolve_builtin_role_ids(workspace_id, owner_account_id)
|
||||
click.echo(f"tenant={workspace_id}")
|
||||
|
||||
if not joins:
|
||||
current_roles_by_account_id: dict[str, set[str]] = {}
|
||||
if not dry_run:
|
||||
current_roles = RBACService.MemberRoles.batch_get(
|
||||
tenant_id=workspace_id,
|
||||
account_id=owner_account_id,
|
||||
member_account_ids=[account_id for account_id, _ in batch],
|
||||
)
|
||||
current_roles_by_account_id = {item.account_id: {role.id for role in item.roles} for item in current_roles}
|
||||
|
||||
replace_jobs: list[tuple[str, str]] = []
|
||||
for member_account_id, legacy_role in batch:
|
||||
resolved_role_id = role_ids_by_legacy_role.get(legacy_role)
|
||||
if not resolved_role_id:
|
||||
raise ValueError(f"Unsupported legacy workspace role: {legacy_role}")
|
||||
|
||||
if dry_run:
|
||||
click.echo(
|
||||
f"tenant={workspace_id} member={member_account_id} "
|
||||
f"legacy_role={legacy_role} -> rbac_role_id={resolved_role_id}"
|
||||
)
|
||||
continue
|
||||
|
||||
if _member_already_has_role(current_roles_by_account_id, member_account_id, resolved_role_id):
|
||||
skipped_count += 1
|
||||
continue
|
||||
|
||||
replace_jobs.append((member_account_id, resolved_role_id))
|
||||
|
||||
if replace_jobs:
|
||||
if workers == 1:
|
||||
for member_account_id, resolved_role_id in replace_jobs:
|
||||
_replace_member_role(workspace_id, owner_account_id, member_account_id, resolved_role_id)
|
||||
migrated_count += 1
|
||||
else:
|
||||
with ThreadPoolExecutor(max_workers=workers) as executor:
|
||||
futures = [
|
||||
executor.submit(
|
||||
_replace_member_role,
|
||||
workspace_id,
|
||||
owner_account_id,
|
||||
member_account_id,
|
||||
resolved_role_id,
|
||||
)
|
||||
for member_account_id, resolved_role_id in replace_jobs
|
||||
]
|
||||
for future in as_completed(futures):
|
||||
future.result()
|
||||
migrated_count += 1
|
||||
|
||||
if scanned_count % 10000 == 0:
|
||||
click.echo(
|
||||
f"progress scanned={scanned_count} migrated={migrated_count} skipped={skipped_count}",
|
||||
err=True,
|
||||
)
|
||||
|
||||
if scanned_count == 0:
|
||||
click.echo(click.style("No workspace members found for migration.", fg="yellow"))
|
||||
return
|
||||
|
||||
owner_account_by_tenant: dict[str, str] = {}
|
||||
resolved_role_ids: dict[tuple[str, str], str] = {}
|
||||
migrated_count = 0
|
||||
|
||||
for join in joins:
|
||||
workspace_id = str(join.tenant_id)
|
||||
member_account_id = str(join.account_id)
|
||||
legacy_role = str(join.role)
|
||||
|
||||
if workspace_id not in owner_account_by_tenant:
|
||||
owner_join = next(
|
||||
(
|
||||
item
|
||||
for item in joins
|
||||
if str(item.tenant_id) == workspace_id and str(item.role) == TenantAccountRole.OWNER.value
|
||||
),
|
||||
None,
|
||||
)
|
||||
if not owner_join:
|
||||
raise ValueError(f"Workspace owner not found for tenant={workspace_id}")
|
||||
owner_account_by_tenant[workspace_id] = str(owner_join.account_id)
|
||||
|
||||
operator_account_id = owner_account_by_tenant[workspace_id]
|
||||
cache_key = (workspace_id, legacy_role)
|
||||
if cache_key not in resolved_role_ids:
|
||||
resolved_role_ids[cache_key] = _resolve_builtin_role_id(workspace_id, operator_account_id, legacy_role)
|
||||
|
||||
resolved_role_id = resolved_role_ids[cache_key]
|
||||
if dry_run:
|
||||
click.echo(
|
||||
f"tenant={workspace_id} member={member_account_id} "
|
||||
f"legacy_role={legacy_role} -> rbac_role_id={resolved_role_id}"
|
||||
click.style(
|
||||
f"Dry run completed. Scanned {scanned_count} members across {tenant_count} tenants. "
|
||||
"No RBAC bindings were written.",
|
||||
fg="yellow",
|
||||
)
|
||||
)
|
||||
else:
|
||||
click.echo(
|
||||
click.style(
|
||||
f"RBAC member-role migration completed. Scanned {scanned_count} members across {tenant_count} tenants, "
|
||||
f"migrated {migrated_count}, skipped {skipped_count} already up-to-date.",
|
||||
fg="green",
|
||||
)
|
||||
)
|
||||
|
||||
if dry_run:
|
||||
continue
|
||||
|
||||
RBACService.MemberRoles.replace(
|
||||
tenant_id=workspace_id,
|
||||
account_id=operator_account_id,
|
||||
member_account_id=member_account_id,
|
||||
role_ids=[resolved_role_id],
|
||||
)
|
||||
migrated_count += 1
|
||||
def _dataset_permission_enum(permission: DatasetPermissionEnum | str | None) -> DatasetPermissionEnum:
|
||||
if permission is None:
|
||||
return DatasetPermissionEnum.ONLY_ME
|
||||
try:
|
||||
return DatasetPermissionEnum(permission)
|
||||
except ValueError as exc:
|
||||
raise ValueError(f"Unsupported legacy dataset permission: {permission}") from exc
|
||||
|
||||
|
||||
def _rbac_dataset_scope_for_legacy_permission(permission: DatasetPermissionEnum) -> RBACResourceWhitelistScope:
|
||||
if permission is DatasetPermissionEnum.ALL_TEAM:
|
||||
return RBACResourceWhitelistScope.ALL
|
||||
if permission in {DatasetPermissionEnum.ONLY_ME, DatasetPermissionEnum.PARTIAL_TEAM}:
|
||||
return RBACResourceWhitelistScope.SPECIFIC
|
||||
raise ValueError(f"Unsupported legacy dataset permission: {permission}")
|
||||
|
||||
|
||||
def _emit_dataset_permission_migration_event(payload: dict[str, object]) -> None:
|
||||
click.echo(json.dumps(payload, sort_keys=True))
|
||||
|
||||
|
||||
@click.command(
|
||||
"rbac-migrate-dataset-permissions",
|
||||
help=(
|
||||
"Migrate legacy dataset permission scopes and partial members into RBAC dataset access bindings. "
|
||||
"Side effect: replacing each dataset whitelist clears existing per-user policy bindings; "
|
||||
"the command then recreates legacy partial-member default bindings."
|
||||
),
|
||||
)
|
||||
@click.option("--tenant-id", help="Only migrate datasets in a single workspace.")
|
||||
@click.option("--dataset-id", help="Only migrate a single dataset.")
|
||||
@click.option("--batch-size", default=500, show_default=True, type=click.IntRange(min=1))
|
||||
@click.option(
|
||||
"--dry-run/--apply",
|
||||
default=True,
|
||||
show_default=True,
|
||||
help="Preview the migration without writing RBAC bindings. Use --apply to write changes.",
|
||||
)
|
||||
def migrate_dataset_permissions_to_rbac(
|
||||
tenant_id: str | None,
|
||||
dataset_id: str | None,
|
||||
batch_size: int,
|
||||
dry_run: bool,
|
||||
) -> None:
|
||||
"""Backfill RBAC dataset access config from legacy `Dataset.permission`.
|
||||
|
||||
Legacy mapping:
|
||||
- all_team_members -> RBAC dataset whitelist scope "all"
|
||||
- partial_members -> RBAC dataset whitelist scope "specific" plus each partial member gets the
|
||||
virtual default policy
|
||||
- only_me -> RBAC dataset whitelist scope "specific" with no member policy bindings
|
||||
|
||||
The command replaces each dataset's RBAC whitelist scope first. RBAC clears
|
||||
existing per-user policy bindings during that replace, then this command
|
||||
recreates the legacy partial-member default bindings. Re-running it is
|
||||
therefore idempotent for a dataset's current legacy configuration.
|
||||
"""
|
||||
click.echo(click.style("Starting RBAC dataset permission migration.", fg="green"))
|
||||
|
||||
scanned_count = 0
|
||||
scope_migrated_count = 0
|
||||
user_policy_migrated_count = 0
|
||||
partial_dataset_count = 0
|
||||
|
||||
last_dataset_id: str | None = None
|
||||
while True:
|
||||
with session_factory.create_session() as session:
|
||||
stmt = (
|
||||
select(Dataset.id, Dataset.tenant_id, Dataset.permission, Dataset.created_by)
|
||||
.order_by(Dataset.id.asc())
|
||||
.limit(batch_size)
|
||||
)
|
||||
if tenant_id:
|
||||
stmt = stmt.where(Dataset.tenant_id == tenant_id)
|
||||
if dataset_id:
|
||||
stmt = stmt.where(Dataset.id == dataset_id)
|
||||
if last_dataset_id:
|
||||
stmt = stmt.where(Dataset.id > last_dataset_id)
|
||||
|
||||
dataset_rows = list(session.execute(stmt).all())
|
||||
if not dataset_rows:
|
||||
break
|
||||
|
||||
dataset_ids = [str(row.id) for row in dataset_rows]
|
||||
partial_members_by_dataset_id: dict[str, list[str]] = {item: [] for item in dataset_ids}
|
||||
permission_rows = session.execute(
|
||||
select(DatasetPermission.dataset_id, DatasetPermission.account_id).where(
|
||||
DatasetPermission.dataset_id.in_(dataset_ids)
|
||||
)
|
||||
).all()
|
||||
for row in permission_rows:
|
||||
partial_members_by_dataset_id[str(row.dataset_id)].append(str(row.account_id))
|
||||
|
||||
for dataset in dataset_rows:
|
||||
workspace_id = str(dataset.tenant_id)
|
||||
current_dataset_id = str(dataset.id)
|
||||
operator_account_id = str(dataset.created_by)
|
||||
permission_value = _dataset_permission_enum(dataset.permission)
|
||||
scope = _rbac_dataset_scope_for_legacy_permission(permission_value)
|
||||
partial_member_ids = sorted(set(partial_members_by_dataset_id[current_dataset_id]))
|
||||
should_bind_partial_members = permission_value is DatasetPermissionEnum.PARTIAL_TEAM
|
||||
|
||||
click.echo(
|
||||
f"tenant={workspace_id} dataset={current_dataset_id} "
|
||||
f"operator={operator_account_id} "
|
||||
f"legacy_permission={permission_value} -> rbac_scope={scope} "
|
||||
f"partial_members={len(partial_member_ids) if should_bind_partial_members else 0}"
|
||||
)
|
||||
|
||||
scanned_count += 1
|
||||
replace_whitelist_payload = ReplaceMemberBindings(scope=scope)
|
||||
if dry_run:
|
||||
_emit_dataset_permission_migration_event(
|
||||
{
|
||||
"event": "dataset_permission_migration_proposed_change",
|
||||
"action": "replace_whitelist",
|
||||
"dry_run": True,
|
||||
"tenant_id": workspace_id,
|
||||
"dataset_id": current_dataset_id,
|
||||
"operator_account_id": operator_account_id,
|
||||
"before": {
|
||||
"legacy_dataset_permission": permission_value.value,
|
||||
"legacy_partial_member_ids": partial_member_ids if should_bind_partial_members else [],
|
||||
},
|
||||
"after": {
|
||||
"rbac_whitelist_scope": scope.value,
|
||||
},
|
||||
"call": {
|
||||
"method": "RBACService.DatasetAccess.replace_whitelist",
|
||||
"kwargs": {
|
||||
"tenant_id": workspace_id,
|
||||
"account_id": operator_account_id,
|
||||
"dataset_id": current_dataset_id,
|
||||
"payload": replace_whitelist_payload.model_dump(mode="json"),
|
||||
},
|
||||
},
|
||||
}
|
||||
)
|
||||
if not dry_run:
|
||||
RBACService.DatasetAccess.replace_whitelist(
|
||||
tenant_id=workspace_id,
|
||||
account_id=operator_account_id,
|
||||
dataset_id=current_dataset_id,
|
||||
payload=replace_whitelist_payload,
|
||||
)
|
||||
scope_migrated_count += 1
|
||||
|
||||
if should_bind_partial_members:
|
||||
partial_dataset_count += 1
|
||||
for member_account_id in partial_member_ids:
|
||||
replace_user_access_policies_payload = ReplaceUserAccessPolicies(
|
||||
access_policy_ids=[_RBAC_DEFAULT_ACCESS_POLICY_ID],
|
||||
)
|
||||
if dry_run:
|
||||
_emit_dataset_permission_migration_event(
|
||||
{
|
||||
"event": "dataset_permission_migration_proposed_change",
|
||||
"action": "replace_user_access_policies",
|
||||
"dry_run": True,
|
||||
"tenant_id": workspace_id,
|
||||
"dataset_id": current_dataset_id,
|
||||
"operator_account_id": operator_account_id,
|
||||
"target_account_id": member_account_id,
|
||||
"before": {
|
||||
"legacy_dataset_permission": permission_value.value,
|
||||
"legacy_partial_member_id": member_account_id,
|
||||
},
|
||||
"after": {
|
||||
"rbac_user_access_policy_ids": [_RBAC_DEFAULT_ACCESS_POLICY_ID],
|
||||
},
|
||||
"call": {
|
||||
"method": "RBACService.DatasetAccess.replace_user_access_policies",
|
||||
"kwargs": {
|
||||
"tenant_id": workspace_id,
|
||||
"account_id": operator_account_id,
|
||||
"dataset_id": current_dataset_id,
|
||||
"target_account_id": member_account_id,
|
||||
"payload": replace_user_access_policies_payload.model_dump(
|
||||
mode="json", exclude_unset=True
|
||||
),
|
||||
},
|
||||
},
|
||||
}
|
||||
)
|
||||
continue
|
||||
RBACService.DatasetAccess.replace_user_access_policies(
|
||||
tenant_id=workspace_id,
|
||||
account_id=operator_account_id,
|
||||
dataset_id=current_dataset_id,
|
||||
target_account_id=member_account_id,
|
||||
payload=replace_user_access_policies_payload,
|
||||
)
|
||||
user_policy_migrated_count += 1
|
||||
|
||||
last_dataset_id = dataset_ids[-1]
|
||||
|
||||
if dataset_id:
|
||||
break
|
||||
|
||||
if scanned_count == 0:
|
||||
click.echo(click.style("No datasets found for migration.", fg="yellow"))
|
||||
return
|
||||
|
||||
if dry_run:
|
||||
click.echo(click.style("Dry run completed. No RBAC bindings were written.", fg="yellow"))
|
||||
click.echo(
|
||||
click.style(
|
||||
f"Dry run completed. Scanned {scanned_count} datasets; "
|
||||
f"{partial_dataset_count} partial-member datasets would be migrated.",
|
||||
fg="yellow",
|
||||
)
|
||||
)
|
||||
else:
|
||||
click.echo(click.style(f"RBAC member-role migration completed. Migrated {migrated_count} members.", fg="green"))
|
||||
click.echo(
|
||||
click.style(
|
||||
"RBAC dataset permission migration completed. "
|
||||
f"Scanned {scanned_count} datasets, migrated {scope_migrated_count} scopes, "
|
||||
f"wrote {user_policy_migrated_count} user default-policy bindings.",
|
||||
fg="green",
|
||||
)
|
||||
)
|
||||
|
||||
@@ -34,6 +34,12 @@ class EnterpriseFeatureConfig(BaseSettings):
|
||||
default=False,
|
||||
)
|
||||
|
||||
ENTERPRISE_RBAC_REQUEST_TIMEOUT: int = Field(
|
||||
ge=1,
|
||||
description="Maximum timeout in seconds for inner RBAC requests.",
|
||||
default=30,
|
||||
)
|
||||
|
||||
|
||||
class EnterpriseTelemetryConfig(BaseSettings):
|
||||
"""
|
||||
|
||||
@@ -43,6 +43,7 @@ from controllers.console.wraps import (
|
||||
from core.ops.ops_trace_manager import OpsTraceManager
|
||||
from core.rag.entities import PreProcessingRule, Rule, Segmentation
|
||||
from core.rag.retrieval.retrieval_methods import RetrievalMethod
|
||||
from core.rbac import RBACResourceWhitelistScope
|
||||
from core.trigger.constants import TRIGGER_NODE_TYPES
|
||||
from extensions.ext_database import db
|
||||
from fields.base import ResponseModel
|
||||
@@ -69,6 +70,7 @@ from services.entities.knowledge_entities.knowledge_entities import (
|
||||
WeightVectorSetting,
|
||||
)
|
||||
from services.feature_service import FeatureService
|
||||
from tasks.initialize_created_app_rbac_access_task import initialize_created_app_rbac_access_task
|
||||
|
||||
ALLOW_CREATE_APP_MODES = ["chat", "agent-chat", "advanced-chat", "workflow", "completion"]
|
||||
|
||||
@@ -666,6 +668,14 @@ class AppListApi(Resource):
|
||||
|
||||
app_service = AppService()
|
||||
app = app_service.create_app(current_tenant_id, params, current_user)
|
||||
if dify_config.RBAC_ENABLED:
|
||||
enterprise_rbac_service.RBACService.AppAccess.replace_whitelist(
|
||||
tenant_id=str(current_tenant_id),
|
||||
account_id=current_user.id,
|
||||
app_id=str(app.id),
|
||||
payload=enterprise_rbac_service.ReplaceMemberBindings(scope=RBACResourceWhitelistScope.ALL),
|
||||
)
|
||||
initialize_created_app_rbac_access_task.delay(current_tenant_id, current_user.id, app_id=app.id)
|
||||
permission_keys_map = enterprise_rbac_service.RBACService.AppPermissions.batch_get(
|
||||
str(current_tenant_id),
|
||||
current_user.id,
|
||||
|
||||
@@ -13,6 +13,7 @@ from controllers.console.wraps import (
|
||||
RBACPermission,
|
||||
RBACResourceScope,
|
||||
account_initialization_required,
|
||||
edit_permission_required,
|
||||
rbac_permission_required,
|
||||
setup_required,
|
||||
)
|
||||
@@ -95,9 +96,12 @@ class TraceAppConfigApi(Resource):
|
||||
console_ns.models[TraceAppConfigResponse.__name__],
|
||||
)
|
||||
@console_ns.response(400, "Invalid request parameters or configuration already exists")
|
||||
@console_ns.response(403, "Insufficient permissions")
|
||||
@setup_required
|
||||
@login_required
|
||||
@account_initialization_required
|
||||
@edit_permission_required
|
||||
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TRACING_CONFIG)
|
||||
@get_app_model
|
||||
def post(self, app_model: App):
|
||||
"""Create a new trace app configuration"""
|
||||
@@ -125,9 +129,12 @@ class TraceAppConfigApi(Resource):
|
||||
console_ns.models[TraceAppConfigResponse.__name__],
|
||||
)
|
||||
@console_ns.response(400, "Invalid request parameters or configuration not found")
|
||||
@console_ns.response(403, "Insufficient permissions")
|
||||
@setup_required
|
||||
@login_required
|
||||
@account_initialization_required
|
||||
@edit_permission_required
|
||||
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TRACING_CONFIG)
|
||||
@get_app_model
|
||||
def patch(self, app_model: App):
|
||||
"""Update an existing trace app configuration"""
|
||||
@@ -149,9 +156,12 @@ class TraceAppConfigApi(Resource):
|
||||
@console_ns.doc(params=query_params_from_model(TraceProviderQuery))
|
||||
@console_ns.response(204, "Tracing configuration deleted successfully")
|
||||
@console_ns.response(400, "Invalid request parameters or configuration not found")
|
||||
@console_ns.response(403, "Insufficient permissions")
|
||||
@setup_required
|
||||
@login_required
|
||||
@account_initialization_required
|
||||
@edit_permission_required
|
||||
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TRACING_CONFIG)
|
||||
@get_app_model
|
||||
def delete(self, app_model: App):
|
||||
"""Delete an existing trace app configuration"""
|
||||
|
||||
@@ -8,11 +8,8 @@ from controllers.common.schema import register_response_schema_models, register_
|
||||
from controllers.console import console_ns
|
||||
from controllers.console.app.wraps import get_app_model
|
||||
from controllers.console.wraps import (
|
||||
RBACPermission,
|
||||
RBACResourceScope,
|
||||
account_initialization_required,
|
||||
edit_permission_required,
|
||||
rbac_permission_required,
|
||||
setup_required,
|
||||
with_current_tenant_id,
|
||||
with_current_user,
|
||||
@@ -223,7 +220,6 @@ class WorkflowCommentListApi(Resource):
|
||||
@setup_required
|
||||
@account_initialization_required
|
||||
@with_current_tenant_id
|
||||
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT)
|
||||
@get_app_model()
|
||||
def get(self, current_tenant_id: str, app_model: App):
|
||||
"""Get all comments for a workflow."""
|
||||
@@ -240,7 +236,6 @@ class WorkflowCommentListApi(Resource):
|
||||
@setup_required
|
||||
@account_initialization_required
|
||||
@edit_permission_required
|
||||
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT)
|
||||
@with_current_user
|
||||
@with_current_tenant_id
|
||||
@get_app_model()
|
||||
@@ -273,7 +268,6 @@ class WorkflowCommentDetailApi(Resource):
|
||||
@setup_required
|
||||
@account_initialization_required
|
||||
@with_current_tenant_id
|
||||
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT)
|
||||
@get_app_model()
|
||||
def get(self, current_tenant_id: str, app_model: App, comment_id: str):
|
||||
"""Get a specific workflow comment."""
|
||||
@@ -292,7 +286,6 @@ class WorkflowCommentDetailApi(Resource):
|
||||
@setup_required
|
||||
@account_initialization_required
|
||||
@edit_permission_required
|
||||
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT)
|
||||
@with_current_user
|
||||
@with_current_tenant_id
|
||||
@get_app_model()
|
||||
@@ -321,7 +314,6 @@ class WorkflowCommentDetailApi(Resource):
|
||||
@setup_required
|
||||
@account_initialization_required
|
||||
@edit_permission_required
|
||||
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT)
|
||||
@with_current_user
|
||||
@with_current_tenant_id
|
||||
@get_app_model()
|
||||
@@ -349,7 +341,6 @@ class WorkflowCommentResolveApi(Resource):
|
||||
@setup_required
|
||||
@account_initialization_required
|
||||
@edit_permission_required
|
||||
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT)
|
||||
@with_current_user
|
||||
@with_current_tenant_id
|
||||
@get_app_model()
|
||||
@@ -378,7 +369,6 @@ class WorkflowCommentReplyApi(Resource):
|
||||
@setup_required
|
||||
@account_initialization_required
|
||||
@edit_permission_required
|
||||
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT)
|
||||
@with_current_user
|
||||
@with_current_tenant_id
|
||||
@get_app_model()
|
||||
@@ -414,7 +404,6 @@ class WorkflowCommentReplyDetailApi(Resource):
|
||||
@setup_required
|
||||
@account_initialization_required
|
||||
@edit_permission_required
|
||||
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT)
|
||||
@with_current_user
|
||||
@with_current_tenant_id
|
||||
@get_app_model()
|
||||
@@ -447,7 +436,6 @@ class WorkflowCommentReplyDetailApi(Resource):
|
||||
@setup_required
|
||||
@account_initialization_required
|
||||
@edit_permission_required
|
||||
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT)
|
||||
@with_current_user
|
||||
@with_current_tenant_id
|
||||
@get_app_model()
|
||||
@@ -483,7 +471,6 @@ class WorkflowCommentMentionUsersApi(Resource):
|
||||
@setup_required
|
||||
@account_initialization_required
|
||||
@with_current_user
|
||||
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT)
|
||||
@get_app_model()
|
||||
def get(self, current_user: Account, app_model: App):
|
||||
"""Get all users in current tenant for mentions."""
|
||||
|
||||
@@ -19,7 +19,8 @@ from controllers.console.app.error import AppNotFoundError
|
||||
from core.db.session_factory import session_factory
|
||||
from extensions.ext_database import db
|
||||
from libs.login import current_account_with_tenant
|
||||
from models import App, AppMode
|
||||
from models import App, AppMode, TrialApp
|
||||
from services.recommended_app_service import RecommendedAppService
|
||||
|
||||
|
||||
def _load_app_model(session: Session, app_id: str) -> App | None:
|
||||
@@ -41,7 +42,10 @@ def _load_app_model_from_scoped_session(app_id: str) -> App | None:
|
||||
|
||||
|
||||
def _load_app_model_with_trial(app_id: str) -> App | None:
|
||||
app_model = db.session.scalar(select(App).where(App.id == app_id, App.status == "normal").limit(1))
|
||||
"""Load a normal app through its trial registration without applying current-tenant scope."""
|
||||
app_model = db.session.scalar(
|
||||
select(App).join(TrialApp, TrialApp.app_id == App.id).where(App.id == app_id, App.status == "normal").limit(1)
|
||||
)
|
||||
return app_model
|
||||
|
||||
|
||||
@@ -195,6 +199,8 @@ def get_app_model_with_trial[**P, R](
|
||||
*,
|
||||
mode: AppMode | list[AppMode] | None = None,
|
||||
) -> Callable[P, R] | Callable[[Callable[P, R]], Callable[P, R]]:
|
||||
"""Inject an app registered for trial or available from the recommended catalog."""
|
||||
|
||||
def decorator(view_func: Callable[P, R]) -> Callable[P, R]:
|
||||
@wraps(view_func)
|
||||
def decorated_view(*args: P.args, **kwargs: P.kwargs) -> R:
|
||||
@@ -207,6 +213,8 @@ def get_app_model_with_trial[**P, R](
|
||||
del kwargs["app_id"]
|
||||
|
||||
app_model = _load_app_model_with_trial(app_id)
|
||||
if app_model is None:
|
||||
app_model = RecommendedAppService.get_app(app_id, session=db.session())
|
||||
|
||||
if not app_model:
|
||||
raise AppNotFoundError()
|
||||
|
||||
@@ -23,9 +23,9 @@ from libs.password import valid_password
|
||||
from models import Account
|
||||
from services.account_service import AccountService
|
||||
from services.billing_service import BillingService
|
||||
from services.errors.account import AccountRegisterError
|
||||
from services.errors.account import AccountRegisterError, SeatsLimitExceededError
|
||||
|
||||
from ..error import AccountInFreezeError, EmailSendIpLimitError
|
||||
from ..error import AccountInFreezeError, EmailSendIpLimitError, SeatsLimitExceeded
|
||||
from ..wraps import email_password_login_enabled, email_register_enabled, setup_required
|
||||
|
||||
|
||||
@@ -208,5 +208,7 @@ class EmailRegisterResetApi(Resource):
|
||||
timezone=timezone,
|
||||
session=db.session,
|
||||
)
|
||||
except SeatsLimitExceededError:
|
||||
raise SeatsLimitExceeded()
|
||||
except AccountRegisterError:
|
||||
raise AccountInFreezeError()
|
||||
|
||||
@@ -25,6 +25,7 @@ from controllers.console.error import (
|
||||
AccountNotFound,
|
||||
EmailSendIpLimitError,
|
||||
NotAllowedCreateWorkspace,
|
||||
SeatsLimitExceeded,
|
||||
WorkspacesLimitExceeded,
|
||||
)
|
||||
from controllers.console.wraps import (
|
||||
@@ -51,7 +52,7 @@ from models.account import Account
|
||||
from services.account_service import AccountService, InvitationDetailDict, RegisterService, TenantService
|
||||
from services.billing_service import BillingService
|
||||
from services.entities.auth_entities import LoginFailureReason, LoginPayloadBase
|
||||
from services.errors.account import AccountRegisterError
|
||||
from services.errors.account import AccountRegisterError, SeatsLimitExceededError
|
||||
from services.errors.workspace import WorkSpaceNotAllowedCreateError, WorkspacesLimitExceededError
|
||||
from services.feature_service import FeatureService
|
||||
|
||||
@@ -317,6 +318,8 @@ class EmailCodeLoginApi(Resource):
|
||||
)
|
||||
except WorkSpaceNotAllowedCreateError:
|
||||
raise NotAllowedCreateWorkspace()
|
||||
except SeatsLimitExceededError:
|
||||
raise SeatsLimitExceeded()
|
||||
except AccountRegisterError:
|
||||
_log_console_login_failure(email=user_email, reason=LoginFailureReason.ACCOUNT_IN_FREEZE)
|
||||
raise AccountInFreezeError()
|
||||
|
||||
@@ -25,7 +25,7 @@ from libs.token import (
|
||||
from models import Account, AccountStatus
|
||||
from services.account_service import AccountService, RegisterService, TenantService
|
||||
from services.billing_service import BillingService
|
||||
from services.errors.account import AccountNotFoundError, AccountRegisterError
|
||||
from services.errors.account import AccountNotFoundError, AccountRegisterError, SeatsLimitExceededError
|
||||
from services.errors.workspace import WorkSpaceNotAllowedCreateError, WorkSpaceNotFoundError
|
||||
from services.feature_service import FeatureService
|
||||
|
||||
@@ -38,6 +38,7 @@ class OAuthLoginQuery(BaseModel):
|
||||
invite_token: str | None = Field(default=None, description="Optional invitation token")
|
||||
timezone: str | None = Field(default=None, description="Preferred timezone")
|
||||
language: str | None = Field(default=None, description="Preferred interface language")
|
||||
redirect_url: str | None = Field(default=None, description="Relative page to resume after login")
|
||||
|
||||
|
||||
class OAuthCallbackQuery(BaseModel):
|
||||
@@ -87,6 +88,36 @@ def _validated_language(value: str | None) -> str | None:
|
||||
return None
|
||||
|
||||
|
||||
def _url_origin(url: str) -> tuple[str, str, int] | None:
|
||||
parsed_url = urllib.parse.urlsplit(url)
|
||||
if parsed_url.scheme not in {"http", "https"} or parsed_url.hostname is None:
|
||||
return None
|
||||
|
||||
try:
|
||||
port = parsed_url.port
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
if port is None:
|
||||
port = 443 if parsed_url.scheme == "https" else 80
|
||||
return parsed_url.scheme, parsed_url.hostname, port
|
||||
|
||||
|
||||
def _get_redirect_target(redirect_url: str | None) -> str:
|
||||
if not redirect_url:
|
||||
return dify_config.CONSOLE_WEB_URL
|
||||
|
||||
parsed_url = urllib.parse.urlsplit(redirect_url)
|
||||
normalized_path = redirect_url.lstrip().replace("\\", "/")
|
||||
if not parsed_url.scheme and not parsed_url.netloc and not normalized_path.startswith("//"):
|
||||
return redirect_url
|
||||
|
||||
redirect_origin = _url_origin(redirect_url)
|
||||
if redirect_origin is not None and redirect_origin == _url_origin(dify_config.CONSOLE_WEB_URL):
|
||||
return redirect_url
|
||||
return dify_config.CONSOLE_WEB_URL
|
||||
|
||||
|
||||
def _preferred_interface_language(language: str | None = None) -> str:
|
||||
if language:
|
||||
return language
|
||||
@@ -109,6 +140,7 @@ class OAuthLogin(Resource):
|
||||
invite_token = request.args.get("invite_token") or None
|
||||
timezone = _validated_timezone(request.args.get("timezone") or None)
|
||||
language = _validated_language(request.args.get("language") or None)
|
||||
redirect_url = request.args.get("redirect_url") or None
|
||||
OAUTH_PROVIDERS = get_oauth_providers()
|
||||
with current_app.app_context():
|
||||
oauth_provider = OAUTH_PROVIDERS.get(provider)
|
||||
@@ -119,6 +151,7 @@ class OAuthLogin(Resource):
|
||||
invite_token=invite_token,
|
||||
timezone=timezone,
|
||||
language=language,
|
||||
redirect_url=redirect_url,
|
||||
)
|
||||
return redirect(auth_url)
|
||||
|
||||
@@ -144,6 +177,7 @@ class OAuthCallback(Resource):
|
||||
invite_token = oauth_state.get("invite_token")
|
||||
timezone = _validated_timezone(oauth_state.get("timezone"))
|
||||
language = _validated_language(oauth_state.get("language"))
|
||||
redirect_url = oauth_state.get("redirect_url")
|
||||
|
||||
if not code:
|
||||
return {"error": "Authorization code is required"}, 400
|
||||
@@ -182,6 +216,8 @@ class OAuthCallback(Resource):
|
||||
f"{dify_config.CONSOLE_WEB_URL}/signin"
|
||||
"?message=Workspace not found, please contact system admin to invite you to join in a workspace."
|
||||
)
|
||||
except SeatsLimitExceededError:
|
||||
return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin?message=Licensed seats limit exceeded.")
|
||||
except AccountRegisterError as e:
|
||||
return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin?message={e.description}")
|
||||
|
||||
@@ -210,9 +246,9 @@ class OAuthCallback(Resource):
|
||||
ip_address=extract_remote_ip(request),
|
||||
)
|
||||
|
||||
base_url = dify_config.CONSOLE_WEB_URL
|
||||
query_char = "&" if "?" in base_url else "?"
|
||||
target_url = f"{base_url}{query_char}oauth_new_user={str(oauth_new_user).lower()}"
|
||||
target_url = _get_redirect_target(redirect_url)
|
||||
query_char = "&" if "?" in target_url else "?"
|
||||
target_url = f"{target_url}{query_char}oauth_new_user={str(oauth_new_user).lower()}"
|
||||
response = redirect(target_url)
|
||||
|
||||
set_access_token_to_cookie(request, response, token_pair.access_token)
|
||||
|
||||
@@ -50,6 +50,8 @@ from models.provider_ids import ModelProviderID
|
||||
from services.api_token_service import ApiTokenCache
|
||||
from services.dataset_service import DatasetPermissionService, DatasetService, DocumentService
|
||||
from services.enterprise import rbac_service as enterprise_rbac_service
|
||||
from services.enterprise.rbac_service import RBACResourceWhitelistScope, ReplaceMemberBindings
|
||||
from tasks.initialize_created_app_rbac_access_task import initialize_created_app_rbac_access_task
|
||||
|
||||
register_response_schema_models(console_ns, ApiBaseUrlResponse, SimpleResultResponse, UsageCheckResponse)
|
||||
|
||||
@@ -565,6 +567,16 @@ class DatasetListApi(Resource):
|
||||
except services.errors.dataset.DatasetNameDuplicateError:
|
||||
raise DatasetNameDuplicateError()
|
||||
|
||||
if dify_config.RBAC_ENABLED:
|
||||
if permission == DatasetPermissionEnum.ALL_TEAM:
|
||||
enterprise_rbac_service.RBACService.DatasetAccess.replace_whitelist(
|
||||
current_tenant_id,
|
||||
current_user.id,
|
||||
dataset.id,
|
||||
ReplaceMemberBindings(scope=RBACResourceWhitelistScope.ALL),
|
||||
)
|
||||
initialize_created_app_rbac_access_task.delay(current_tenant_id, current_user.id, dataset_id=dataset.id)
|
||||
|
||||
permission_keys_map = enterprise_rbac_service.RBACService.DatasetPermissions.batch_get(
|
||||
str(current_tenant_id),
|
||||
current_user.id,
|
||||
|
||||
@@ -58,6 +58,12 @@ class WorkspacesLimitExceeded(BaseHTTPException):
|
||||
code = 400
|
||||
|
||||
|
||||
class SeatsLimitExceeded(BaseHTTPException):
|
||||
error_code = "limit_exceeded"
|
||||
description = "Unable to create account because the licensed seats limit was exceeded"
|
||||
code = 400
|
||||
|
||||
|
||||
class AccountBannedError(BaseHTTPException):
|
||||
error_code = "account_banned"
|
||||
description = "Account is banned."
|
||||
|
||||
@@ -36,7 +36,7 @@ from libs.login import current_account_with_tenant, login_required
|
||||
from models.account import Account, TenantAccountJoin, TenantAccountRole
|
||||
from services.account_service import AccountService, RegisterService, TenantService
|
||||
from services.enterprise import rbac_service as enterprise_rbac_service
|
||||
from services.errors.account import AccountAlreadyInTenantError
|
||||
from services.errors.account import AccountAlreadyInTenantError, SeatsLimitExceededError
|
||||
from services.feature_service import FeatureService
|
||||
|
||||
|
||||
@@ -291,6 +291,14 @@ class MemberInviteEmailApi(Resource):
|
||||
"message": "Account already in workspace.",
|
||||
}
|
||||
)
|
||||
except SeatsLimitExceededError:
|
||||
invitation_results.append(
|
||||
MemberInviteFailedResponse(
|
||||
status="failed",
|
||||
email=invitee_email,
|
||||
message="Licensed seats limit exceeded.",
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
invitation_results.append({"status": "failed", "email": invitee_email, "message": str(e)})
|
||||
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from enum import StrEnum
|
||||
from typing import Any
|
||||
|
||||
from flask import request
|
||||
@@ -14,9 +13,11 @@ from controllers.common.schema import register_response_schema_models
|
||||
from controllers.console import console_ns
|
||||
from controllers.console.wraps import RBACPermission, RBACResourceScope, rbac_permission_required
|
||||
from core.db.session_factory import session_factory
|
||||
from core.rbac import RBACResourceWhitelistScope
|
||||
from libs.login import current_account_with_tenant, login_required
|
||||
from models import Account
|
||||
from services.enterprise import rbac_service as svc
|
||||
from tasks.initialize_created_app_rbac_access_task import initialize_created_app_rbac_access_task
|
||||
|
||||
|
||||
class _RBACRoleList(svc.Paginated[svc.RBACRole]):
|
||||
@@ -511,14 +512,8 @@ class RBACAccessPolicyBindingUnlockApi(Resource):
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _AccessScope(StrEnum):
|
||||
ALL = "all"
|
||||
SPECIFIC = "specific"
|
||||
ONLY_ME = "only_me"
|
||||
|
||||
|
||||
class _ResourceAccessScopeRequest(BaseModel):
|
||||
scope: _AccessScope
|
||||
scope: RBACResourceWhitelistScope
|
||||
|
||||
|
||||
class _ReplaceBindingsRequest(BaseModel):
|
||||
@@ -584,14 +579,15 @@ class RBACAppWhitelistApi(Resource):
|
||||
def put(self, app_id):
|
||||
tenant_id, account_id = _current_ids()
|
||||
request = _payload(_ResourceAccessScopeRequest)
|
||||
return _dump(
|
||||
svc.RBACService.AppAccess.replace_whitelist(
|
||||
tenant_id,
|
||||
account_id,
|
||||
str(app_id),
|
||||
svc.ReplaceMemberBindings(scope=request.scope.value),
|
||||
)
|
||||
result = svc.RBACService.AppAccess.replace_whitelist(
|
||||
tenant_id,
|
||||
account_id,
|
||||
str(app_id),
|
||||
svc.ReplaceMemberBindings(scope=request.scope.value),
|
||||
)
|
||||
if dify_config.RBAC_ENABLED and request.scope is RBACResourceWhitelistScope.ALL:
|
||||
initialize_created_app_rbac_access_task.delay(tenant_id, account_id, str(app_id))
|
||||
return _dump(result)
|
||||
|
||||
|
||||
@console_ns.route("/workspaces/current/rbac/apps/<uuid:app_id>/user-access-policies")
|
||||
@@ -616,8 +612,8 @@ class RBACAppUserAccessPolicyAssignmentApi(Resource):
|
||||
svc.RBACService.AppAccess.replace_user_access_policies(
|
||||
tenant_id,
|
||||
account_id,
|
||||
str(app_id),
|
||||
str(target_account_id),
|
||||
app_id,
|
||||
target_account_id,
|
||||
payload,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -48,6 +48,7 @@ from services.errors.account import (
|
||||
MemberNotInTenantError,
|
||||
NoPermissionError,
|
||||
RoleAlreadyAssignedError,
|
||||
SeatsLimitExceededError,
|
||||
)
|
||||
from services.feature_service import FeatureService
|
||||
|
||||
@@ -190,6 +191,8 @@ class WorkspaceMembersApi(Resource):
|
||||
raise BadRequest(str(exc))
|
||||
except NoPermissionError as exc:
|
||||
raise BadRequest(str(exc))
|
||||
except SeatsLimitExceededError:
|
||||
raise BadRequest("licensed seats limit exceeded")
|
||||
except AccountRegisterError as exc:
|
||||
raise BadRequest(str(exc))
|
||||
|
||||
|
||||
@@ -43,6 +43,7 @@ from models.dataset import DatasetPermissionEnum
|
||||
from models.enums import TagType
|
||||
from models.provider_ids import ModelProviderID
|
||||
from services.dataset_service import DatasetPermissionService, DatasetService, DocumentService
|
||||
from services.enterprise.rbac_service import RBACResourceWhitelistScope, RBACService, ReplaceMemberBindings
|
||||
from services.entities.knowledge_entities.knowledge_entities import (
|
||||
ExternalRetrievalModel,
|
||||
KnowledgeProvider,
|
||||
@@ -58,6 +59,7 @@ from services.tag_service import (
|
||||
from services.tag_service import (
|
||||
UpdateTagPayload as UpdateTagServicePayload,
|
||||
)
|
||||
from tasks.initialize_created_app_rbac_access_task import initialize_created_app_rbac_access_task
|
||||
|
||||
register_enum_models(service_api_ns, DatasetPermissionEnum)
|
||||
|
||||
@@ -523,6 +525,15 @@ class DatasetListApi(DatasetApiResource):
|
||||
except services.errors.dataset.DatasetNameDuplicateError:
|
||||
raise DatasetNameDuplicateError()
|
||||
|
||||
if payload.permission == DatasetPermissionEnum.ALL_TEAM and dify_config.RBAC_ENABLED:
|
||||
RBACService.DatasetAccess.replace_whitelist(
|
||||
tenant_id,
|
||||
current_user.id,
|
||||
dataset.id,
|
||||
ReplaceMemberBindings(scope=RBACResourceWhitelistScope.ALL),
|
||||
)
|
||||
initialize_created_app_rbac_access_task.delay(tenant_id, current_user.id, dataset_id=dataset.id)
|
||||
|
||||
return _dump_service_dataset_detail(dataset), 200
|
||||
|
||||
|
||||
|
||||
@@ -12,6 +12,7 @@ from flask_restx import Resource
|
||||
from flask_restx.utils import merge
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
from werkzeug.exceptions import Forbidden, NotFound, Unauthorized
|
||||
|
||||
from configs import dify_config
|
||||
@@ -269,8 +270,8 @@ def cloud_edition_billing_rate_limit_check[**P, R](
|
||||
subscription_plan=knowledge_rate_limit.subscription_plan,
|
||||
operation="knowledge",
|
||||
)
|
||||
db.session.add(rate_limit_log)
|
||||
db.session.commit()
|
||||
with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as session:
|
||||
session.add(rate_limit_log)
|
||||
raise Forbidden(
|
||||
"Sorry, you have reached the knowledge base request rate limit of your subscription."
|
||||
)
|
||||
|
||||
@@ -47,6 +47,7 @@ from core.repositories import DifyCoreRepositoryFactory
|
||||
from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository
|
||||
from extensions.ext_database import db
|
||||
from factories import file_factory
|
||||
from graphon.filters import ResponseStreamFilter
|
||||
from graphon.graph_engine.layers import GraphEngineLayer
|
||||
from graphon.model_runtime.errors.invoke import InvokeAuthorizationError
|
||||
from graphon.runtime import GraphRuntimeState
|
||||
@@ -268,6 +269,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
||||
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
|
||||
graph_runtime_state: GraphRuntimeState,
|
||||
pause_state_config: PauseStateLayerConfig | None = None,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
):
|
||||
"""
|
||||
Resume a paused advanced chat execution.
|
||||
@@ -297,6 +299,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
||||
stream=application_generate_entity.stream,
|
||||
pause_state_config=pause_state_config,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
response_stream_filter=response_stream_filter,
|
||||
)
|
||||
|
||||
def single_iteration_generate(
|
||||
@@ -491,6 +494,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
||||
pause_state_config: PauseStateLayerConfig | None = None,
|
||||
graph_runtime_state: GraphRuntimeState | None = None,
|
||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
|
||||
"""
|
||||
Generate App response.
|
||||
@@ -538,12 +542,14 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
||||
)
|
||||
|
||||
graph_layers: list[GraphEngineLayer] = list(graph_engine_layers)
|
||||
resolved_response_stream_filter = response_stream_filter or ResponseStreamFilter()
|
||||
if pause_state_config is not None:
|
||||
graph_layers.append(
|
||||
PauseStatePersistenceLayer(
|
||||
session_factory=pause_state_config.session_factory,
|
||||
generate_entity=application_generate_entity,
|
||||
state_owner_user_id=pause_state_config.state_owner_user_id,
|
||||
response_stream_filter=resolved_response_stream_filter,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -564,6 +570,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
||||
"workflow_node_execution_repository": workflow_node_execution_repository,
|
||||
"graph_engine_layers": tuple(graph_layers),
|
||||
"graph_runtime_state": graph_runtime_state,
|
||||
"response_stream_filter": resolved_response_stream_filter,
|
||||
},
|
||||
)
|
||||
|
||||
@@ -603,6 +610,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
||||
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
|
||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||
graph_runtime_state: GraphRuntimeState | None = None,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
):
|
||||
"""
|
||||
Generate worker in a new thread.
|
||||
@@ -662,6 +670,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
||||
workflow_node_execution_repository=workflow_node_execution_repository,
|
||||
graph_engine_layers=graph_engine_layers,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
response_stream_filter=response_stream_filter,
|
||||
)
|
||||
|
||||
try:
|
||||
|
||||
@@ -39,6 +39,7 @@ from extensions.ext_database import db
|
||||
from extensions.ext_redis import redis_client
|
||||
from extensions.otel import WorkflowAppRunnerHandler, trace_span
|
||||
from graphon.enums import WorkflowType
|
||||
from graphon.filters import ResponseStreamFilter
|
||||
from graphon.graph_engine.command_channels import RedisChannel
|
||||
from graphon.graph_engine.layers import GraphEngineLayer
|
||||
from graphon.runtime import GraphRuntimeState, VariablePool
|
||||
@@ -73,6 +74,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
|
||||
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
|
||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||
graph_runtime_state: GraphRuntimeState | None = None,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
):
|
||||
super().__init__(
|
||||
queue_manager=queue_manager,
|
||||
@@ -90,6 +92,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
|
||||
self._workflow_execution_repository = workflow_execution_repository
|
||||
self._workflow_node_execution_repository = workflow_node_execution_repository
|
||||
self._resume_graph_runtime_state = graph_runtime_state
|
||||
self._response_stream_filter = response_stream_filter
|
||||
|
||||
@trace_span(WorkflowAppRunnerHandler)
|
||||
def run(self):
|
||||
@@ -227,6 +230,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
|
||||
variable_pool=variable_pool,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
command_channel=command_channel,
|
||||
response_stream_filter=self._response_stream_filter,
|
||||
)
|
||||
|
||||
self._queue_manager.graph_runtime_state = graph_runtime_state
|
||||
|
||||
@@ -42,6 +42,7 @@ from core.repositories import DifyCoreRepositoryFactory
|
||||
from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository
|
||||
from extensions.ext_database import db
|
||||
from factories import file_factory
|
||||
from graphon.filters import ResponseStreamFilter
|
||||
from graphon.graph_engine.layers import GraphEngineLayer
|
||||
from graphon.model_runtime.errors.invoke import InvokeAuthorizationError
|
||||
from graphon.runtime import GraphRuntimeState
|
||||
@@ -280,6 +281,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||
pause_state_config: PauseStateLayerConfig | None = None,
|
||||
variable_loader: VariableLoader = DUMMY_VARIABLE_LOADER,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
|
||||
"""
|
||||
Resume a paused workflow execution using the persisted runtime state.
|
||||
@@ -310,6 +312,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
graph_engine_layers=graph_engine_layers,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
pause_state_config=pause_state_config,
|
||||
response_stream_filter=response_stream_filter,
|
||||
)
|
||||
|
||||
def _generate(
|
||||
@@ -328,6 +331,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||
graph_runtime_state: GraphRuntimeState | None = None,
|
||||
pause_state_config: PauseStateLayerConfig | None = None,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
|
||||
"""
|
||||
Generate App response.
|
||||
@@ -356,12 +360,14 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
app_mode=app_model.mode,
|
||||
)
|
||||
|
||||
resolved_response_stream_filter = response_stream_filter or ResponseStreamFilter()
|
||||
if pause_state_config is not None:
|
||||
graph_layers.append(
|
||||
PauseStatePersistenceLayer(
|
||||
session_factory=pause_state_config.session_factory,
|
||||
generate_entity=application_generate_entity,
|
||||
state_owner_user_id=pause_state_config.state_owner_user_id,
|
||||
response_stream_filter=resolved_response_stream_filter,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -384,6 +390,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
"workflow_node_execution_repository": workflow_node_execution_repository,
|
||||
"graph_engine_layers": tuple(graph_layers),
|
||||
"graph_runtime_state": graph_runtime_state,
|
||||
"response_stream_filter": resolved_response_stream_filter,
|
||||
},
|
||||
)
|
||||
|
||||
@@ -590,6 +597,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
root_node_id: str | None = None,
|
||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||
graph_runtime_state: GraphRuntimeState | None = None,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Generate worker in a new thread.
|
||||
@@ -638,6 +646,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
root_node_id=root_node_id,
|
||||
graph_engine_layers=graph_engine_layers,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
response_stream_filter=response_stream_filter,
|
||||
)
|
||||
|
||||
try:
|
||||
|
||||
@@ -18,6 +18,7 @@ from core.workflow.workflow_entry import WorkflowEntry
|
||||
from extensions.ext_redis import redis_client
|
||||
from extensions.otel import WorkflowAppRunnerHandler, trace_span
|
||||
from graphon.enums import WorkflowType
|
||||
from graphon.filters import ResponseStreamFilter
|
||||
from graphon.graph_engine.command_channels import RedisChannel
|
||||
from graphon.graph_engine.layers import GraphEngineLayer
|
||||
from graphon.runtime import GraphRuntimeState, VariablePool
|
||||
@@ -46,6 +47,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
||||
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
|
||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||
graph_runtime_state: GraphRuntimeState | None = None,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
):
|
||||
super().__init__(
|
||||
queue_manager=queue_manager,
|
||||
@@ -60,6 +62,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
||||
self._workflow_execution_repository = workflow_execution_repository
|
||||
self._workflow_node_execution_repository = workflow_node_execution_repository
|
||||
self._resume_graph_runtime_state = graph_runtime_state
|
||||
self._response_stream_filter = response_stream_filter
|
||||
|
||||
@trace_span(WorkflowAppRunnerHandler)
|
||||
def run(self):
|
||||
@@ -163,6 +166,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
||||
variable_pool=variable_pool,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
command_channel=command_channel,
|
||||
response_stream_filter=self._response_stream_filter,
|
||||
)
|
||||
|
||||
persistence_layer = WorkflowPersistenceLayer(
|
||||
|
||||
@@ -7,6 +7,7 @@ from sqlalchemy.orm import Session, sessionmaker
|
||||
|
||||
from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity, WorkflowAppGenerateEntity
|
||||
from core.workflow.system_variables import SystemVariableKey, get_system_text
|
||||
from graphon.filters import ResponseStreamFilter
|
||||
from graphon.graph_engine.layers import GraphEngineLayer
|
||||
from graphon.graph_events import GraphEngineEvent, GraphRunPausedEvent
|
||||
from models.model import AppMode
|
||||
@@ -41,6 +42,10 @@ class WorkflowResumptionContext(BaseModel):
|
||||
# Only workflow / chatflow could be paused.
|
||||
generate_entity: _GenerateEntityUnion
|
||||
serialized_graph_runtime_state: str
|
||||
# Optional so that a workflow run paused before this field existed still
|
||||
# loads: it just degrades to fresh-filter behavior on resume for that one
|
||||
# stale run.
|
||||
serialized_response_stream_filter_state: str | None = None
|
||||
|
||||
def dumps(self) -> str:
|
||||
return self.model_dump_json()
|
||||
@@ -52,6 +57,12 @@ class WorkflowResumptionContext(BaseModel):
|
||||
def get_generate_entity(self) -> WorkflowAppGenerateEntity | AdvancedChatAppGenerateEntity:
|
||||
return self.generate_entity.entity
|
||||
|
||||
def get_response_stream_filter(self) -> ResponseStreamFilter:
|
||||
response_stream_filter = ResponseStreamFilter()
|
||||
if self.serialized_response_stream_filter_state is not None:
|
||||
response_stream_filter.loads(self.serialized_response_stream_filter_state)
|
||||
return response_stream_filter
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PauseStateLayerConfig:
|
||||
@@ -67,11 +78,17 @@ class PauseStatePersistenceLayer(GraphEngineLayer):
|
||||
session_factory: Engine | sessionmaker[Session],
|
||||
generate_entity: WorkflowAppGenerateEntity | AdvancedChatAppGenerateEntity,
|
||||
state_owner_user_id: str,
|
||||
response_stream_filter: ResponseStreamFilter,
|
||||
):
|
||||
"""Create a PauseStatePersistenceLayer.
|
||||
|
||||
The `state_owner_user_id` is used when creating state file for pause.
|
||||
It generally should id of the creator of workflow.
|
||||
|
||||
`response_stream_filter` must be the exact same instance that
|
||||
`WorkflowEntry` is using to stream this run's events — this layer
|
||||
dumps its state on pause, and a different instance would silently
|
||||
persist the wrong (empty) filter state.
|
||||
"""
|
||||
if isinstance(session_factory, Engine):
|
||||
session_factory = sessionmaker(session_factory)
|
||||
@@ -79,6 +96,7 @@ class PauseStatePersistenceLayer(GraphEngineLayer):
|
||||
self._session_maker = session_factory
|
||||
self._state_owner_user_id = state_owner_user_id
|
||||
self._generate_entity = generate_entity
|
||||
self._response_stream_filter = response_stream_filter
|
||||
|
||||
def _get_repo(self) -> APIWorkflowRunRepository:
|
||||
return DifyAPIRepositoryFactory.create_api_workflow_run_repository(self._session_maker)
|
||||
@@ -119,6 +137,7 @@ class PauseStatePersistenceLayer(GraphEngineLayer):
|
||||
state = WorkflowResumptionContext(
|
||||
serialized_graph_runtime_state=self.graph_runtime_state.dumps(),
|
||||
generate_entity=entity_wrapper,
|
||||
serialized_response_stream_filter_state=self._response_stream_filter.dumps(),
|
||||
)
|
||||
|
||||
workflow_run_id = get_system_text(
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
from collections.abc import Generator, Iterable, Mapping
|
||||
from typing import Any
|
||||
|
||||
from configs import dify_config
|
||||
from core.callback_handler.agent_tool_callback_handler import DifyAgentCallbackHandler, print_text
|
||||
from core.ops.ops_trace_manager import TraceQueueManager
|
||||
from core.tools.entities.tool_entities import ToolInvokeMessage
|
||||
@@ -19,8 +20,9 @@ class DifyWorkflowCallbackHandler(DifyAgentCallbackHandler):
|
||||
trace_manager: TraceQueueManager | None = None,
|
||||
) -> Generator[ToolInvokeMessage, None, None]:
|
||||
for tool_output in tool_outputs:
|
||||
print_text("\n[on_tool_execution]\n", color=self.color)
|
||||
print_text("Tool: " + tool_name + "\n", color=self.color)
|
||||
print_text("Outputs: " + tool_output.model_dump_json()[:1000] + "\n", color=self.color)
|
||||
print_text("\n")
|
||||
if dify_config.DEBUG:
|
||||
print_text("\n[on_tool_execution]\n", color=self.color)
|
||||
print_text("Tool: " + tool_name + "\n", color=self.color)
|
||||
print_text("Outputs: " + tool_output.model_dump_json()[:1000] + "\n", color=self.color)
|
||||
print_text("\n")
|
||||
yield tool_output
|
||||
|
||||
@@ -13,7 +13,7 @@ from core.helper.code_executor.jinja2.jinja2_transformer import Jinja2TemplateTr
|
||||
from core.helper.code_executor.python3.python3_transformer import Python3TemplateTransformer
|
||||
from core.helper.code_executor.template_transformer import TemplateTransformer
|
||||
from core.helper.http_client_pooling import get_pooled_http_client
|
||||
from graphon.nodes.code.entities import CodeLanguage
|
||||
from graphon.nodes.code.entities import CodeLanguage as CodeLanguage # noqa: PLC0414
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
code_execution_endpoint_url = URL(str(dify_config.CODE_EXECUTION_ENDPOINT))
|
||||
@@ -133,7 +133,9 @@ class CodeExecutor:
|
||||
return response_code.data.stdout or ""
|
||||
|
||||
@classmethod
|
||||
def execute_workflow_code_template(cls, language: CodeLanguage, code: str, inputs: Mapping[str, Any]):
|
||||
def execute_workflow_code_template(
|
||||
cls, language: CodeLanguage, code: str, inputs: Mapping[str, Any]
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Execute code
|
||||
:param language: code language
|
||||
|
||||
@@ -11,7 +11,7 @@ class Jinja2TemplateTransformer(TemplateTransformer):
|
||||
|
||||
@classmethod
|
||||
@override
|
||||
def transform_response(cls, response: str):
|
||||
def transform_response(cls, response: str) -> dict[str, Any]:
|
||||
"""
|
||||
Transform response to dict
|
||||
:param response: response
|
||||
|
||||
@@ -36,14 +36,14 @@ class TemplateTransformer(ABC):
|
||||
return runner_script, preload_script
|
||||
|
||||
@classmethod
|
||||
def extract_result_str_from_response(cls, response: str):
|
||||
def extract_result_str_from_response(cls, response: str) -> str:
|
||||
result = re.search(rf"{cls._result_tag}(.*){cls._result_tag}", response, re.DOTALL)
|
||||
if not result:
|
||||
raise ValueError(f"Failed to parse result: no result tag found in response. Response: {response[:200]}...")
|
||||
return result.group(1)
|
||||
|
||||
@classmethod
|
||||
def transform_response(cls, response: str) -> Mapping[str, Any]:
|
||||
def transform_response(cls, response: str) -> dict[str, Any]:
|
||||
"""
|
||||
Transform response to dict
|
||||
:param response: response
|
||||
@@ -71,7 +71,7 @@ class TemplateTransformer(ABC):
|
||||
return result
|
||||
|
||||
@classmethod
|
||||
def _post_process_result(cls, result: dict[Any, Any]) -> dict[Any, Any]:
|
||||
def _post_process_result(cls, result: dict[str, Any]) -> dict[str, Any]:
|
||||
"""
|
||||
Post-process the result to convert scientific notation strings back to numbers
|
||||
"""
|
||||
@@ -89,7 +89,7 @@ class TemplateTransformer(ABC):
|
||||
return [convert_scientific_notation(v) for v in value]
|
||||
return value
|
||||
|
||||
return convert_scientific_notation(result)
|
||||
return {key: convert_scientific_notation(value) for key, value in result.items()}
|
||||
|
||||
@classmethod
|
||||
@abstractmethod
|
||||
|
||||
@@ -24,7 +24,7 @@ def upload_dsl(dsl_file_bytes: bytes, filename: str = "template.yaml") -> str:
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
claim_code = data.get("data", {}).get("claim_code")
|
||||
if not claim_code:
|
||||
if not isinstance(claim_code, str) or not claim_code:
|
||||
raise ValueError("Creators Platform did not return a valid claim_code")
|
||||
return claim_code
|
||||
|
||||
|
||||
@@ -10,18 +10,21 @@ def is_credential_exists(credential_id: str, credential_type: "PluginCredentialT
|
||||
"""
|
||||
Check if the credential still exists in the database.
|
||||
|
||||
Uses the configured SQLAlchemy session factory instead of Flask-SQLAlchemy's
|
||||
``db.engine`` because workflow graph node construction may run without an
|
||||
active Flask application context.
|
||||
|
||||
:param credential_id: The credential ID to check
|
||||
:param credential_type: The type of credential (MODEL or TOOL)
|
||||
:return: True if credential exists, False otherwise
|
||||
"""
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from extensions.ext_database import db
|
||||
from core.db import session_factory
|
||||
from models.provider import ProviderCredential, ProviderModelCredential
|
||||
from models.tools import BuiltinToolProvider
|
||||
|
||||
with Session(db.engine) as session:
|
||||
with session_factory.create_session() as session:
|
||||
if credential_type == PluginCredentialType.MODEL:
|
||||
# Check both pre-defined and custom model credentials using a single UNION query
|
||||
stmt = (
|
||||
@@ -42,7 +45,7 @@ def is_credential_exists(credential_id: str, credential_type: "PluginCredentialT
|
||||
|
||||
def runtime_check_credential_policy_compliance(
|
||||
credential_id: str, provider: str, credential_type: "PluginCredentialType", check_existence: bool = True
|
||||
):
|
||||
) -> None:
|
||||
if dify_config.ENTERPRISE_DISABLE_RUNTIME_CREDENTIAL_CHECK:
|
||||
return
|
||||
check_credential_policy_compliance(
|
||||
|
||||
@@ -1,4 +1,7 @@
|
||||
def download_with_size_limit(url, max_download_size: int, **kwargs):
|
||||
from typing import Any
|
||||
|
||||
|
||||
def download_with_size_limit(url: str, max_download_size: int, **kwargs: Any) -> bytes:
|
||||
from core.file import remote_fetcher
|
||||
|
||||
response = remote_fetcher.make_request("GET", url, follow_redirects=True, **kwargs)
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
import base64
|
||||
|
||||
from Crypto.PublicKey import RSA
|
||||
|
||||
from libs import rsa
|
||||
|
||||
|
||||
@@ -11,13 +13,13 @@ def obfuscated_token(token: str) -> str:
|
||||
return token[:6] + "*" * 12 + token[-2:]
|
||||
|
||||
|
||||
def full_mask_token(token_length=20):
|
||||
def full_mask_token(token_length: int = 20) -> str:
|
||||
return "*" * token_length
|
||||
|
||||
|
||||
def encrypt_token(tenant_id: str, token: str):
|
||||
from extensions.ext_database import db
|
||||
def encrypt_token(tenant_id: str, token: str) -> str:
|
||||
from models.account import Tenant
|
||||
from models.engine import db
|
||||
|
||||
if not (tenant := db.session.get(Tenant, tenant_id)):
|
||||
raise ValueError(f"Tenant with id {tenant_id} not found")
|
||||
@@ -30,15 +32,15 @@ def decrypt_token(tenant_id: str, token: str) -> str:
|
||||
return rsa.decrypt(base64.b64decode(token), tenant_id)
|
||||
|
||||
|
||||
def batch_decrypt_token(tenant_id: str, tokens: list[str]):
|
||||
def batch_decrypt_token(tenant_id: str, tokens: list[str]) -> list[str]:
|
||||
rsa_key, cipher_rsa = rsa.get_decrypt_decoding(tenant_id)
|
||||
|
||||
return [rsa.decrypt_token_with_decoding(base64.b64decode(token), rsa_key, cipher_rsa) for token in tokens]
|
||||
|
||||
|
||||
def get_decrypt_decoding(tenant_id: str):
|
||||
def get_decrypt_decoding(tenant_id: str) -> tuple[RSA.RsaKey, object]:
|
||||
return rsa.get_decrypt_decoding(tenant_id)
|
||||
|
||||
|
||||
def decrypt_token_with_decoding(token: str, rsa_key, cipher_rsa):
|
||||
def decrypt_token_with_decoding(token: str, rsa_key: RSA.RsaKey, cipher_rsa: object) -> str:
|
||||
return rsa.decrypt_token_with_decoding(base64.b64decode(token), rsa_key, cipher_rsa)
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import logging
|
||||
from collections.abc import Sequence
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from yarl import URL
|
||||
@@ -19,7 +20,7 @@ def get_plugin_pkg_url(plugin_unique_identifier: str) -> str:
|
||||
return str((marketplace_api_url / "api/v1/plugins/download").with_query(unique_identifier=plugin_unique_identifier))
|
||||
|
||||
|
||||
def download_plugin_pkg(plugin_unique_identifier: str):
|
||||
def download_plugin_pkg(plugin_unique_identifier: str) -> bytes:
|
||||
return download_with_size_limit(get_plugin_pkg_url(plugin_unique_identifier), dify_config.PLUGIN_MAX_PACKAGE_SIZE)
|
||||
|
||||
|
||||
@@ -39,7 +40,7 @@ def batch_fetch_plugin_manifests(plugin_ids: list[str]) -> Sequence[MarketplaceP
|
||||
return [MarketplacePluginDeclaration.model_validate(plugin) for plugin in response.json()["data"]["plugins"]]
|
||||
|
||||
|
||||
def batch_fetch_plugin_by_ids(plugin_ids: list[str]) -> list[dict]:
|
||||
def batch_fetch_plugin_by_ids(plugin_ids: list[str]) -> list[dict[str, Any]]:
|
||||
if not plugin_ids:
|
||||
return []
|
||||
|
||||
@@ -53,10 +54,19 @@ def batch_fetch_plugin_by_ids(plugin_ids: list[str]) -> list[dict]:
|
||||
response.raise_for_status()
|
||||
|
||||
data = response.json()
|
||||
return data.get("data", {}).get("plugins", [])
|
||||
plugins = data.get("data", {}).get("plugins", [])
|
||||
if not isinstance(plugins, list):
|
||||
raise ValueError("Marketplace did not return a valid plugins list")
|
||||
|
||||
result: list[dict[str, Any]] = []
|
||||
for plugin in plugins:
|
||||
if not isinstance(plugin, dict) or not all(isinstance(key, str) for key in plugin):
|
||||
raise ValueError("Marketplace did not return a valid plugins list")
|
||||
result.append(plugin)
|
||||
return result
|
||||
|
||||
|
||||
def record_install_plugin_event(plugin_unique_identifier: str):
|
||||
def record_install_plugin_event(plugin_unique_identifier: str) -> None:
|
||||
url = str(marketplace_api_url / "api/v1/stats/plugins/install_count")
|
||||
response = httpx.post(url, json={"unique_identifier": plugin_unique_identifier}, timeout=MARKETPLACE_TIMEOUT)
|
||||
response.raise_for_status()
|
||||
|
||||
@@ -34,7 +34,7 @@ class ProviderCredentialsCache:
|
||||
else:
|
||||
return None
|
||||
|
||||
def set(self, credentials: dict[str, Any]):
|
||||
def set(self, credentials: dict[str, Any]) -> None:
|
||||
"""
|
||||
Cache model provider credentials.
|
||||
|
||||
@@ -43,7 +43,7 @@ class ProviderCredentialsCache:
|
||||
"""
|
||||
redis_client.setex(self.cache_key, 86400, json.dumps(credentials))
|
||||
|
||||
def delete(self):
|
||||
def delete(self) -> None:
|
||||
"""
|
||||
Delete cached model provider credentials.
|
||||
|
||||
|
||||
@@ -20,17 +20,18 @@ def import_module_from_source[T: (str, bytes)](
|
||||
raise Exception(f"Failed to load module {module_name} from {py_file_path!r}")
|
||||
else:
|
||||
# Refer to: https://docs.python.org/3/library/importlib.html#importing-a-source-file-directly
|
||||
# FIXME: mypy does not support the type of spec.loader
|
||||
spec = importlib.util.spec_from_file_location(module_name, py_file_path) # type: ignore[assignment]
|
||||
if not spec or not spec.loader:
|
||||
new_spec = importlib.util.spec_from_file_location(module_name, py_file_path)
|
||||
if not new_spec or not new_spec.loader:
|
||||
raise Exception(f"Failed to load module {module_name} from {py_file_path!r}")
|
||||
if use_lazy_loader:
|
||||
# Refer to: https://docs.python.org/3/library/importlib.html#implementing-lazy-imports
|
||||
spec.loader = importlib.util.LazyLoader(spec.loader)
|
||||
new_spec.loader = importlib.util.LazyLoader(new_spec.loader)
|
||||
spec = new_spec
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
if not existed_spec:
|
||||
sys.modules[module_name] = module
|
||||
spec.loader.exec_module(module)
|
||||
if spec.loader is not None:
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
except Exception as e:
|
||||
logger.exception("Failed to load module %s from script file '%s'", module_name, repr(py_file_path))
|
||||
|
||||
@@ -9,11 +9,11 @@ from extensions.ext_redis import redis_client
|
||||
class ProviderCredentialsCache(ABC):
|
||||
"""Base class for provider credentials cache"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
def __init__(self, **kwargs: Any) -> None:
|
||||
self.cache_key = self._generate_cache_key(**kwargs)
|
||||
|
||||
@abstractmethod
|
||||
def _generate_cache_key(self, **kwargs) -> str:
|
||||
def _generate_cache_key(self, **kwargs: Any) -> str:
|
||||
"""Generate cache key based on subclass implementation"""
|
||||
pass
|
||||
|
||||
@@ -28,11 +28,11 @@ class ProviderCredentialsCache(ABC):
|
||||
return None
|
||||
return None
|
||||
|
||||
def set(self, config: dict[str, Any]):
|
||||
def set(self, config: dict[str, Any]) -> None:
|
||||
"""Cache provider credentials"""
|
||||
redis_client.setex(self.cache_key, 86400, json.dumps(config))
|
||||
|
||||
def delete(self):
|
||||
def delete(self) -> None:
|
||||
"""Delete cached provider credentials"""
|
||||
redis_client.delete(self.cache_key)
|
||||
|
||||
@@ -48,7 +48,7 @@ class SingletonProviderCredentialsCache(ProviderCredentialsCache):
|
||||
)
|
||||
|
||||
@override
|
||||
def _generate_cache_key(self, **kwargs) -> str:
|
||||
def _generate_cache_key(self, **kwargs: Any) -> str:
|
||||
tenant_id = kwargs["tenant_id"]
|
||||
provider_type = kwargs["provider_type"]
|
||||
identity_name = kwargs["provider_identity"]
|
||||
@@ -63,7 +63,7 @@ class ToolProviderCredentialsCache(ProviderCredentialsCache):
|
||||
super().__init__(tenant_id=tenant_id, provider=provider, credential_id=credential_id)
|
||||
|
||||
@override
|
||||
def _generate_cache_key(self, **kwargs) -> str:
|
||||
def _generate_cache_key(self, **kwargs: Any) -> str:
|
||||
tenant_id = kwargs["tenant_id"]
|
||||
provider = kwargs["provider"]
|
||||
credential_id = kwargs["credential_id"]
|
||||
@@ -77,10 +77,10 @@ class NoOpProviderCredentialCache:
|
||||
"""Get cached provider credentials"""
|
||||
return None
|
||||
|
||||
def set(self, config: dict[str, Any]):
|
||||
def set(self, config: dict[str, Any]) -> None:
|
||||
"""Cache provider credentials"""
|
||||
pass
|
||||
|
||||
def delete(self):
|
||||
def delete(self) -> None:
|
||||
"""Delete cached provider credentials"""
|
||||
pass
|
||||
|
||||
@@ -125,5 +125,7 @@ class ProviderConfigEncrypter:
|
||||
return data
|
||||
|
||||
|
||||
def create_provider_encrypter(tenant_id: str, config: list[BasicProviderConfig], cache: ProviderConfigCache):
|
||||
def create_provider_encrypter(
|
||||
tenant_id: str, config: list[BasicProviderConfig], cache: ProviderConfigCache
|
||||
) -> tuple[ProviderConfigEncrypter, ProviderConfigCache]:
|
||||
return ProviderConfigEncrypter(tenant_id=tenant_id, config=config, provider_config_cache=cache), cache
|
||||
|
||||
@@ -37,11 +37,11 @@ class ToolParameterCache:
|
||||
else:
|
||||
return None
|
||||
|
||||
def set(self, parameters: dict[str, Any]):
|
||||
def set(self, parameters: dict[str, Any]) -> None:
|
||||
"""Cache model provider credentials."""
|
||||
redis_client.setex(self.cache_key, 86400, json.dumps(parameters))
|
||||
|
||||
def delete(self):
|
||||
def delete(self) -> None:
|
||||
"""
|
||||
Delete cached model provider credentials.
|
||||
|
||||
|
||||
@@ -61,7 +61,7 @@ def get_external_trace_id(request: Any) -> str | None:
|
||||
return None
|
||||
|
||||
|
||||
def extract_external_trace_id_from_args(args: Mapping[str, Any]):
|
||||
def extract_external_trace_id_from_args(args: Mapping[str, Any]) -> dict[str, Any]:
|
||||
"""
|
||||
Extract 'external_trace_id' from args.
|
||||
|
||||
|
||||
@@ -76,7 +76,11 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation):
|
||||
if not user_id:
|
||||
user = EndUserService.get_or_create_end_user(app)
|
||||
else:
|
||||
user = cls._get_user(user_id, app)
|
||||
try:
|
||||
user = cls._get_user(user_id, app)
|
||||
except ValueError:
|
||||
# Plugins such as WeCom Bot pass external sender IDs rather than EndUser UUIDs.
|
||||
user = EndUserService.get_or_create_end_user(app, user_id=user_id)
|
||||
|
||||
conversation_id = conversation_id or ""
|
||||
|
||||
@@ -226,6 +230,13 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation):
|
||||
EndUser.app_id == app.id,
|
||||
)
|
||||
user = session.scalar(stmt)
|
||||
if not user:
|
||||
stmt = select(EndUser).where(
|
||||
EndUser.session_id == user_id,
|
||||
EndUser.tenant_id == app.tenant_id,
|
||||
EndUser.app_id == app.id,
|
||||
)
|
||||
user = session.scalar(stmt)
|
||||
if not user:
|
||||
stmt = select(Account).where(
|
||||
Account.id == user_id,
|
||||
|
||||
@@ -228,7 +228,7 @@ class CredentialType(enum.StrEnum):
|
||||
OAUTH2 = "oauth2"
|
||||
UNAUTHORIZED = "unauthorized"
|
||||
|
||||
def get_name(self):
|
||||
def get_name(self) -> str:
|
||||
if self == CredentialType.API_KEY:
|
||||
return "API KEY"
|
||||
elif self == CredentialType.OAUTH2:
|
||||
|
||||
@@ -4,6 +4,8 @@ This module owns plugin daemon management calls that are shared by API services
|
||||
and core runtimes. Plugin model provider discovery is cached here, alongside
|
||||
plugin install, uninstall, and upgrade invalidation, so all cache mutations for
|
||||
plugin-owned provider metadata stay tenant-scoped and in one place.
|
||||
Provider cache payloads may be stored as prefixed zstd bytes; readers also
|
||||
accept legacy plain JSON payloads for rolling upgrades and existing Redis keys.
|
||||
|
||||
The console plugin list also normalizes endpoint setup counters against live
|
||||
endpoint records. Some plugin daemon builds return stale ``endpoints_*``
|
||||
@@ -14,12 +16,15 @@ metadata.
|
||||
|
||||
import logging
|
||||
import time
|
||||
from collections.abc import Mapping, Sequence
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from contextlib import contextmanager
|
||||
from mimetypes import guess_type
|
||||
from typing import ClassVar
|
||||
from typing import Literal, Protocol
|
||||
|
||||
import zstandard
|
||||
from pydantic import BaseModel, TypeAdapter, ValidationError
|
||||
from redis import RedisError
|
||||
from redis.exceptions import LockError
|
||||
from sqlalchemy import delete, select, update
|
||||
from sqlalchemy.orm import Session
|
||||
from yarl import URL
|
||||
@@ -67,14 +72,18 @@ logger = logging.getLogger(__name__)
|
||||
_provider_entities_adapter: TypeAdapter[list[ProviderEntity]] = TypeAdapter(list[ProviderEntity])
|
||||
|
||||
|
||||
class PluginService:
|
||||
_plugin_model_providers_memory_cache: ClassVar[dict[str, tuple[int, float, tuple[ProviderEntity, ...]]]] = {}
|
||||
class _RedisLock(Protocol):
|
||||
def acquire(self, *, blocking: bool = True, blocking_timeout: float | None = None) -> bool: ...
|
||||
|
||||
def release(self) -> None: ...
|
||||
|
||||
|
||||
class PluginService:
|
||||
class LatestPluginCache(BaseModel):
|
||||
plugin_id: str
|
||||
version: str
|
||||
unique_identifier: str
|
||||
status: str
|
||||
status: Literal["active", "deleted"]
|
||||
deprecated_reason: str
|
||||
alternative_plugin_id: str
|
||||
|
||||
@@ -82,6 +91,12 @@ class PluginService:
|
||||
REDIS_TTL = 60 * 5 # 5 minutes
|
||||
PLUGIN_MODEL_PROVIDERS_REDIS_KEY_PREFIX = "plugin_model_providers:tenant_id:"
|
||||
PLUGIN_MODEL_PROVIDERS_GENERATION_REDIS_KEY_PREFIX = "plugin_model_providers_generation:tenant_id:"
|
||||
PLUGIN_MODEL_PROVIDERS_LOCK_REDIS_KEY_PREFIX = "plugin_model_providers_refresh_lock:tenant_id:"
|
||||
PLUGIN_MODEL_PROVIDERS_LOCK_TTL = 30
|
||||
PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_TIMEOUT = 2.0
|
||||
PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_INTERVAL = 0.05
|
||||
PLUGIN_MODEL_PROVIDERS_CACHE_COMPRESSION_PREFIX = b"\x00dify-plugin-model-providers-zstd-v1:"
|
||||
PLUGIN_MODEL_PROVIDERS_CACHE_COMPRESSION_MIN_BYTES = 64 * 1024
|
||||
PLUGIN_INSTALL_TASK_TERMINAL_STATUSES = (PluginInstallTaskStatus.Success, PluginInstallTaskStatus.Failed)
|
||||
# Mirror the detail-panel endpoint query size so list reconciliation and
|
||||
# the visible endpoint drawer exercise the same daemon pagination path.
|
||||
@@ -98,6 +113,10 @@ class PluginService:
|
||||
def _get_plugin_model_providers_generation_cache_key(cls, tenant_id: str) -> str:
|
||||
return f"{cls.PLUGIN_MODEL_PROVIDERS_GENERATION_REDIS_KEY_PREFIX}{tenant_id}"
|
||||
|
||||
@classmethod
|
||||
def _get_plugin_model_providers_lock_key(cls, tenant_id: str, generation: int) -> str:
|
||||
return f"{cls.PLUGIN_MODEL_PROVIDERS_LOCK_REDIS_KEY_PREFIX}{tenant_id}:generation:{generation}"
|
||||
|
||||
@staticmethod
|
||||
def _get_provider_short_name_alias(provider: PluginModelProviderEntity) -> str:
|
||||
"""
|
||||
@@ -129,8 +148,25 @@ class PluginService:
|
||||
return declaration
|
||||
|
||||
@classmethod
|
||||
def _copy_provider_entities(cls, providers: Sequence[ProviderEntity]) -> tuple[ProviderEntity, ...]:
|
||||
return tuple(provider.model_copy(deep=True) for provider in providers)
|
||||
def _encode_plugin_model_providers_cache_payload(cls, payload: bytes) -> bytes:
|
||||
if len(payload) < cls.PLUGIN_MODEL_PROVIDERS_CACHE_COMPRESSION_MIN_BYTES:
|
||||
return payload
|
||||
|
||||
return cls.PLUGIN_MODEL_PROVIDERS_CACHE_COMPRESSION_PREFIX + zstandard.compress(payload, level=1)
|
||||
|
||||
@classmethod
|
||||
def _decode_plugin_model_providers_cache_payload(cls, payload: bytes | bytearray | str) -> bytes | bytearray | str:
|
||||
if isinstance(payload, str):
|
||||
return payload
|
||||
|
||||
prefix = cls.PLUGIN_MODEL_PROVIDERS_CACHE_COMPRESSION_PREFIX
|
||||
if not payload.startswith(prefix):
|
||||
return payload
|
||||
|
||||
try:
|
||||
return zstandard.decompress(payload[len(prefix) :])
|
||||
except zstandard.ZstdError as exc:
|
||||
raise ValueError("Invalid compressed plugin model providers cache payload.") from exc
|
||||
|
||||
@classmethod
|
||||
def _load_plugin_model_providers_generation(cls, tenant_id: str) -> int | None:
|
||||
@@ -163,76 +199,35 @@ class PluginService:
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def _load_in_memory_plugin_model_providers(
|
||||
cls, memory_cache_key: str, generation: int
|
||||
) -> tuple[ProviderEntity, ...] | None:
|
||||
cached_entry = cls._plugin_model_providers_memory_cache.get(memory_cache_key)
|
||||
if cached_entry is None:
|
||||
return None
|
||||
def _load_cached_plugin_model_providers_for_generation(
|
||||
cls, tenant_id: str, generation: int | None
|
||||
) -> tuple[tuple[ProviderEntity, ...] | None, bool]:
|
||||
if generation is None:
|
||||
return None, False
|
||||
|
||||
cached_generation, expires_at, providers = cached_entry
|
||||
if cached_generation != generation or time.monotonic() >= expires_at:
|
||||
cls._plugin_model_providers_memory_cache.pop(memory_cache_key, None)
|
||||
return None
|
||||
|
||||
return cls._copy_provider_entities(providers)
|
||||
|
||||
@classmethod
|
||||
def _store_in_memory_plugin_model_providers(
|
||||
cls, memory_cache_key: str, generation: int, providers: Sequence[ProviderEntity]
|
||||
) -> None:
|
||||
ttl = dify_config.PLUGIN_MODEL_PROVIDERS_CACHE_TTL
|
||||
if ttl <= 0:
|
||||
cls._plugin_model_providers_memory_cache.pop(memory_cache_key, None)
|
||||
return
|
||||
|
||||
cls._plugin_model_providers_memory_cache[memory_cache_key] = (
|
||||
generation,
|
||||
time.monotonic() + ttl,
|
||||
cls._copy_provider_entities(providers),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _load_cached_plugin_model_providers(
|
||||
cls, tenant_id: str, *, client: PluginModelClient | None = None
|
||||
) -> tuple[ProviderEntity, ...] | None:
|
||||
generation = cls._load_plugin_model_providers_generation(tenant_id)
|
||||
if generation is not None:
|
||||
in_memory_cached_providers = cls._load_in_memory_plugin_model_providers(tenant_id, generation)
|
||||
if in_memory_cached_providers is not None:
|
||||
return in_memory_cached_providers
|
||||
|
||||
cache_keys = []
|
||||
if generation is not None:
|
||||
cache_keys.append(cls._get_plugin_model_providers_cache_key(tenant_id, generation))
|
||||
if generation == 0:
|
||||
cache_keys.append(cls._get_plugin_model_providers_cache_key(tenant_id))
|
||||
|
||||
if not cache_keys:
|
||||
return None
|
||||
cache_keys = [cls._get_plugin_model_providers_cache_key(tenant_id, generation)]
|
||||
|
||||
try:
|
||||
cached_provider_entries = redis_client.mget(cache_keys)
|
||||
except (RedisError, RuntimeError):
|
||||
except (LockError, RedisError, RuntimeError):
|
||||
logger.warning("Failed to read cached plugin model providers for tenant %s.", tenant_id, exc_info=True)
|
||||
return None
|
||||
return None, False
|
||||
|
||||
if len(cached_provider_entries) != len(cache_keys):
|
||||
logger.warning(
|
||||
"Unexpected cached plugin model providers response size for tenant %s.",
|
||||
tenant_id,
|
||||
)
|
||||
return None
|
||||
return None, False
|
||||
|
||||
for cache_key, cached_providers in zip(cache_keys, cached_provider_entries):
|
||||
if not cached_providers:
|
||||
continue
|
||||
|
||||
try:
|
||||
providers = tuple(_provider_entities_adapter.validate_json(cached_providers))
|
||||
if generation is not None:
|
||||
cls._store_in_memory_plugin_model_providers(tenant_id, generation, providers)
|
||||
return providers
|
||||
payload = cls._decode_plugin_model_providers_cache_payload(cached_providers)
|
||||
providers = tuple(_provider_entities_adapter.validate_json(payload))
|
||||
return providers, True
|
||||
except (TypeError, ValueError, ValidationError):
|
||||
logger.warning(
|
||||
"Invalid cached plugin model providers for tenant %s; deleting cache key %s.",
|
||||
@@ -249,7 +244,7 @@ class PluginService:
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
return None
|
||||
return None, True
|
||||
|
||||
@classmethod
|
||||
def _store_cached_plugin_model_providers(
|
||||
@@ -257,15 +252,94 @@ class PluginService:
|
||||
) -> None:
|
||||
cache_key = cls._get_plugin_model_providers_cache_key(tenant_id, generation)
|
||||
try:
|
||||
payload = _provider_entities_adapter.dump_json(list(providers)).decode("utf-8")
|
||||
payload = cls._encode_plugin_model_providers_cache_payload(
|
||||
_provider_entities_adapter.dump_json(list(providers))
|
||||
)
|
||||
redis_client.setex(cache_key, dify_config.PLUGIN_MODEL_PROVIDERS_CACHE_TTL, payload)
|
||||
except (RedisError, RuntimeError):
|
||||
logger.warning("Failed to cache plugin model providers for tenant %s.", tenant_id, exc_info=True)
|
||||
|
||||
@classmethod
|
||||
@contextmanager
|
||||
def _plugin_model_providers_refresh_lock(
|
||||
cls, tenant_id: str, generation: int, *, wait_timeout: float
|
||||
) -> Iterator[bool]:
|
||||
lock_key = cls._get_plugin_model_providers_lock_key(tenant_id, generation)
|
||||
try:
|
||||
refresh_lock: _RedisLock = redis_client.lock(
|
||||
lock_key,
|
||||
timeout=cls.PLUGIN_MODEL_PROVIDERS_LOCK_TTL,
|
||||
sleep=cls.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_INTERVAL,
|
||||
)
|
||||
except (RedisError, RuntimeError):
|
||||
logger.warning(
|
||||
"Failed to create plugin model providers refresh lock for tenant %s.",
|
||||
tenant_id,
|
||||
exc_info=True,
|
||||
)
|
||||
yield False
|
||||
return
|
||||
|
||||
try:
|
||||
lock_acquired = refresh_lock.acquire(blocking=True, blocking_timeout=wait_timeout)
|
||||
except LockError:
|
||||
logger.warning(
|
||||
"Provider refresh lock timed out; direct daemon fallback. tenant_id=%s generation=%s",
|
||||
tenant_id,
|
||||
generation,
|
||||
exc_info=True,
|
||||
)
|
||||
yield False
|
||||
return
|
||||
except (RedisError, RuntimeError):
|
||||
# Redis failures should not block provider discovery; callers fetch directly from the daemon.
|
||||
logger.warning(
|
||||
"Failed to acquire plugin model providers refresh lock for tenant %s.",
|
||||
tenant_id,
|
||||
exc_info=True,
|
||||
)
|
||||
yield False
|
||||
return
|
||||
|
||||
if not lock_acquired:
|
||||
logger.warning(
|
||||
"Provider refresh lock timed out; direct daemon fallback. tenant_id=%s generation=%s",
|
||||
tenant_id,
|
||||
generation,
|
||||
)
|
||||
yield False
|
||||
return
|
||||
|
||||
try:
|
||||
yield True
|
||||
finally:
|
||||
try:
|
||||
refresh_lock.release()
|
||||
except (LockError, RedisError, RuntimeError):
|
||||
# Release failures must not hide the daemon result or the original exception.
|
||||
logger.warning(
|
||||
"Failed to release plugin model providers refresh lock for tenant %s generation %s.",
|
||||
tenant_id,
|
||||
generation,
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _fetch_and_cache_plugin_model_providers(
|
||||
cls, tenant_id: str, client: PluginModelClient | None, *, refresh_generation: int | None
|
||||
) -> tuple[ProviderEntity, ...]:
|
||||
model_client = client or PluginModelClient()
|
||||
providers = tuple(
|
||||
cls._to_provider_entity(provider) for provider in model_client.fetch_model_providers(tenant_id)
|
||||
)
|
||||
generation = cls._load_plugin_model_providers_generation(tenant_id)
|
||||
if generation is not None and generation == refresh_generation:
|
||||
cls._store_cached_plugin_model_providers(tenant_id, generation, providers)
|
||||
return providers
|
||||
|
||||
@classmethod
|
||||
def invalidate_plugin_model_providers_cache(cls, tenant_id: str) -> None:
|
||||
"""Invalidate tenant-scoped provider metadata across Redis and worker-local mirrors."""
|
||||
cls._plugin_model_providers_memory_cache.pop(tenant_id, None)
|
||||
"""Invalidate tenant-scoped provider metadata stored in Redis."""
|
||||
cache_key = cls._get_plugin_model_providers_cache_key(tenant_id)
|
||||
generation_key = cls._get_plugin_model_providers_generation_cache_key(tenant_id)
|
||||
try:
|
||||
@@ -287,21 +361,68 @@ class PluginService:
|
||||
are intentionally owned by this service so tenant isolation and cache
|
||||
expiry are handled in one place.
|
||||
"""
|
||||
cached_providers = cls._load_cached_plugin_model_providers(tenant_id, client=client)
|
||||
if cached_providers is not None:
|
||||
return cached_providers
|
||||
deadline = time.monotonic() + cls.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_TIMEOUT
|
||||
|
||||
model_client = client or PluginModelClient()
|
||||
providers = tuple(
|
||||
cls._to_provider_entity(provider) for provider in model_client.fetch_model_providers(tenant_id)
|
||||
)
|
||||
if not providers:
|
||||
return providers
|
||||
generation = cls._load_plugin_model_providers_generation(tenant_id)
|
||||
if generation is not None:
|
||||
cls._store_in_memory_plugin_model_providers(tenant_id, generation, providers)
|
||||
cls._store_cached_plugin_model_providers(tenant_id, generation, providers)
|
||||
return providers
|
||||
while True:
|
||||
generation = cls._load_plugin_model_providers_generation(tenant_id)
|
||||
cached_providers, cache_available = cls._load_cached_plugin_model_providers_for_generation(
|
||||
tenant_id, generation
|
||||
)
|
||||
if cached_providers is not None:
|
||||
return cached_providers
|
||||
|
||||
if generation is None or not cache_available:
|
||||
return cls._fetch_and_cache_plugin_model_providers(
|
||||
tenant_id,
|
||||
client,
|
||||
refresh_generation=generation,
|
||||
)
|
||||
|
||||
wait_timeout = deadline - time.monotonic()
|
||||
if wait_timeout < 0:
|
||||
logger.warning(
|
||||
"Provider refresh lock timed out; direct daemon fallback. tenant_id=%s generation=%s",
|
||||
tenant_id,
|
||||
generation,
|
||||
)
|
||||
return cls._fetch_and_cache_plugin_model_providers(
|
||||
tenant_id,
|
||||
client,
|
||||
refresh_generation=generation,
|
||||
)
|
||||
|
||||
with cls._plugin_model_providers_refresh_lock(
|
||||
tenant_id,
|
||||
generation,
|
||||
wait_timeout=wait_timeout,
|
||||
) as lock_acquired:
|
||||
if not lock_acquired:
|
||||
return cls._fetch_and_cache_plugin_model_providers(
|
||||
tenant_id,
|
||||
client,
|
||||
refresh_generation=generation,
|
||||
)
|
||||
|
||||
latest_generation = cls._load_plugin_model_providers_generation(tenant_id)
|
||||
cached_providers, cache_available = cls._load_cached_plugin_model_providers_for_generation(
|
||||
tenant_id, latest_generation
|
||||
)
|
||||
if cached_providers is not None:
|
||||
return cached_providers
|
||||
if latest_generation is None or not cache_available:
|
||||
return cls._fetch_and_cache_plugin_model_providers(
|
||||
tenant_id,
|
||||
client,
|
||||
refresh_generation=latest_generation,
|
||||
)
|
||||
if latest_generation != generation:
|
||||
continue
|
||||
|
||||
return cls._fetch_and_cache_plugin_model_providers(
|
||||
tenant_id,
|
||||
client,
|
||||
refresh_generation=generation,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def fetch_latest_plugin_version(plugin_ids: Sequence[str]) -> Mapping[str, LatestPluginCache | None]:
|
||||
|
||||
@@ -1030,6 +1030,10 @@ class DatasetRetrieval:
|
||||
):
|
||||
"""
|
||||
Persist dataset query audit rows for retrieval requests.
|
||||
|
||||
Query audit logging is a side effect of retrieval. Keep it in an
|
||||
independent transaction so failures or commits here do not affect the
|
||||
request/workflow transaction that called the retriever.
|
||||
"""
|
||||
if not query and not attachment_ids:
|
||||
return
|
||||
@@ -1041,6 +1045,9 @@ class DatasetRetrieval:
|
||||
app_id,
|
||||
)
|
||||
return
|
||||
created_by_role = self._resolve_creator_user_role(user_from)
|
||||
if created_by_role is None:
|
||||
return
|
||||
dataset_queries = []
|
||||
for dataset_id in dataset_ids:
|
||||
contents = []
|
||||
@@ -1055,13 +1062,16 @@ class DatasetRetrieval:
|
||||
content=json.dumps(contents),
|
||||
source=DatasetQuerySource.APP,
|
||||
source_app_id=app_id,
|
||||
created_by_role=CreatorUserRole(user_from),
|
||||
created_by_role=created_by_role,
|
||||
created_by=created_by,
|
||||
)
|
||||
dataset_queries.append(dataset_query)
|
||||
if dataset_queries:
|
||||
db.session.add_all(dataset_queries)
|
||||
db.session.commit()
|
||||
|
||||
if not dataset_queries:
|
||||
return
|
||||
|
||||
with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as session:
|
||||
session.add_all(dataset_queries)
|
||||
|
||||
def _retriever(
|
||||
self,
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
from core.rbac.entities import RBACPermission, RBACResourceScope
|
||||
from core.rbac.entities import RBACPermission, RBACResourceScope, RBACResourceWhitelistScope
|
||||
|
||||
__all__ = ["RBACPermission", "RBACResourceScope"]
|
||||
__all__ = ["RBACPermission", "RBACResourceScope", "RBACResourceWhitelistScope"]
|
||||
|
||||
@@ -13,6 +13,14 @@ class RBACResourceScope(StrEnum):
|
||||
WORKSPACE = "workspace"
|
||||
|
||||
|
||||
class RBACResourceWhitelistScope(StrEnum):
|
||||
"""Whitelist scopes accepted by RBAC app and dataset access config APIs."""
|
||||
|
||||
ALL = "all"
|
||||
SPECIFIC = "specific"
|
||||
ONLY_ME = "only_me"
|
||||
|
||||
|
||||
class RBACPermission(StrEnum):
|
||||
"""Permission points (RBAC scenes) checked by ``rbac_permission_required``.
|
||||
|
||||
|
||||
@@ -102,16 +102,17 @@ class ApiTool(Tool):
|
||||
elif not isinstance(credentials["api_key_value"], str):
|
||||
raise ToolProviderCredentialValidationError("api_key_value must be a string")
|
||||
|
||||
api_key_value = credentials["api_key_value"]
|
||||
if "api_key_header_prefix" in credentials:
|
||||
api_key_header_prefix = credentials["api_key_header_prefix"]
|
||||
if api_key_header_prefix == "basic" and credentials["api_key_value"]:
|
||||
credentials["api_key_value"] = f"Basic {credentials['api_key_value']}"
|
||||
elif api_key_header_prefix == "bearer" and credentials["api_key_value"]:
|
||||
credentials["api_key_value"] = f"Bearer {credentials['api_key_value']}"
|
||||
if api_key_header_prefix == "basic" and api_key_value:
|
||||
api_key_value = f"Basic {api_key_value}"
|
||||
elif api_key_header_prefix == "bearer" and api_key_value:
|
||||
api_key_value = f"Bearer {api_key_value}"
|
||||
elif api_key_header_prefix == "custom":
|
||||
pass
|
||||
|
||||
headers[api_key_header] = credentials["api_key_value"]
|
||||
headers[api_key_header] = api_key_value
|
||||
|
||||
elif credentials["auth_type"] == "api_key_query":
|
||||
# For query parameter authentication, we don't add anything to headers
|
||||
|
||||
@@ -7,6 +7,7 @@ from datetime import UTC, datetime
|
||||
from mimetypes import guess_type
|
||||
from typing import Any, Union, cast
|
||||
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
from yarl import URL
|
||||
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
@@ -338,47 +339,49 @@ class ToolEngine:
|
||||
user_id: str,
|
||||
) -> list[str]:
|
||||
"""
|
||||
Create message file
|
||||
Create message files produced by a tool call.
|
||||
|
||||
Tool file persistence is a side effect of agent execution. Use an
|
||||
independent transaction so this helper never commits or closes the
|
||||
caller's request-scoped session.
|
||||
|
||||
:return: message file ids
|
||||
"""
|
||||
result = []
|
||||
|
||||
for message in tool_messages:
|
||||
if "image" in message.mimetype:
|
||||
file_type = FileType.IMAGE
|
||||
elif "video" in message.mimetype:
|
||||
file_type = FileType.VIDEO
|
||||
elif "audio" in message.mimetype:
|
||||
file_type = FileType.AUDIO
|
||||
elif "text" in message.mimetype or "pdf" in message.mimetype:
|
||||
file_type = FileType.DOCUMENT
|
||||
else:
|
||||
file_type = FileType.CUSTOM
|
||||
with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as session:
|
||||
for message in tool_messages:
|
||||
# extract tool file id from url
|
||||
tool_file_id = message.url.split("/")[-1].split(".")[0]
|
||||
message_file = MessageFile(
|
||||
message_id=agent_message.id,
|
||||
type=ToolEngine._resolve_tool_file_type(message),
|
||||
transfer_method=FileTransferMethod.TOOL_FILE,
|
||||
belongs_to=MessageFileBelongsTo.ASSISTANT,
|
||||
url=message.url,
|
||||
upload_file_id=tool_file_id,
|
||||
created_by_role=(
|
||||
CreatorUserRole.ACCOUNT
|
||||
if invoke_from in {InvokeFrom.EXPLORE, InvokeFrom.DEBUGGER}
|
||||
else CreatorUserRole.END_USER
|
||||
),
|
||||
created_by=user_id,
|
||||
)
|
||||
|
||||
# extract tool file id from url
|
||||
tool_file_id = message.url.split("/")[-1].split(".")[0]
|
||||
message_file = MessageFile(
|
||||
message_id=agent_message.id,
|
||||
type=file_type,
|
||||
transfer_method=FileTransferMethod.TOOL_FILE,
|
||||
belongs_to=MessageFileBelongsTo.ASSISTANT,
|
||||
url=message.url,
|
||||
upload_file_id=tool_file_id,
|
||||
created_by_role=(
|
||||
CreatorUserRole.ACCOUNT
|
||||
if invoke_from in {InvokeFrom.EXPLORE, InvokeFrom.DEBUGGER}
|
||||
else CreatorUserRole.END_USER
|
||||
),
|
||||
created_by=user_id,
|
||||
)
|
||||
|
||||
db.session.add(message_file)
|
||||
db.session.commit()
|
||||
db.session.refresh(message_file)
|
||||
|
||||
result.append(message_file.id)
|
||||
|
||||
db.session.close()
|
||||
session.add(message_file)
|
||||
result.append(message_file.id)
|
||||
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _resolve_tool_file_type(message: ToolInvokeMessageBinary) -> FileType:
|
||||
if "image" in message.mimetype:
|
||||
return FileType.IMAGE
|
||||
elif "video" in message.mimetype:
|
||||
return FileType.VIDEO
|
||||
elif "audio" in message.mimetype:
|
||||
return FileType.AUDIO
|
||||
elif "text" in message.mimetype or "pdf" in message.mimetype:
|
||||
return FileType.DOCUMENT
|
||||
else:
|
||||
return FileType.CUSTOM
|
||||
|
||||
@@ -5,10 +5,10 @@ from json import loads as json_loads
|
||||
from json.decoder import JSONDecodeError
|
||||
from typing import Any, TypedDict
|
||||
|
||||
import httpx
|
||||
from flask import has_request_context, request
|
||||
from yaml import YAMLError, safe_load
|
||||
|
||||
from core.helper import ssrf_proxy
|
||||
from core.tools.entities.common_entities import I18nObject
|
||||
from core.tools.entities.tool_bundle import ApiToolBundle
|
||||
from core.tools.entities.tool_entities import ApiProviderSchemaType, ToolParameter
|
||||
@@ -376,7 +376,7 @@ class ApiBasedToolSchemaParser:
|
||||
raise ToolNotSupportedError("Only openapi is supported now.")
|
||||
|
||||
# get openapi yaml
|
||||
response = httpx.get(
|
||||
response = ssrf_proxy.get(
|
||||
api_url, headers={"User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) "}, timeout=5
|
||||
)
|
||||
|
||||
|
||||
@@ -46,18 +46,26 @@ logger = logging.getLogger(__name__)
|
||||
_file_access_controller = DatabaseFileAccessController()
|
||||
|
||||
|
||||
def iter_dify_graph_engine_events(engine: GraphEngine) -> Generator[GraphEngineEvent, None, None]:
|
||||
def iter_dify_graph_engine_events(
|
||||
engine: GraphEngine,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
) -> Generator[GraphEngineEvent, None, None]:
|
||||
"""
|
||||
Apply Dify's response streaming compatibility filter to GraphEngine events.
|
||||
|
||||
Graphon v0.5.0 emits raw variable stream chunks and requires callers to opt
|
||||
into the legacy response-ordered stream behavior that Dify exposes to its
|
||||
workflow runners and tests.
|
||||
|
||||
``response_stream_filter``, when supplied, must be the same instance a
|
||||
caller intends to persist on pause (see ``PauseStatePersistenceLayer``) so
|
||||
the filter's ``paths_map`` reflects everything the engine has actually
|
||||
streamed for this run.
|
||||
"""
|
||||
yield from filter_graph_events(
|
||||
engine.run(),
|
||||
context=GraphEventFilterContext.from_engine(engine),
|
||||
filters=[ResponseStreamFilter()],
|
||||
filters=[response_stream_filter or ResponseStreamFilter()],
|
||||
)
|
||||
|
||||
|
||||
@@ -167,6 +175,7 @@ class WorkflowEntry:
|
||||
variable_pool: VariablePool,
|
||||
graph_runtime_state: GraphRuntimeState,
|
||||
command_channel: CommandChannel | None = None,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Init workflow entry
|
||||
@@ -183,6 +192,8 @@ class WorkflowEntry:
|
||||
:param variable_pool: variable pool
|
||||
:param graph_runtime_state: pre-created graph runtime state
|
||||
:param command_channel: command channel for external control (optional, defaults to InMemoryChannel)
|
||||
:param response_stream_filter: pre-restored filter for resumed runs (optional, defaults to a fresh
|
||||
ResponseStreamFilter for runs with no prior pause)
|
||||
:param thread_pool_id: thread pool id
|
||||
"""
|
||||
# check call depth
|
||||
@@ -195,6 +206,7 @@ class WorkflowEntry:
|
||||
command_channel = InMemoryChannel()
|
||||
|
||||
self.command_channel = command_channel
|
||||
self._response_stream_filter = response_stream_filter or ResponseStreamFilter()
|
||||
execution_context = capture_current_context()
|
||||
graph_runtime_state.execution_context = execution_context
|
||||
self._child_engine_builder = _WorkflowChildEngineBuilder(tenant_id=tenant_id)
|
||||
@@ -240,7 +252,7 @@ class WorkflowEntry:
|
||||
|
||||
try:
|
||||
# Preserve Dify's response-stream semantics on top of Graphon 0.5.0.
|
||||
generator = iter_dify_graph_engine_events(graph_engine)
|
||||
generator = iter_dify_graph_engine_events(graph_engine, self._response_stream_filter)
|
||||
yield from generator
|
||||
except GenerateTaskStoppedError:
|
||||
pass
|
||||
|
||||
@@ -35,10 +35,10 @@ if [[ "${MODE}" == "worker" ]]; then
|
||||
if [[ -z "${CELERY_QUEUES}" ]]; then
|
||||
if [[ "${EDITION}" == "CLOUD" ]]; then
|
||||
# Cloud edition: separate queues for dataset and trigger tasks
|
||||
DEFAULT_QUEUES="api_token,dataset,dataset_summary,priority_dataset,priority_pipeline,pipeline,mail,ops_trace,app_deletion,plugin,workflow_storage,conversation,workflow_professional,workflow_team,workflow_sandbox,schedule_poller,schedule_executor,triggered_workflow_dispatcher,trigger_refresh_publisher,trigger_refresh_executor,retention,workflow_based_app_execution"
|
||||
DEFAULT_QUEUES="api_token,dataset,dataset_summary,priority_dataset,priority_pipeline,pipeline,mail,ops_trace,app_deletion,app_rbac,plugin,workflow_storage,conversation,workflow_professional,workflow_team,workflow_sandbox,schedule_poller,schedule_executor,triggered_workflow_dispatcher,trigger_refresh_publisher,trigger_refresh_executor,retention,workflow_based_app_execution"
|
||||
else
|
||||
# Community edition (SELF_HOSTED): dataset, pipeline and workflow have separate queues
|
||||
DEFAULT_QUEUES="api_token,dataset,dataset_summary,priority_dataset,priority_pipeline,pipeline,mail,ops_trace,app_deletion,plugin,workflow_storage,conversation,workflow,schedule_poller,schedule_executor,triggered_workflow_dispatcher,trigger_refresh_publisher,trigger_refresh_executor,retention,workflow_based_app_execution"
|
||||
DEFAULT_QUEUES="api_token,dataset,dataset_summary,priority_dataset,priority_pipeline,pipeline,mail,ops_trace,app_deletion,app_rbac,plugin,workflow_storage,conversation,workflow,schedule_poller,schedule_executor,triggered_workflow_dispatcher,trigger_refresh_publisher,trigger_refresh_executor,retention,workflow_based_app_execution"
|
||||
fi
|
||||
else
|
||||
DEFAULT_QUEUES="${CELERY_QUEUES}"
|
||||
|
||||
@@ -153,6 +153,7 @@ def init_app(app: DifyApp) -> Celery:
|
||||
"tasks.trigger_processing_tasks", # async trigger processing
|
||||
"tasks.generate_summary_index_task", # summary index generation
|
||||
"tasks.regenerate_summary_index_task", # summary index regeneration
|
||||
"tasks.initialize_created_app_rbac_access_task", # app access initialization
|
||||
"tasks.app_generate.resume_agent_app_task", # ENG-635: Agent v2 chat ask_human resume
|
||||
]
|
||||
day = dify_config.CELERY_BEAT_SCHEDULER_TIME
|
||||
|
||||
@@ -27,6 +27,7 @@ def init_app(app: DifyApp):
|
||||
install_plugins,
|
||||
install_rag_pipeline_plugins,
|
||||
migrate_data_for_plugin,
|
||||
migrate_dataset_permissions_to_rbac,
|
||||
migrate_member_roles_to_rbac,
|
||||
migrate_oss,
|
||||
migration_data_wizard,
|
||||
@@ -56,6 +57,7 @@ def init_app(app: DifyApp):
|
||||
upgrade_db,
|
||||
fix_app_site_missing,
|
||||
migrate_data_for_plugin,
|
||||
migrate_dataset_permissions_to_rbac,
|
||||
migrate_member_roles_to_rbac,
|
||||
backfill_plugin_auto_upgrade,
|
||||
extract_plugins,
|
||||
|
||||
+19
-2
@@ -34,6 +34,7 @@ class OAuthState(TypedDict, total=False):
|
||||
invite_token: str
|
||||
timezone: str
|
||||
language: str
|
||||
redirect_url: str
|
||||
|
||||
|
||||
class GitHubEmailRecord(TypedDict, total=False):
|
||||
@@ -72,6 +73,7 @@ def encode_oauth_state(
|
||||
invite_token: str | None = None,
|
||||
timezone: str | None = None,
|
||||
language: str | None = None,
|
||||
redirect_url: str | None = None,
|
||||
) -> str | None:
|
||||
state: OAuthState = {}
|
||||
if invite_token:
|
||||
@@ -80,6 +82,8 @@ def encode_oauth_state(
|
||||
state["timezone"] = timezone
|
||||
if language:
|
||||
state["language"] = language
|
||||
if redirect_url:
|
||||
state["redirect_url"] = redirect_url
|
||||
if not state:
|
||||
return None
|
||||
|
||||
@@ -122,6 +126,7 @@ class OAuth:
|
||||
invite_token: str | None = None,
|
||||
timezone: str | None = None,
|
||||
language: str | None = None,
|
||||
redirect_url: str | None = None,
|
||||
) -> str:
|
||||
raise NotImplementedError()
|
||||
|
||||
@@ -151,13 +156,19 @@ class GitHubOAuth(OAuth):
|
||||
invite_token: str | None = None,
|
||||
timezone: str | None = None,
|
||||
language: str | None = None,
|
||||
redirect_url: str | None = None,
|
||||
) -> str:
|
||||
params = {
|
||||
"client_id": self.client_id,
|
||||
"redirect_uri": self.redirect_uri,
|
||||
"scope": "user:email", # Request only basic user information
|
||||
}
|
||||
state = encode_oauth_state(invite_token=invite_token, timezone=timezone, language=language)
|
||||
state = encode_oauth_state(
|
||||
invite_token=invite_token,
|
||||
timezone=timezone,
|
||||
language=language,
|
||||
redirect_url=redirect_url,
|
||||
)
|
||||
if state:
|
||||
params["state"] = state
|
||||
return f"{self._AUTH_URL}?{urllib.parse.urlencode(params)}"
|
||||
@@ -248,6 +259,7 @@ class GoogleOAuth(OAuth):
|
||||
invite_token: str | None = None,
|
||||
timezone: str | None = None,
|
||||
language: str | None = None,
|
||||
redirect_url: str | None = None,
|
||||
) -> str:
|
||||
params = {
|
||||
"client_id": self.client_id,
|
||||
@@ -255,7 +267,12 @@ class GoogleOAuth(OAuth):
|
||||
"redirect_uri": self.redirect_uri,
|
||||
"scope": "openid email",
|
||||
}
|
||||
state = encode_oauth_state(invite_token=invite_token, timezone=timezone, language=language)
|
||||
state = encode_oauth_state(
|
||||
invite_token=invite_token,
|
||||
timezone=timezone,
|
||||
language=language,
|
||||
redirect_url=redirect_url,
|
||||
)
|
||||
if state:
|
||||
params["state"] = state
|
||||
return f"{self._AUTH_URL}?{urllib.parse.urlencode(params)}"
|
||||
|
||||
@@ -214,6 +214,18 @@ class EndUserType(StrEnum):
|
||||
SERVICE_API = "service-api"
|
||||
TRIGGER = "trigger"
|
||||
|
||||
@classmethod
|
||||
@override
|
||||
def _missing_(cls, value):
|
||||
# Legacy rows persisted the service-api type with an underscore before it
|
||||
# was normalized to the hyphenated value. The
|
||||
# `4f7b2c8d9a10_normalize_legacy_end_user_type` migration rewrites those
|
||||
# rows, but tolerate the old value here as well so an unmigrated end user
|
||||
# keeps loading instead of failing enum validation on every request.
|
||||
if value == "service_api":
|
||||
return cls.SERVICE_API
|
||||
return super()._missing_(value)
|
||||
|
||||
|
||||
class DocumentDocType(StrEnum):
|
||||
"""Document doc_type classification"""
|
||||
|
||||
@@ -2868,6 +2868,7 @@ Delete an existing tracing configuration for an application
|
||||
| ---- | ----------- |
|
||||
| 204 | Tracing configuration deleted successfully |
|
||||
| 400 | Invalid request parameters or configuration not found |
|
||||
| 403 | Insufficient permissions |
|
||||
|
||||
### [GET] /apps/{app_id}/trace-config
|
||||
Get tracing configuration for an application
|
||||
@@ -2909,6 +2910,7 @@ Update an existing tracing configuration for an application
|
||||
| ---- | ----------- | ------ |
|
||||
| 200 | Tracing configuration updated successfully | **application/json**: [TraceAppConfigResponse](#traceappconfigresponse)<br> |
|
||||
| 400 | Invalid request parameters or configuration not found | |
|
||||
| 403 | Insufficient permissions | |
|
||||
|
||||
### [POST] /apps/{app_id}/trace-config
|
||||
**Create a new trace app configuration**
|
||||
@@ -2933,6 +2935,7 @@ Create a new tracing configuration for an application
|
||||
| ---- | ----------- | ------ |
|
||||
| 201 | Tracing configuration created successfully | **application/json**: [TraceAppConfigResponse](#traceappconfigresponse)<br> |
|
||||
| 400 | Invalid request parameters or configuration already exists | |
|
||||
| 403 | Insufficient permissions | |
|
||||
|
||||
### [POST] /apps/{app_id}/trigger-enable
|
||||
**Update app trigger (enable/disable)**
|
||||
@@ -7000,6 +7003,7 @@ Initiate OAuth login process
|
||||
| provider | path | OAuth provider name (github/google) | Yes | string |
|
||||
| invite_token | query | Optional invitation token | No | string |
|
||||
| language | query | Preferred interface language | No | string |
|
||||
| redirect_url | query | Relative page to resume after login | No | string |
|
||||
| timezone | query | Preferred timezone | No | string |
|
||||
|
||||
#### Responses
|
||||
@@ -16859,6 +16863,7 @@ Enum class for large language model mode.
|
||||
| Name | Type | Description | Required |
|
||||
| ---- | ---- | ----------- | -------- |
|
||||
| expired_at | string | | Yes |
|
||||
| seats | [LicenseLimitationModel](#licenselimitationmodel) | | Yes |
|
||||
| status | [LicenseStatus](#licensestatus) | | Yes |
|
||||
| workspaces | [LicenseLimitationModel](#licenselimitationmodel) | | Yes |
|
||||
|
||||
@@ -17529,6 +17534,7 @@ Coarse node-level status used by Inspector to pick a banner.
|
||||
| ---- | ---- | ----------- | -------- |
|
||||
| invite_token | string | Optional invitation token | No |
|
||||
| language | string | Preferred interface language | No |
|
||||
| redirect_url | string | Relative page to resume after login | No |
|
||||
| timezone | string | Preferred timezone | No |
|
||||
|
||||
#### OAuthProviderAccountResponse
|
||||
@@ -18795,6 +18801,13 @@ Model class for provider quota configuration.
|
||||
| ---- | ---- | ----------- | -------- |
|
||||
| url | string | URL to fetch | Yes |
|
||||
|
||||
#### ReplaceUserAccessPolicies
|
||||
|
||||
| Name | Type | Description | Required |
|
||||
| ---- | ---- | ----------- | -------- |
|
||||
| access_policy_ids | [ string ] | | No |
|
||||
| account_ids | [ string ] | | No |
|
||||
|
||||
#### ReplaceUserAccessPoliciesResponse
|
||||
|
||||
| Name | Type | Description | Required |
|
||||
|
||||
@@ -1337,6 +1337,7 @@ Parsed multipart form fields for HITL uploads.
|
||||
| Name | Type | Description | Required |
|
||||
| ---- | ---- | ----------- | -------- |
|
||||
| expired_at | string | | Yes |
|
||||
| seats | [LicenseLimitationModel](#licenselimitationmodel) | | Yes |
|
||||
| status | [LicenseStatus](#licensestatus) | | Yes |
|
||||
| workspaces | [LicenseLimitationModel](#licenselimitationmodel) | | Yes |
|
||||
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
import uuid
|
||||
from enum import StrEnum
|
||||
from typing import Any, override
|
||||
@@ -17,6 +18,8 @@ from models.dataset import Dataset
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
METADATA_KEY_PATTERN = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
|
||||
|
||||
|
||||
class MyScaleConfig(BaseModel):
|
||||
host: str
|
||||
@@ -102,7 +105,10 @@ class MyScaleVector(BaseVector):
|
||||
|
||||
@override
|
||||
def text_exists(self, id: str) -> bool:
|
||||
results = self._client.query(f"SELECT id FROM {self._config.database}.{self._collection_name} WHERE id='{id}'")
|
||||
results = self._client.query(
|
||||
f"SELECT id FROM {self._config.database}.{self._collection_name} WHERE id={{id:String}}",
|
||||
parameters={"id": id},
|
||||
)
|
||||
return results.row_count > 0
|
||||
|
||||
@override
|
||||
@@ -110,20 +116,26 @@ class MyScaleVector(BaseVector):
|
||||
if not ids:
|
||||
return
|
||||
self._client.command(
|
||||
f"DELETE FROM {self._config.database}.{self._collection_name} WHERE id IN {str(tuple(ids))}"
|
||||
f"DELETE FROM {self._config.database}.{self._collection_name} WHERE id IN {{ids:Array(String)}}",
|
||||
parameters={"ids": ids},
|
||||
)
|
||||
|
||||
@override
|
||||
def get_ids_by_metadata_field(self, key: str, value: str):
|
||||
self._validate_metadata_key(key)
|
||||
rows = self._client.query(
|
||||
f"SELECT DISTINCT id FROM {self._config.database}.{self._collection_name} WHERE metadata.{key}='{value}'"
|
||||
f"SELECT DISTINCT id FROM {self._config.database}.{self._collection_name} "
|
||||
f"WHERE metadata.{key}={{value:String}}",
|
||||
parameters={"value": value},
|
||||
).result_rows
|
||||
return [row[0] for row in rows]
|
||||
|
||||
@override
|
||||
def delete_by_metadata_field(self, key: str, value: str):
|
||||
self._validate_metadata_key(key)
|
||||
self._client.command(
|
||||
f"DELETE FROM {self._config.database}.{self._collection_name} WHERE metadata.{key}='{value}'"
|
||||
f"DELETE FROM {self._config.database}.{self._collection_name} WHERE metadata.{key}={{value:String}}",
|
||||
parameters={"value": value},
|
||||
)
|
||||
|
||||
@override
|
||||
@@ -132,22 +144,29 @@ class MyScaleVector(BaseVector):
|
||||
|
||||
@override
|
||||
def search_by_full_text(self, query: str, **kwargs: Any) -> list[Document]:
|
||||
return self._search(f"TextSearch('enable_nlq=false')(text, '{query}')", SortOrder.DESC, **kwargs)
|
||||
return self._search(
|
||||
"TextSearch('enable_nlq=false')(text, {query:String})",
|
||||
SortOrder.DESC,
|
||||
parameters={"query": query},
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def _search(self, dist: str, order: SortOrder, **kwargs: Any) -> list[Document]:
|
||||
def _search(
|
||||
self, dist: str, order: SortOrder, parameters: dict[str, Any] | None = None, **kwargs: Any
|
||||
) -> list[Document]:
|
||||
top_k = kwargs.get("top_k", 4)
|
||||
if not isinstance(top_k, int) or top_k <= 0:
|
||||
raise ValueError("top_k must be a positive integer")
|
||||
score_threshold = float(kwargs.get("score_threshold") or 0.0)
|
||||
where_str = (
|
||||
f"WHERE dist < {1 - score_threshold}"
|
||||
if self._metric.upper() == "COSINE" and order == SortOrder.ASC and score_threshold > 0.0
|
||||
else ""
|
||||
)
|
||||
where_conditions = []
|
||||
query_parameters = dict(parameters or {})
|
||||
if self._metric.upper() == "COSINE" and order == SortOrder.ASC and score_threshold > 0.0:
|
||||
where_conditions.append(f"dist < {1 - score_threshold}")
|
||||
document_ids_filter = kwargs.get("document_ids_filter")
|
||||
if document_ids_filter:
|
||||
document_ids = ", ".join(f"'{id}'" for id in document_ids_filter)
|
||||
where_str = f"{where_str} AND metadata['document_id'] in ({document_ids})"
|
||||
where_conditions.append("metadata['document_id'] IN {document_ids_filter:Array(String)}")
|
||||
query_parameters["document_ids_filter"] = document_ids_filter
|
||||
where_str = f"WHERE {' AND '.join(where_conditions)}" if where_conditions else ""
|
||||
sql = f"""
|
||||
SELECT text, vector, metadata, {dist} as dist FROM {self._config.database}.{self._collection_name}
|
||||
{where_str} ORDER BY dist {order.value} LIMIT {top_k}
|
||||
@@ -159,12 +178,17 @@ class MyScaleVector(BaseVector):
|
||||
vector=r["vector"],
|
||||
metadata=r["metadata"],
|
||||
)
|
||||
for r in self._client.query(sql).named_results()
|
||||
for r in self._client.query(sql, parameters=query_parameters).named_results()
|
||||
]
|
||||
except Exception:
|
||||
logger.exception("Vector search operation failed")
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
def _validate_metadata_key(key: str) -> None:
|
||||
if not METADATA_KEY_PATTERN.match(key):
|
||||
raise ValueError("metadata key must be a valid identifier")
|
||||
|
||||
@override
|
||||
def delete(self):
|
||||
self._client.command(f"DROP TABLE IF EXISTS {self._config.database}.{self._collection_name}")
|
||||
|
||||
@@ -181,14 +181,41 @@ def test_text_exists_and_metadata_operations(myscale_module):
|
||||
vector = myscale_module.MyScaleVector("collection_1", _config(myscale_module))
|
||||
vector._client.query.return_value = SimpleNamespace(row_count=1, result_rows=[("id-1",), ("id-2",)])
|
||||
|
||||
assert vector.text_exists("id-1") is True
|
||||
assert vector.get_ids_by_metadata_field("document_id", "doc-1") == ["id-1", "id-2"]
|
||||
assert vector.text_exists("id-1' OR '1'='1") is True
|
||||
text_exists_call = vector._client.query.call_args
|
||||
assert "id={id:String}" in text_exists_call.args[0]
|
||||
assert "id-1' OR '1'='1" not in text_exists_call.args[0]
|
||||
assert text_exists_call.kwargs["parameters"] == {"id": "id-1' OR '1'='1"}
|
||||
|
||||
assert vector.get_ids_by_metadata_field("document_id", "doc-1' OR '1'='1") == ["id-1", "id-2"]
|
||||
metadata_query_call = vector._client.query.call_args
|
||||
assert "metadata.document_id={value:String}" in metadata_query_call.args[0]
|
||||
assert "doc-1' OR '1'='1" not in metadata_query_call.args[0]
|
||||
assert metadata_query_call.kwargs["parameters"] == {"value": "doc-1' OR '1'='1"}
|
||||
|
||||
vector.delete_by_ids(["id-1", "id-2"])
|
||||
vector.delete_by_metadata_field("document_id", "doc-1")
|
||||
delete_ids_call = vector._client.command.call_args
|
||||
assert "id IN {ids:Array(String)}" in delete_ids_call.args[0]
|
||||
assert delete_ids_call.kwargs["parameters"] == {"ids": ["id-1", "id-2"]}
|
||||
|
||||
vector.delete_by_metadata_field("document_id", "doc-1' OR '1'='1")
|
||||
delete_metadata_call = vector._client.command.call_args
|
||||
assert "metadata.document_id={value:String}" in delete_metadata_call.args[0]
|
||||
assert "doc-1' OR '1'='1" not in delete_metadata_call.args[0]
|
||||
assert delete_metadata_call.kwargs["parameters"] == {"value": "doc-1' OR '1'='1"}
|
||||
assert vector._client.command.call_count >= 2
|
||||
|
||||
|
||||
def test_metadata_operations_reject_invalid_key(myscale_module):
|
||||
vector = myscale_module.MyScaleVector("collection_1", _config(myscale_module))
|
||||
|
||||
with pytest.raises(ValueError, match="metadata key must be a valid identifier"):
|
||||
vector.get_ids_by_metadata_field("document_id) OR 1=1 --", "doc-1")
|
||||
|
||||
with pytest.raises(ValueError, match="metadata key must be a valid identifier"):
|
||||
vector.delete_by_metadata_field("document_id) OR 1=1 --", "doc-1")
|
||||
|
||||
|
||||
def test_search_delegation_methods(myscale_module):
|
||||
vector = myscale_module.MyScaleVector("collection_1", _config(myscale_module))
|
||||
vector._search = MagicMock(return_value=["result"])
|
||||
@@ -199,6 +226,28 @@ def test_search_delegation_methods(myscale_module):
|
||||
assert result_vector == ["result"]
|
||||
assert result_text == ["result"]
|
||||
assert vector._search.call_count == 2
|
||||
vector._search.assert_any_call(
|
||||
"TextSearch('enable_nlq=false')(text, {query:String})",
|
||||
myscale_module.SortOrder.DESC,
|
||||
parameters={"query": "hello"},
|
||||
top_k=2,
|
||||
)
|
||||
|
||||
|
||||
def test_search_by_full_text_uses_query_parameters(myscale_module):
|
||||
vector = myscale_module.MyScaleVector("collection_1", _config(myscale_module))
|
||||
vector._client.query.return_value = SimpleNamespace(
|
||||
named_results=lambda: [{"text": "doc", "vector": [0.1], "metadata": {"doc_id": "1"}}]
|
||||
)
|
||||
payload = "x') AS dist FROM dify.collection_1 UNION ALL SELECT secret FROM users --"
|
||||
|
||||
docs = vector.search_by_full_text(payload, top_k=2)
|
||||
|
||||
assert len(docs) == 1
|
||||
sql = vector._client.query.call_args.args[0]
|
||||
assert payload not in sql
|
||||
assert "TextSearch('enable_nlq=false')(text, {query:String})" in sql
|
||||
assert vector._client.query.call_args.kwargs["parameters"] == {"query": payload}
|
||||
|
||||
|
||||
def test_search_with_document_filter_and_exception(myscale_module):
|
||||
@@ -215,7 +264,8 @@ def test_search_with_document_filter_and_exception(myscale_module):
|
||||
)
|
||||
assert len(docs) == 1
|
||||
sql = vector._client.query.call_args.args[0]
|
||||
assert "metadata['document_id'] in ('doc-1', 'doc-2')" in sql
|
||||
assert "WHERE metadata['document_id'] IN {document_ids_filter:Array(String)}" in sql
|
||||
assert vector._client.query.call_args.kwargs["parameters"] == {"document_ids_filter": ["doc-1", "doc-2"]}
|
||||
|
||||
vector._client.query.side_effect = RuntimeError("boom")
|
||||
assert vector._search("distance(vector, [0.1])", myscale_module.SortOrder.ASC, top_k=1) == []
|
||||
|
||||
@@ -33,7 +33,6 @@ from models.dataset import Dataset, DatasetCollectionBinding
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from qdrant_client.conversions import common_types
|
||||
from qdrant_client.http import models as rest
|
||||
|
||||
type DictFilter = dict[str, str | int | bool | dict | list]
|
||||
type MetadataFilter = DictFilter | common_types.Filter
|
||||
|
||||
-1
@@ -41,7 +41,6 @@ from models.enums import TidbAuthBindingStatus
|
||||
if TYPE_CHECKING:
|
||||
from qdrant_client import grpc # noqa
|
||||
from qdrant_client.conversions import common_types
|
||||
from qdrant_client.http import models as rest
|
||||
|
||||
type DictFilter = dict[str, str | int | bool | dict | list]
|
||||
type MetadataFilter = DictFilter | common_types.Filter
|
||||
|
||||
+3
-1
@@ -42,11 +42,12 @@ dependencies = [
|
||||
"opentelemetry-propagator-b3>=1.41.1,<2.0.0",
|
||||
"readabilipy==0.3.0",
|
||||
"resend>=2.27.0,<3.0.0",
|
||||
"zstandard==0.25.0",
|
||||
# Emerging: newer and fast-moving, use compatible pins
|
||||
"fastopenapi[flask]==0.7.0",
|
||||
"graphon==0.5.3",
|
||||
"httpx-sse==0.4.3",
|
||||
"json-repair==0.59.4",
|
||||
"json-repair==0.60.1",
|
||||
]
|
||||
# Before adding new dependency, consider place it in
|
||||
# alphabet order (a-z) and suitable group.
|
||||
@@ -107,6 +108,7 @@ override-dependencies = [
|
||||
"litellm>=1.83.10,<2.0.0",
|
||||
"pyarrow>=23.0.1,<24.0.0",
|
||||
"cryptography>=49.0.0,<50.0.0",
|
||||
"setuptools>=80.10.2,<81",
|
||||
]
|
||||
|
||||
[dependency-groups]
|
||||
|
||||
@@ -10,6 +10,7 @@ import json
|
||||
import logging
|
||||
import secrets
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from hashlib import sha256
|
||||
from typing import Any, NotRequired, TypedDict, cast
|
||||
@@ -64,6 +65,7 @@ from services.errors.account import (
|
||||
MemberNotInTenantError,
|
||||
NoPermissionError,
|
||||
RoleAlreadyAssignedError,
|
||||
SeatsLimitExceededError,
|
||||
TenantNotFoundError,
|
||||
)
|
||||
from services.errors.workspace import WorkSpaceNotAllowedCreateError, WorkspacesLimitExceededError
|
||||
@@ -433,6 +435,13 @@ class AccountService:
|
||||
|
||||
raise AccountNotFound()
|
||||
|
||||
# A licensed seat is one Account row, deployment-wide; joining an existing
|
||||
# account into another workspace does not pass through here and costs no seat.
|
||||
# is_authenticated=True: server-side enforcement needs the full license payload,
|
||||
# which the enterprise fill withholds from unauthenticated (browser-facing) calls.
|
||||
if not FeatureService.get_system_features(is_authenticated=True).license.seats.is_available():
|
||||
raise SeatsLimitExceededError("licensed seats limit exceeded")
|
||||
|
||||
if dify_config.BILLING_ENABLED and BillingService.is_email_in_freeze(email):
|
||||
raise AccountRegisterError(
|
||||
description=(
|
||||
@@ -1576,6 +1585,25 @@ class TenantService:
|
||||
|
||||
return updated_accounts
|
||||
|
||||
@staticmethod
|
||||
def iter_member_account_id_batches(tenant_id: str, batch_size: int, *, session: Session) -> Iterator[list[str]]:
|
||||
"""Yield workspace member account ids in bounded, ordered batches."""
|
||||
offset = 0
|
||||
while True:
|
||||
stmt = (
|
||||
select(TenantAccountJoin.account_id)
|
||||
.where(TenantAccountJoin.tenant_id == tenant_id)
|
||||
.order_by(TenantAccountJoin.id)
|
||||
.offset(offset)
|
||||
.limit(batch_size)
|
||||
)
|
||||
account_ids = list(session.scalars(stmt).all())
|
||||
if not account_ids:
|
||||
return
|
||||
|
||||
yield account_ids
|
||||
offset += batch_size
|
||||
|
||||
@staticmethod
|
||||
def get_dataset_operator_members(tenant: Tenant, *, session: scoped_session | Session) -> list[Account]:
|
||||
"""Get dataset admin members"""
|
||||
@@ -1981,6 +2009,10 @@ class RegisterService:
|
||||
session.rollback()
|
||||
logger.exception("Register failed")
|
||||
raise AccountRegisterError("Workspace is not allowed to create.")
|
||||
except SeatsLimitExceededError:
|
||||
session.rollback()
|
||||
logger.exception("Register failed")
|
||||
raise
|
||||
except AccountRegisterError as are:
|
||||
session.rollback()
|
||||
logger.exception("Register failed")
|
||||
|
||||
@@ -50,9 +50,26 @@ class QuotaReleaseResult(TypedDict):
|
||||
released: int
|
||||
|
||||
|
||||
class QuotaBalanceResult(TypedDict):
|
||||
available: int
|
||||
reserved: int
|
||||
quota: int
|
||||
usage: int
|
||||
|
||||
|
||||
class QuotaConsumeCappedResult(TypedDict):
|
||||
deducted: int
|
||||
available: int
|
||||
reserved: int
|
||||
quota: int
|
||||
usage: int
|
||||
|
||||
|
||||
_quota_reserve_adapter = TypeAdapter(QuotaReserveResult)
|
||||
_quota_commit_adapter = TypeAdapter(QuotaCommitResult)
|
||||
_quota_release_adapter = TypeAdapter(QuotaReleaseResult)
|
||||
_quota_balance_adapter = TypeAdapter(QuotaBalanceResult)
|
||||
_quota_consume_capped_adapter = TypeAdapter(QuotaConsumeCappedResult)
|
||||
|
||||
|
||||
class _TenantFeatureQuota(TypedDict):
|
||||
@@ -176,6 +193,7 @@ class DismissNotificationDict(TypedDict):
|
||||
|
||||
class BillingService:
|
||||
base_url = os.environ.get("BILLING_API_URL", "BILLING_API_URL")
|
||||
quota_base_url = os.environ.get("BILLING_QUOTA_API_URL") or base_url
|
||||
secret_key = os.environ.get("BILLING_API_SECRET_KEY", "BILLING_API_SECRET_KEY")
|
||||
|
||||
compliance_download_rate_limiter = RateLimiter("compliance_download_rate_limiter", 4, 60)
|
||||
@@ -215,12 +233,18 @@ class BillingService:
|
||||
def get_quota_info(cls, tenant_id: str) -> TenantFeatureQuotaInfo:
|
||||
params = {"tenant_id": tenant_id}
|
||||
return _tenant_feature_quota_info_adapter.validate_python(
|
||||
cls._send_request("GET", "/quota/info", params=params)
|
||||
cls._send_quota_request("GET", "/quota/info", params=params)
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def quota_reserve(
|
||||
cls, tenant_id: str, feature_key: str, request_id: str, amount: int = 1, meta: dict | None = None
|
||||
cls,
|
||||
tenant_id: str,
|
||||
feature_key: str,
|
||||
request_id: str,
|
||||
amount: int = 1,
|
||||
meta: dict | None = None,
|
||||
bucket: str = "",
|
||||
) -> QuotaReserveResult:
|
||||
"""Reserve quota before task execution."""
|
||||
payload: dict = {
|
||||
@@ -229,13 +253,21 @@ class BillingService:
|
||||
"request_id": request_id,
|
||||
"amount": amount,
|
||||
}
|
||||
if bucket:
|
||||
payload["bucket"] = bucket
|
||||
if meta:
|
||||
payload["meta"] = meta
|
||||
return _quota_reserve_adapter.validate_python(cls._send_request("POST", "/quota/reserve", json=payload))
|
||||
return _quota_reserve_adapter.validate_python(cls._send_quota_request("POST", "/quota/reserve", json=payload))
|
||||
|
||||
@classmethod
|
||||
def quota_commit(
|
||||
cls, tenant_id: str, feature_key: str, reservation_id: str, actual_amount: int, meta: dict | None = None
|
||||
cls,
|
||||
tenant_id: str,
|
||||
feature_key: str,
|
||||
reservation_id: str,
|
||||
actual_amount: int,
|
||||
meta: dict | None = None,
|
||||
bucket: str = "",
|
||||
) -> QuotaCommitResult:
|
||||
"""Commit a reservation with actual consumption."""
|
||||
payload: dict = {
|
||||
@@ -244,23 +276,57 @@ class BillingService:
|
||||
"reservation_id": reservation_id,
|
||||
"actual_amount": actual_amount,
|
||||
}
|
||||
if bucket:
|
||||
payload["bucket"] = bucket
|
||||
if meta:
|
||||
payload["meta"] = meta
|
||||
return _quota_commit_adapter.validate_python(cls._send_request("POST", "/quota/commit", json=payload))
|
||||
return _quota_commit_adapter.validate_python(cls._send_quota_request("POST", "/quota/commit", json=payload))
|
||||
|
||||
@classmethod
|
||||
def quota_release(cls, tenant_id: str, feature_key: str, reservation_id: str) -> QuotaReleaseResult:
|
||||
def quota_release(
|
||||
cls, tenant_id: str, feature_key: str, reservation_id: str, bucket: str = ""
|
||||
) -> QuotaReleaseResult:
|
||||
"""Release a reservation (cancel, return frozen quota)."""
|
||||
return _quota_release_adapter.validate_python(
|
||||
cls._send_request(
|
||||
"POST",
|
||||
"/quota/release",
|
||||
json={
|
||||
"tenant_id": tenant_id,
|
||||
"feature_key": feature_key,
|
||||
"reservation_id": reservation_id,
|
||||
},
|
||||
)
|
||||
payload = {
|
||||
"tenant_id": tenant_id,
|
||||
"feature_key": feature_key,
|
||||
"reservation_id": reservation_id,
|
||||
}
|
||||
if bucket:
|
||||
payload["bucket"] = bucket
|
||||
return _quota_release_adapter.validate_python(cls._send_quota_request("POST", "/quota/release", json=payload))
|
||||
|
||||
@classmethod
|
||||
def quota_get_balance(cls, tenant_id: str, feature_key: str, bucket: str = "") -> QuotaBalanceResult:
|
||||
"""Get quota balance for a feature bucket."""
|
||||
params = {"tenant_id": tenant_id, "feature_key": feature_key}
|
||||
if bucket:
|
||||
params["bucket"] = bucket
|
||||
return _quota_balance_adapter.validate_python(cls._send_quota_request("GET", "/quota/balance", params=params))
|
||||
|
||||
@classmethod
|
||||
def quota_consume_capped(
|
||||
cls,
|
||||
tenant_id: str,
|
||||
feature_key: str,
|
||||
request_id: str,
|
||||
amount: int,
|
||||
meta: dict | None = None,
|
||||
bucket: str = "",
|
||||
) -> QuotaConsumeCappedResult:
|
||||
"""Consume up to the available quota and return the actual deducted amount."""
|
||||
payload: dict = {
|
||||
"tenant_id": tenant_id,
|
||||
"feature_key": feature_key,
|
||||
"request_id": request_id,
|
||||
"amount": amount,
|
||||
}
|
||||
if bucket:
|
||||
payload["bucket"] = bucket
|
||||
if meta:
|
||||
payload["meta"] = meta
|
||||
return _quota_consume_capped_adapter.validate_python(
|
||||
cls._send_quota_request("POST", "/quota/consume-capped", json=payload)
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@@ -334,6 +400,12 @@ class BillingService:
|
||||
params = {"tenant_id": tenant_id, "feature_key": feature_key}
|
||||
return cls._send_request("GET", "/billing/tenant_feature_plan/usage", params=params)
|
||||
|
||||
@classmethod
|
||||
def _send_quota_request(
|
||||
cls, method: Literal["GET", "POST", "DELETE", "PUT"], endpoint: str, json=None, params=None
|
||||
):
|
||||
return cls._send_request(method, endpoint, json=json, params=params, base_url=cls.quota_base_url)
|
||||
|
||||
@classmethod
|
||||
@retry(
|
||||
wait=wait_fixed(2),
|
||||
@@ -341,10 +413,17 @@ class BillingService:
|
||||
retry=retry_if_exception_type(httpx.RequestError),
|
||||
reraise=True,
|
||||
)
|
||||
def _send_request(cls, method: Literal["GET", "POST", "DELETE", "PUT"], endpoint: str, json=None, params=None):
|
||||
def _send_request(
|
||||
cls,
|
||||
method: Literal["GET", "POST", "DELETE", "PUT"],
|
||||
endpoint: str,
|
||||
json=None,
|
||||
params=None,
|
||||
base_url: str | None = None,
|
||||
):
|
||||
headers = {"Content-Type": "application/json", "Billing-Api-Secret-Key": cls.secret_key}
|
||||
|
||||
url = f"{cls.base_url}{endpoint}"
|
||||
url = f"{base_url or cls.base_url}{endpoint}"
|
||||
response = _http_client.request(method, url, json=json, params=params, headers=headers, follow_redirects=True)
|
||||
if method == "GET" and response.status_code != httpx.codes.OK:
|
||||
raise ValueError("Unable to retrieve billing information. Please try again later or contact support.")
|
||||
|
||||
@@ -7,6 +7,8 @@ from piling up database transactions while preserving cross-tenant concurrency.
|
||||
|
||||
import logging
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from uuid import uuid4
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -21,11 +23,37 @@ from models.enums import ProviderQuotaType
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
FEATURE_KEY_CREDIT_POOL = "credit_pool"
|
||||
CREDIT_POOL_TENANT_LOCK_TIMEOUT_SECONDS = 10
|
||||
CREDIT_POOL_TENANT_LOCK_BLOCKING_TIMEOUT_SECONDS = 5
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CreditPoolBalance:
|
||||
tenant_id: str
|
||||
pool_type: str
|
||||
quota_limit: int
|
||||
quota_used: int
|
||||
|
||||
@property
|
||||
def remaining_credits(self) -> int:
|
||||
if self.quota_limit == -1:
|
||||
return -1
|
||||
return max(0, self.quota_limit - self.quota_used)
|
||||
|
||||
def has_sufficient_credits(self, required_credits: int) -> bool:
|
||||
return self.quota_limit == -1 or self.remaining_credits >= required_credits
|
||||
|
||||
|
||||
class CreditPoolService:
|
||||
@staticmethod
|
||||
def _normalize_pool_type(pool_type: str | ProviderQuotaType) -> str:
|
||||
return pool_type.value if isinstance(pool_type, ProviderQuotaType) else str(pool_type)
|
||||
|
||||
@staticmethod
|
||||
def _use_billing_quota() -> bool:
|
||||
return bool(dify_config.BILLING_ENABLED)
|
||||
|
||||
@staticmethod
|
||||
def _get_tenant_lock_key(tenant_id: str) -> str:
|
||||
return f"credit_pool:tenant:{tenant_id}:deduct_lock"
|
||||
@@ -79,14 +107,32 @@ class CreditPoolService:
|
||||
return credit_pool
|
||||
|
||||
@classmethod
|
||||
def get_pool(cls, tenant_id: str, pool_type: str = "trial") -> TenantCreditPool | None:
|
||||
def get_pool(
|
||||
cls, tenant_id: str, pool_type: str | ProviderQuotaType = "trial"
|
||||
) -> TenantCreditPool | CreditPoolBalance | None:
|
||||
"""get tenant credit pool"""
|
||||
normalized_pool_type = cls._normalize_pool_type(pool_type)
|
||||
if cls._use_billing_quota():
|
||||
from services.billing_service import BillingService
|
||||
|
||||
balance = BillingService.quota_get_balance(
|
||||
tenant_id=tenant_id,
|
||||
feature_key=FEATURE_KEY_CREDIT_POOL,
|
||||
bucket=normalized_pool_type,
|
||||
)
|
||||
return CreditPoolBalance(
|
||||
tenant_id=tenant_id,
|
||||
pool_type=normalized_pool_type,
|
||||
quota_limit=balance["quota"],
|
||||
quota_used=balance["usage"],
|
||||
)
|
||||
|
||||
with session_factory.get_session_maker().begin() as session:
|
||||
return session.scalar(
|
||||
select(TenantCreditPool)
|
||||
.where(
|
||||
TenantCreditPool.tenant_id == tenant_id,
|
||||
TenantCreditPool.pool_type == pool_type,
|
||||
TenantCreditPool.pool_type == normalized_pool_type,
|
||||
)
|
||||
.limit(1)
|
||||
)
|
||||
@@ -102,7 +148,7 @@ class CreditPoolService:
|
||||
pool = cls.get_pool(tenant_id, pool_type)
|
||||
if not pool:
|
||||
return False
|
||||
return pool.remaining_credits >= credits_required
|
||||
return pool.has_sufficient_credits(credits_required)
|
||||
|
||||
@classmethod
|
||||
def check_and_deduct_credits(
|
||||
@@ -114,10 +160,54 @@ class CreditPoolService:
|
||||
"""Deduct exactly the requested credits or raise without mutating the pool."""
|
||||
if credits_required <= 0:
|
||||
return 0
|
||||
normalized_pool_type = cls._normalize_pool_type(pool_type)
|
||||
|
||||
if cls._use_billing_quota():
|
||||
from services.billing_service import BillingService
|
||||
|
||||
request_id = str(uuid4())
|
||||
result = BillingService.quota_reserve(
|
||||
tenant_id=tenant_id,
|
||||
feature_key=FEATURE_KEY_CREDIT_POOL,
|
||||
bucket=normalized_pool_type,
|
||||
request_id=request_id,
|
||||
amount=credits_required,
|
||||
meta={"source": "credit_pool.check_and_deduct"},
|
||||
)
|
||||
reservation_id = result.get("reservation_id", "")
|
||||
if not reservation_id:
|
||||
raise QuotaExceededError("Insufficient credits remaining")
|
||||
try:
|
||||
BillingService.quota_commit(
|
||||
tenant_id=tenant_id,
|
||||
feature_key=FEATURE_KEY_CREDIT_POOL,
|
||||
bucket=normalized_pool_type,
|
||||
reservation_id=reservation_id,
|
||||
actual_amount=credits_required,
|
||||
meta={"source": "credit_pool.check_and_deduct"},
|
||||
)
|
||||
except Exception:
|
||||
try:
|
||||
BillingService.quota_release(
|
||||
tenant_id=tenant_id,
|
||||
feature_key=FEATURE_KEY_CREDIT_POOL,
|
||||
bucket=normalized_pool_type,
|
||||
reservation_id=reservation_id,
|
||||
)
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Failed to release reserved credit pool quota, tenant_id=%s, pool_type=%s, reservation_id=%s",
|
||||
tenant_id,
|
||||
normalized_pool_type,
|
||||
reservation_id,
|
||||
exc_info=True,
|
||||
)
|
||||
raise
|
||||
return credits_required
|
||||
|
||||
def deduct() -> int:
|
||||
with session_factory.get_session_maker().begin() as session:
|
||||
pool = cls._get_locked_pool(session=session, tenant_id=tenant_id, pool_type=pool_type)
|
||||
pool = cls._get_locked_pool(session=session, tenant_id=tenant_id, pool_type=normalized_pool_type)
|
||||
if not pool:
|
||||
raise QuotaExceededError("Credit pool not found")
|
||||
|
||||
@@ -148,12 +238,26 @@ class CreditPoolService:
|
||||
"""Deduct up to the available balance and return the actual deducted credits."""
|
||||
if credits_required <= 0:
|
||||
return 0
|
||||
normalized_pool_type = cls._normalize_pool_type(pool_type)
|
||||
|
||||
if cls._use_billing_quota():
|
||||
from services.billing_service import BillingService
|
||||
|
||||
result = BillingService.quota_consume_capped(
|
||||
tenant_id=tenant_id,
|
||||
feature_key=FEATURE_KEY_CREDIT_POOL,
|
||||
bucket=normalized_pool_type,
|
||||
request_id=str(uuid4()),
|
||||
amount=credits_required,
|
||||
meta={"source": "credit_pool.deduct_capped"},
|
||||
)
|
||||
return result["deducted"]
|
||||
|
||||
def deduct() -> int:
|
||||
with session_factory.get_session_maker().begin() as session:
|
||||
pool = cls._get_locked_pool(session=session, tenant_id=tenant_id, pool_type=pool_type)
|
||||
pool = cls._get_locked_pool(session=session, tenant_id=tenant_id, pool_type=normalized_pool_type)
|
||||
if not pool:
|
||||
logger.warning("Credit pool not found, tenant_id=%s, pool_type=%s", tenant_id, pool_type)
|
||||
logger.warning("Credit pool not found, tenant_id=%s, pool_type=%s", tenant_id, normalized_pool_type)
|
||||
return 0
|
||||
|
||||
deducted_credits = min(credits_required, pool.remaining_credits)
|
||||
|
||||
@@ -12,6 +12,7 @@ from sqlalchemy.exc import SQLAlchemyError
|
||||
|
||||
from configs import dify_config
|
||||
from core.db.session_factory import session_factory
|
||||
from core.rbac import RBACResourceWhitelistScope
|
||||
from models import TenantAccountJoin, TenantAccountRole
|
||||
from services.enterprise.base import EnterpriseRequest
|
||||
|
||||
@@ -254,6 +255,7 @@ class ResourceUserAccessPoliciesResponse(_RBACModel):
|
||||
|
||||
class ReplaceUserAccessPolicies(_RBACModel):
|
||||
access_policy_ids: list[str] = Field(default_factory=list)
|
||||
account_ids: list[str] = Field(default_factory=list)
|
||||
|
||||
@field_validator("access_policy_ids", mode="before")
|
||||
@classmethod
|
||||
@@ -364,7 +366,6 @@ _LEGACY_WORKSPACE_ADMIN_KEYS: list[str] = [
|
||||
]
|
||||
|
||||
_LEGACY_WORKSPACE_EDITOR_KEYS: list[str] = [
|
||||
"workspace.member.manage",
|
||||
"api_extension.manage",
|
||||
"plugin.install",
|
||||
"credential.use",
|
||||
@@ -435,6 +436,7 @@ _LEGACY_APP_EDITOR_KEYS: list[str] = [
|
||||
"app.acl.delete",
|
||||
"app.acl.release_and_version",
|
||||
"app.acl.monitor",
|
||||
"app.acl.log_and_annotation",
|
||||
"app.acl.access_config",
|
||||
]
|
||||
|
||||
@@ -649,17 +651,18 @@ class ReplaceRoleBindings(_RBACModel):
|
||||
|
||||
|
||||
class ReplaceMemberBindings(_RBACModel):
|
||||
scope: str = "specific"
|
||||
scope: RBACResourceWhitelistScope = RBACResourceWhitelistScope.SPECIFIC
|
||||
|
||||
@field_validator("scope")
|
||||
@classmethod
|
||||
def _normalize_scope(cls, value: Any) -> str:
|
||||
def _normalize_scope(cls, value: Any) -> RBACResourceWhitelistScope:
|
||||
scope = str(value or "").strip().lower()
|
||||
if scope in {"", "specific"}:
|
||||
return "specific"
|
||||
if scope in {"all", "only_me"}:
|
||||
return scope
|
||||
raise ValueError(f"invalid scope: {value}")
|
||||
if scope == "":
|
||||
return RBACResourceWhitelistScope.SPECIFIC
|
||||
try:
|
||||
return RBACResourceWhitelistScope(scope)
|
||||
except ValueError as exc:
|
||||
raise ValueError(f"invalid scope: {value}") from exc
|
||||
|
||||
|
||||
class DeleteMemberBindings(_RBACModel):
|
||||
@@ -743,6 +746,7 @@ def _inner_call(
|
||||
account_id=account_id,
|
||||
json=json,
|
||||
params=params,
|
||||
timeout=dify_config.ENTERPRISE_RBAC_REQUEST_TIMEOUT,
|
||||
)
|
||||
|
||||
|
||||
@@ -1122,16 +1126,17 @@ class RBACService:
|
||||
tenant_id: str,
|
||||
account_id: str | None,
|
||||
app_id: str,
|
||||
target_account_id: str,
|
||||
target_account_id: str | None,
|
||||
payload: ReplaceUserAccessPolicies,
|
||||
) -> ReplaceUserAccessPoliciesResponse:
|
||||
request_data = payload.model_dump(mode="json")
|
||||
data = _inner_call(
|
||||
"PUT",
|
||||
f"{_INNER_PREFIX}/apps/user-access-policies",
|
||||
tenant_id=tenant_id,
|
||||
account_id=account_id,
|
||||
params={"app_id": app_id, "account_id": target_account_id},
|
||||
json=payload.model_dump(mode="json"),
|
||||
json=request_data,
|
||||
)
|
||||
return ReplaceUserAccessPoliciesResponse.model_validate(data or {})
|
||||
|
||||
@@ -1291,7 +1296,7 @@ class RBACService:
|
||||
tenant_id: str,
|
||||
account_id: str | None,
|
||||
dataset_id: str,
|
||||
target_account_id: str,
|
||||
target_account_id: str | None,
|
||||
payload: ReplaceUserAccessPolicies,
|
||||
) -> ReplaceUserAccessPoliciesResponse:
|
||||
data = _inner_call(
|
||||
@@ -1300,7 +1305,7 @@ class RBACService:
|
||||
tenant_id=tenant_id,
|
||||
account_id=account_id,
|
||||
params={"dataset_id": dataset_id, "account_id": target_account_id},
|
||||
json=payload.model_dump(mode="json"),
|
||||
json=payload.model_dump(mode="json", exclude_unset=True),
|
||||
)
|
||||
return ReplaceUserAccessPoliciesResponse.model_validate(data or {})
|
||||
|
||||
|
||||
@@ -37,6 +37,10 @@ class AccountAlreadyInTenantError(BaseServiceError):
|
||||
pass
|
||||
|
||||
|
||||
class SeatsLimitExceededError(BaseServiceError):
|
||||
pass
|
||||
|
||||
|
||||
class InvalidActionError(BaseServiceError):
|
||||
pass
|
||||
|
||||
|
||||
@@ -79,6 +79,7 @@ class LicenseModel(FeatureResponseModel):
|
||||
status: LicenseStatus = LicenseStatus.NONE
|
||||
expired_at: str = ""
|
||||
workspaces: LicenseLimitationModel = LicenseLimitationModel(enabled=False, size=0, limit=0)
|
||||
seats: LicenseLimitationModel = LicenseLimitationModel(enabled=False, size=0, limit=0)
|
||||
|
||||
|
||||
class BrandingModel(FeatureResponseModel):
|
||||
@@ -457,6 +458,11 @@ class FeatureService:
|
||||
features.license.workspaces.limit = workspaces_info.get("limit", 0)
|
||||
features.license.workspaces.size = workspaces_info.get("used", 0)
|
||||
|
||||
if seats_info := license_info.get("licensedSeats"):
|
||||
features.license.seats.enabled = seats_info.get("enabled", False)
|
||||
features.license.seats.limit = seats_info.get("limit", 0)
|
||||
features.license.seats.size = seats_info.get("used", 0)
|
||||
|
||||
if "PluginInstallationPermission" in enterprise_info:
|
||||
plugin_installation_info = enterprise_info["PluginInstallationPermission"]
|
||||
features.plugin_installation_permission.plugin_installation_scope = plugin_installation_info[
|
||||
|
||||
@@ -1,17 +1,28 @@
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import scoped_session
|
||||
from sqlalchemy.orm import Session, scoped_session
|
||||
|
||||
from configs import dify_config
|
||||
from models.model import AccountTrialAppRecord, TrialApp
|
||||
from models.model import AccountTrialAppRecord, App, TrialApp
|
||||
from services.feature_service import FeatureService
|
||||
from services.recommend_app.recommend_app_factory import RecommendAppRetrievalFactory
|
||||
|
||||
|
||||
class RecommendedAppService:
|
||||
@classmethod
|
||||
def get_recommended_apps_and_categories(cls, session: scoped_session, language: str):
|
||||
def get_app(cls, app_id: str, *, session: Session) -> App | None:
|
||||
"""Return a normal app only when it belongs to the recommended catalog."""
|
||||
mode = dify_config.HOSTED_FETCH_APP_TEMPLATES_MODE
|
||||
retrieval_instance = RecommendAppRetrievalFactory.get_recommend_app_factory(mode)()
|
||||
recommended_app_detail = retrieval_instance.get_recommend_app_detail(app_id, session=session)
|
||||
if recommended_app_detail is None:
|
||||
return None
|
||||
|
||||
return session.scalar(select(App).where(App.id == app_id, App.status == "normal").limit(1))
|
||||
|
||||
@classmethod
|
||||
def get_recommended_apps_and_categories(cls, language: str, *, session: Session):
|
||||
"""
|
||||
Get recommended apps and categories.
|
||||
:param language: language
|
||||
|
||||
@@ -2,11 +2,11 @@ import json
|
||||
import logging
|
||||
from typing import Any, TypedDict, cast
|
||||
|
||||
from httpx import get
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from core.entities.provider_entities import ProviderConfig
|
||||
from core.helper import ssrf_proxy
|
||||
from core.tools.__base.tool_runtime import ToolRuntime
|
||||
from core.tools.custom_tool.provider import ApiToolProviderController
|
||||
from core.tools.entities.api_entities import ToolApiEntity, ToolProviderApiEntity
|
||||
@@ -212,7 +212,7 @@ class ApiToolManageService:
|
||||
}
|
||||
|
||||
try:
|
||||
response = get(url, headers=headers, timeout=10)
|
||||
response = ssrf_proxy.get(url, headers=headers, timeout=10)
|
||||
if response.status_code != 200:
|
||||
raise ValueError(f"Got status code {response.status_code}")
|
||||
schema = response.text
|
||||
|
||||
@@ -25,6 +25,7 @@ from core.repositories import DifyCoreRepositoryFactory
|
||||
from extensions.ext_database import db
|
||||
from graphon.entities import WorkflowStartReason
|
||||
from graphon.enums import WorkflowExecutionStatus
|
||||
from graphon.filters import ResponseStreamFilter
|
||||
from graphon.runtime import GraphRuntimeState
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from libs.flask_utils import set_login_user
|
||||
@@ -486,6 +487,7 @@ def _resume_app_execution(payload: dict[str, Any]) -> None:
|
||||
generate_entity = resumption_context.get_generate_entity()
|
||||
|
||||
graph_runtime_state = GraphRuntimeState.from_snapshot(resumption_context.serialized_graph_runtime_state)
|
||||
response_stream_filter = resumption_context.get_response_stream_filter()
|
||||
|
||||
conversation = None
|
||||
message = None
|
||||
@@ -562,6 +564,7 @@ def _resume_app_execution(payload: dict[str, Any]) -> None:
|
||||
message=message,
|
||||
generate_entity=generate_entity,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
response_stream_filter=response_stream_filter,
|
||||
session_factory=session_factory,
|
||||
pause_state_config=pause_config,
|
||||
workflow_run_id=workflow_run_id,
|
||||
@@ -574,6 +577,7 @@ def _resume_app_execution(payload: dict[str, Any]) -> None:
|
||||
user=user,
|
||||
generate_entity=generate_entity,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
response_stream_filter=response_stream_filter,
|
||||
session_factory=session_factory,
|
||||
pause_state_config=pause_config,
|
||||
workflow_run_id=workflow_run_id,
|
||||
@@ -592,6 +596,7 @@ def _resume_advanced_chat(
|
||||
message: Message,
|
||||
generate_entity: AdvancedChatAppGenerateEntity,
|
||||
graph_runtime_state: GraphRuntimeState,
|
||||
response_stream_filter: ResponseStreamFilter,
|
||||
session_factory: sessionmaker,
|
||||
pause_state_config: PauseStateLayerConfig,
|
||||
workflow_run_id: str,
|
||||
@@ -631,6 +636,7 @@ def _resume_advanced_chat(
|
||||
workflow_node_execution_repository=workflow_node_execution_repository,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
pause_state_config=pause_state_config,
|
||||
response_stream_filter=response_stream_filter,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Failed to resume chatflow execution for workflow run %s", workflow_run_id)
|
||||
@@ -654,6 +660,7 @@ def _resume_workflow(
|
||||
user: Account | EndUser,
|
||||
generate_entity: WorkflowAppGenerateEntity,
|
||||
graph_runtime_state: GraphRuntimeState,
|
||||
response_stream_filter: ResponseStreamFilter,
|
||||
session_factory: sessionmaker,
|
||||
pause_state_config: PauseStateLayerConfig,
|
||||
workflow_run_id: str,
|
||||
@@ -693,6 +700,7 @@ def _resume_workflow(
|
||||
workflow_execution_repository=workflow_execution_repository,
|
||||
workflow_node_execution_repository=workflow_node_execution_repository,
|
||||
pause_state_config=pause_state_config,
|
||||
response_stream_filter=response_stream_filter,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Failed to resume workflow execution for workflow run %s", workflow_run_id)
|
||||
|
||||
@@ -232,6 +232,7 @@ def resume_workflow_execution(task_data_dict: dict[str, Any]) -> None:
|
||||
return
|
||||
|
||||
graph_runtime_state = GraphRuntimeState.from_snapshot(resumption_context.serialized_graph_runtime_state)
|
||||
response_stream_filter = resumption_context.get_response_stream_filter()
|
||||
|
||||
with session_factory() as session:
|
||||
workflow = session.scalar(select(Workflow).where(Workflow.id == workflow_run.workflow_id))
|
||||
@@ -294,6 +295,7 @@ def resume_workflow_execution(task_data_dict: dict[str, Any]) -> None:
|
||||
workflow_node_execution_repository=workflow_node_execution_repository,
|
||||
graph_engine_layers=graph_engine_layers,
|
||||
pause_state_config=pause_config,
|
||||
response_stream_filter=response_stream_filter,
|
||||
)
|
||||
workflow_run_repo.delete_workflow_pause(pause_entity)
|
||||
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
"""Initialize default RBAC access for existing workspace members after app creation."""
|
||||
|
||||
import logging
|
||||
|
||||
from celery import shared_task
|
||||
|
||||
from configs import dify_config
|
||||
from extensions.ext_database import db
|
||||
from services.account_service import TenantService
|
||||
from services.enterprise import rbac_service as enterprise_rbac_service
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
APP_RBAC_ACCOUNT_POLICY_BATCH_SIZE = 500
|
||||
APP_RBAC_DEFAULT_ACCESS_POLICY_ID = "default"
|
||||
APP_RBAC_QUEUE = "app_rbac"
|
||||
|
||||
|
||||
@shared_task(queue=APP_RBAC_QUEUE, bind=True, max_retries=3, default_retry_delay=60)
|
||||
def initialize_created_app_rbac_access_task(
|
||||
self, tenant_id: str, account_id: str, app_id: str | None = None, dataset_id: str | None = None
|
||||
) -> None:
|
||||
"""Grant the default app policy to current workspace members.
|
||||
|
||||
App scope is persisted synchronously before this task is queued. Replacing
|
||||
member policies is idempotent, so retrying the whole synchronization is safe
|
||||
when the enterprise RBAC service is temporarily unavailable.
|
||||
"""
|
||||
if not dify_config.RBAC_ENABLED:
|
||||
return
|
||||
|
||||
try:
|
||||
for account_ids in TenantService.iter_member_account_id_batches(
|
||||
tenant_id,
|
||||
APP_RBAC_ACCOUNT_POLICY_BATCH_SIZE,
|
||||
session=db.session(),
|
||||
):
|
||||
if app_id is not None:
|
||||
enterprise_rbac_service.RBACService.AppAccess.replace_user_access_policies(
|
||||
tenant_id=tenant_id,
|
||||
account_id=account_id,
|
||||
app_id=app_id,
|
||||
target_account_id=None,
|
||||
payload=enterprise_rbac_service.ReplaceUserAccessPolicies(
|
||||
access_policy_ids=[APP_RBAC_DEFAULT_ACCESS_POLICY_ID],
|
||||
account_ids=account_ids,
|
||||
),
|
||||
)
|
||||
elif dataset_id is not None:
|
||||
enterprise_rbac_service.RBACService.DatasetAccess.replace_user_access_policies(
|
||||
tenant_id=tenant_id,
|
||||
account_id=account_id,
|
||||
dataset_id=dataset_id,
|
||||
target_account_id=None,
|
||||
payload=enterprise_rbac_service.ReplaceUserAccessPolicies(
|
||||
access_policy_ids=[APP_RBAC_DEFAULT_ACCESS_POLICY_ID],
|
||||
account_ids=account_ids,
|
||||
),
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.exception(
|
||||
"Failed to initialize app RBAC access; retrying: tenant_id=%s app_id=%s attempt=%s",
|
||||
tenant_id,
|
||||
app_id,
|
||||
self.request.retries + 1,
|
||||
)
|
||||
raise self.retry(exc=exc)
|
||||
@@ -701,16 +701,29 @@ def _delete_records(query_sql: str, params: dict[str, Any], delete_func: Callabl
|
||||
if not rows:
|
||||
break
|
||||
|
||||
success_count = 0
|
||||
for i in rows:
|
||||
record_id = str(i.id)
|
||||
try:
|
||||
delete_func(session, record_id)
|
||||
logger.info(click.style(f"Deleted {name} {record_id}", fg="green"))
|
||||
session.commit()
|
||||
success_count += 1
|
||||
except Exception:
|
||||
logger.exception("Error occurred while deleting %s %s", name, record_id)
|
||||
# continue with next record even if one deletion fails
|
||||
session.rollback()
|
||||
break
|
||||
session.commit()
|
||||
continue
|
||||
|
||||
rs.close()
|
||||
|
||||
# If we couldn't delete ANY records in this batch, we must break out of the while loop
|
||||
# to prevent an infinite loop where we keep fetching the same failing records.
|
||||
if success_count == 0:
|
||||
logger.warning(
|
||||
click.style(
|
||||
f"Failed to delete any {name} in the current batch. Stopping to prevent infinite loop.",
|
||||
fg="yellow",
|
||||
)
|
||||
)
|
||||
break
|
||||
|
||||
+217
@@ -0,0 +1,217 @@
|
||||
"""Regression test: if-else branch + human_input pause + downstream answer nodes.
|
||||
|
||||
Reproduces https://github.com/langgenius/dify/issues/38525 at the
|
||||
iter_dify_graph_engine_events layer: without a restored ResponseStreamFilter,
|
||||
answer nodes downstream of a pre-pause branch never unlock for streaming on
|
||||
resume, even though the graph executes correctly.
|
||||
"""
|
||||
|
||||
from datetime import timedelta
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from core.workflow.nodes.human_input.callback import DifyHITLCallback
|
||||
from core.workflow.nodes.human_input.entities import HumanInputNodeData, UserActionConfig
|
||||
from core.workflow.nodes.human_input.enums import HumanInputFormStatus
|
||||
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom
|
||||
from core.repositories.human_input_repository import HumanInputFormEntity, HumanInputFormRepository
|
||||
from core.workflow.system_variables import build_system_variables
|
||||
from core.workflow.workflow_entry import iter_dify_graph_engine_events
|
||||
from graphon.filters import GraphEventFilterContext, ResponseStreamFilter, filter_graph_events
|
||||
from graphon.graph import Graph
|
||||
from graphon.graph_engine import GraphEngine, GraphEngineConfig
|
||||
from graphon.graph_engine.command_channels import InMemoryChannel
|
||||
from graphon.graph_events import GraphRunPausedEvent, GraphRunSucceededEvent, NodeRunStreamChunkEvent
|
||||
from graphon.nodes.answer.answer_node import AnswerNode
|
||||
from graphon.nodes.answer.entities import AnswerNodeData
|
||||
from graphon.nodes.human_input.human_input_node import HumanInputNode
|
||||
from graphon.nodes.if_else.entities import IfElseNodeData
|
||||
from graphon.nodes.if_else.if_else_node import IfElseNode
|
||||
from graphon.nodes.start.entities import StartNodeData
|
||||
from graphon.nodes.start.start_node import StartNode
|
||||
from graphon.runtime import GraphRuntimeState, VariablePool
|
||||
from graphon.utils.condition.entities import Condition
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from tests.workflow_test_utils import build_test_graph_init_params
|
||||
|
||||
WORKFLOW_EXECUTION_ID = "wf-exec-38525"
|
||||
|
||||
|
||||
def _mock_repo_paused() -> HumanInputFormRepository:
|
||||
repo = MagicMock(spec=HumanInputFormRepository)
|
||||
form = MagicMock(spec=HumanInputFormEntity)
|
||||
form.id = "form-1"
|
||||
form.submission_token = "token-1"
|
||||
form.recipients = []
|
||||
form.rendered_content = "rendered"
|
||||
form.submitted = False
|
||||
repo.create_form.return_value = form
|
||||
repo.get_form.return_value = None
|
||||
return repo
|
||||
|
||||
|
||||
def _mock_repo_resumed(action_id: str = "continue") -> HumanInputFormRepository:
|
||||
repo = MagicMock(spec=HumanInputFormRepository)
|
||||
form = MagicMock(spec=HumanInputFormEntity)
|
||||
form.id = "form-1"
|
||||
form.submission_token = "token-1"
|
||||
form.recipients = []
|
||||
form.rendered_content = "rendered"
|
||||
form.submitted = True
|
||||
form.selected_action_id = action_id
|
||||
form.submitted_data = {}
|
||||
form.status = HumanInputFormStatus.WAITING
|
||||
form.expiration_time = naive_utc_now() + timedelta(hours=1)
|
||||
repo.get_form.return_value = form
|
||||
return repo
|
||||
|
||||
|
||||
def _build_graph(runtime_state: GraphRuntimeState, form_repository: HumanInputFormRepository) -> Graph:
|
||||
params = build_test_graph_init_params(
|
||||
workflow_id="wf",
|
||||
graph_config={"nodes": [], "edges": []},
|
||||
user_from=UserFrom.ACCOUNT,
|
||||
invoke_from=InvokeFrom.DEBUGGER,
|
||||
)
|
||||
|
||||
start_node = StartNode(
|
||||
node_id="start",
|
||||
data=StartNodeData(title="start", variables=[]),
|
||||
graph_init_params=params,
|
||||
graph_runtime_state=runtime_state,
|
||||
)
|
||||
|
||||
if_else_node = IfElseNode(
|
||||
node_id="if_else",
|
||||
data=IfElseNodeData(
|
||||
title="if-else",
|
||||
cases=[
|
||||
IfElseNodeData.Case(
|
||||
case_id="true",
|
||||
logical_operator="and",
|
||||
conditions=[
|
||||
Condition(
|
||||
variable_selector=["start", "category"],
|
||||
comparison_operator="is",
|
||||
value="fruit",
|
||||
)
|
||||
],
|
||||
)
|
||||
],
|
||||
),
|
||||
graph_init_params=params,
|
||||
graph_runtime_state=runtime_state,
|
||||
)
|
||||
|
||||
human_data = HumanInputNodeData(
|
||||
title="human",
|
||||
form_content="Awaiting human input",
|
||||
inputs=[],
|
||||
user_actions=[UserActionConfig(id="continue", title="Continue")],
|
||||
)
|
||||
human_node = HumanInputNode(
|
||||
node_id="human_input",
|
||||
data=human_data,
|
||||
graph_init_params=params,
|
||||
graph_runtime_state=runtime_state,
|
||||
hitl_callback=DifyHITLCallback(form_repository=form_repository, node_data=human_data),
|
||||
)
|
||||
|
||||
answer_false_node = AnswerNode(
|
||||
node_id="answer_false",
|
||||
data=AnswerNodeData(title="answer_false", answer="unreachable branch"),
|
||||
graph_init_params=params,
|
||||
graph_runtime_state=runtime_state,
|
||||
)
|
||||
|
||||
answer_after_pause = AnswerNode(
|
||||
node_id="answer_after_pause",
|
||||
data=AnswerNodeData(title="answer_after_pause", answer="Post-branch answer chunk 1"),
|
||||
graph_init_params=params,
|
||||
graph_runtime_state=runtime_state,
|
||||
)
|
||||
|
||||
answer_after_pause_2 = AnswerNode(
|
||||
node_id="answer_after_pause_2",
|
||||
data=AnswerNodeData(title="answer_after_pause_2", answer="Post-branch answer chunk 2"),
|
||||
graph_init_params=params,
|
||||
graph_runtime_state=runtime_state,
|
||||
)
|
||||
|
||||
return (
|
||||
Graph.new()
|
||||
.add_root(start_node)
|
||||
.add_node(if_else_node, from_node_id="start")
|
||||
.add_node(human_node, from_node_id="if_else", source_handle="true")
|
||||
.add_node(answer_false_node, from_node_id="if_else", source_handle="false")
|
||||
.add_node(answer_after_pause, from_node_id="human_input", source_handle="continue")
|
||||
.add_node(answer_after_pause_2, from_node_id="answer_after_pause")
|
||||
.build()
|
||||
)
|
||||
|
||||
|
||||
def _build_runtime_state() -> GraphRuntimeState:
|
||||
variable_pool = VariablePool.from_bootstrap(
|
||||
system_variables=build_system_variables(
|
||||
workflow_execution_id=WORKFLOW_EXECUTION_ID,
|
||||
app_id="app",
|
||||
workflow_id="wf",
|
||||
user_id="user",
|
||||
),
|
||||
user_inputs={},
|
||||
conversation_variables=[],
|
||||
)
|
||||
variable_pool.add(("start", "category"), "fruit") # drives the if-else "true" branch
|
||||
return GraphRuntimeState(variable_pool=variable_pool, start_at=0.0)
|
||||
|
||||
|
||||
def test_if_else_human_input_pause_resume_answer_chunks_survive_resume() -> None:
|
||||
# ---- Phase 1: run to GraphRunPausedEvent ----
|
||||
runtime_state_1 = _build_runtime_state()
|
||||
graph_1 = _build_graph(runtime_state_1, _mock_repo_paused())
|
||||
engine_1 = GraphEngine(
|
||||
workflow_id="wf",
|
||||
graph=graph_1,
|
||||
graph_runtime_state=runtime_state_1,
|
||||
command_channel=InMemoryChannel(),
|
||||
config=GraphEngineConfig(),
|
||||
)
|
||||
filter_1 = ResponseStreamFilter()
|
||||
phase1_events = list(
|
||||
filter_graph_events(
|
||||
engine_1.run(),
|
||||
context=GraphEventFilterContext.from_engine(engine_1),
|
||||
filters=[filter_1],
|
||||
)
|
||||
)
|
||||
|
||||
assert any(isinstance(e, GraphRunPausedEvent) for e in phase1_events)
|
||||
phase1_chunks = [e for e in phase1_events if isinstance(e, NodeRunStreamChunkEvent)]
|
||||
assert not any(e.node_id in ("answer_after_pause", "answer_after_pause_2") for e in phase1_chunks)
|
||||
|
||||
response_filter_snapshot = filter_1.dumps()
|
||||
runtime_snapshot = runtime_state_1.dumps()
|
||||
|
||||
# ---- Phase 2: rebuild engine + filter from snapshots, resume to completion ----
|
||||
runtime_state_2 = GraphRuntimeState.from_snapshot(runtime_snapshot)
|
||||
graph_2 = _build_graph(runtime_state_2, _mock_repo_resumed(action_id="continue"))
|
||||
engine_2 = GraphEngine(
|
||||
workflow_id="wf",
|
||||
graph=graph_2,
|
||||
graph_runtime_state=runtime_state_2,
|
||||
command_channel=InMemoryChannel(),
|
||||
config=GraphEngineConfig(),
|
||||
)
|
||||
filter_2 = ResponseStreamFilter()
|
||||
filter_2.loads(response_filter_snapshot)
|
||||
|
||||
phase2_events = list(iter_dify_graph_engine_events(engine_2, filter_2))
|
||||
|
||||
assert any(isinstance(e, GraphRunSucceededEvent) for e in phase2_events)
|
||||
|
||||
phase2_chunks = [e for e in phase2_events if isinstance(e, NodeRunStreamChunkEvent)]
|
||||
answer_1_chunks = [e for e in phase2_chunks if e.node_id == "answer_after_pause"]
|
||||
answer_2_chunks = [e for e in phase2_chunks if e.node_id == "answer_after_pause_2"]
|
||||
|
||||
assert answer_1_chunks, "answer_after_pause produced no stream chunks after resume"
|
||||
assert answer_2_chunks, "answer_after_pause_2 produced no stream chunks after resume"
|
||||
@@ -105,6 +105,7 @@ class TestOAuthLogin:
|
||||
invite_token=expected_token,
|
||||
timezone=None,
|
||||
language=None,
|
||||
redirect_url=None,
|
||||
)
|
||||
mock_redirect.assert_called_once_with("https://github.com/login/oauth/authorize?...")
|
||||
|
||||
@@ -127,6 +128,7 @@ class TestOAuthLogin:
|
||||
invite_token=None,
|
||||
timezone="Asia/Shanghai",
|
||||
language=None,
|
||||
redirect_url=None,
|
||||
)
|
||||
mock_redirect.assert_called_once_with("https://github.com/login/oauth/authorize?...")
|
||||
|
||||
@@ -149,6 +151,7 @@ class TestOAuthLogin:
|
||||
invite_token=None,
|
||||
timezone=None,
|
||||
language="zh-Hans",
|
||||
redirect_url=None,
|
||||
)
|
||||
mock_redirect.assert_called_once_with("https://github.com/login/oauth/authorize?...")
|
||||
|
||||
|
||||
+19
@@ -20,6 +20,7 @@ providing more reliable and realistic test scenarios than mocks.
|
||||
import json
|
||||
import uuid
|
||||
from time import time
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import Engine, delete, select
|
||||
@@ -35,6 +36,7 @@ from core.workflow.system_variables import build_system_variables
|
||||
from extensions.ext_storage import storage
|
||||
from graphon.entities.pause_reason import SchedulingPause
|
||||
from graphon.enums import WorkflowExecutionStatus
|
||||
from graphon.filters import GraphEventFilterContext, ResponseStreamFilter
|
||||
from graphon.graph_engine.entities.commands import GraphEngineCommand
|
||||
from graphon.graph_engine.layers.base import GraphEngineLayerNotInitializedError
|
||||
from graphon.graph_events import GraphRunPausedEvent
|
||||
@@ -49,6 +51,22 @@ from services.file_service import FileService
|
||||
from services.workflow_run_service import WorkflowRunService
|
||||
|
||||
|
||||
def _create_initialized_response_stream_filter() -> ResponseStreamFilter:
|
||||
"""Build a `ResponseStreamFilter` that has already run `initialize()`.
|
||||
|
||||
`ResponseStreamFilter.dumps()` raises `RuntimeError` unless the filter has
|
||||
processed a `GraphEventFilterContext` first. In production this always
|
||||
happens before any event (including `GraphRunPausedEvent`) reaches
|
||||
`PauseStatePersistenceLayer.on_event`, so tests that exercise `on_event`
|
||||
or a subsequent `dumps()` call need a filter in that same state. A
|
||||
nodeless graph is enough to satisfy the precondition.
|
||||
"""
|
||||
response_stream_filter = ResponseStreamFilter()
|
||||
context = GraphEventFilterContext(graph=Mock(nodes={}), runtime_state=Mock())
|
||||
response_stream_filter.initialize(context)
|
||||
return response_stream_filter
|
||||
|
||||
|
||||
class _TestCommandChannelImpl:
|
||||
"""Real implementation of CommandChannel for testing."""
|
||||
|
||||
@@ -295,6 +313,7 @@ class TestPauseStatePersistenceLayerTestContainers:
|
||||
session_factory=self.session.get_bind(),
|
||||
state_owner_user_id=owner_id,
|
||||
generate_entity=entity,
|
||||
response_stream_filter=_create_initialized_response_stream_filter(),
|
||||
)
|
||||
|
||||
def test_complete_pause_flow_with_real_dependencies(self, db_session_with_containers: Session):
|
||||
|
||||
@@ -17,6 +17,7 @@ from services.errors.account import (
|
||||
AccountPasswordError,
|
||||
AccountRegisterError,
|
||||
CurrentPasswordIncorrectError,
|
||||
SeatsLimitExceededError,
|
||||
TenantNotFoundError,
|
||||
)
|
||||
from services.errors.workspace import WorkSpaceNotAllowedCreateError, WorkspacesLimitExceededError
|
||||
@@ -477,6 +478,32 @@ class TestAccountService:
|
||||
session=db_session_with_containers,
|
||||
)
|
||||
|
||||
def test_create_account_seats_limit_exceeded(
|
||||
self, db_session_with_containers: Session, mock_external_service_dependencies
|
||||
):
|
||||
"""
|
||||
Test account creation when the licensed seats limit is exceeded.
|
||||
"""
|
||||
fake = Faker()
|
||||
email = fake.email()
|
||||
name = fake.name()
|
||||
password = generate_valid_password(fake)
|
||||
# Setup mocks
|
||||
mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True
|
||||
mock_external_service_dependencies[
|
||||
"feature_service"
|
||||
].get_system_features.return_value.license.seats.is_available.return_value = False
|
||||
mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = False
|
||||
|
||||
with pytest.raises(SeatsLimitExceededError):
|
||||
AccountService.create_account(
|
||||
email=email,
|
||||
name=name,
|
||||
interface_language="en-US",
|
||||
password=password,
|
||||
session=db_session_with_containers,
|
||||
)
|
||||
|
||||
def test_link_account_integrate_new_provider(
|
||||
self, db_session_with_containers: Session, mock_external_service_dependencies
|
||||
):
|
||||
|
||||
+71
-2
@@ -9,11 +9,16 @@ import pytest
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from extensions.ext_database import db
|
||||
from models.model import AccountTrialAppRecord, TrialApp
|
||||
from models.model import AccountTrialAppRecord, App, AppMode, TrialApp
|
||||
from services import recommended_app_service as service_module
|
||||
from services.recommended_app_service import RecommendedAppService
|
||||
|
||||
pytestmark = pytest.mark.parametrize(
|
||||
"sqlite_session",
|
||||
[(TrialApp, AccountTrialAppRecord, App)],
|
||||
indirect=True,
|
||||
)
|
||||
|
||||
|
||||
class RecommendedAppPayload(TypedDict, total=False):
|
||||
id: str
|
||||
@@ -103,6 +108,37 @@ def _mock_factory_for_apps(
|
||||
return retrieval_instance, builtin_instance
|
||||
|
||||
|
||||
def _mock_factory_for_app_detail(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
*,
|
||||
result: RecommendedAppPayload | None,
|
||||
) -> MagicMock:
|
||||
retrieval_instance = MagicMock()
|
||||
retrieval_instance.get_recommend_app_detail.return_value = result
|
||||
retrieval_factory = MagicMock(return_value=retrieval_instance)
|
||||
monkeypatch.setattr(service_module.dify_config, "HOSTED_FETCH_APP_TEMPLATES_MODE", "remote", raising=False)
|
||||
monkeypatch.setattr(
|
||||
service_module.RecommendAppRetrievalFactory,
|
||||
"get_recommend_app_factory",
|
||||
MagicMock(return_value=retrieval_factory),
|
||||
)
|
||||
return retrieval_instance
|
||||
|
||||
|
||||
def _persist_app(session: Session, *, name: str) -> App:
|
||||
app = App(
|
||||
tenant_id=str(uuid.uuid4()),
|
||||
name=name,
|
||||
mode=AppMode.CHAT,
|
||||
enable_site=True,
|
||||
enable_api=True,
|
||||
)
|
||||
app.id = str(uuid.uuid4())
|
||||
session.add(app)
|
||||
session.commit()
|
||||
return app
|
||||
|
||||
|
||||
# ── Pure logic tests: get_recommended_apps_and_categories ──────────────
|
||||
|
||||
|
||||
@@ -202,6 +238,39 @@ class TestRecommendedAppServiceGetApps:
|
||||
mock_factory_class.get_recommend_app_factory.assert_called_with(mode)
|
||||
|
||||
|
||||
# ── Database-backed tests: get_app ─────────────────────────────────────
|
||||
|
||||
|
||||
class TestRecommendedAppServiceGetApp:
|
||||
def test_returns_normal_recommended_app(self, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None:
|
||||
app = _persist_app(sqlite_session, name="Recommended App")
|
||||
|
||||
retrieval_instance = _mock_factory_for_app_detail(
|
||||
monkeypatch,
|
||||
result=RecommendedAppPayload(id=app.id),
|
||||
)
|
||||
feature_lookup = MagicMock(side_effect=AssertionError("get_app must not inspect trial features"))
|
||||
monkeypatch.setattr(service_module.FeatureService, "get_system_features", feature_lookup)
|
||||
|
||||
result = RecommendedAppService.get_app(app.id, session=sqlite_session)
|
||||
|
||||
assert result is app
|
||||
retrieval_instance.get_recommend_app_detail.assert_called_once_with(app.id, session=sqlite_session)
|
||||
feature_lookup.assert_not_called()
|
||||
|
||||
def test_returns_none_when_app_is_not_recommended(
|
||||
self, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
|
||||
) -> None:
|
||||
app = _persist_app(sqlite_session, name="Private App")
|
||||
|
||||
retrieval_instance = _mock_factory_for_app_detail(monkeypatch, result=None)
|
||||
|
||||
result = RecommendedAppService.get_app(app.id, session=sqlite_session)
|
||||
|
||||
assert result is None
|
||||
retrieval_instance.get_recommend_app_detail.assert_called_once_with(app.id, session=sqlite_session)
|
||||
|
||||
|
||||
# ── Pure logic tests: get_recommend_app_detail ─────────────────────────
|
||||
|
||||
|
||||
|
||||
@@ -377,8 +377,8 @@ def test_import_options_prompts_explain_secrets_reuse_and_conflicts(monkeypatch)
|
||||
assert prompt_calls[0][0] == "Import ID strategy. Enter one of: preserve-id, generate-new-id"
|
||||
assert prompt_calls[0][1]["default"] == "preserve-id"
|
||||
assert prompt_calls[0][1]["show_default"] is True
|
||||
assert prompt_calls[0][1]["type"].choices == ID_STRATEGY_CHOICES
|
||||
assert list(prompt_calls[0][1]["type"].choices) == ID_STRATEGY_CHOICES
|
||||
assert prompt_calls[1][0] == "Import conflict strategy. Enter one of: fail, skip, update"
|
||||
assert prompt_calls[1][1]["default"] == "update"
|
||||
assert prompt_calls[1][1]["show_default"] is True
|
||||
assert prompt_calls[1][1]["type"].choices == CONFLICT_STRATEGY_CHOICES
|
||||
assert list(prompt_calls[1][1]["type"].choices) == CONFLICT_STRATEGY_CHOICES
|
||||
|
||||
@@ -336,6 +336,174 @@ def _insert_load_balancing_model_config(
|
||||
)
|
||||
|
||||
|
||||
def test_data_migrate_group_registers_dataset_permission_rbac_migration(command_module) -> None:
|
||||
command = command_module.data_migrate.commands["rbac-migrate-dataset-permissions"]
|
||||
|
||||
assert command is command_module.migrate_dataset_permissions_to_rbac
|
||||
assert "operator_account_id" not in {param.name for param in command.params}
|
||||
|
||||
|
||||
def test_dataset_permission_rbac_migration_help_mentions_binding_clear_side_effect(command_module) -> None:
|
||||
result = CliRunner().invoke(
|
||||
command_module.data_migrate,
|
||||
["rbac-migrate-dataset-permissions", "--help"],
|
||||
)
|
||||
|
||||
assert result.exit_code == 0
|
||||
normalized_output = " ".join(result.output.split())
|
||||
assert "clears existing per-user policy bindings" in normalized_output
|
||||
assert "recreates legacy partial-member default bindings" in normalized_output
|
||||
|
||||
|
||||
def test_dataset_permission_rbac_migration_maps_legacy_permissions_to_enum_scopes() -> None:
|
||||
rbac_module = importlib.import_module("commands.rbac")
|
||||
|
||||
assert (
|
||||
rbac_module._rbac_dataset_scope_for_legacy_permission(rbac_module.DatasetPermissionEnum.ALL_TEAM)
|
||||
is rbac_module.RBACResourceWhitelistScope.ALL
|
||||
)
|
||||
assert (
|
||||
rbac_module._rbac_dataset_scope_for_legacy_permission(rbac_module.DatasetPermissionEnum.PARTIAL_TEAM)
|
||||
is rbac_module.RBACResourceWhitelistScope.SPECIFIC
|
||||
)
|
||||
assert rbac_module._dataset_permission_enum("partial_members") is rbac_module.DatasetPermissionEnum.PARTIAL_TEAM
|
||||
assert rbac_module._dataset_permission_enum(None) is rbac_module.DatasetPermissionEnum.ONLY_ME
|
||||
|
||||
|
||||
def test_dataset_permission_rbac_migration_uses_dataset_creator_as_operator(
|
||||
command_module,
|
||||
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], [], []]
|
||||
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()
|
||||
|
||||
def fake_replace_whitelist(**kwargs):
|
||||
assert session_closed is True
|
||||
calls.append(kwargs)
|
||||
|
||||
monkeypatch.setattr(rbac_module, "session_factory", FakeSessionFactory)
|
||||
monkeypatch.setattr(rbac_module.RBACService.DatasetAccess, "replace_whitelist", fake_replace_whitelist)
|
||||
|
||||
command_module.migrate_dataset_permissions_to_rbac.callback(
|
||||
tenant_id=None,
|
||||
dataset_id=None,
|
||||
batch_size=500,
|
||||
dry_run=False,
|
||||
)
|
||||
|
||||
assert calls[0]["tenant_id"] == "tenant-1"
|
||||
assert calls[0]["account_id"] == "creator-account-1"
|
||||
assert calls[0]["dataset_id"] == "dataset-1"
|
||||
assert calls[0]["payload"].scope is rbac_module.RBACResourceWhitelistScope.SPECIFIC
|
||||
|
||||
|
||||
def test_dataset_permission_rbac_migration_dry_run_outputs_structured_proposed_changes(
|
||||
command_module,
|
||||
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",
|
||||
)
|
||||
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)
|
||||
monkeypatch.setattr(
|
||||
rbac_module.RBACService.DatasetAccess,
|
||||
"replace_whitelist",
|
||||
lambda **kwargs: pytest.fail("dry-run must not replace whitelist"),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rbac_module.RBACService.DatasetAccess,
|
||||
"replace_user_access_policies",
|
||||
lambda **kwargs: pytest.fail("dry-run must not replace user access policies"),
|
||||
)
|
||||
|
||||
result = CliRunner().invoke(
|
||||
command_module.data_migrate,
|
||||
["rbac-migrate-dataset-permissions", "--dry-run"],
|
||||
)
|
||||
|
||||
assert result.exit_code == 0
|
||||
events = [json.loads(line) for line in result.output.splitlines() if line.startswith("{")]
|
||||
assert [event["action"] for event in events] == ["replace_whitelist", "replace_user_access_policies"]
|
||||
assert events[0]["before"] == {
|
||||
"legacy_dataset_permission": "partial_members",
|
||||
"legacy_partial_member_ids": ["member-account-1"],
|
||||
}
|
||||
assert events[0]["after"] == {"rbac_whitelist_scope": "specific"}
|
||||
assert events[0]["call"] == {
|
||||
"method": "RBACService.DatasetAccess.replace_whitelist",
|
||||
"kwargs": {
|
||||
"tenant_id": "tenant-1",
|
||||
"account_id": "creator-account-1",
|
||||
"dataset_id": "dataset-1",
|
||||
"payload": {"scope": "specific"},
|
||||
},
|
||||
}
|
||||
assert events[1]["target_account_id"] == "member-account-1"
|
||||
assert events[1]["after"] == {"rbac_user_access_policy_ids": ["default"]}
|
||||
assert events[1]["call"]["kwargs"]["payload"] == {"access_policy_ids": ["default"]}
|
||||
|
||||
|
||||
def test_data_migrate_command_defaults_output_to_stdout_stream(
|
||||
command_module,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
|
||||
@@ -543,6 +543,7 @@ def test_app_create_api_attaches_permission_keys(app, app_module):
|
||||
|
||||
with app.test_request_context("/apps", method="POST", json={}):
|
||||
with pytest.MonkeyPatch.context() as monkeypatch:
|
||||
monkeypatch.setattr(app_module.dify_config, "RBAC_ENABLED", True)
|
||||
app_module.console_ns.payload = {
|
||||
"name": "Created App",
|
||||
"description": "Summary",
|
||||
@@ -558,11 +559,25 @@ def test_app_create_api_attaches_permission_keys(app, app_module):
|
||||
"batch_get",
|
||||
lambda tenant_id, account_id, app_ids: {"app-new": ["app.acl.view_layout", "app.acl.edit"]},
|
||||
)
|
||||
initialize_rbac_task = MagicMock()
|
||||
monkeypatch.setattr(
|
||||
app_module,
|
||||
"initialize_created_app_rbac_access_task",
|
||||
initialize_rbac_task,
|
||||
)
|
||||
replace_whitelist = MagicMock()
|
||||
monkeypatch.setattr(
|
||||
app_module.enterprise_rbac_service.RBACService.AppAccess,
|
||||
"replace_whitelist",
|
||||
replace_whitelist,
|
||||
)
|
||||
|
||||
resp, status = method(app_module.AppListApi(), "tenant-1", SimpleNamespace(id="acct-1"))
|
||||
|
||||
assert status == 201
|
||||
assert resp["permission_keys"] == ["app.acl.view_layout", "app.acl.edit"]
|
||||
assert replace_whitelist.call_args.kwargs["payload"].scope is app_module.RBACResourceWhitelistScope.ALL
|
||||
initialize_rbac_task.delay.assert_called_once_with("tenant-1", "acct-1", app_id="app-new")
|
||||
|
||||
|
||||
def test_app_list_api_attaches_permission_keys(app, app_module):
|
||||
|
||||
@@ -0,0 +1,154 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from contextlib import nullcontext
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, PropertyMock, patch
|
||||
|
||||
import pytest
|
||||
from flask import Flask
|
||||
from werkzeug.exceptions import Forbidden
|
||||
|
||||
from controllers.common import wraps as common_wraps
|
||||
from controllers.console import console_ns
|
||||
from controllers.console import wraps as console_wraps
|
||||
from controllers.console.app import ops_trace as ops_trace_module
|
||||
from controllers.console.app import wraps as app_wraps
|
||||
from libs import login as login_lib
|
||||
from models.account import Account, AccountStatus, TenantAccountRole
|
||||
|
||||
|
||||
def _make_account(role: TenantAccountRole) -> Account:
|
||||
account = Account(name="tester", email="tester@example.com")
|
||||
account.id = "account-123" # type: ignore[assignment]
|
||||
account.status = AccountStatus.ACTIVE
|
||||
account.role = role
|
||||
account._current_tenant = SimpleNamespace(id="tenant-123") # type: ignore[assignment]
|
||||
account._get_current_object = lambda: account # type: ignore[attr-defined]
|
||||
return account
|
||||
|
||||
|
||||
def _make_app() -> SimpleNamespace:
|
||||
return SimpleNamespace(id="app-123", tenant_id="tenant-123", status="normal", mode="chat")
|
||||
|
||||
|
||||
def _patch_console_guards(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
account: Account,
|
||||
app_model: SimpleNamespace,
|
||||
*,
|
||||
rbac_enabled: bool = False,
|
||||
) -> None:
|
||||
monkeypatch.setattr(login_lib.dify_config, "LOGIN_DISABLED", True)
|
||||
monkeypatch.setattr(login_lib.dify_config, "RBAC_ENABLED", rbac_enabled)
|
||||
monkeypatch.setattr(console_wraps.dify_config, "EDITION", "CLOUD")
|
||||
monkeypatch.setattr(login_lib, "current_user", account)
|
||||
monkeypatch.setattr(login_lib, "current_account_with_tenant", lambda: (account, account.current_tenant_id))
|
||||
monkeypatch.setattr(console_wraps, "current_account_with_tenant", lambda: (account, account.current_tenant_id))
|
||||
monkeypatch.setattr(common_wraps, "current_account_with_tenant", lambda: (account, account.current_tenant_id))
|
||||
monkeypatch.setattr(app_wraps, "_load_app_model_from_scoped_session", lambda _app_id: app_model)
|
||||
|
||||
|
||||
def _patch_payload(payload: dict[str, object] | None):
|
||||
if payload is None:
|
||||
return nullcontext()
|
||||
return patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("method_name", "path", "payload", "service_method_name", "service_result"),
|
||||
[
|
||||
(
|
||||
"post",
|
||||
"/console/api/apps/app-123/trace-config",
|
||||
{"tracing_provider": "mlflow", "tracing_config": {"endpoint": "https://trace.example.com"}},
|
||||
"create_tracing_app_config",
|
||||
{"id": "trace-config-1"},
|
||||
),
|
||||
(
|
||||
"patch",
|
||||
"/console/api/apps/app-123/trace-config",
|
||||
{"tracing_provider": "mlflow", "tracing_config": {"endpoint": "https://trace.example.com"}},
|
||||
"update_tracing_app_config",
|
||||
True,
|
||||
),
|
||||
(
|
||||
"delete",
|
||||
"/console/api/apps/app-123/trace-config?tracing_provider=mlflow",
|
||||
None,
|
||||
"delete_tracing_app_config",
|
||||
True,
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_trace_config_mutations_require_edit_permission(
|
||||
app: Flask,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
method_name: str,
|
||||
path: str,
|
||||
payload: dict[str, object] | None,
|
||||
service_method_name: str,
|
||||
service_result: object,
|
||||
) -> None:
|
||||
app.config.setdefault("RESTX_MASK_HEADER", "X-Fields")
|
||||
account = _make_account(TenantAccountRole.NORMAL)
|
||||
_patch_console_guards(monkeypatch, account, _make_app())
|
||||
service_mock = MagicMock(return_value=service_result)
|
||||
monkeypatch.setattr(ops_trace_module.OpsService, service_method_name, service_mock)
|
||||
|
||||
with app.test_request_context(path, method=method_name.upper(), json=payload):
|
||||
with _patch_payload(payload):
|
||||
with pytest.raises(Forbidden):
|
||||
getattr(ops_trace_module.TraceAppConfigApi(), method_name)(app_id="app-123")
|
||||
|
||||
service_mock.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("method_name", "path", "payload", "service_method_name", "service_result"),
|
||||
[
|
||||
(
|
||||
"post",
|
||||
"/console/api/apps/app-123/trace-config",
|
||||
{"tracing_provider": "mlflow", "tracing_config": {"endpoint": "https://trace.example.com"}},
|
||||
"create_tracing_app_config",
|
||||
{"id": "trace-config-1"},
|
||||
),
|
||||
(
|
||||
"patch",
|
||||
"/console/api/apps/app-123/trace-config",
|
||||
{"tracing_provider": "mlflow", "tracing_config": {"endpoint": "https://trace.example.com"}},
|
||||
"update_tracing_app_config",
|
||||
True,
|
||||
),
|
||||
(
|
||||
"delete",
|
||||
"/console/api/apps/app-123/trace-config?tracing_provider=mlflow",
|
||||
None,
|
||||
"delete_tracing_app_config",
|
||||
True,
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_trace_config_mutations_require_rbac_permission(
|
||||
app: Flask,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
method_name: str,
|
||||
path: str,
|
||||
payload: dict[str, object] | None,
|
||||
service_method_name: str,
|
||||
service_result: object,
|
||||
) -> None:
|
||||
app.config.setdefault("RESTX_MASK_HEADER", "X-Fields")
|
||||
account = _make_account(TenantAccountRole.NORMAL)
|
||||
_patch_console_guards(monkeypatch, account, _make_app(), rbac_enabled=True)
|
||||
monkeypatch.setattr(common_wraps.db, "session", SimpleNamespace(scalar=lambda _stmt: "other-account"))
|
||||
monkeypatch.setattr(common_wraps.RBACService.CheckAccess, "check", MagicMock(return_value=False))
|
||||
service_mock = MagicMock(return_value=service_result)
|
||||
monkeypatch.setattr(ops_trace_module.OpsService, service_method_name, service_mock)
|
||||
|
||||
with app.test_request_context(path, method=method_name.upper(), json=payload):
|
||||
with _patch_payload(payload):
|
||||
with pytest.raises(Forbidden):
|
||||
getattr(ops_trace_module.TraceAppConfigApi(), method_name)(app_id="app-123")
|
||||
|
||||
service_mock.assert_not_called()
|
||||
@@ -1,12 +1,15 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import Select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from controllers.console.app import wraps as wraps_module
|
||||
from controllers.console.app.error import AppNotFoundError
|
||||
from models.model import AppMode
|
||||
from models.model import App, AppMode, TrialApp
|
||||
|
||||
|
||||
class FakeSession:
|
||||
@@ -117,6 +120,66 @@ def test_get_app_model_rejects_wrong_mode(monkeypatch: pytest.MonkeyPatch) -> No
|
||||
handler(app_id="app-1")
|
||||
|
||||
|
||||
def test_get_app_model_with_trial_requires_trial_app_registration(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
app_model = SimpleNamespace(id="app-1", mode=AppMode.CHAT.value, status="normal", tenant_id="t1")
|
||||
|
||||
def scalar(statement: Select[tuple[App]]) -> object | None:
|
||||
has_trial_app_join = any(
|
||||
from_clause.is_derived_from(TrialApp.__table__) for from_clause in statement.get_final_froms()
|
||||
)
|
||||
return None if has_trial_app_join else app_model
|
||||
|
||||
scoped_session = MagicMock()
|
||||
scoped_session.scalar.side_effect = scalar
|
||||
recommended_get_app = MagicMock(return_value=None)
|
||||
monkeypatch.setattr(wraps_module.db, "session", scoped_session)
|
||||
monkeypatch.setattr(wraps_module.RecommendedAppService, "get_app", recommended_get_app)
|
||||
|
||||
@wraps_module.get_app_model_with_trial
|
||||
def handler(app_model):
|
||||
return app_model.id
|
||||
|
||||
with pytest.raises(AppNotFoundError):
|
||||
handler(app_id="app-1")
|
||||
|
||||
recommended_get_app.assert_called_once_with("app-1", session=scoped_session.return_value)
|
||||
|
||||
|
||||
def test_get_app_model_with_trial_falls_back_to_recommended_app(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
app_model = SimpleNamespace(id="app-1", mode=AppMode.CHAT.value, status="normal", tenant_id="t1")
|
||||
session = MagicMock(spec=Session)
|
||||
scoped_session = MagicMock(return_value=session)
|
||||
trial_app_loader = MagicMock(return_value=None)
|
||||
recommended_get_app = MagicMock(return_value=app_model)
|
||||
monkeypatch.setattr(wraps_module.db, "session", scoped_session)
|
||||
monkeypatch.setattr(wraps_module, "_load_app_model_with_trial", trial_app_loader)
|
||||
monkeypatch.setattr(wraps_module.RecommendedAppService, "get_app", recommended_get_app)
|
||||
|
||||
@wraps_module.get_app_model_with_trial
|
||||
def handler(app_model):
|
||||
return app_model.id
|
||||
|
||||
assert handler(app_id="app-1") == "app-1"
|
||||
trial_app_loader.assert_called_once_with("app-1")
|
||||
recommended_get_app.assert_called_once_with("app-1", session=session)
|
||||
|
||||
|
||||
def test_get_app_model_with_trial_prefers_trial_registration(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
app_model = SimpleNamespace(id="app-1", mode=AppMode.CHAT.value, status="normal", tenant_id="t1")
|
||||
trial_app_loader = MagicMock(return_value=app_model)
|
||||
recommended_get_app = MagicMock()
|
||||
monkeypatch.setattr(wraps_module, "_load_app_model_with_trial", trial_app_loader)
|
||||
monkeypatch.setattr(wraps_module.RecommendedAppService, "get_app", recommended_get_app)
|
||||
|
||||
@wraps_module.get_app_model_with_trial
|
||||
def handler(app_model):
|
||||
return app_model.id
|
||||
|
||||
assert handler(app_id="app-1") == "app-1"
|
||||
trial_app_loader.assert_called_once_with("app-1")
|
||||
recommended_get_app.assert_not_called()
|
||||
|
||||
|
||||
def test_get_app_model_requires_app_id() -> None:
|
||||
@wraps_module.get_app_model
|
||||
def handler(app_model):
|
||||
|
||||
@@ -26,10 +26,11 @@ from controllers.console.auth.login import EmailCodeLoginApi, LoginApi, LogoutAp
|
||||
from controllers.console.error import (
|
||||
AccountBannedError,
|
||||
AccountInFreezeError,
|
||||
SeatsLimitExceeded,
|
||||
WorkspacesLimitExceeded,
|
||||
)
|
||||
from services.entities.auth_entities import LoginFailureReason
|
||||
from services.errors.account import AccountLoginError, AccountPasswordError
|
||||
from services.errors.account import AccountLoginError, AccountPasswordError, SeatsLimitExceededError
|
||||
|
||||
|
||||
def encode_password(password: str) -> str:
|
||||
@@ -487,6 +488,45 @@ class TestLoginApi:
|
||||
assert warn_records[0].args[0] == "user@example.com"
|
||||
assert warn_records[0].args[1] == LoginFailureReason.ACCOUNT_BANNED
|
||||
|
||||
@patch("controllers.console.wraps.db")
|
||||
@patch("controllers.console.auth.login.db")
|
||||
@patch("controllers.console.auth.login.AccountService.create_account_and_tenant")
|
||||
@patch("controllers.console.auth.login.AccountService.get_email_code_login_data")
|
||||
@patch("controllers.console.auth.login.AccountService.revoke_email_code_login_token")
|
||||
@patch("controllers.console.auth.login._get_account_with_case_fallback")
|
||||
def test_email_code_login_fails_when_seats_limit_exceeded(
|
||||
self,
|
||||
mock_get_account: MagicMock,
|
||||
mock_revoke_token: MagicMock,
|
||||
mock_get_token_data: MagicMock,
|
||||
mock_create_account: MagicMock,
|
||||
mock_login_db: MagicMock,
|
||||
mock_db: MagicMock,
|
||||
app: Flask,
|
||||
):
|
||||
"""
|
||||
Test email-code login failure when creating the account would exceed the licensed seats.
|
||||
|
||||
Verifies that:
|
||||
- the new-account path is taken when no account exists for the email
|
||||
- the service-layer SeatsLimitExceededError is translated to the SeatsLimitExceeded HTTP error
|
||||
"""
|
||||
# Arrange: valid token, no existing account -> account-creation path
|
||||
mock_get_token_data.return_value = {"email": "User@Example.com", "code": "123456"}
|
||||
mock_get_account.return_value = None
|
||||
mock_create_account.side_effect = SeatsLimitExceededError("licensed seats limit exceeded")
|
||||
|
||||
# Act & Assert
|
||||
with app.test_request_context(
|
||||
"/email-code-login/validity",
|
||||
method="POST",
|
||||
json={"email": "User@Example.com", "code": encode_code("123456"), "token": "token-123"},
|
||||
):
|
||||
with pytest.raises(SeatsLimitExceeded):
|
||||
EmailCodeLoginApi().post()
|
||||
|
||||
mock_create_account.assert_called_once()
|
||||
|
||||
|
||||
class TestLogoutApi:
|
||||
"""Test cases for the LogoutApi endpoint."""
|
||||
|
||||
@@ -0,0 +1,93 @@
|
||||
import urllib.parse
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from flask import Flask
|
||||
|
||||
from controllers.console.auth.oauth import OAuthCallback, OAuthLogin
|
||||
from libs.oauth import OAuthUserInfo, encode_oauth_state
|
||||
from models.account import AccountStatus
|
||||
|
||||
REDIRECT_URL = "/apps?category=workflow"
|
||||
CONSOLE_WEB_URL = "https://console.example.com"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def app() -> Flask:
|
||||
app = Flask(__name__)
|
||||
app.config["TESTING"] = True
|
||||
return app
|
||||
|
||||
|
||||
def test_oauth_login_passes_relative_redirect_url_through(app: Flask) -> None:
|
||||
oauth_provider = MagicMock()
|
||||
oauth_provider.get_authorization_url.return_value = "https://accounts.google.com/o/oauth2/v2/auth?state=..."
|
||||
query = urllib.parse.urlencode({"redirect_url": REDIRECT_URL})
|
||||
|
||||
with (
|
||||
patch("controllers.console.auth.oauth.get_oauth_providers", return_value={"google": oauth_provider}),
|
||||
app.test_request_context(f"/oauth/login/google?{query}"),
|
||||
):
|
||||
response = OAuthLogin().get("google")
|
||||
|
||||
oauth_provider.get_authorization_url.assert_called_once_with(
|
||||
invite_token=None,
|
||||
timezone=None,
|
||||
language=None,
|
||||
redirect_url=REDIRECT_URL,
|
||||
)
|
||||
assert response.status_code == 302
|
||||
assert response.headers["Location"] == "https://accounts.google.com/o/oauth2/v2/auth?state=..."
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("redirect_url", "expected_target_url"),
|
||||
[
|
||||
(REDIRECT_URL, REDIRECT_URL),
|
||||
(f"{CONSOLE_WEB_URL}{REDIRECT_URL}", f"{CONSOLE_WEB_URL}{REDIRECT_URL}"),
|
||||
("https://console.example.com.malicious.example/apps", CONSOLE_WEB_URL),
|
||||
("//malicious.example.com/apps", CONSOLE_WEB_URL),
|
||||
("///malicious.example.com/apps", CONSOLE_WEB_URL),
|
||||
(r"\\malicious.example.com/apps", CONSOLE_WEB_URL),
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize("oauth_new_user", [False, True])
|
||||
def test_oauth_callback_validates_redirect_url_and_appends_new_user_flag(
|
||||
app: Flask,
|
||||
redirect_url: str,
|
||||
expected_target_url: str,
|
||||
oauth_new_user: bool,
|
||||
) -> None:
|
||||
oauth_provider = MagicMock()
|
||||
oauth_provider.get_access_token.return_value = "google-access-token"
|
||||
oauth_provider.get_user_info.return_value = OAuthUserInfo(
|
||||
id="google-user-123",
|
||||
name="Test User",
|
||||
email="test@example.com",
|
||||
)
|
||||
account = MagicMock()
|
||||
account.status = AccountStatus.ACTIVE
|
||||
token_pair = MagicMock()
|
||||
token_pair.access_token = "dify-access-token"
|
||||
token_pair.refresh_token = "dify-refresh-token"
|
||||
token_pair.csrf_token = "dify-csrf-token"
|
||||
state = encode_oauth_state(redirect_url=redirect_url)
|
||||
|
||||
with (
|
||||
patch("controllers.console.auth.oauth.get_oauth_providers", return_value={"google": oauth_provider}),
|
||||
patch("controllers.console.auth.oauth.dify_config.CONSOLE_WEB_URL", CONSOLE_WEB_URL),
|
||||
patch("controllers.console.auth.oauth._generate_account", return_value=(account, oauth_new_user)),
|
||||
patch("controllers.console.auth.oauth.TenantService.create_owner_tenant_if_not_exist"),
|
||||
patch("controllers.console.auth.oauth.AccountService.login", return_value=token_pair),
|
||||
patch("controllers.console.auth.oauth.set_access_token_to_cookie"),
|
||||
patch("controllers.console.auth.oauth.set_refresh_token_to_cookie"),
|
||||
patch("controllers.console.auth.oauth.set_csrf_token_to_cookie"),
|
||||
app.test_request_context(f"/oauth/authorize/google?code=test-code&state={state}"),
|
||||
):
|
||||
response = OAuthCallback().get("google")
|
||||
|
||||
assert response.status_code == 302
|
||||
query_char = "&" if "?" in expected_target_url else "?"
|
||||
assert response.headers["Location"] == (
|
||||
f"{expected_target_url}{query_char}oauth_new_user={str(oauth_new_user).lower()}"
|
||||
)
|
||||
@@ -33,6 +33,7 @@ def test_oauth_login_passes_language_and_timezone_to_authorization_url(
|
||||
invite_token=None,
|
||||
timezone="Asia/Shanghai",
|
||||
language="zh-Hans",
|
||||
redirect_url=None,
|
||||
)
|
||||
mock_redirect.assert_called_once_with("https://github.com/login/oauth/authorize?state=...")
|
||||
|
||||
|
||||
@@ -136,7 +136,7 @@ class TestPydanticModels:
|
||||
|
||||
def test_resource_access_scope_defaults_empty_account_ids(self):
|
||||
parsed = rbac_mod._ResourceAccessScopeRequest.model_validate({"scope": "specific"})
|
||||
assert parsed.scope is rbac_mod._AccessScope.SPECIFIC
|
||||
assert parsed.scope is rbac_mod.RBACResourceWhitelistScope.SPECIFIC
|
||||
|
||||
def test_resource_access_scope_coerce_null_account_ids(self):
|
||||
rbac_mod._ResourceAccessScopeRequest.model_validate({"scope": "all"})
|
||||
@@ -262,6 +262,25 @@ class TestPaginationMapping:
|
||||
|
||||
|
||||
class TestResourceAccessScopeBindings:
|
||||
def test_app_whitelist_all_schedules_member_policy_sync(self, app):
|
||||
with (
|
||||
app.test_request_context(
|
||||
"/workspaces/current/rbac/apps/app-1/whitelist",
|
||||
method="PUT",
|
||||
json={"scope": "all"},
|
||||
),
|
||||
patch("controllers.console.workspace.rbac._current_ids", return_value=("tenant-1", "acct-actor")),
|
||||
patch("controllers.console.workspace.rbac.dify_config.RBAC_ENABLED", True),
|
||||
patch(
|
||||
"controllers.console.workspace.rbac.svc.RBACService.AppAccess.replace_whitelist",
|
||||
return_value=rbac_mod.svc.ResourceWhitelist(),
|
||||
),
|
||||
patch("controllers.console.workspace.rbac.initialize_created_app_rbac_access_task") as mock_sync_task,
|
||||
):
|
||||
inspect.unwrap(rbac_mod.RBACAppWhitelistApi.put)(rbac_mod.RBACAppWhitelistApi(), "app-1")
|
||||
|
||||
mock_sync_task.delay.assert_called_once_with("tenant-1", "acct-actor", "app-1")
|
||||
|
||||
def test_app_user_access_policy_assignment_forwards_ids(self, app):
|
||||
with (
|
||||
app.test_request_context(
|
||||
|
||||
@@ -3,7 +3,7 @@ Unit tests for Service API wraps (authentication decorators)
|
||||
"""
|
||||
|
||||
import uuid
|
||||
from unittest.mock import Mock, patch
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
from flask import Flask
|
||||
@@ -469,7 +469,10 @@ 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")
|
||||
def test_rejects_over_rate_limit(self, mock_db, mock_get_rate_limit, mock_validate_token, app: Flask):
|
||||
@patch("controllers.service_api.wraps.sessionmaker")
|
||||
def test_rejects_over_rate_limit(
|
||||
self, mock_sessionmaker, mock_db, mock_get_rate_limit, mock_validate_token, app: Flask
|
||||
):
|
||||
"""Test that Forbidden is raised when over rate limit."""
|
||||
# Arrange
|
||||
mock_validate_token.return_value = Mock(tenant_id="tenant123")
|
||||
@@ -479,6 +482,10 @@ class TestCloudEditionBillingRateLimitCheck:
|
||||
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:
|
||||
mock_redis.zcard.return_value = 15 # Over limit
|
||||
@@ -492,6 +499,9 @@ 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()
|
||||
|
||||
|
||||
class TestValidateDatasetToken:
|
||||
|
||||
@@ -14,7 +14,8 @@ from core.app.layers.pause_state_persist_layer import (
|
||||
_WorkflowGenerateEntityWrapper,
|
||||
)
|
||||
from core.workflow.system_variables import SystemVariableKey
|
||||
from graphon.entities.pause_reason import SchedulingPause
|
||||
from graphon.entities.pause_reason import HitlRequired, SchedulingPause
|
||||
from graphon.filters import GraphEventFilterContext, ResponseStreamFilter
|
||||
from graphon.graph_engine.entities.commands import GraphEngineCommand
|
||||
from graphon.graph_engine.layers.base import GraphEngineLayerNotInitializedError
|
||||
from graphon.graph_events import (
|
||||
@@ -29,6 +30,22 @@ from models.model import AppMode
|
||||
from repositories.factory import DifyAPIRepositoryFactory
|
||||
|
||||
|
||||
def _create_initialized_response_stream_filter() -> ResponseStreamFilter:
|
||||
"""Build a `ResponseStreamFilter` that has already run `initialize()`.
|
||||
|
||||
`ResponseStreamFilter.dumps()` raises `RuntimeError` unless the filter has
|
||||
processed a `GraphEventFilterContext` first. In production this always
|
||||
happens before any event (including `GraphRunPausedEvent`) reaches
|
||||
`PauseStatePersistenceLayer.on_event`, so tests that exercise `on_event`
|
||||
or a subsequent `dumps()` call need a filter in that same state. A
|
||||
nodeless graph is enough to satisfy the precondition.
|
||||
"""
|
||||
response_stream_filter = ResponseStreamFilter()
|
||||
context = GraphEventFilterContext(graph=Mock(nodes={}), runtime_state=Mock())
|
||||
response_stream_filter.initialize(context)
|
||||
return response_stream_filter
|
||||
|
||||
|
||||
class TestDataFactory:
|
||||
"""Factory helpers for constructing graph events used in tests."""
|
||||
|
||||
@@ -200,6 +217,7 @@ class TestPauseStatePersistenceLayer:
|
||||
session_factory=session_factory,
|
||||
state_owner_user_id=state_owner_user_id,
|
||||
generate_entity=self._create_generate_entity(),
|
||||
response_stream_filter=ResponseStreamFilter(),
|
||||
)
|
||||
|
||||
assert layer._session_maker is session_factory
|
||||
@@ -214,6 +232,7 @@ class TestPauseStatePersistenceLayer:
|
||||
session_factory=session_factory,
|
||||
state_owner_user_id="owner",
|
||||
generate_entity=self._create_generate_entity(),
|
||||
response_stream_filter=ResponseStreamFilter(),
|
||||
)
|
||||
|
||||
graph_runtime_state = MockReadOnlyGraphRuntimeState()
|
||||
@@ -231,6 +250,7 @@ class TestPauseStatePersistenceLayer:
|
||||
session_factory=session_factory,
|
||||
state_owner_user_id="owner-123",
|
||||
generate_entity=generate_entity,
|
||||
response_stream_filter=_create_initialized_response_stream_filter(),
|
||||
)
|
||||
|
||||
mock_repo = Mock()
|
||||
@@ -263,12 +283,71 @@ class TestPauseStatePersistenceLayer:
|
||||
|
||||
assert isinstance(pause_reasons, list)
|
||||
|
||||
def test_on_event_enriches_hitl_pause_reasons_before_persisting(self, monkeypatch: pytest.MonkeyPatch):
|
||||
session_factory = Mock(name="session_factory")
|
||||
generate_entity = self._create_generate_entity(workflow_execution_id="run-123")
|
||||
layer = PauseStatePersistenceLayer(
|
||||
session_factory=session_factory,
|
||||
state_owner_user_id="owner-123",
|
||||
generate_entity=generate_entity,
|
||||
response_stream_filter=_create_initialized_response_stream_filter(),
|
||||
)
|
||||
|
||||
mock_repo = Mock()
|
||||
mock_factory = Mock(return_value=mock_repo)
|
||||
mock_form_repository = Mock(name="form_repository")
|
||||
enriched_reason = HumanInputRequired(
|
||||
form_id="form-123",
|
||||
form_content="Rendered content",
|
||||
inputs=[],
|
||||
actions=[],
|
||||
node_id="node-123",
|
||||
node_title="Ask for approval",
|
||||
)
|
||||
enrich_mock = Mock(return_value=[enriched_reason])
|
||||
monkeypatch.setattr(DifyAPIRepositoryFactory, "create_api_workflow_run_repository", mock_factory)
|
||||
monkeypatch.setattr(
|
||||
pause_layer_module,
|
||||
"HumanInputFormSubmissionRepository",
|
||||
Mock(return_value=mock_form_repository),
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
pause_layer_module,
|
||||
"enrich_graph_pause_reasons",
|
||||
enrich_mock,
|
||||
raising=False,
|
||||
)
|
||||
|
||||
graph_runtime_state = MockReadOnlyGraphRuntimeState(
|
||||
workflow_execution_id="run-123",
|
||||
)
|
||||
command_channel = MockCommandChannel()
|
||||
layer.initialize(graph_runtime_state, command_channel)
|
||||
|
||||
raw_reason = HitlRequired(
|
||||
session_id="session-123",
|
||||
node_id="node-123",
|
||||
node_title="Ask for approval",
|
||||
)
|
||||
event = GraphRunPausedEvent(reasons=[raw_reason], outputs={})
|
||||
|
||||
layer.on_event(event)
|
||||
|
||||
enrich_mock.assert_called_once_with(
|
||||
reasons=[raw_reason],
|
||||
form_repository=mock_form_repository,
|
||||
variable_pool=graph_runtime_state.variable_pool,
|
||||
)
|
||||
assert mock_repo.create_workflow_pause.call_args.kwargs["pause_reasons"] == [enriched_reason]
|
||||
|
||||
def test_on_event_ignores_non_paused_events(self, monkeypatch: pytest.MonkeyPatch):
|
||||
session_factory = Mock(name="session_factory")
|
||||
layer = PauseStatePersistenceLayer(
|
||||
session_factory=session_factory,
|
||||
state_owner_user_id="owner-123",
|
||||
generate_entity=self._create_generate_entity(),
|
||||
response_stream_filter=ResponseStreamFilter(),
|
||||
)
|
||||
|
||||
mock_repo = Mock()
|
||||
@@ -297,6 +376,7 @@ class TestPauseStatePersistenceLayer:
|
||||
session_factory=session_factory,
|
||||
state_owner_user_id="owner-123",
|
||||
generate_entity=self._create_generate_entity(),
|
||||
response_stream_filter=ResponseStreamFilter(),
|
||||
)
|
||||
|
||||
event = TestDataFactory.create_graph_run_paused_event()
|
||||
@@ -310,6 +390,7 @@ class TestPauseStatePersistenceLayer:
|
||||
session_factory=session_factory,
|
||||
state_owner_user_id="owner-123",
|
||||
generate_entity=self._create_generate_entity(),
|
||||
response_stream_filter=_create_initialized_response_stream_filter(),
|
||||
)
|
||||
|
||||
mock_repo = Mock()
|
||||
@@ -409,3 +490,53 @@ def test_workflow_resumption_context_dumps_loads_roundtrip(state: WorkflowResump
|
||||
restored_entity = loaded.get_generate_entity()
|
||||
assert isinstance(restored_entity, type(state.generate_entity.entity))
|
||||
assert restored_entity.extras["trace_session_id"] == "session-1"
|
||||
|
||||
|
||||
def test_on_event_persists_response_stream_filter_dump(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
session_factory = Mock(name="session_factory")
|
||||
generate_entity = TestPauseStatePersistenceLayer._create_generate_entity(workflow_execution_id="run-123")
|
||||
response_stream_filter = _create_initialized_response_stream_filter()
|
||||
layer = PauseStatePersistenceLayer(
|
||||
session_factory=session_factory,
|
||||
state_owner_user_id="owner-123",
|
||||
generate_entity=generate_entity,
|
||||
response_stream_filter=response_stream_filter,
|
||||
)
|
||||
|
||||
mock_repo = Mock()
|
||||
mock_factory = Mock(return_value=mock_repo)
|
||||
monkeypatch.setattr(DifyAPIRepositoryFactory, "create_api_workflow_run_repository", mock_factory)
|
||||
|
||||
graph_runtime_state = MockReadOnlyGraphRuntimeState(workflow_execution_id="run-123")
|
||||
layer.initialize(graph_runtime_state, MockCommandChannel())
|
||||
|
||||
event = TestDataFactory.create_graph_run_paused_event()
|
||||
layer.on_event(event)
|
||||
|
||||
serialized_state = mock_repo.create_workflow_pause.call_args.kwargs["state"]
|
||||
resumption_context = WorkflowResumptionContext.loads(serialized_state)
|
||||
assert resumption_context.serialized_response_stream_filter_state == response_stream_filter.dumps()
|
||||
|
||||
|
||||
def test_get_response_stream_filter_restores_dumped_state() -> None:
|
||||
original = _create_initialized_response_stream_filter()
|
||||
context = WorkflowResumptionContext(
|
||||
serialized_graph_runtime_state=json.dumps({"state": "workflow"}),
|
||||
generate_entity=_WorkflowGenerateEntityWrapper(entity=TestPauseStatePersistenceLayer._create_generate_entity()),
|
||||
serialized_response_stream_filter_state=original.dumps(),
|
||||
)
|
||||
|
||||
restored = context.get_response_stream_filter()
|
||||
|
||||
assert restored.dumps() == original.dumps()
|
||||
|
||||
|
||||
def test_get_response_stream_filter_defaults_when_state_missing() -> None:
|
||||
context = WorkflowResumptionContext(
|
||||
serialized_graph_runtime_state=json.dumps({"state": "workflow"}),
|
||||
generate_entity=_WorkflowGenerateEntityWrapper(entity=TestPauseStatePersistenceLayer._create_generate_entity()),
|
||||
)
|
||||
|
||||
restored = context.get_response_stream_filter()
|
||||
|
||||
assert isinstance(restored, ResponseStreamFilter)
|
||||
|
||||
@@ -32,8 +32,16 @@ def mock_print_text(mocker: MockerFixture):
|
||||
return mocker.patch("core.callback_handler.workflow_tool_callback_handler.print_text")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def enable_debug(mocker: MockerFixture):
|
||||
"""Force DEBUG on so the handler emits its verbose stdout traces."""
|
||||
mocker.patch("core.callback_handler.workflow_tool_callback_handler.dify_config.DEBUG", True)
|
||||
|
||||
|
||||
class TestDifyWorkflowCallbackHandler:
|
||||
def test_on_tool_execution_single_output_success(self, handler: DifyWorkflowCallbackHandler, mock_print_text):
|
||||
def test_on_tool_execution_single_output_success(
|
||||
self, handler: DifyWorkflowCallbackHandler, mock_print_text, enable_debug
|
||||
):
|
||||
# Arrange
|
||||
tool_name = "test_tool"
|
||||
tool_inputs = {"a": 1}
|
||||
@@ -63,7 +71,9 @@ class TestDifyWorkflowCallbackHandler:
|
||||
]
|
||||
)
|
||||
|
||||
def test_on_tool_execution_multiple_outputs(self, handler: DifyWorkflowCallbackHandler, mock_print_text):
|
||||
def test_on_tool_execution_multiple_outputs(
|
||||
self, handler: DifyWorkflowCallbackHandler, mock_print_text, enable_debug
|
||||
):
|
||||
# Arrange
|
||||
tool_name = "multi_tool"
|
||||
outputs = [
|
||||
@@ -101,6 +111,29 @@ class TestDifyWorkflowCallbackHandler:
|
||||
assert results == []
|
||||
mock_print_text.assert_not_called()
|
||||
|
||||
def test_on_tool_execution_skips_print_when_debug_disabled(
|
||||
self, handler: DifyWorkflowCallbackHandler, mock_print_text, mocker: MockerFixture
|
||||
):
|
||||
"""When DEBUG is off, outputs are still yielded but nothing is printed
|
||||
and model_dump_json() is never invoked."""
|
||||
# Arrange
|
||||
mocker.patch("core.callback_handler.workflow_tool_callback_handler.dify_config.DEBUG", False)
|
||||
message = MagicMock()
|
||||
|
||||
# Act
|
||||
results = list(
|
||||
handler.on_tool_execution(
|
||||
tool_name="quiet_tool",
|
||||
tool_inputs={},
|
||||
tool_outputs=[message],
|
||||
)
|
||||
)
|
||||
|
||||
# Assert
|
||||
assert results == [message]
|
||||
mock_print_text.assert_not_called()
|
||||
message.model_dump_json.assert_not_called()
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("invalid_outputs", "expected_exception"),
|
||||
[
|
||||
@@ -110,7 +143,7 @@ class TestDifyWorkflowCallbackHandler:
|
||||
],
|
||||
)
|
||||
def test_on_tool_execution_invalid_outputs_type(
|
||||
self, handler: DifyWorkflowCallbackHandler, invalid_outputs, expected_exception
|
||||
self, handler: DifyWorkflowCallbackHandler, invalid_outputs, expected_exception, enable_debug
|
||||
):
|
||||
# Arrange
|
||||
tool_name = "invalid_tool"
|
||||
@@ -125,7 +158,9 @@ class TestDifyWorkflowCallbackHandler:
|
||||
)
|
||||
)
|
||||
|
||||
def test_on_tool_execution_long_json_truncation(self, handler: DifyWorkflowCallbackHandler, mock_print_text):
|
||||
def test_on_tool_execution_long_json_truncation(
|
||||
self, handler: DifyWorkflowCallbackHandler, mock_print_text, enable_debug
|
||||
):
|
||||
# Arrange
|
||||
tool_name = "long_json_tool"
|
||||
long_json = "x" * 1500
|
||||
@@ -147,7 +182,9 @@ class TestDifyWorkflowCallbackHandler:
|
||||
color="blue",
|
||||
)
|
||||
|
||||
def test_on_tool_execution_model_dump_json_exception(self, handler: DifyWorkflowCallbackHandler, mock_print_text):
|
||||
def test_on_tool_execution_model_dump_json_exception(
|
||||
self, handler: DifyWorkflowCallbackHandler, mock_print_text, enable_debug
|
||||
):
|
||||
# Arrange
|
||||
tool_name = "exception_tool"
|
||||
bad_message = MagicMock()
|
||||
@@ -167,7 +204,7 @@ class TestDifyWorkflowCallbackHandler:
|
||||
assert mock_print_text.call_count >= 2
|
||||
|
||||
def test_on_tool_execution_none_message_id_and_trace_manager(
|
||||
self, handler: DifyWorkflowCallbackHandler, mock_print_text
|
||||
self, handler: DifyWorkflowCallbackHandler, mock_print_text, enable_debug
|
||||
):
|
||||
# Arrange
|
||||
tool_name = "optional_params_tool"
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
import json
|
||||
from base64 import b64decode
|
||||
from collections.abc import Mapping
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
@@ -44,7 +43,7 @@ def test_serialize_inputs_encodes_payload() -> None:
|
||||
def test_transform_response_parses_json_result_and_converts_scientific_notation() -> None:
|
||||
response = '<<RESULT>>{"value": "1e+3", "nested": {"x": "2E-2"}, "arr": ["3e+1"]}<<RESULT>>'
|
||||
|
||||
result: Mapping[str, Any] = _DummyTransformer.transform_response(response)
|
||||
result: dict[str, Any] = _DummyTransformer.transform_response(response)
|
||||
|
||||
assert result == {"value": 1000.0, "nested": {"x": 0.02}, "arr": [30.0]}
|
||||
|
||||
|
||||
@@ -46,6 +46,18 @@ class TestUploadDSL:
|
||||
with pytest.raises(ValueError, match="claim_code"):
|
||||
upload_dsl(b"app: demo")
|
||||
|
||||
@patch("core.helper.creators.httpx.post")
|
||||
def test_raises_on_non_string_claim_code(self, mock_post):
|
||||
mock_response = MagicMock(spec=httpx.Response)
|
||||
mock_response.json.return_value = {"data": {"claim_code": 123}}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
from core.helper.creators import upload_dsl
|
||||
|
||||
with pytest.raises(ValueError, match="claim_code"):
|
||||
upload_dsl(b"app: demo")
|
||||
|
||||
@patch("core.helper.creators.httpx.post")
|
||||
def test_raises_on_http_error(self, mock_post):
|
||||
mock_response = MagicMock(spec=httpx.Response)
|
||||
|
||||
@@ -114,10 +114,11 @@ def test_is_credential_exists_by_type(
|
||||
scalar_result: str | None,
|
||||
expected: bool,
|
||||
) -> None:
|
||||
mocker.patch("extensions.ext_database.db", new=SimpleNamespace(engine=object()))
|
||||
session_cls = mocker.patch("sqlalchemy.orm.Session")
|
||||
session = session_cls.return_value.__enter__.return_value
|
||||
session = mocker.MagicMock()
|
||||
session.scalar.return_value = scalar_result
|
||||
session_context = mocker.MagicMock()
|
||||
session_context.__enter__.return_value = session
|
||||
mocker.patch("core.db.session_factory.create_session", return_value=session_context)
|
||||
|
||||
result = is_credential_exists("cred-1", credential_type)
|
||||
|
||||
@@ -128,11 +129,33 @@ def test_is_credential_exists_by_type(
|
||||
def test_is_credential_exists_returns_false_for_unknown_type(
|
||||
mocker: MockerFixture,
|
||||
) -> None:
|
||||
mocker.patch("extensions.ext_database.db", new=SimpleNamespace(engine=object()))
|
||||
session_cls = mocker.patch("sqlalchemy.orm.Session")
|
||||
session = session_cls.return_value.__enter__.return_value
|
||||
session = mocker.MagicMock()
|
||||
session_context = mocker.MagicMock()
|
||||
session_context.__enter__.return_value = session
|
||||
mocker.patch("core.db.session_factory.create_session", return_value=session_context)
|
||||
|
||||
result = is_credential_exists("cred-1", cast(PluginCredentialType, "unknown"))
|
||||
|
||||
assert result is False
|
||||
session.scalar.assert_not_called()
|
||||
|
||||
|
||||
def test_is_credential_exists_uses_configured_session_factory_without_flask_app_context(
|
||||
mocker: MockerFixture,
|
||||
) -> None:
|
||||
class RaisingDB:
|
||||
@property
|
||||
def engine(self):
|
||||
raise RuntimeError("Working outside of application context.")
|
||||
|
||||
session = mocker.MagicMock()
|
||||
session.scalar.return_value = "model-credential"
|
||||
session_context = mocker.MagicMock()
|
||||
session_context.__enter__.return_value = session
|
||||
create_session = mocker.patch("core.db.session_factory.create_session", return_value=session_context)
|
||||
mocker.patch("extensions.ext_database.db", new=RaisingDB())
|
||||
|
||||
result = is_credential_exists("cred-1", PluginCredentialType.MODEL)
|
||||
|
||||
assert result is True
|
||||
create_session.assert_called_once_with()
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from pytest_mock import MockerFixture
|
||||
|
||||
from core.helper.marketplace import (
|
||||
@@ -53,6 +54,16 @@ def test_batch_fetch_plugin_by_ids_returns_plugins_from_response(mocker: MockerF
|
||||
response.raise_for_status.assert_called_once()
|
||||
|
||||
|
||||
def test_batch_fetch_plugin_by_ids_rejects_invalid_plugins_response(mocker: MockerFixture) -> None:
|
||||
response = MagicMock()
|
||||
response.json.return_value = {"data": {"plugins": ["p1"]}}
|
||||
response.raise_for_status.return_value = None
|
||||
mocker.patch("core.helper.marketplace.httpx.post", return_value=response)
|
||||
|
||||
with pytest.raises(ValueError, match="plugins list"):
|
||||
batch_fetch_plugin_by_ids(["p1"])
|
||||
|
||||
|
||||
def test_batch_fetch_plugin_manifests_returns_empty_for_empty_input(mocker: MockerFixture) -> None:
|
||||
post_mock = mocker.patch("core.helper.marketplace.httpx.post")
|
||||
|
||||
|
||||
@@ -346,14 +346,32 @@ class TestPluginAppBackwardsInvocation:
|
||||
assert "end_users.app_id" in compiled
|
||||
assert stmt.compile().params == {"id_1": "uid", "tenant_id_1": "tenant-1", "app_id_1": "app-1"}
|
||||
|
||||
def test_get_user_returns_end_user_by_session_id(self, mocker: MockerFixture):
|
||||
session = self.patch_create_session(mocker, side_effect=[None, MagicMock(id="session-user")])
|
||||
app = SimpleNamespace(id="app-1", tenant_id="tenant-1")
|
||||
|
||||
user = PluginAppBackwardsInvocation._get_user("wecom-sender-1", app)
|
||||
|
||||
assert user.id == "session-user"
|
||||
stmt = session.scalar.call_args_list[1].args[0]
|
||||
compiled = str(stmt.compile(dialect=postgresql.dialect()))
|
||||
assert "end_users.session_id" in compiled
|
||||
assert "end_users.tenant_id" in compiled
|
||||
assert "end_users.app_id" in compiled
|
||||
assert stmt.compile().params == {
|
||||
"session_id_1": "wecom-sender-1",
|
||||
"tenant_id_1": "tenant-1",
|
||||
"app_id_1": "app-1",
|
||||
}
|
||||
|
||||
def test_get_user_falls_back_to_account_user(self, mocker: MockerFixture):
|
||||
session = self.patch_create_session(mocker, side_effect=[None, MagicMock(id="account-user")])
|
||||
session = self.patch_create_session(mocker, side_effect=[None, None, MagicMock(id="account-user")])
|
||||
app = SimpleNamespace(id="app-1", tenant_id="tenant-1")
|
||||
|
||||
user = PluginAppBackwardsInvocation._get_user("uid", app)
|
||||
|
||||
assert user.id == "account-user"
|
||||
stmt = session.scalar.call_args_list[1].args[0]
|
||||
stmt = session.scalar.call_args_list[2].args[0]
|
||||
compiled = str(stmt.compile(dialect=postgresql.dialect()))
|
||||
assert "accounts.id" in compiled
|
||||
assert "tenant_account_joins.account_id" in compiled
|
||||
@@ -361,12 +379,41 @@ class TestPluginAppBackwardsInvocation:
|
||||
assert stmt.compile().params == {"id_1": "uid", "tenant_id_1": "tenant-1"}
|
||||
|
||||
def test_get_user_raises_when_user_not_found(self, mocker: MockerFixture):
|
||||
self.patch_create_session(mocker, side_effect=[None, None])
|
||||
self.patch_create_session(mocker, side_effect=[None, None, None])
|
||||
app = SimpleNamespace(id="app-1", tenant_id="tenant-1")
|
||||
|
||||
with pytest.raises(ValueError, match="user not found"):
|
||||
PluginAppBackwardsInvocation._get_user("uid", app)
|
||||
|
||||
def test_invoke_app_creates_end_user_for_unknown_external_user_id(self, mocker: MockerFixture):
|
||||
app = MagicMock(mode=AppMode.WORKFLOW)
|
||||
end_user = MagicMock()
|
||||
workflow = MagicMock()
|
||||
mocker.patch.object(PluginAppBackwardsInvocation, "_get_app", return_value=app)
|
||||
mocker.patch.object(PluginAppBackwardsInvocation, "_get_workflow", return_value=workflow)
|
||||
mocker.patch.object(PluginAppBackwardsInvocation, "_get_user", side_effect=ValueError("user not found"))
|
||||
get_or_create = mocker.patch(
|
||||
"core.plugin.backwards_invocation.app.EndUserService.get_or_create_end_user",
|
||||
return_value=end_user,
|
||||
)
|
||||
route = mocker.patch.object(PluginAppBackwardsInvocation, "invoke_workflow_app", return_value={"ok": True})
|
||||
|
||||
result = PluginAppBackwardsInvocation.invoke_app(
|
||||
MagicMock(),
|
||||
app_id="app",
|
||||
user_id="wecom-sender-1",
|
||||
tenant_id="tenant",
|
||||
conversation_id="",
|
||||
query=None,
|
||||
stream=True,
|
||||
inputs={},
|
||||
files=[],
|
||||
)
|
||||
|
||||
assert result == {"ok": True}
|
||||
get_or_create.assert_called_once_with(app, user_id="wecom-sender-1")
|
||||
assert route.call_args.args[2] is end_user
|
||||
|
||||
def test_get_app_returns_app(self, mocker: MockerFixture):
|
||||
app_obj = MagicMock(id="app")
|
||||
self.patch_create_session(mocker, return_value=app_obj)
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
import datetime
|
||||
import uuid
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import Mock, patch, sentinel
|
||||
from unittest.mock import MagicMock, Mock, patch, sentinel
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -44,12 +44,34 @@ class _FakeRedis:
|
||||
def delete(self, key: str) -> None:
|
||||
self._values.pop(key, None)
|
||||
|
||||
def lock(
|
||||
self,
|
||||
key: str,
|
||||
*,
|
||||
timeout: int,
|
||||
sleep: float,
|
||||
) -> "_FakeRedisLock":
|
||||
return _FakeRedisLock(self, key)
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def clear_plugin_model_provider_memory_cache() -> None:
|
||||
PluginService._plugin_model_providers_memory_cache.clear()
|
||||
yield
|
||||
PluginService._plugin_model_providers_memory_cache.clear()
|
||||
|
||||
class _FakeRedisLock:
|
||||
def __init__(self, redis: _FakeRedis, key: str) -> None:
|
||||
self._redis = redis
|
||||
self._key = key
|
||||
self._acquired = False
|
||||
|
||||
def acquire(self, *, blocking: bool = True, blocking_timeout: float | None = None) -> bool:
|
||||
if self._key in self._redis._values:
|
||||
return False
|
||||
|
||||
self._redis._values[self._key] = "locked"
|
||||
self._acquired = True
|
||||
return True
|
||||
|
||||
def release(self) -> None:
|
||||
if self._acquired:
|
||||
self._redis.delete(self._key)
|
||||
self._acquired = False
|
||||
|
||||
|
||||
def _build_model_schema() -> AIModelEntity:
|
||||
@@ -413,12 +435,13 @@ class TestPluginModelRuntime:
|
||||
"redis_client",
|
||||
SimpleNamespace(
|
||||
get=Mock(return_value=None),
|
||||
mget=Mock(return_value=[None, None]),
|
||||
mget=Mock(return_value=[None]),
|
||||
delete=Mock(),
|
||||
setex=Mock(),
|
||||
lock=Mock(return_value=MagicMock()),
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(plugin_service_module.dify_config, "PLUGIN_MODEL_PROVIDERS_CACHE_TTL", 300)
|
||||
monkeypatch.setattr(plugin_service_module.dify_config, "PLUGIN_MODEL_PROVIDERS_CACHE_TTL", 0)
|
||||
runtime = PluginModelRuntime(tenant_id="tenant", user_id="user", client=client, plugin_service=PluginService)
|
||||
|
||||
runtime.fetch_model_providers()
|
||||
|
||||
@@ -3888,7 +3888,17 @@ class TestDatasetRetrievalAdditionalHelpers:
|
||||
trace_manager.add_trace_task.assert_not_called()
|
||||
|
||||
def test_on_query(self, retrieval: DatasetRetrieval) -> None:
|
||||
with patch("core.rag.retrieval.dataset_retrieval.db.session") as mock_session:
|
||||
db_mock = Mock()
|
||||
audit_session = MagicMock()
|
||||
session_factory = MagicMock()
|
||||
session_factory.begin.return_value.__enter__.return_value = audit_session
|
||||
|
||||
with (
|
||||
patch("core.rag.retrieval.dataset_retrieval.db", db_mock),
|
||||
patch(
|
||||
"core.rag.retrieval.dataset_retrieval.sessionmaker", return_value=session_factory
|
||||
) as sessionmaker_mock,
|
||||
):
|
||||
retrieval._on_query(
|
||||
query=None,
|
||||
attachment_ids=None,
|
||||
@@ -3897,7 +3907,7 @@ class TestDatasetRetrievalAdditionalHelpers:
|
||||
user_from="account",
|
||||
user_id="u1",
|
||||
)
|
||||
mock_session.add_all.assert_not_called()
|
||||
audit_session.add_all.assert_not_called()
|
||||
|
||||
retrieval._on_query(
|
||||
query="python",
|
||||
@@ -3907,11 +3917,22 @@ class TestDatasetRetrievalAdditionalHelpers:
|
||||
user_from="account",
|
||||
user_id="u1",
|
||||
)
|
||||
mock_session.add_all.assert_called()
|
||||
mock_session.commit.assert_called()
|
||||
sessionmaker_mock.assert_called_once_with(bind=db_mock.engine, expire_on_commit=False)
|
||||
audit_session.add_all.assert_called_once()
|
||||
added_queries = audit_session.add_all.call_args.args[0]
|
||||
assert len(added_queries) == 2
|
||||
db_mock.session.commit.assert_not_called()
|
||||
|
||||
def test_on_query_normalizes_workflow_end_user_role(self, retrieval: DatasetRetrieval) -> None:
|
||||
with patch("core.rag.retrieval.dataset_retrieval.db.session") as mock_session:
|
||||
db_mock = Mock()
|
||||
audit_session = MagicMock()
|
||||
session_factory = MagicMock()
|
||||
session_factory.begin.return_value.__enter__.return_value = audit_session
|
||||
|
||||
with (
|
||||
patch("core.rag.retrieval.dataset_retrieval.db", db_mock),
|
||||
patch("core.rag.retrieval.dataset_retrieval.sessionmaker", return_value=session_factory),
|
||||
):
|
||||
retrieval._on_query(
|
||||
query="python",
|
||||
attachment_ids=None,
|
||||
@@ -3921,12 +3942,11 @@ class TestDatasetRetrievalAdditionalHelpers:
|
||||
user_id="u1",
|
||||
)
|
||||
|
||||
mock_session.add_all.assert_called_once()
|
||||
added_queries = mock_session.add_all.call_args.args[0]
|
||||
audit_session.add_all.assert_called_once()
|
||||
added_queries = audit_session.add_all.call_args.args[0]
|
||||
|
||||
assert len(added_queries) == 1
|
||||
assert added_queries[0].created_by_role == CreatorUserRole.END_USER
|
||||
mock_session.commit.assert_called_once()
|
||||
|
||||
def test_handle_invoke_result(self, retrieval: DatasetRetrieval) -> None:
|
||||
usage = LLMUsage.empty_usage()
|
||||
|
||||
@@ -84,6 +84,10 @@ def test_assembling_request_auth_header_assembly():
|
||||
assert headers["Authorization"] == "Bearer abc"
|
||||
|
||||
tool.runtime.credentials = {"auth_type": "api_key_header", "api_key_header_prefix": "basic", "api_key_value": "abc"}
|
||||
headers = tool.assembling_request(parameters={})
|
||||
assert headers["Authorization"] == "Basic abc"
|
||||
assert tool.runtime.credentials["api_key_value"] == "abc"
|
||||
|
||||
headers = tool.assembling_request(parameters={})
|
||||
assert headers["Authorization"] == "Basic abc"
|
||||
|
||||
|
||||
@@ -3,7 +3,7 @@ from __future__ import annotations
|
||||
from collections.abc import Generator
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
from unittest.mock import Mock, patch
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -131,18 +131,25 @@ def test_create_message_files_and_invoke_generator():
|
||||
created.append(obj)
|
||||
return obj
|
||||
|
||||
with patch("core.tools.tool_engine.MessageFile", side_effect=_message_file_factory):
|
||||
with patch("core.tools.tool_engine.db") as mock_db:
|
||||
ids = ToolEngine._create_message_files(
|
||||
tool_messages=binaries,
|
||||
agent_message=SimpleNamespace(id="msg-1"),
|
||||
invoke_from=InvokeFrom.DEBUGGER,
|
||||
user_id="user-1",
|
||||
)
|
||||
file_session = MagicMock()
|
||||
session_factory = MagicMock()
|
||||
session_factory.begin.return_value.__enter__.return_value = file_session
|
||||
with (
|
||||
patch("core.tools.tool_engine.MessageFile", side_effect=_message_file_factory),
|
||||
patch("core.tools.tool_engine.db") as mock_db,
|
||||
patch("core.tools.tool_engine.sessionmaker", return_value=session_factory) as mock_sessionmaker,
|
||||
):
|
||||
ids = ToolEngine._create_message_files(
|
||||
tool_messages=binaries,
|
||||
agent_message=SimpleNamespace(id="msg-1"),
|
||||
invoke_from=InvokeFrom.DEBUGGER,
|
||||
user_id="user-1",
|
||||
)
|
||||
|
||||
assert ids == ["mf-1", "mf-2"]
|
||||
assert mock_db.session.add.call_count == 2
|
||||
mock_db.session.close.assert_called_once()
|
||||
mock_sessionmaker.assert_called_once_with(bind=mock_db.engine, expire_on_commit=False)
|
||||
assert file_session.add.call_count == 2
|
||||
mock_db.session.close.assert_not_called()
|
||||
|
||||
tool = _build_tool()
|
||||
invoked = list(ToolEngine._invoke(tool, {"a": 1}, user_id="u"))
|
||||
|
||||
@@ -327,7 +327,7 @@ def test_parse_openai_plugin_json_branches(app):
|
||||
def test_parse_openai_plugin_json_http_branches(app):
|
||||
with app.test_request_context():
|
||||
response = type("Resp", (), {"status_code": 500, "text": "", "close": Mock()})()
|
||||
with patch("core.tools.utils.parser.httpx.get", return_value=response):
|
||||
with patch("core.tools.utils.parser.ssrf_proxy.get", return_value=response):
|
||||
with pytest.raises(ToolProviderNotFoundError, match="cannot get openapi yaml"):
|
||||
ApiBasedToolSchemaParser.parse_openai_plugin_json_to_tool_bundle(
|
||||
'{"api": {"url": "https://x", "type": "openapi"}}'
|
||||
@@ -335,7 +335,7 @@ def test_parse_openai_plugin_json_http_branches(app):
|
||||
response.close.assert_called_once()
|
||||
|
||||
success_response = type("Resp", (), {"status_code": 200, "text": "openapi: 3.0.0", "close": Mock()})()
|
||||
with patch("core.tools.utils.parser.httpx.get", return_value=success_response):
|
||||
with patch("core.tools.utils.parser.ssrf_proxy.get", return_value=success_response):
|
||||
with patch(
|
||||
"core.tools.utils.parser.ApiBasedToolSchemaParser.parse_openapi_yaml_to_tool_bundle",
|
||||
return_value=["bundle"],
|
||||
|
||||
@@ -13,6 +13,7 @@ from graphon.entities.base_node_data import BaseNodeData
|
||||
from graphon.enums import NodeType, WorkflowNodeExecutionStatus
|
||||
from graphon.errors import WorkflowNodeRunFailedError
|
||||
from graphon.file import File, FileTransferMethod, FileType
|
||||
from graphon.filters import ResponseStreamFilter
|
||||
from graphon.graph import Graph
|
||||
from graphon.graph_events import GraphRunFailedEvent
|
||||
from graphon.model_runtime.entities.llm_entities import LLMMode, LLMUsage
|
||||
@@ -241,6 +242,37 @@ class TestWorkflowChildEngineBuilder:
|
||||
)
|
||||
|
||||
|
||||
def _build_minimal_workflow_entry(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
*,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
) -> workflow_entry.WorkflowEntry:
|
||||
"""Construct a minimal WorkflowEntry with GraphEngine construction mocked out."""
|
||||
graph_engine = MagicMock()
|
||||
graph_runtime_state = SimpleNamespace(execution_context=None)
|
||||
|
||||
monkeypatch.setattr(workflow_entry, "capture_current_context", lambda: sentinel.execution_context)
|
||||
monkeypatch.setattr(workflow_entry, "GraphEngine", MagicMock(return_value=graph_engine))
|
||||
monkeypatch.setattr(workflow_entry, "GraphEngineConfig", MagicMock(return_value=sentinel.graph_engine_config))
|
||||
monkeypatch.setattr(workflow_entry, "InMemoryChannel", MagicMock(return_value=sentinel.command_channel))
|
||||
monkeypatch.setattr(workflow_entry, "LLMQuotaLayer", MagicMock(return_value=sentinel.llm_quota_layer))
|
||||
|
||||
return workflow_entry.WorkflowEntry(
|
||||
tenant_id="tenant-id",
|
||||
app_id="app-id",
|
||||
workflow_id="workflow-id",
|
||||
graph_config={"nodes": [], "edges": []},
|
||||
graph=sentinel.graph,
|
||||
user_id="user-id",
|
||||
user_from=UserFrom.ACCOUNT,
|
||||
invoke_from=InvokeFrom.DEBUGGER,
|
||||
call_depth=0,
|
||||
variable_pool=sentinel.variable_pool,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
response_stream_filter=response_stream_filter,
|
||||
)
|
||||
|
||||
|
||||
class TestWorkflowEntryInit:
|
||||
def test_rejects_call_depth_above_limit(self):
|
||||
call_depth = workflow_entry.dify_config.WORKFLOW_CALL_MAX_DEPTH + 1
|
||||
@@ -329,12 +361,24 @@ class TestWorkflowEntryInit:
|
||||
((observability_layer,), {}),
|
||||
]
|
||||
|
||||
def test_workflow_entry_stores_supplied_response_stream_filter(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
supplied_filter = ResponseStreamFilter()
|
||||
entry = _build_minimal_workflow_entry(monkeypatch, response_stream_filter=supplied_filter)
|
||||
|
||||
assert entry._response_stream_filter is supplied_filter
|
||||
|
||||
def test_workflow_entry_defaults_to_fresh_response_stream_filter(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
entry = _build_minimal_workflow_entry(monkeypatch, response_stream_filter=None)
|
||||
|
||||
assert isinstance(entry._response_stream_filter, ResponseStreamFilter)
|
||||
|
||||
|
||||
class TestWorkflowEntryRun:
|
||||
def test_run_swallows_generate_task_stopped_errors(self):
|
||||
entry = object.__new__(workflow_entry.WorkflowEntry)
|
||||
entry.graph_engine = MagicMock()
|
||||
entry.graph_engine.run.side_effect = GenerateTaskStoppedError()
|
||||
entry._response_stream_filter = ResponseStreamFilter()
|
||||
|
||||
assert list(entry.run()) == []
|
||||
|
||||
@@ -373,6 +417,7 @@ class TestWorkflowEntryRun:
|
||||
def test_run_delegates_to_dify_event_iterator(self):
|
||||
entry = object.__new__(workflow_entry.WorkflowEntry)
|
||||
entry.graph_engine = sentinel.graph_engine
|
||||
entry._response_stream_filter = sentinel.response_stream_filter
|
||||
|
||||
with patch.object(
|
||||
workflow_entry,
|
||||
@@ -382,12 +427,13 @@ class TestWorkflowEntryRun:
|
||||
events = list(entry.run())
|
||||
|
||||
assert events == [sentinel.filtered_event]
|
||||
iter_dify_graph_engine_events.assert_called_once_with(sentinel.graph_engine)
|
||||
iter_dify_graph_engine_events.assert_called_once_with(sentinel.graph_engine, sentinel.response_stream_filter)
|
||||
|
||||
def test_run_emits_failed_event_for_unexpected_errors(self):
|
||||
entry = object.__new__(workflow_entry.WorkflowEntry)
|
||||
entry.graph_engine = MagicMock()
|
||||
entry.graph_engine.run.side_effect = RuntimeError("boom")
|
||||
entry._response_stream_filter = ResponseStreamFilter()
|
||||
|
||||
events = list(entry.run())
|
||||
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user