refactor(api): clarify DSL import and plugin migration boundaries (#38483)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
WH-2099
2026-07-07 02:06:49 +00:00
committed by GitHub
co-authored by autofix-ci[bot]
parent b9c7199d34
commit 0a3426ea38
12 changed files with 329 additions and 133 deletions
+10 -3
View File
@@ -39,6 +39,7 @@ from libs.datetime_utils import naive_utc_now
from models import Account, App, AppMode
from models.model import AppModelConfig, AppModelConfigDict, IconType
from models.workflow import Workflow
from services.dsl_content import DSL_MAX_SIZE, dsl_content_size
from services.dsl_version import check_version_compatibility
from services.entities.dsl_entities import CheckDependenciesResult, ImportMode, ImportStatus
from services.errors.app import WorkflowNotFoundError
@@ -51,7 +52,6 @@ logger = logging.getLogger(__name__)
IMPORT_INFO_REDIS_KEY_PREFIX = "app_import_info:"
CHECK_DEPENDENCIES_REDIS_KEY_PREFIX = "app_check_dependencies:"
IMPORT_INFO_REDIS_EXPIRY = 10 * 60 # 10 minutes
DSL_MAX_SIZE = 10 * 1024 * 1024 # 10MB
CURRENT_DSL_VERSION = CURRENT_APP_DSL_VERSION
@@ -131,15 +131,16 @@ class AppDslService:
yaml_url = yaml_url.replace("/blob/", "/")
response = remote_fetcher.make_request("GET", yaml_url.strip(), follow_redirects=True, timeout=(10, 10))
response.raise_for_status()
content = response.content.decode()
raw_content = response.content
if len(content) > DSL_MAX_SIZE:
if dsl_content_size(raw_content) > DSL_MAX_SIZE:
return Import(
id=import_id,
status=ImportStatus.FAILED,
error="File size exceeds the limit of 10MB",
)
content = raw_content.decode("utf-8")
if not content:
return Import(
id=import_id,
@@ -160,6 +161,12 @@ class AppDslService:
error="yaml_content is required when import_mode is yaml-content",
)
content = yaml_content
if dsl_content_size(content) > DSL_MAX_SIZE:
return Import(
id=import_id,
status=ImportStatus.FAILED,
error="File size exceeds the limit of 10MB",
)
# Process YAML content
try:
+9
View File
@@ -0,0 +1,9 @@
"""Shared DSL content size and decoding rules."""
DSL_MAX_SIZE = 10 * 1024 * 1024 # 10MB
def dsl_content_size(content: str | bytes) -> int:
if isinstance(content, bytes):
return len(content)
return len(content.encode("utf-8"))
+65 -51
View File
@@ -307,9 +307,9 @@ class PluginMigration:
return result
@classmethod
def _fetch_plugin_unique_identifier(cls, plugin_id: str) -> str | None:
def _fetch_latest_package_identifier(cls, plugin_id: str) -> str | None:
"""
Fetch plugin unique identifier using plugin id.
Fetch the latest marketplace package identifier using a plugin id.
"""
if not dify_config.MARKETPLACE_ENABLED:
return None
@@ -328,7 +328,7 @@ class PluginMigration:
@classmethod
def extract_unique_plugins(cls, extracted_plugins: str) -> ExtractedPluginsDict:
plugins: dict[str, str] = {}
package_identifier_by_plugin_id: dict[str, str] = {}
plugin_ids = []
plugin_not_exist = []
logger.info("Extracting unique plugins from %s", extracted_plugins)
@@ -341,19 +341,19 @@ class PluginMigration:
def fetch_plugin(plugin_id):
try:
unique_identifier = cls._fetch_plugin_unique_identifier(plugin_id)
if unique_identifier:
plugins[plugin_id] = unique_identifier
latest_package_identifier = cls._fetch_latest_package_identifier(plugin_id)
if latest_package_identifier:
package_identifier_by_plugin_id[plugin_id] = latest_package_identifier
else:
plugin_not_exist.append(plugin_id)
except Exception:
logger.exception("Failed to fetch plugin unique identifier for %s", plugin_id)
logger.exception("Failed to fetch latest package identifier for %s", plugin_id)
plugin_not_exist.append(plugin_id)
with ThreadPoolExecutor(max_workers=10) as executor:
list(tqdm.tqdm(executor.map(fetch_plugin, plugin_ids), total=len(plugin_ids)))
return {"plugins": plugins, "plugin_not_exist": plugin_not_exist}
return {"plugins": package_identifier_by_plugin_id, "plugin_not_exist": plugin_not_exist}
@classmethod
def install_plugins(cls, extracted_plugins: str, output_file: str, workers: int = 100):
@@ -362,17 +362,22 @@ class PluginMigration:
"""
manager = PluginInstaller()
plugins = cls.extract_unique_plugins(extracted_plugins)
extracted = cls.extract_unique_plugins(extracted_plugins)
package_identifier_by_plugin_id = extracted["plugins"]
not_installed = []
plugin_install_failed = []
# use a fake tenant id to install all the plugins
fake_tenant_id = uuid4().hex
logger.info("Installing %s plugin instances for fake tenant %s", len(plugins["plugins"]), fake_tenant_id)
logger.info(
"Installing %s plugin instances for fake tenant %s",
len(package_identifier_by_plugin_id),
fake_tenant_id,
)
thread_pool = ThreadPoolExecutor(max_workers=workers)
response = cls.handle_plugin_instance_install(fake_tenant_id, plugins["plugins"])
response = cls.handle_plugin_instance_install(fake_tenant_id, package_identifier_by_plugin_id)
if response.get("failed"):
plugin_install_failed.extend(response.get("failed", []))
@@ -384,21 +389,21 @@ class PluginMigration:
# at most 64 plugins one batch
for i in range(0, len(plugin_ids), 64):
batch_plugin_ids = plugin_ids[i : i + 64]
batch_plugin_identifiers = [
plugins["plugins"][plugin_id]
batch_package_identifiers = [
package_identifier_by_plugin_id[plugin_id]
for plugin_id in batch_plugin_ids
if plugin_id not in installed_plugins_ids and plugin_id in plugins["plugins"]
if plugin_id not in installed_plugins_ids and plugin_id in package_identifier_by_plugin_id
]
if batch_plugin_identifiers:
if batch_package_identifiers:
manager.install_from_identifiers(
tenant_id,
batch_plugin_identifiers,
batch_package_identifiers,
PluginInstallationSource.Marketplace,
metas=[
{
"plugin_unique_identifier": identifier,
"plugin_unique_identifier": package_identifier,
}
for identifier in batch_plugin_identifiers
for package_identifier in batch_package_identifiers
],
)
PluginService.invalidate_plugin_model_providers_cache(tenant_id)
@@ -412,10 +417,8 @@ class PluginMigration:
tenant_id = data["tenant_id"]
plugin_ids = data["plugins"]
plugin_not_exist: list[str] = []
# get plugin unique identifier
for plugin_id in plugin_ids:
unique_identifier = plugins.get(plugin_id)
if unique_identifier:
if plugin_id not in package_identifier_by_plugin_id:
plugin_not_exist.append(plugin_id)
if plugin_not_exist:
@@ -459,36 +462,44 @@ class PluginMigration:
"""
manager = PluginInstaller()
plugins = cls.extract_unique_plugins(extracted_plugins)
extracted = cls.extract_unique_plugins(extracted_plugins)
package_identifier_by_plugin_id = extracted["plugins"]
plugin_install_failed = []
# use a fake tenant id to install all the plugins
fake_tenant_id = uuid4().hex
logger.info("Installing %s plugin instances for fake tenant %s", len(plugins["plugins"]), fake_tenant_id)
logger.info(
"Installing %s plugin instances for fake tenant %s",
len(package_identifier_by_plugin_id),
fake_tenant_id,
)
thread_pool = ThreadPoolExecutor(max_workers=workers)
response = cls.handle_plugin_instance_install(fake_tenant_id, plugins["plugins"])
response = cls.handle_plugin_instance_install(fake_tenant_id, package_identifier_by_plugin_id)
if response.get("failed"):
plugin_install_failed.extend(response.get("failed", []))
def install(
tenant_id: str, plugin_ids: dict[str, str], total_success_tenant: int, total_failed_tenant: int
tenant_id: str,
package_identifier_by_plugin_id: dict[str, str],
total_success_tenant: int,
total_failed_tenant: int,
) -> None:
logger.info("Installing %s plugins for tenant %s", len(plugin_ids), tenant_id)
logger.info("Installing %s plugins for tenant %s", len(package_identifier_by_plugin_id), tenant_id)
try:
# fetch plugin already installed
installed_plugins = manager.list_plugins(tenant_id)
installed_plugins_ids = [plugin.plugin_id for plugin in installed_plugins]
# at most 64 plugins one batch
for i in range(0, len(plugin_ids), 64):
batch_plugin_ids = list(plugin_ids.keys())[i : i + 64]
batch_plugin_identifiers = [
plugin_ids[plugin_id]
for i in range(0, len(package_identifier_by_plugin_id), 64):
batch_plugin_ids = list(package_identifier_by_plugin_id.keys())[i : i + 64]
batch_package_identifiers = [
package_identifier_by_plugin_id[plugin_id]
for plugin_id in batch_plugin_ids
if plugin_id not in installed_plugins_ids and plugin_id in plugin_ids
if plugin_id not in installed_plugins_ids and plugin_id in package_identifier_by_plugin_id
]
PluginService.install_from_marketplace_pkg(tenant_id, batch_plugin_identifiers)
PluginService.install_from_marketplace_pkg(tenant_id, batch_package_identifiers)
total_success_tenant += 1
except Exception:
@@ -510,7 +521,7 @@ class PluginMigration:
thread_pool.submit(
install,
tenant_id,
plugins.get("plugins", {}),
package_identifier_by_plugin_id,
total_success_tenant,
total_failed_tenant,
)
@@ -542,12 +553,12 @@ class PluginMigration:
@classmethod
def handle_plugin_instance_install(
cls, tenant_id: str, plugin_identifiers_map: Mapping[str, str]
cls, tenant_id: str, package_identifier_by_plugin_id: Mapping[str, str]
) -> PluginInstallResultDict:
"""
Install plugins for a tenant.
"""
if plugin_identifiers_map and not dify_config.MARKETPLACE_ENABLED:
if package_identifier_by_plugin_id and not dify_config.MARKETPLACE_ENABLED:
raise ValueError(
"Marketplace disabled in offline mode; cannot bulk-install plugins. "
"Pre-upload plugin packages via Console first."
@@ -558,17 +569,17 @@ class PluginMigration:
thread_pool = ThreadPoolExecutor(max_workers=10)
futures = []
for plugin_id, plugin_identifier in plugin_identifiers_map.items():
for plugin_id, package_identifier in package_identifier_by_plugin_id.items():
def download_and_upload(tenant_id, plugin_id, plugin_identifier):
plugin_package = marketplace.download_plugin_pkg(plugin_identifier)
def download_and_upload(tenant_id, plugin_id, package_identifier):
plugin_package = marketplace.download_plugin_pkg(package_identifier)
if not plugin_package:
raise Exception(f"Failed to download plugin {plugin_identifier}")
raise Exception(f"Failed to download plugin {package_identifier}")
# upload
manager.upload_pkg(tenant_id, plugin_package, verify_signature=True)
futures.append(thread_pool.submit(download_and_upload, tenant_id, plugin_id, plugin_identifier))
futures.append(thread_pool.submit(download_and_upload, tenant_id, plugin_id, package_identifier))
# Wait for all downloads to complete
for future in futures:
@@ -578,33 +589,33 @@ class PluginMigration:
success = []
failed = []
reverse_map = {v: k for k, v in plugin_identifiers_map.items()}
plugin_id_by_package_identifier = {v: k for k, v in package_identifier_by_plugin_id.items()}
# at most 8 plugins one batch
for i in range(0, len(plugin_identifiers_map), 8):
batch_plugin_ids = list(plugin_identifiers_map.keys())[i : i + 8]
batch_plugin_identifiers = [plugin_identifiers_map[plugin_id] for plugin_id in batch_plugin_ids]
for i in range(0, len(package_identifier_by_plugin_id), 8):
batch_plugin_ids = list(package_identifier_by_plugin_id.keys())[i : i + 8]
batch_package_identifiers = [package_identifier_by_plugin_id[plugin_id] for plugin_id in batch_plugin_ids]
try:
response = manager.install_from_identifiers(
tenant_id=tenant_id,
identifiers=batch_plugin_identifiers,
identifiers=batch_package_identifiers,
source=PluginInstallationSource.Marketplace,
metas=[
{
"plugin_unique_identifier": identifier,
"plugin_unique_identifier": package_identifier,
}
for identifier in batch_plugin_identifiers
for package_identifier in batch_package_identifiers
],
)
PluginService.invalidate_plugin_model_providers_cache(tenant_id)
except Exception:
# add to failed
failed.extend(batch_plugin_identifiers)
failed.extend(batch_plugin_ids)
continue
if response.all_installed:
success.extend(batch_plugin_identifiers)
success.extend(batch_plugin_ids)
continue
task_id = response.task_id
@@ -614,10 +625,13 @@ class PluginMigration:
if status.status in [PluginInstallTaskStatus.Failed, PluginInstallTaskStatus.Success]:
PluginService.invalidate_plugin_model_providers_cache(tenant_id)
for plugin in status.plugins:
plugin_id = plugin_id_by_package_identifier.get(
plugin.plugin_unique_identifier, plugin.plugin_unique_identifier.split(":", 1)[0]
)
if plugin.status == PluginInstallTaskStatus.Success:
success.append(reverse_map[plugin.plugin_unique_identifier])
success.append(plugin_id)
else:
failed.append(reverse_map[plugin.plugin_unique_identifier])
failed.append(plugin_id)
logger.error(
"Failed to install plugin %s, error: %s",
plugin.plugin_unique_identifier,
@@ -36,6 +36,7 @@ from models import Account
from models.dataset import Dataset, DatasetCollectionBinding, Pipeline
from models.enums import CollectionBindingType, DatasetRuntimeMode
from models.workflow import Workflow, WorkflowType
from services.dsl_content import DSL_MAX_SIZE, dsl_content_size
from services.dsl_version import check_version_compatibility
from services.entities.dsl_entities import CheckDependenciesResult, ImportMode, ImportStatus
from services.entities.knowledge_entities.rag_pipeline_entities import (
@@ -50,7 +51,6 @@ logger = logging.getLogger(__name__)
IMPORT_INFO_REDIS_KEY_PREFIX = "app_import_info:"
CHECK_DEPENDENCIES_REDIS_KEY_PREFIX = "app_check_dependencies:"
IMPORT_INFO_REDIS_EXPIRY = 10 * 60 # 10 minutes
DSL_MAX_SIZE = 10 * 1024 * 1024 # 10MB
CURRENT_DSL_VERSION = "0.1.0"
@@ -127,15 +127,16 @@ class RagPipelineDslService:
yaml_url = yaml_url.replace("/blob/", "/")
response = remote_fetcher.make_request("GET", yaml_url.strip(), follow_redirects=True, timeout=(10, 10))
response.raise_for_status()
content = response.content.decode()
raw_content = response.content
if len(content) > DSL_MAX_SIZE:
if dsl_content_size(raw_content) > DSL_MAX_SIZE:
return RagPipelineImportInfo(
id=import_id,
status=ImportStatus.FAILED,
error="File size exceeds the limit of 10MB",
)
content = raw_content.decode("utf-8")
if not content:
return RagPipelineImportInfo(
id=import_id,
@@ -156,6 +157,12 @@ class RagPipelineDslService:
error="yaml_content is required when import_mode is yaml-content",
)
content = yaml_content
if dsl_content_size(content) > DSL_MAX_SIZE:
return RagPipelineImportInfo(
id=import_id,
status=ImportStatus.FAILED,
error="File size exceeds the limit of 10MB",
)
# Process YAML content
try:
@@ -269,11 +269,13 @@ class RagPipelineTransformService:
installed_plugins_ids = [plugin.plugin_id for plugin in installed_plugins]
dependencies = pipeline_yaml.get("dependencies", [])
need_install_plugin_unique_identifiers = []
package_identifiers_to_install = []
for dependency in dependencies:
if dependency.get("type") == "marketplace":
plugin_unique_identifier = dependency.get("value", {}).get("plugin_unique_identifier")
plugin_id = plugin_unique_identifier.split(":")[0]
package_identifier = dependency.get("value", {}).get("plugin_unique_identifier")
if not package_identifier:
continue
plugin_id = package_identifier.split(":", 1)[0]
if plugin_id not in installed_plugins_ids:
if not dify_config.MARKETPLACE_ENABLED:
logger.warning(
@@ -282,12 +284,12 @@ class RagPipelineTransformService:
plugin_id,
)
continue
plugin_unique_identifier = plugin_migration._fetch_plugin_unique_identifier(plugin_id) # type: ignore
if plugin_unique_identifier:
need_install_plugin_unique_identifiers.append(plugin_unique_identifier)
if need_install_plugin_unique_identifiers:
logger.debug("Installing missing pipeline plugins %s", need_install_plugin_unique_identifiers)
PluginService.install_from_marketplace_pkg(tenant_id, need_install_plugin_unique_identifiers)
latest_package_identifier = plugin_migration._fetch_latest_package_identifier(plugin_id) # type: ignore
if latest_package_identifier:
package_identifiers_to_install.append(latest_package_identifier)
if package_identifiers_to_install:
logger.debug("Installing missing pipeline plugins %s", package_identifiers_to_install)
PluginService.install_from_marketplace_pkg(tenant_id, package_identifiers_to_install)
def _transform_to_empty_pipeline(self, dataset: Dataset, session: Session):
pipeline = Pipeline(
+10 -44
View File
@@ -3,12 +3,10 @@ import logging
import uuid
from collections.abc import Mapping
from datetime import UTC, datetime
from enum import StrEnum
from urllib.parse import urlparse
import yaml
from packaging import version
from pydantic import BaseModel, Field
from pydantic import BaseModel
from sqlalchemy import select
from sqlalchemy.orm import Session
@@ -20,6 +18,9 @@ from graphon.model_runtime.utils.encoders import jsonable_encoder
from models import Account
from models.snippet import CustomizedSnippet, SnippetType
from models.workflow import Workflow
from services.dsl_content import DSL_MAX_SIZE, dsl_content_size
from services.dsl_version import check_version_compatibility
from services.entities.dsl_entities import CheckDependenciesResult, ImportMode, ImportStatus
from services.plugin.dependencies_analysis import DependenciesAnalysisService
from services.snippet_service import SNIPPET_FORBIDDEN_NODE_TYPES, SnippetService
@@ -28,22 +29,9 @@ logger = logging.getLogger(__name__)
IMPORT_INFO_REDIS_KEY_PREFIX = "snippet_import_info:"
CHECK_DEPENDENCIES_REDIS_KEY_PREFIX = "snippet_check_dependencies:"
IMPORT_INFO_REDIS_EXPIRY = 10 * 60 # 10 minutes
DSL_MAX_SIZE = 10 * 1024 * 1024 # 10MB
CURRENT_DSL_VERSION = "0.1.0"
class ImportMode(StrEnum):
YAML_CONTENT = "yaml-content"
YAML_URL = "yaml-url"
class ImportStatus(StrEnum):
COMPLETED = "completed"
COMPLETED_WITH_WARNINGS = "completed-with-warnings"
PENDING = "pending"
FAILED = "failed"
class SnippetImportInfo(BaseModel):
id: str
status: ImportStatus
@@ -53,32 +41,9 @@ class SnippetImportInfo(BaseModel):
error: str = ""
class CheckDependenciesResult(BaseModel):
leaked_dependencies: list[PluginDependency] = Field(default_factory=list)
def _check_version_compatibility(imported_version: str) -> ImportStatus:
"""Determine import status based on version comparison"""
try:
current_ver = version.parse(CURRENT_DSL_VERSION)
imported_ver = version.parse(imported_version)
except version.InvalidVersion:
return ImportStatus.FAILED
# If imported version is newer than current, always return PENDING
if imported_ver > current_ver:
return ImportStatus.PENDING
# If imported version is older than current's major, return PENDING
if imported_ver.major < current_ver.major:
return ImportStatus.PENDING
# If imported version is older than current's minor, return COMPLETED_WITH_WARNINGS
if imported_ver.minor < current_ver.minor:
return ImportStatus.COMPLETED_WITH_WARNINGS
# If imported version equals or is older than current's micro, return COMPLETED
return ImportStatus.COMPLETED
"""Determine import status based on version comparison."""
return check_version_compatibility(imported_version, CURRENT_DSL_VERSION)
class SnippetPendingData(BaseModel):
@@ -145,13 +110,14 @@ class SnippetDslService:
status=ImportStatus.FAILED,
error=f"Failed to fetch YAML from URL: {response.status_code}",
)
content = response.text
if len(content) > DSL_MAX_SIZE:
raw_content = response.content
if dsl_content_size(raw_content) > DSL_MAX_SIZE:
return SnippetImportInfo(
id=import_id,
status=ImportStatus.FAILED,
error=f"YAML content size exceeds maximum limit of {DSL_MAX_SIZE} bytes",
)
content = raw_content.decode("utf-8")
except Exception as e:
logger.exception("Failed to fetch YAML from URL")
return SnippetImportInfo(
@@ -167,7 +133,7 @@ class SnippetDslService:
error="yaml_content is required when import_mode is yaml-content",
)
content = yaml_content
if len(content) > DSL_MAX_SIZE:
if dsl_content_size(content) > DSL_MAX_SIZE:
return SnippetImportInfo(
id=import_id,
status=ImportStatus.FAILED,