diff --git a/api/controllers/console/workspace/model_providers.py b/api/controllers/console/workspace/model_providers.py index adebe1a4a91..5ff114aa4e0 100644 --- a/api/controllers/console/workspace/model_providers.py +++ b/api/controllers/console/workspace/model_providers.py @@ -4,9 +4,11 @@ from typing import Any, Literal from flask import request, send_file from flask_restx import Resource from pydantic import BaseModel, Field, field_validator +from sqlalchemy.orm import Session from controllers.common.fields import SimpleResultResponse, ValidationResultResponse from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models +from controllers.common.session import with_session from controllers.console import console_ns from controllers.console.wraps import ( RBACPermission, @@ -26,8 +28,13 @@ from libs.helper import dump_response, uuid_value from libs.login import login_required from models import Account from services.billing_service import BillingService -from services.entities.model_provider_entities import ProviderResponse +from services.entities.model_provider_entities import ( + ModelProviderPluginSummaryResponse, + ModelProviderSummaryResponse, + ProviderResponse, +) from services.model_provider_service import ModelProviderService +from services.workspace_service import WorkspaceService class ParserModelList(BaseModel): @@ -91,6 +98,22 @@ class ModelProviderListResponse(ResponseModel): data: list[ProviderResponse] +class ModelProviderSummaryListResponse(ResponseModel): + data: list[ModelProviderSummaryResponse] + plugins: dict[str, ModelProviderPluginSummaryResponse] + + +class ModelProviderCreditsResponse(ResponseModel): + pool_type: Literal["paid", "trial"] | None + quota_limit: int | None = Field(description="Credit limit for the effective pool; -1 means unlimited.") + quota_used: int | None + remaining_credits: int | None = Field(description="Remaining credits; -1 means unlimited.") + is_unlimited: bool + is_exhausted: bool + exhausted_at: int | None + next_credit_reset_date: int | None + + class ProviderCredentialsResponse(ResponseModel): credentials: dict[str, Any] | None = None @@ -114,6 +137,8 @@ register_response_schema_models( console_ns, SimpleResultResponse, ModelProviderListResponse, + ModelProviderSummaryListResponse, + ModelProviderCreditsResponse, ProviderCredentialsResponse, ValidationResultResponse, ModelProviderPaymentCheckoutUrlResponse, @@ -140,6 +165,40 @@ class ModelProviderListApi(Resource): return ModelProviderListResponse(data=provider_list).model_dump(mode="json") +@console_ns.route("/workspaces/current/model-providers/summary") +class ModelProviderSummaryListApi(Resource): + @console_ns.response( + 200, + "Model provider summaries retrieved successfully", + console_ns.models[ModelProviderSummaryListResponse.__name__], + ) + @setup_required + @login_required + @account_initialization_required + @with_current_tenant_id + def get(self, tenant_id: str): + providers, plugins = ModelProviderService().get_provider_summary_list(tenant_id=tenant_id) + return dump_response( + ModelProviderSummaryListResponse, + {"data": providers, "plugins": plugins}, + ) + + +@console_ns.route("/workspaces/current/model-providers/credits") +class ModelProviderCreditsApi(Resource): + @console_ns.response( + 200, "Model provider credits retrieved successfully", console_ns.models[ModelProviderCreditsResponse.__name__] + ) + @setup_required + @login_required + @account_initialization_required + @with_current_tenant_id + @with_session(write=False) + def get(self, session: Session, tenant_id: str): + credit_pool = WorkspaceService.get_effective_credit_pool(tenant_id, session=session) + return dump_response(ModelProviderCreditsResponse, credit_pool) + + @console_ns.route("/workspaces/current/model-providers//credentials") class ModelProviderCredentialApi(Resource): @console_ns.doc(params=query_params_from_model(ParserCredentialId)) diff --git a/api/controllers/console/workspace/plugin.py b/api/controllers/console/workspace/plugin.py index 36ba2a621e1..e762d83caf5 100644 --- a/api/controllers/console/workspace/plugin.py +++ b/api/controllers/console/workspace/plugin.py @@ -1,5 +1,5 @@ import io -from collections.abc import Mapping +from collections.abc import Mapping, Sequence from datetime import datetime from typing import Any, Literal, TypedDict @@ -43,6 +43,7 @@ from core.plugin.entities.plugin_daemon import PluginDecodeResponse, PluginInsta from core.plugin.impl.exc import PluginDaemonClientSideError from core.plugin.plugin_service import PluginService from core.tools.builtin_tool.providers._positions import BuiltinToolProviderSort +from core.tools.entities.api_entities import ToolProviderApiEntity from core.tools.entities.common_entities import I18nObject from core.tools.entities.tool_entities import ToolProviderType from core.tools.tool_manager import ToolManager @@ -90,9 +91,21 @@ class ParserList(BaseModel): page_size: int = Field(default=256, ge=1, le=256, description="Page size (1-256)") +type PluginCategoryListLanguage = Literal["en_US", "zh_Hans", "ja_JP", "pt_BR"] + + class PluginCategoryListQuery(BaseModel): page: int = Field(default=1, ge=1, description="Page number") page_size: int = Field(default=256, ge=1, le=256, description="Page size (1-256)") + query: str = Field(default="", max_length=256, description="Case-insensitive search query") + tags: list[str] = Field(default_factory=list, max_length=128, description="Match any plugin tag") + language: Literal["en_US", "zh_Hans", "ja_JP", "pt_BR"] = Field( + default="en_US", description="Language used for localized label and description search" + ) + + +class PluginInstalledIdsQuery(BaseModel): + category: PluginCategory = Field(description="Plugin category to include") class ParserLatest(BaseModel): @@ -326,6 +339,10 @@ class PluginListResponse(ResponseModel): total: int +class PluginInstalledIdsResponse(ResponseModel): + plugin_ids: list[str] + + class PluginVersionsResponse(ResponseModel): versions: Mapping[str, PluginService.LatestPluginCache | None] @@ -385,6 +402,7 @@ register_schema_models( console_ns, ParserList, PluginCategoryListQuery, + PluginInstalledIdsQuery, PluginAutoUpgradeSettingsPayload, PluginPermissionSettingsPayload, ParserLatest, @@ -421,6 +439,7 @@ register_response_schema_models( PluginDebuggingKeyResponse, PluginDynamicOptionsResponse, PluginInstallationsResponse, + PluginInstalledIdsResponse, PluginInstallTaskStartResponse, PluginListResponse, PluginManifestResponse, @@ -478,7 +497,39 @@ def _read_upload_content(file: FileStorage, max_size: int) -> bytes: return content -def _list_hardcoded_builtin_tool_providers(tenant_id: str) -> list[dict[str, Any]]: +def _localized_builtin_tool_text(value: I18nObject, language: PluginCategoryListLanguage) -> str: + return value.to_dict()[language] or value.en_US + + +def _builtin_tool_provider_matches_filters( + provider: ToolProviderApiEntity, + *, + query: str, + tags: Sequence[str], + language: PluginCategoryListLanguage, +) -> bool: + if tags and not any(tag in provider.labels for tag in tags): + return False + if not query: + return True + + lower_query = query.lower() + candidates = ( + provider.name, + _localized_builtin_tool_text(provider.label, language), + _localized_builtin_tool_text(provider.description, language), + ) + return any(lower_query in candidate.lower() for candidate in candidates) + + +def _list_hardcoded_builtin_tool_providers( + tenant_id: str, + *, + query: str = "", + tags: Sequence[str] = (), + language: PluginCategoryListLanguage = "en_US", +) -> list[dict[str, Any]]: + """List builtin providers using the same search and tag semantics as category plugins.""" db_builtin_providers = { str(ToolProviderID(provider.provider)): provider for provider in ToolManager.list_default_builtin_providers(tenant_id) @@ -499,6 +550,13 @@ def _list_hardcoded_builtin_tool_providers(tenant_id: str) -> list[dict[str, Any db_provider=db_builtin_providers.get(provider.entity.identity.name), decrypt_credentials=False, ) + if not _builtin_tool_provider_matches_filters( + user_provider, + query=query, + tags=tags, + language=language, + ): + continue ToolTransformService.repack_provider(tenant_id=tenant_id, provider=user_provider) builtin_providers.append(user_provider) @@ -553,7 +611,9 @@ class PluginCategoryListApi(Resource): @account_initialization_required @with_current_tenant_id def get(self, tenant_id: str, category: str): - args = PluginCategoryListQuery.model_validate(request.args.to_dict(flat=True)) + args = PluginCategoryListQuery.model_validate( + {**request.args.to_dict(flat=True), "tags": request.args.getlist("tags")} + ) try: plugin_category = PluginCategory(category) @@ -561,13 +621,26 @@ class PluginCategoryListApi(Resource): return {"code": "invalid_param", "message": "invalid plugin category"}, 400 try: - plugins = PluginService.list_by_category(tenant_id, plugin_category, args.page, args.page_size) + plugins = PluginService.list_by_category( + tenant_id, + plugin_category, + args.page, + args.page_size, + query=args.query, + tags=args.tags, + language=args.language, + ) except PluginDaemonClientSideError as e: return {"code": "plugin_error", "message": e.description}, 400 builtin_tools = [] if plugin_category == PluginCategory.Tool: - builtin_tools = _list_hardcoded_builtin_tool_providers(tenant_id) + builtin_tools = _list_hardcoded_builtin_tool_providers( + tenant_id, + query=args.query, + tags=args.tags, + language=args.language, + ) return dump_response( PluginCategoryListResponse, @@ -579,6 +652,24 @@ class PluginCategoryListApi(Resource): ) +@console_ns.route("/workspaces/current/plugin/installed-ids") +class PluginInstalledIdsApi(Resource): + @console_ns.doc(params=query_params_from_model(PluginInstalledIdsQuery)) + @console_ns.response(200, "Success", console_ns.models[PluginInstalledIdsResponse.__name__]) + @setup_required + @login_required + @account_initialization_required + @with_current_tenant_id + def get(self, tenant_id: str): + args = PluginInstalledIdsQuery.model_validate(request.args.to_dict(flat=True)) + try: + plugin_ids = PluginService.list_installed_plugin_ids(tenant_id, args.category) + except PluginDaemonClientSideError as e: + return {"code": "plugin_error", "message": e.description}, 400 + + return dump_response(PluginInstalledIdsResponse, {"plugin_ids": plugin_ids}) + + @console_ns.route("/workspaces/current/plugin/list/latest-versions") class PluginListLatestVersionsApi(Resource): @console_ns.expect(console_ns.models[ParserLatest.__name__]) diff --git a/api/core/plugin/entities/plugin_daemon.py b/api/core/plugin/entities/plugin_daemon.py index 6b40ab4aaf7..521884e21c0 100644 --- a/api/core/plugin/entities/plugin_daemon.py +++ b/api/core/plugin/entities/plugin_daemon.py @@ -102,6 +102,19 @@ class PluginModelProviderEntity(BaseModel): declaration: ProviderEntity = Field(description="The declaration of the model provider.") +class PluginModelProviderBinding(BaseModel): + """Lightweight installation metadata for one model provider.""" + + provider: str + installation_id: str + plugin_id: str + plugin_unique_identifier: str + runtime_type: str + source: PluginInstallationSource + version: str + verified: bool = False + + class PluginTextEmbeddingNumTokensResponse(BaseModel): """ Response for number of tokens. @@ -215,6 +228,10 @@ class PluginListResponse(BaseModel): total: int +class PluginInstalledIdsDaemonResponse(BaseModel): + plugin_ids: list[str] + + class PluginListWithoutTotalResponse(BaseModel): list: list[PluginEntity] has_more: bool diff --git a/api/core/plugin/impl/model.py b/api/core/plugin/impl/model.py index ee1bd11901e..c38399bdcff 100644 --- a/api/core/plugin/impl/model.py +++ b/api/core/plugin/impl/model.py @@ -6,6 +6,7 @@ from core.plugin.entities.plugin_daemon import ( PluginBasicBooleanResponse, PluginDaemonInnerError, PluginLLMNumTokensResponse, + PluginModelProviderBinding, PluginModelProviderEntity, PluginModelSchemaEntity, PluginStringResultResponse, @@ -48,6 +49,14 @@ class PluginModelClient(BasePluginClient): ) return response + def fetch_model_provider_bindings(self, tenant_id: str) -> Sequence[PluginModelProviderBinding]: + """Fetch only model-provider installation identities from the daemon.""" + return self._request_with_plugin_daemon_response( + "GET", + f"plugin/{tenant_id}/management/models/bindings", + list[PluginModelProviderBinding], + ) + def get_model_schema( self, tenant_id: str, diff --git a/api/core/plugin/impl/plugin.py b/api/core/plugin/impl/plugin.py index 34e8d315d86..8ac33b50297 100644 --- a/api/core/plugin/impl/plugin.py +++ b/api/core/plugin/impl/plugin.py @@ -14,6 +14,7 @@ from core.plugin.entities.plugin import ( ) from core.plugin.entities.plugin_daemon import ( PluginDecodeResponse, + PluginInstalledIdsDaemonResponse, PluginInstallTask, PluginInstallTaskStartResponse, PluginListResponse, @@ -68,6 +69,16 @@ class PluginInstaller(BasePluginClient): ) return result.list + def list_installed_plugin_ids(self, tenant_id: str, category: PluginCategory) -> list[str]: + """List all currently installed plugin IDs in one category.""" + result = self._request_with_plugin_daemon_response( + "GET", + f"plugin/{tenant_id}/management/installation/ids", + PluginInstalledIdsDaemonResponse, + params={"category": category.value}, + ) + return result.plugin_ids + def list_plugins_with_total(self, tenant_id: str, page: int, page_size: int) -> PluginListResponse: return self._request_with_plugin_daemon_response( "GET", @@ -77,13 +88,28 @@ class PluginInstaller(BasePluginClient): ) def list_plugins_by_category( - self, tenant_id: str, category: PluginCategory, page: int, page_size: int + self, + tenant_id: str, + category: PluginCategory, + page: int, + page_size: int, + *, + query: str = "", + tags: Sequence[str] = (), + language: str = "en_US", ) -> PluginListWithoutTotalResponse: return self._request_with_plugin_daemon_response( "GET", f"plugin/{tenant_id}/management/{category.value}/list", PluginListWithoutTotalResponse, - params={"page": page, "page_size": page_size, "response_type": "paged"}, + params={ + "page": page, + "page_size": page_size, + "response_type": "paged", + "query": query, + "tags": list(tags), + "language": language, + }, ) def upload_pkg( diff --git a/api/core/plugin/plugin_service.py b/api/core/plugin/plugin_service.py index 8ef50a96925..410287e604f 100644 --- a/api/core/plugin/plugin_service.py +++ b/api/core/plugin/plugin_service.py @@ -48,6 +48,7 @@ from core.plugin.entities.plugin_daemon import ( PluginInstallTaskStatus, PluginListResponse, PluginListWithoutTotalResponse, + PluginModelProviderBinding, PluginModelProviderDeclaration, PluginModelProviderEntity, PluginVerification, @@ -81,6 +82,12 @@ class _RedisLock(Protocol): def release(self) -> None: ... +class _ModelPluginIdentity(Protocol): + plugin_id: str + plugin_unique_identifier: str + source: PluginInstallationSource + + class PluginService: class LatestPluginCache(BaseModel): plugin_id: str @@ -301,7 +308,7 @@ class PluginService: logger.warning("Failed to cache plugin model providers for tenant %s.", tenant_id, exc_info=True) @classmethod - def _get_remote_model_plugin_cache_marker(cls, plugins: Sequence[PluginEntity]) -> str | None: + def _get_remote_model_plugin_cache_marker(cls, plugins: Sequence[_ModelPluginIdentity]) -> str | None: remote_model_plugins = sorted( f"{plugin.plugin_id}:{plugin.plugin_unique_identifier}" for plugin in plugins @@ -386,7 +393,7 @@ class PluginService: def _should_invalidate_model_provider_cache_for_remote_model_plugins( cls, tenant_id: str, - plugins: Sequence[PluginEntity], + plugins: Sequence[_ModelPluginIdentity], ) -> bool: remote_model_plugin_marker = cls._get_remote_model_plugin_cache_marker(plugins) cached_remote_model_plugin_marker = cls._load_cached_remote_model_plugin_marker(tenant_id) @@ -715,6 +722,26 @@ class PluginService: plugins = manager.list_plugins(tenant_id) return plugins + @staticmethod + def list_installed_plugin_ids(tenant_id: str, category: PluginCategory) -> Sequence[str]: + """List all currently installed plugin IDs in one category through the daemon's lightweight query.""" + manager = PluginInstaller() + return manager.list_installed_plugin_ids(tenant_id, category) + + @staticmethod + def list_model_provider_bindings( + tenant_id: str, *, client: PluginModelClient | None = None + ) -> Sequence[PluginModelProviderBinding]: + """Return fresh model bindings and reconcile remote-debug provider metadata before it is read.""" + model_client = client or PluginModelClient() + bindings = model_client.fetch_model_provider_bindings(tenant_id) + if PluginService._should_invalidate_model_provider_cache_for_remote_model_plugins(tenant_id, bindings): + PluginService.invalidate_plugin_model_providers_cache(tenant_id) + + marker = PluginService._get_remote_model_plugin_cache_marker(bindings) + PluginService._store_cached_remote_model_plugin_marker(tenant_id, marker) + return bindings + @staticmethod def list_with_total(tenant_id: str, user_id: str, page: int, page_size: int) -> PluginListResponse: """List tenant plugins with endpoint counts reconciled from live records. @@ -731,17 +758,33 @@ class PluginService: @staticmethod def list_by_category( - tenant_id: str, category: PluginCategory, page: int, page_size: int + tenant_id: str, + category: PluginCategory, + page: int, + page_size: int, + *, + query: str = "", + tags: Sequence[str] = (), + language: str = "en_US", ) -> PluginListWithoutTotalResponse: """ List plugins in one category with a has-more cursor signal and without calculating total. - The daemon scans tenant installations in the existing list order and stops once it finds one extra match. - This keeps pagination usable before category is persisted on installation rows. + The daemon applies category, search, and tag filters before pagination, then stops once it finds one extra + match. Only a complete, unfiltered first page may reconcile the model-provider cache; the unpaginated model + binding read is the authoritative marker source for larger result sets. """ manager = PluginInstaller() - plugins = manager.list_plugins_by_category(tenant_id, category, page, page_size) - if category == PluginCategory.Model: + plugins = manager.list_plugins_by_category( + tenant_id, + category, + page, + page_size, + query=query, + tags=tags, + language=language, + ) + if category == PluginCategory.Model and page == 1 and not plugins.has_more and not query and not tags: should_invalidate_model_provider_cache = ( PluginService._should_invalidate_model_provider_cache_for_remote_model_plugins( tenant_id, diff --git a/api/openapi/markdown/console-openapi.md b/api/openapi/markdown/console-openapi.md index bca7cd812b8..67e1a4f9e99 100644 --- a/api/openapi/markdown/console-openapi.md +++ b/api/openapi/markdown/console-openapi.md @@ -10597,6 +10597,20 @@ Update a plugin endpoint | ---- | ----------- | ------ | | 200 | Model providers retrieved successfully | **application/json**: [ModelProviderListResponse](#modelproviderlistresponse)
| +### [GET] /workspaces/current/model-providers/credits +#### Responses + +| Code | Description | Schema | +| ---- | ----------- | ------ | +| 200 | Model provider credits retrieved successfully | **application/json**: [ModelProviderCreditsResponse](#modelprovidercreditsresponse)
| + +### [GET] /workspaces/current/model-providers/summary +#### Responses + +| Code | Description | Schema | +| ---- | ----------- | ------ | +| 200 | Model provider summaries retrieved successfully | **application/json**: [ModelProviderSummaryListResponse](#modelprovidersummarylistresponse)
| + ### [GET] /workspaces/current/model-providers/{provider}/checkout-url #### Parameters @@ -11142,6 +11156,19 @@ Returns permission flags that control workspace features like member invitations | ---- | ----------- | ------ | | 200 | Success | **application/json**: [PluginInstallTaskStartResponse](#plugininstalltaskstartresponse)
| +### [GET] /workspaces/current/plugin/installed-ids +#### Parameters + +| Name | Located in | Description | Required | Schema | +| ---- | ---------- | ----------- | -------- | ------ | +| category | query | Plugin category to include | Yes | string,
**Available values:** "agent-strategy", "datasource", "extension", "model", "tool", "trigger" | + +#### Responses + +| Code | Description | Schema | +| ---- | ----------- | ------ | +| 200 | Success | **application/json**: [PluginInstalledIdsResponse](#plugininstalledidsresponse)
| + ### [GET] /workspaces/current/plugin/list #### Parameters @@ -11400,8 +11427,11 @@ Returns permission flags that control workspace features like member invitations | Name | Located in | Description | Required | Schema | | ---- | ---------- | ----------- | -------- | ------ | +| language | query | Language used for localized label and description search | No | string,
**Available values:** "en_US", "ja_JP", "pt_BR", "zh_Hans",
**Default:** en_US | | page | query | Page number | No | integer,
**Default:** 1 | | page_size | query | Page size (1-256) | No | integer,
**Default:** 256 | +| query | query | Case-insensitive search query | No | string | +| tags | query | Match any plugin tag | No | [ string ] | | category | path | | Yes | string | #### Responses @@ -19381,6 +19411,30 @@ Enum class for model property key. | ---- | ---- | ----------- | -------- | | ModelPropertyKey | string | Enum class for model property key. | | +#### ModelProviderCreditsResponse + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| exhausted_at | integer | | Yes | +| is_exhausted | boolean | | Yes | +| is_unlimited | boolean | | Yes | +| next_credit_reset_date | integer | | Yes | +| pool_type | string | | Yes | +| quota_limit | integer | Credit limit for the effective pool; -1 means unlimited. | Yes | +| quota_used | integer | | Yes | +| remaining_credits | integer | Remaining credits; -1 means unlimited. | Yes | + +#### ModelProviderCustomConfigurationSummaryResponse + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| available_credentials | [ [CredentialConfiguration](#credentialconfiguration) ] | | Yes | +| current_credential_id | string | | No | +| current_credential_name | string | | No | +| current_credential_usable | boolean | | Yes | +| has_custom_models | boolean | Whether custom model configuration exists, including saved model credentials. | Yes | +| status | [CustomConfigurationStatus](#customconfigurationstatus) | | Yes | + #### ModelProviderListResponse | Name | Type | Description | Required | @@ -19393,6 +19447,49 @@ Enum class for model property key. | ---- | ---- | ----------- | -------- | | payment_link | string | | Yes | +#### ModelProviderPluginSummaryResponse + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| installation_id | string | | Yes | +| plugin_id | string | | Yes | +| plugin_unique_identifier | string | | Yes | +| runtime_type | string | | Yes | +| source | [PluginInstallationSource](#plugininstallationsource) | | Yes | +| version | string | | Yes | + +#### ModelProviderSummaryListResponse + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| data | [ [ModelProviderSummaryResponse](#modelprovidersummaryresponse) ] | | Yes | +| plugins | object | | Yes | + +#### ModelProviderSummaryResponse + +Fields required to render the collapsed model-provider list. + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| configurate_methods | [ [ConfigurateMethod](#configuratemethod) ] | | Yes | +| custom_configuration | [ModelProviderCustomConfigurationSummaryResponse](#modelprovidercustomconfigurationsummaryresponse) | | Yes | +| description | [I18nObject](#i18nobject) | | No | +| icon_small | [I18nObject](#i18nobject) | | No | +| icon_small_dark | [I18nObject](#i18nobject) | | No | +| is_configured | boolean | | Yes | +| label | [I18nObject](#i18nobject) | | Yes | +| plugin_id | string | | Yes | +| preferred_provider_type | [ProviderType](#providertype) | | Yes | +| provider | string | | Yes | +| supported_model_types | [ [ModelType](#modeltype) ] | | Yes | +| system_configuration | [ModelProviderSystemConfigurationSummaryResponse](#modelprovidersystemconfigurationsummaryresponse) | | Yes | + +#### ModelProviderSystemConfigurationSummaryResponse + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| enabled | boolean | | Yes | + #### ModelSelectorScope | Name | Type | Description | Required | @@ -20382,8 +20479,11 @@ Shared permission levels for resources (datasets, credentials, etc.) | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | +| language | string,
**Available values:** "en_US", "ja_JP", "pt_BR", "zh_Hans",
**Default:** en_US | Language used for localized label and description search
*Enum:* `"en_US"`, `"ja_JP"`, `"pt_BR"`, `"zh_Hans"` | No | | page | integer,
**Default:** 1 | Page number | No | | page_size | integer,
**Default:** 256 | Page size (1-256) | No | +| query | string | Case-insensitive search query | No | +| tags | [ string ] | Match any plugin tag | No | #### PluginCategoryListResponse @@ -20584,6 +20684,18 @@ Shared permission levels for resources (datasets, credentials, etc.) | ---- | ---- | ----------- | -------- | | plugins | [ [PluginInstallationItemResponse](#plugininstallationitemresponse) ] | | Yes | +#### PluginInstalledIdsQuery + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| category | [PluginCategory](#plugincategory) | Plugin category to include | Yes | + +#### PluginInstalledIdsResponse + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| plugin_ids | [ string ] | | Yes | + #### PluginListResponse | Name | Type | Description | Required | diff --git a/api/services/entities/model_provider_entities.py b/api/services/entities/model_provider_entities.py index 020dc4a2ea9..8a9e8dd66b3 100644 --- a/api/services/entities/model_provider_entities.py +++ b/api/services/entities/model_provider_entities.py @@ -17,6 +17,7 @@ from core.entities.provider_entities import ( QuotaConfiguration, UnaddedModelConfiguration, ) +from core.plugin.entities.plugin import PluginInstallationSource from graphon.model_runtime.entities.common_entities import I18nObject from graphon.model_runtime.entities.model_entities import ( FetchFrom, @@ -69,6 +70,67 @@ class SystemConfigurationResponse(BaseModel): quota_configurations: list[QuotaConfiguration] = [] +class ModelProviderCustomConfigurationSummaryResponse(BaseModel): + status: CustomConfigurationStatus + has_custom_models: bool = Field( + description="Whether custom model configuration exists, including saved model credentials." + ) + available_credentials: list[CredentialConfiguration] + current_credential_id: str | None = None + current_credential_name: str | None = None + current_credential_usable: bool + + +class ModelProviderSystemConfigurationSummaryResponse(BaseModel): + enabled: bool + + +class ModelProviderPluginSummaryResponse(BaseModel): + installation_id: str + plugin_id: str + plugin_unique_identifier: str + runtime_type: str + source: PluginInstallationSource + version: str + + +class ModelProviderSummaryResponse(BaseModel): + """Fields required to render the collapsed model-provider list.""" + + tenant_id: str = Field(exclude=True) + provider: str + plugin_id: str + label: I18nObject + description: I18nObject | None = None + icon_small: I18nObject | None = None + icon_small_dark: I18nObject | None = None + supported_model_types: Sequence[ModelType] + configurate_methods: list[ConfigurateMethod] + preferred_provider_type: ProviderType + is_configured: bool + custom_configuration: ModelProviderCustomConfigurationSummaryResponse + system_configuration: ModelProviderSystemConfigurationSummaryResponse + + model_config = ConfigDict(protected_namespaces=()) + + @model_validator(mode="after") + def build_icon_urls(self): + url_prefix = ( + dify_config.CONSOLE_API_URL + f"/console/api/workspaces/{self.tenant_id}/model-providers/{self.provider}" + ) + if self.icon_small is not None: + self.icon_small = I18nObject( + en_US=f"{url_prefix}/icon_small/en_US", + zh_Hans=f"{url_prefix}/icon_small/zh_Hans", + ) + if self.icon_small_dark is not None: + self.icon_small_dark = I18nObject( + en_US=f"{url_prefix}/icon_small_dark/en_US", + zh_Hans=f"{url_prefix}/icon_small_dark/zh_Hans", + ) + return self + + class ProviderResponse(BaseModel): """ Model class for provider response. diff --git a/api/services/model_provider_service.py b/api/services/model_provider_service.py index 7c34afd42e1..e30f25bfea6 100644 --- a/api/services/model_provider_service.py +++ b/api/services/model_provider_service.py @@ -1,18 +1,42 @@ import logging +from collections import defaultdict +from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any +from sqlalchemy import and_, select + if TYPE_CHECKING: from models.account import Account +from configs import dify_config +from core.db.session_factory import session_factory from core.entities.model_entities import ModelWithProviderEntity, ProviderModelWithStatusEntity +from core.entities.provider_entities import CredentialConfiguration +from core.helper.position_helper import is_filtered +from core.plugin.entities.plugin import PluginInstallationSource +from core.plugin.entities.plugin_daemon import PluginModelProviderBinding from core.plugin.impl.model_runtime_factory import create_plugin_model_provider_factory, create_plugin_provider_manager +from core.plugin.plugin_service import PluginService from core.provider_manager import ProviderManager +from extensions import ext_hosting_provider from graphon.model_runtime.entities.model_entities import ModelType, ParameterRule -from models.provider import ProviderType +from models.provider import ( + Provider, + ProviderCredential, + ProviderModel, + ProviderModelCredential, + ProviderType, + TenantPreferredModelProvider, +) +from models.provider_ids import ModelProviderID from services.entities.model_provider_entities import ( CustomConfigurationResponse, CustomConfigurationStatus, DefaultModelResponse, + ModelProviderCustomConfigurationSummaryResponse, + ModelProviderPluginSummaryResponse, + ModelProviderSummaryResponse, + ModelProviderSystemConfigurationSummaryResponse, ModelWithProviderEntityResponse, ProviderResponse, ProviderWithModelsResponse, @@ -24,6 +48,17 @@ from services.errors.app_model_config import ProviderNotFoundError logger = logging.getLogger(__name__) +@dataclass(slots=True) +class _ProviderSummaryState: + has_custom_provider: bool = False + available_credentials: list[CredentialConfiguration] = field(default_factory=list) + has_custom_models: bool = False + current_credential_id: str | None = None + current_credential_name: str | None = None + current_credential_usable: bool = False + preferred_provider_type: ProviderType | None = None + + class ModelProviderService: """ Model Provider Service @@ -132,6 +167,249 @@ class ModelProviderService: return provider_responses + @staticmethod + def _load_provider_summary_states(tenant_id: str) -> dict[str, _ProviderSummaryState]: + """Load only the workspace columns required by the collapsed provider list.""" + with session_factory.create_session() as session: + custom_provider_rows = session.execute( + select( + Provider.provider_name, + Provider.credential_id, + ProviderCredential.provider_name.label("credential_provider_name"), + ProviderCredential.credential_name, + ) + .outerjoin( + ProviderCredential, + and_( + ProviderCredential.id == Provider.credential_id, + ProviderCredential.tenant_id == tenant_id, + ), + ) + .where( + Provider.tenant_id == tenant_id, + Provider.provider_type == ProviderType.CUSTOM, + Provider.is_valid.is_(True), + ) + ).all() + credential_rows = session.execute( + select( + ProviderCredential.id, + ProviderCredential.provider_name, + ProviderCredential.credential_name, + ) + .where(ProviderCredential.tenant_id == tenant_id) + .order_by( + ProviderCredential.created_at.desc(), + ProviderCredential.id.desc(), + ) + ).all() + custom_model_rows = session.execute( + select(ProviderModel.provider_name.label("provider_name")) + .where( + ProviderModel.tenant_id == tenant_id, + ProviderModel.is_valid.is_(True), + ) + .union( + select(ProviderModelCredential.provider_name.label("provider_name")).where( + ProviderModelCredential.tenant_id == tenant_id + ) + ) + ).all() + preferred_provider_rows = session.execute( + select( + TenantPreferredModelProvider.provider_name, + TenantPreferredModelProvider.preferred_provider_type, + ).where(TenantPreferredModelProvider.tenant_id == tenant_id) + ).all() + + states: defaultdict[str, _ProviderSummaryState] = defaultdict(_ProviderSummaryState) + for credential in credential_rows: + provider_name = str(ModelProviderID(credential.provider_name)) + states[provider_name].available_credentials.append( + CredentialConfiguration( + credential_id=credential.id, + credential_name=credential.credential_name, + ) + ) + + selected_provider_priorities: dict[str, bool] = {} + for provider in custom_provider_rows: + provider_name = str(ModelProviderID(provider.provider_name)) + state = states[provider_name] + state.has_custom_provider = True + + is_canonical_row = provider.provider_name == provider_name + if provider_name in selected_provider_priorities and not is_canonical_row: + continue + selected_provider_priorities[provider_name] = is_canonical_row + state.current_credential_id = provider.credential_id + if ( + provider.credential_provider_name is not None + and str(ModelProviderID(provider.credential_provider_name)) == provider_name + ): + state.current_credential_name = provider.credential_name + state.current_credential_usable = True + else: + state.current_credential_name = None + state.current_credential_usable = False + + for model in custom_model_rows: + states[str(ModelProviderID(model.provider_name))].has_custom_models = True + + preferred_provider_priorities: dict[str, bool] = {} + for preferred_provider in preferred_provider_rows: + provider_name = str(ModelProviderID(preferred_provider.provider_name)) + is_canonical_row = preferred_provider.provider_name == provider_name + if provider_name in preferred_provider_priorities and not is_canonical_row: + continue + preferred_provider_priorities[provider_name] = is_canonical_row + states[provider_name].preferred_provider_type = preferred_provider.preferred_provider_type + + return dict(states) + + @staticmethod + def _has_system_provider_hosting_configuration(provider: str) -> bool: + configuration = ext_hosting_provider.hosting_configuration.provider_map.get(provider) + return bool(configuration and configuration.enabled and configuration.quotas) + + @staticmethod + def _select_binding( + current_binding: PluginModelProviderBinding | None, + candidate_binding: PluginModelProviderBinding, + ) -> PluginModelProviderBinding: + """Prefer a remote-debug runtime when one shadows an installed plugin.""" + if current_binding is None: + return candidate_binding + if ( + candidate_binding.source == PluginInstallationSource.Remote + and current_binding.source != PluginInstallationSource.Remote + ): + return candidate_binding + return current_binding + + @staticmethod + def _get_preferred_provider_type( + state: _ProviderSummaryState, + *, + custom_present: bool, + system_enabled: bool, + ) -> ProviderType: + if state.preferred_provider_type is not None: + return state.preferred_provider_type + if dify_config.EDITION == "CLOUD" and system_enabled: + return ProviderType.SYSTEM + if custom_present: + return ProviderType.CUSTOM + if system_enabled: + return ProviderType.SYSTEM + return ProviderType.CUSTOM + + def get_provider_summary_list( + self, tenant_id: str + ) -> tuple[list[ModelProviderSummaryResponse], dict[str, ModelProviderPluginSummaryResponse]]: + """Build the complete first-screen provider projection without assembling provider configurations.""" + # Read bindings first: remote-debug identity changes invalidate provider metadata + # before the provider cache is consulted. + bindings = PluginService.list_model_provider_bindings(tenant_id) + provider_entities = PluginService.fetch_plugin_model_providers(tenant_id=tenant_id) + states = self._load_provider_summary_states(tenant_id) + + bindings_by_provider: dict[str, PluginModelProviderBinding] = {} + for binding in bindings: + provider_name = ( + str(ModelProviderID(binding.provider)) + if binding.provider.count("/") == 2 + else str(ModelProviderID(f"{binding.plugin_id}/{binding.provider}")) + ) + bindings_by_provider[provider_name] = self._select_binding( + bindings_by_provider.get(provider_name), + binding, + ) + + provider_summaries: list[ModelProviderSummaryResponse] = [] + emitted_provider_names: set[str] = set() + for provider_entity in provider_entities: + if is_filtered( + include_set=dify_config.POSITION_PROVIDER_INCLUDES_SET, + exclude_set=dify_config.POSITION_PROVIDER_EXCLUDES_SET, + data=provider_entity, + name_func=lambda provider: provider.provider, + ): + continue + + provider_id = ModelProviderID(provider_entity.provider) + provider_name = str(provider_id) + if provider_name in emitted_provider_names: + continue + emitted_provider_names.add(provider_name) + + state = states.get(provider_name, _ProviderSummaryState()) + custom_configured = ( + state.has_custom_provider and bool(state.available_credentials) + ) or state.has_custom_models + custom_present = state.has_custom_provider or state.has_custom_models + provider_binding = bindings_by_provider.get(provider_name) + system_enabled = bool( + provider_binding + and self._has_system_provider_hosting_configuration(provider_name) + and provider_binding.source != PluginInstallationSource.Package + and provider_binding.verified + ) + preferred_provider_type = self._get_preferred_provider_type( + state, + custom_present=custom_present, + system_enabled=system_enabled, + ) + + provider_summaries.append( + ModelProviderSummaryResponse( + tenant_id=tenant_id, + provider=provider_name, + plugin_id=provider_id.plugin_id, + label=provider_entity.label, + description=provider_entity.description, + icon_small=provider_entity.icon_small, + icon_small_dark=provider_entity.icon_small_dark, + supported_model_types=provider_entity.supported_model_types, + configurate_methods=provider_entity.configurate_methods, + preferred_provider_type=preferred_provider_type, + is_configured=custom_configured or system_enabled, + custom_configuration=ModelProviderCustomConfigurationSummaryResponse( + status=CustomConfigurationStatus.ACTIVE + if custom_configured + else CustomConfigurationStatus.NO_CONFIGURE, + has_custom_models=state.has_custom_models, + available_credentials=state.available_credentials, + current_credential_id=state.current_credential_id, + current_credential_name=state.current_credential_name, + current_credential_usable=state.current_credential_usable, + ), + system_configuration=ModelProviderSystemConfigurationSummaryResponse( + enabled=system_enabled, + ), + ) + ) + + plugin_bindings: dict[str, PluginModelProviderBinding] = {} + for binding in bindings_by_provider.values(): + plugin_bindings[binding.plugin_id] = self._select_binding( + plugin_bindings.get(binding.plugin_id), + binding, + ) + + plugin_summaries = { + plugin_id: ModelProviderPluginSummaryResponse( + installation_id=binding.installation_id, + plugin_id=binding.plugin_id, + plugin_unique_identifier=binding.plugin_unique_identifier, + runtime_type=binding.runtime_type, + source=binding.source, + version=binding.version, + ) + for plugin_id, binding in plugin_bindings.items() + } + return provider_summaries, plugin_summaries + def get_models_by_provider(self, tenant_id: str, provider: str) -> list[ModelWithProviderEntityResponse]: """ get provider models. diff --git a/api/services/workspace_service.py b/api/services/workspace_service.py index 0635f519438..aa09b24278c 100644 --- a/api/services/workspace_service.py +++ b/api/services/workspace_service.py @@ -1,3 +1,6 @@ +from dataclasses import dataclass +from typing import Literal + from flask_login import current_user from sqlalchemy import select from sqlalchemy.orm import Session @@ -7,9 +10,37 @@ from enums.cloud_plan import CloudPlan from enums.deployment_edition import DeploymentEdition from models.account import Tenant, TenantAccountJoin, TenantAccountRole from services.account_service import TenantService +from services.billing_service import BillingService from services.feature_service import FeatureService +@dataclass(frozen=True) +class EffectiveCreditPool: + plan: str | None = None + pool_type: Literal["paid", "trial"] | None = None + quota_limit: int | None = None + quota_used: int | None = None + exhausted_at: int | None = None + next_credit_reset_date: int | None = None + + @property + def remaining_credits(self) -> int | None: + if self.quota_limit is None or self.quota_used is None: + return None + if self.is_unlimited: + return -1 + return max(0, self.quota_limit - self.quota_used) + + @property + def is_unlimited(self) -> bool: + return self.quota_limit == -1 + + @property + def is_exhausted(self) -> bool: + remaining_credits = self.remaining_credits + return not self.is_unlimited and (remaining_credits is None or remaining_credits <= 0) + + def _set_credit_pool_info( tenant_info: dict[str, object], *, quota_limit: int, quota_used: int, exhausted_at: int | None = None ) -> None: @@ -20,6 +51,51 @@ def _set_credit_pool_info( class WorkspaceService: + @classmethod + def get_effective_credit_pool(cls, tenant_id: str, *, session: Session) -> EffectiveCreditPool: + if not dify_config.BILLING_ENABLED: + return EffectiveCreditPool() + + billing_info = BillingService.get_info(tenant_id, exclude_vector_space=True) + subscription_plan: str = billing_info["subscription"]["plan"] + + from services.credit_pool_service import CreditPoolBalance, CreditPoolService + + effective_pool = None + effective_pool_type: Literal["paid", "trial"] = "trial" + if subscription_plan != CloudPlan.SANDBOX: + paid_pool = CreditPoolService.get_pool(tenant_id=tenant_id, pool_type="paid", session=session) + if paid_pool is not None and (paid_pool.quota_limit == -1 or paid_pool.quota_limit > paid_pool.quota_used): + effective_pool = paid_pool + effective_pool_type = "paid" + + if effective_pool is None: + effective_pool = CreditPoolService.get_pool(tenant_id=tenant_id, pool_type="trial", session=session) + + if effective_pool is None: + return EffectiveCreditPool( + plan=subscription_plan if billing_info["enabled"] else None, + next_credit_reset_date=billing_info.get("next_credit_reset_date"), + ) + + exhausted_at = effective_pool.exhausted_at if isinstance(effective_pool, CreditPoolBalance) else None + if not ( + isinstance(exhausted_at, int) + and exhausted_at > 0 + and effective_pool.quota_limit > 0 + and effective_pool.quota_used >= effective_pool.quota_limit + ): + exhausted_at = None + + return EffectiveCreditPool( + plan=subscription_plan if billing_info["enabled"] else None, + pool_type=effective_pool_type, + quota_limit=effective_pool.quota_limit, + quota_used=effective_pool.quota_used, + exhausted_at=exhausted_at, + next_credit_reset_date=billing_info.get("next_credit_reset_date"), + ) + @classmethod def get_tenant_info(cls, tenant: Tenant, session: Session): if not tenant: diff --git a/api/tests/unit_tests/controllers/console/workspace/test_model_providers.py b/api/tests/unit_tests/controllers/console/workspace/test_model_providers.py index af37c20edaa..5fe0f47ecde 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_model_providers.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_model_providers.py @@ -12,12 +12,15 @@ from configs import dify_config from controllers.console.workspace.model_providers import ( ModelProviderCredentialApi, ModelProviderCredentialSwitchApi, + ModelProviderCreditsApi, ModelProviderIconApi, ModelProviderListApi, ModelProviderPaymentCheckoutUrlApi, + ModelProviderSummaryListApi, ModelProviderValidateApi, PreferredProviderTypeUpdateApi, ) +from core.entities.provider_entities import CredentialConfiguration from graphon.model_runtime.entities.common_entities import I18nObject from graphon.model_runtime.entities.model_entities import ModelType from graphon.model_runtime.entities.provider_entities import ConfigurateMethod @@ -27,9 +30,14 @@ from models.provider import ProviderType from services.entities.model_provider_entities import ( CustomConfigurationResponse, CustomConfigurationStatus, + ModelProviderCustomConfigurationSummaryResponse, + ModelProviderPluginSummaryResponse, + ModelProviderSummaryResponse, + ModelProviderSystemConfigurationSummaryResponse, ProviderResponse, SystemConfigurationResponse, ) +from services.workspace_service import EffectiveCreditPool VALID_UUID = "123e4567-e89b-12d3-a456-426614174000" INVALID_UUID = "123" @@ -140,6 +148,136 @@ class TestModelProviderListApi: assert result == {"data": []} +class TestModelProviderSummaryListApi: + def test_get_success(self, app: Flask): + api = ModelProviderSummaryListApi() + method = unwrap(api.get) + provider = ModelProviderSummaryResponse( + tenant_id="tenant1", + provider="langgenius/openai/openai", + plugin_id="langgenius/openai", + label=I18nObject(en_US="OpenAI"), + description=I18nObject(en_US="OpenAI models"), + icon_small=I18nObject(en_US="icon.svg"), + icon_small_dark=None, + supported_model_types=[ModelType.LLM], + configurate_methods=[ConfigurateMethod.PREDEFINED_MODEL], + preferred_provider_type=ProviderType.CUSTOM, + is_configured=True, + custom_configuration=ModelProviderCustomConfigurationSummaryResponse( + status=CustomConfigurationStatus.ACTIVE, + has_custom_models=True, + available_credentials=[ + CredentialConfiguration( + credential_id=VALID_UUID, + credential_name="production", + ), + CredentialConfiguration( + credential_id="223e4567-e89b-12d3-a456-426614174000", + credential_name="backup", + ), + ], + current_credential_id=VALID_UUID, + current_credential_name="production", + current_credential_usable=True, + ), + system_configuration=ModelProviderSystemConfigurationSummaryResponse(enabled=False), + ) + plugin = ModelProviderPluginSummaryResponse( + installation_id="installation-1", + plugin_id="langgenius/openai", + plugin_unique_identifier="langgenius/openai:1.0.0@checksum", + runtime_type="local", + source="marketplace", + version="1.0.0", + ) + + with ( + app.test_request_context("/"), + patch( + "controllers.console.workspace.model_providers.ModelProviderService.get_provider_summary_list", + return_value=([provider], {"langgenius/openai": plugin}), + ) as get_provider_summary_list, + ): + result = method(api, "tenant1") + + get_provider_summary_list.assert_called_once_with(tenant_id="tenant1") + assert result["data"][0]["provider"] == "langgenius/openai/openai" + assert "tenant_id" not in result["data"][0] + assert result["data"][0]["custom_configuration"] == { + "status": "active", + "has_custom_models": True, + "available_credentials": [ + { + "credential_id": VALID_UUID, + "credential_name": "production", + }, + { + "credential_id": "223e4567-e89b-12d3-a456-426614174000", + "credential_name": "backup", + }, + ], + "current_credential_id": VALID_UUID, + "current_credential_name": "production", + "current_credential_usable": True, + } + assert result["plugins"]["langgenius/openai"]["installation_id"] == "installation-1" + + +class TestModelProviderCreditsApi: + def test_get_success(self): + api = ModelProviderCreditsApi() + method = unwrap(api.get) + session = SimpleNamespace() + credit_pool = EffectiveCreditPool( + plan="team", + pool_type="paid", + quota_limit=-1, + quota_used=999, + next_credit_reset_date=1775001600, + ) + + with patch( + "controllers.console.workspace.model_providers.WorkspaceService.get_effective_credit_pool", + return_value=credit_pool, + ) as get_effective_credit_pool: + result = method(api, session, "tenant1") + + get_effective_credit_pool.assert_called_once_with("tenant1", session=session) + assert result == { + "pool_type": "paid", + "quota_limit": -1, + "quota_used": 999, + "remaining_credits": -1, + "is_unlimited": True, + "is_exhausted": False, + "exhausted_at": None, + "next_credit_reset_date": 1775001600, + } + + def test_get_without_effective_pool(self): + api = ModelProviderCreditsApi() + method = unwrap(api.get) + session = SimpleNamespace() + + with patch( + "controllers.console.workspace.model_providers.WorkspaceService.get_effective_credit_pool", + return_value=EffectiveCreditPool(), + ): + result = method(api, session, "tenant1") + + assert result == { + "pool_type": None, + "quota_limit": None, + "quota_used": None, + "remaining_credits": None, + "is_unlimited": False, + "is_exhausted": True, + "exhausted_at": None, + "next_credit_reset_date": None, + } + + class TestModelProviderCredentialApi: def test_get_success(self, app: Flask): api = ModelProviderCredentialApi() diff --git a/api/tests/unit_tests/controllers/console/workspace/test_plugin.py b/api/tests/unit_tests/controllers/console/workspace/test_plugin.py index e92c4651983..7a1984c64d0 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_plugin.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_plugin.py @@ -28,6 +28,7 @@ from controllers.console.workspace.plugin import ( PluginFetchMarketplacePkgApi, PluginFetchPermissionApi, PluginIconApi, + PluginInstalledIdsApi, PluginInstallFromGithubApi, PluginInstallFromMarketplaceApi, PluginInstallFromPkgApi, @@ -41,12 +42,16 @@ from controllers.console.workspace.plugin import ( PluginUploadFromBundleApi, PluginUploadFromGithubApi, PluginUploadFromPkgApi, + _list_hardcoded_builtin_tool_providers, ) from core.plugin.entities.parameters import PluginParameterOption -from core.plugin.entities.plugin import PluginDeclaration, PluginEntity, PluginInstallation +from core.plugin.entities.plugin import PluginCategory, PluginDeclaration, PluginEntity, PluginInstallation from core.plugin.entities.plugin_daemon import PluginInstallTask from core.plugin.impl.exc import PluginDaemonClientSideError from core.plugin.plugin_service import PluginService +from core.tools.entities.api_entities import ToolProviderApiEntity +from core.tools.entities.common_entities import I18nObject +from core.tools.entities.tool_entities import ToolProviderType from models.account import ( Account, TenantAccountRole, @@ -445,7 +450,7 @@ class TestPluginCategoryListApi: mock_list = MagicMock(list=[plugin_item], has_more=True) with ( - app.test_request_context("/?page=2&page_size=10"), + app.test_request_context("/?page=2&page_size=10&query=weather&tags=search&tags=rag&language=zh_Hans"), patch( "controllers.console.workspace.plugin.PluginService.list_by_category", return_value=mock_list ) as list_mock, @@ -456,18 +461,75 @@ class TestPluginCategoryListApi: ): result = method(api, "t1", "tool") - list_mock.assert_called_once() - assert list_mock.call_args.args[0] == "t1" - assert list_mock.call_args.args[1] == "tool" - assert list_mock.call_args.args[2] == 2 - assert list_mock.call_args.args[3] == 10 + list_mock.assert_called_once_with( + "t1", + "tool", + 2, + 10, + query="weather", + tags=["search", "rag"], + language="zh_Hans", + ) assert result["plugins"][0]["id"] == "entity-1" assert result["plugins"][0]["plugin_unique_identifier"] == "test-author/test-plugin:1.0.0@checksum" assert result["builtin_tools"][0]["id"] == "builtin" assert result["builtin_tools"][0]["type"] == "builtin" assert result["has_more"] is True assert "total" not in result - builtin_mock.assert_called_once_with("t1") + builtin_mock.assert_called_once_with( + "t1", + query="weather", + tags=["search", "rag"], + language="zh_Hans", + ) + + def test_builtin_tool_providers_use_the_category_list_filters(self): + search_provider = ToolProviderApiEntity( + id="search-provider", + author="dify", + name="search-provider", + description=I18nObject(en_US="Search provider", zh_Hans="搜索工具"), + icon="icon.svg", + label=I18nObject(en_US="Search", zh_Hans="搜索"), + type=ToolProviderType.BUILT_IN, + labels=["search"], + ) + rag_provider = ToolProviderApiEntity( + id="rag-provider", + author="dify", + name="rag-provider", + description=I18nObject(en_US="RAG provider", zh_Hans="知识库工具"), + icon="icon.svg", + label=I18nObject(en_US="RAG", zh_Hans="知识库"), + type=ToolProviderType.BUILT_IN, + labels=["rag"], + ) + + with ( + patch("controllers.console.workspace.plugin.ToolManager.list_default_builtin_providers", return_value=[]), + patch( + "controllers.console.workspace.plugin.ToolManager.list_hardcoded_providers", + return_value=[MagicMock(), MagicMock()], + ), + patch("controllers.console.workspace.plugin.is_filtered", return_value=False), + patch( + "controllers.console.workspace.plugin.ToolTransformService.builtin_provider_to_user_provider", + side_effect=[search_provider, rag_provider], + ), + patch("controllers.console.workspace.plugin.ToolTransformService.repack_provider"), + patch( + "controllers.console.workspace.plugin.BuiltinToolProviderSort.sort", + side_effect=lambda providers: providers, + ), + ): + result = _list_hardcoded_builtin_tool_providers( + "t1", + query="搜索", + tags=["search", "weather"], + language="zh_Hans", + ) + + assert [provider["id"] for provider in result] == ["search-provider"] def test_non_tool_category_does_not_include_builtin_tools(self, app: Flask): api = PluginCategoryListApi() @@ -739,6 +801,39 @@ class TestPluginListInstallationsFromIdsApi: assert result == ({"code": "plugin_error", "message": "error"}, 400) +class TestPluginInstalledIdsApi: + def test_success(self, app: Flask): + api = PluginInstalledIdsApi() + method = unwrap(api.get) + + with ( + app.test_request_context("/?category=tool"), + patch( + "controllers.console.workspace.plugin.PluginService.list_installed_plugin_ids", + return_value=["langgenius/openai", "langgenius/anthropic"], + ) as list_installed_plugin_ids, + ): + result = method(api, "t1") + + assert result == {"plugin_ids": ["langgenius/openai", "langgenius/anthropic"]} + list_installed_plugin_ids.assert_called_once_with("t1", PluginCategory.Tool) + + def test_daemon_error(self, app: Flask): + api = PluginInstalledIdsApi() + method = unwrap(api.get) + + with ( + app.test_request_context("/?category=tool"), + patch( + "controllers.console.workspace.plugin.PluginService.list_installed_plugin_ids", + side_effect=PluginDaemonClientSideError("error"), + ), + ): + result = method(api, "t1") + + assert result == ({"code": "plugin_error", "message": "error"}, 400) + + class TestPluginUploadFromGithubApi: def test_success(self, app: Flask, user): api = PluginUploadFromGithubApi() diff --git a/api/tests/unit_tests/controllers/test_swagger.py b/api/tests/unit_tests/controllers/test_swagger.py index 3b0bb6e4e63..fbca5ee600c 100644 --- a/api/tests/unit_tests/controllers/test_swagger.py +++ b/api/tests/unit_tests/controllers/test_swagger.py @@ -630,6 +630,11 @@ def test_console_plugin_category_list_exported_schema_uses_typed_items(tmp_path) console_openapi_path = next(path for path in written_paths if path.name == "console-openapi.json") payload = json.loads(console_openapi_path.read_text(encoding="utf-8")) operation = payload["paths"]["/workspaces/current/plugin/{category}/list"]["get"] + parameters = {parameter["name"]: parameter for parameter in operation["parameters"]} + assert parameters["query"]["in"] == "query" + assert parameters["tags"]["in"] == "query" + assert parameters["tags"]["schema"]["type"] == "array" + assert parameters["language"]["in"] == "query" response_ref = operation["responses"]["200"]["content"]["application/json"]["schema"]["$ref"].removeprefix( "#/components/schemas/" ) @@ -657,3 +662,114 @@ def test_console_plugin_category_list_exported_schema_uses_typed_items(tmp_path) builtin_tool_schema = schemas["PluginCategoryBuiltinToolProviderResponse"] for field in ("plugin_unique_identifier", "team_credentials", "type", "tools"): assert field in builtin_tool_schema["properties"] + + +def test_console_installed_plugin_ids_exported_schema_is_lightweight(tmp_path): + from dev.generate_swagger_specs import generate_specs + + written_paths = generate_specs(tmp_path) + console_openapi_path = next(path for path in written_paths if path.name == "console-openapi.json") + payload = json.loads(console_openapi_path.read_text(encoding="utf-8")) + operation = payload["paths"]["/workspaces/current/plugin/installed-ids"]["get"] + parameters = {parameter["name"]: parameter for parameter in operation["parameters"]} + assert parameters["category"]["in"] == "query" + assert parameters["category"]["required"] is True + assert parameters["category"]["schema"]["enum"] == [ + "agent-strategy", + "datasource", + "extension", + "model", + "tool", + "trigger", + ] + response_ref = operation["responses"]["200"]["content"]["application/json"]["schema"]["$ref"].removeprefix( + "#/components/schemas/" + ) + response_schema = payload["components"]["schemas"][response_ref] + + assert response_schema["required"] == ["plugin_ids"] + assert response_schema["properties"] == { + "plugin_ids": { + "items": {"type": "string"}, + "title": "Plugin Ids", + "type": "array", + } + } + + +def test_console_model_provider_summary_exported_schema_is_lightweight(tmp_path): + from dev.generate_swagger_specs import generate_specs + + written_paths = generate_specs(tmp_path) + console_openapi_path = next(path for path in written_paths if path.name == "console-openapi.json") + payload = json.loads(console_openapi_path.read_text(encoding="utf-8")) + operation = payload["paths"]["/workspaces/current/model-providers/summary"]["get"] + assert operation.get("parameters", []) == [] + + response_ref = operation["responses"]["200"]["content"]["application/json"]["schema"]["$ref"].removeprefix( + "#/components/schemas/" + ) + response_schema = payload["components"]["schemas"][response_ref] + assert response_schema["required"] == ["data", "plugins"] + assert response_schema["properties"]["data"]["items"]["$ref"] == ( + "#/components/schemas/ModelProviderSummaryResponse" + ) + assert response_schema["properties"]["plugins"]["additionalProperties"]["$ref"] == ( + "#/components/schemas/ModelProviderPluginSummaryResponse" + ) + + provider_properties = payload["components"]["schemas"]["ModelProviderSummaryResponse"]["properties"] + assert set(provider_properties) == { + "configurate_methods", + "custom_configuration", + "description", + "icon_small", + "icon_small_dark", + "is_configured", + "label", + "plugin_id", + "preferred_provider_type", + "provider", + "supported_model_types", + "system_configuration", + } + assert "provider_credential_schema" not in provider_properties + assert "model_credential_schema" not in provider_properties + + custom_configuration_schema = payload["components"]["schemas"]["ModelProviderCustomConfigurationSummaryResponse"] + custom_configuration_properties = custom_configuration_schema["properties"] + assert set(custom_configuration_schema["required"]) == { + "available_credentials", + "current_credential_usable", + "has_custom_models", + "status", + } + assert set(custom_configuration_properties) == { + "available_credentials", + "current_credential_id", + "current_credential_name", + "current_credential_usable", + "has_custom_models", + "status", + } + assert custom_configuration_properties["available_credentials"]["items"]["$ref"] == ( + "#/components/schemas/CredentialConfiguration" + ) + assert "has_credentials" not in custom_configuration_properties + + credential_properties = payload["components"]["schemas"]["CredentialConfiguration"]["properties"] + assert set(credential_properties) == { + "credential_id", + "credential_name", + } + assert "encrypted_config" not in credential_properties + + plugin_properties = payload["components"]["schemas"]["ModelProviderPluginSummaryResponse"]["properties"] + assert set(plugin_properties) == { + "installation_id", + "plugin_id", + "plugin_unique_identifier", + "runtime_type", + "source", + "version", + } diff --git a/api/tests/unit_tests/core/plugin/impl/test_model_client.py b/api/tests/unit_tests/core/plugin/impl/test_model_client.py index c707b52ccaf..70c26a6cc0e 100644 --- a/api/tests/unit_tests/core/plugin/impl/test_model_client.py +++ b/api/tests/unit_tests/core/plugin/impl/test_model_client.py @@ -28,6 +28,19 @@ class TestPluginModelClient: ) assert request_mock.call_args.kwargs["params"] == {"page": 1, "page_size": 256} + def test_fetch_model_provider_bindings(self, mocker: MockerFixture): + client = PluginModelClient() + request_mock = mocker.patch.object(client, "_request_with_plugin_daemon_response", return_value=["binding-a"]) + + result = client.fetch_model_provider_bindings("tenant-1") + + assert result == ["binding-a"] + assert request_mock.call_args.args[:2] == ( + "GET", + "plugin/tenant-1/management/models/bindings", + ) + assert "params" not in request_mock.call_args.kwargs + def test_get_model_schema(self, mocker: MockerFixture): client = PluginModelClient() schema = SimpleNamespace(name="schema") diff --git a/api/tests/unit_tests/core/plugin/test_plugin_manager.py b/api/tests/unit_tests/core/plugin/test_plugin_manager.py index 1aa019254ff..290c7301bbe 100644 --- a/api/tests/unit_tests/core/plugin/test_plugin_manager.py +++ b/api/tests/unit_tests/core/plugin/test_plugin_manager.py @@ -29,6 +29,7 @@ from core.plugin.entities.plugin import ( ) from core.plugin.entities.plugin_daemon import ( PluginDecodeResponse, + PluginInstalledIdsDaemonResponse, PluginInstallTask, PluginInstallTaskStartResponse, PluginInstallTaskStatus, @@ -132,7 +133,13 @@ class TestPluginDiscovery: plugin_installer, "_request_with_plugin_daemon_response", return_value=mock_response ) as mock_request: result = plugin_installer.list_plugins_by_category( - "test-tenant", category=PluginCategory.Tool, page=2, page_size=10 + "test-tenant", + category=PluginCategory.Tool, + page=2, + page_size=10, + query="weather", + tags=["search", "rag"], + language="zh_Hans", ) mock_request.assert_called_once() @@ -141,6 +148,9 @@ class TestPluginDiscovery: assert call_args.args[2] is PluginListWithoutTotalResponse assert call_args.kwargs["params"]["page"] == 2 assert call_args.kwargs["params"]["page_size"] == 10 + assert call_args.kwargs["params"]["query"] == "weather" + assert call_args.kwargs["params"]["tags"] == ["search", "rag"] + assert call_args.kwargs["params"]["language"] == "zh_Hans" assert result.list == [mock_plugin_entity] assert result.has_more is True @@ -156,6 +166,23 @@ class TestPluginDiscovery: # Assert: Verify empty list is returned assert len(result) == 0 + def test_list_installed_plugin_ids(self, plugin_installer): + """The lightweight ID endpoint is unpaginated and does not request plugin details.""" + mock_response = PluginInstalledIdsDaemonResponse(plugin_ids=["langgenius/openai", "langgenius/anthropic"]) + + with patch.object( + plugin_installer, "_request_with_plugin_daemon_response", return_value=mock_response + ) as mock_request: + result = plugin_installer.list_installed_plugin_ids("test-tenant", PluginCategory.Tool) + + mock_request.assert_called_once_with( + "GET", + "plugin/test-tenant/management/installation/ids", + PluginInstalledIdsDaemonResponse, + params={"category": "tool"}, + ) + assert result == ["langgenius/openai", "langgenius/anthropic"] + def test_fetch_plugin_by_identifier_found(self, plugin_installer): """Test fetching a plugin by its unique identifier when it exists.""" # Arrange: Mock successful fetch diff --git a/api/tests/unit_tests/services/plugin/test_plugin_service.py b/api/tests/unit_tests/services/plugin/test_plugin_service.py index e8c752f2fa7..c82787dc640 100644 --- a/api/tests/unit_tests/services/plugin/test_plugin_service.py +++ b/api/tests/unit_tests/services/plugin/test_plugin_service.py @@ -881,7 +881,105 @@ class TestPluginListEndpointCounts: assert tool_plugin.endpoints_active == 0 +class TestPluginCategoryList: + def test_list_by_category_forwards_search_and_tag_filters(self) -> None: + plugins = SimpleNamespace(list=[], has_more=False) + + with patch(f"{MODULE}.PluginInstaller") as installer_cls: + installer_cls.return_value.list_plugins_by_category.return_value = plugins + + from core.plugin.plugin_service import PluginService + + result = PluginService.list_by_category( + "tenant-1", + PluginCategory.Tool, + 2, + 25, + query="weather", + tags=["search", "rag"], + language="zh_Hans", + ) + + assert result is plugins + installer_cls.return_value.list_plugins_by_category.assert_called_once_with( + "tenant-1", + PluginCategory.Tool, + 2, + 25, + query="weather", + tags=["search", "rag"], + language="zh_Hans", + ) + + def test_filtered_model_category_does_not_reconcile_from_a_partial_result(self) -> None: + plugins = SimpleNamespace(list=[], has_more=False) + + with ( + patch(f"{MODULE}.PluginInstaller") as installer_cls, + patch(f"{MODULE}.PluginService.invalidate_plugin_model_providers_cache") as invalidate_cache, + patch(f"{MODULE}.PluginService._store_cached_remote_model_plugin_marker") as store_marker, + ): + installer_cls.return_value.list_plugins_by_category.return_value = plugins + + from core.plugin.plugin_service import PluginService + + result = PluginService.list_by_category( + "tenant-1", + PluginCategory.Model, + 1, + 100, + query="openai", + tags=[], + language="en_US", + ) + + assert result is plugins + invalidate_cache.assert_not_called() + store_marker.assert_not_called() + + +class TestInstalledPluginIds: + def test_list_installed_plugin_ids_uses_lightweight_daemon_endpoint(self) -> None: + with patch(f"{MODULE}.PluginInstaller") as installer_cls: + installer_cls.return_value.list_installed_plugin_ids.return_value = [ + "langgenius/openai", + "langgenius/anthropic", + ] + + from core.plugin.plugin_service import PluginService + + result = PluginService.list_installed_plugin_ids("tenant-1", PluginCategory.Tool) + + assert result == ["langgenius/openai", "langgenius/anthropic"] + installer_cls.return_value.list_installed_plugin_ids.assert_called_once_with("tenant-1", PluginCategory.Tool) + + class TestPluginModelProviderCacheInvalidation: + def test_list_model_provider_bindings_reconciles_remote_provider_cache(self) -> None: + """The summary binding read owns the remote marker once the full category list leaves the first-load path.""" + remote_binding = _build_remote_model_plugin() + client = MagicMock() + client.fetch_model_provider_bindings.return_value = [remote_binding] + remote_plugin_marker = "langgenius/debug-model:langgenius/debug-model:1.0.0" + + with ( + patch( + f"{MODULE}.PluginService._should_invalidate_model_provider_cache_for_remote_model_plugins", + return_value=True, + ) as should_invalidate, + patch(f"{MODULE}.PluginService.invalidate_plugin_model_providers_cache") as invalidate_cache, + patch(f"{MODULE}.PluginService._store_cached_remote_model_plugin_marker") as store_marker, + ): + from core.plugin.plugin_service import PluginService + + result = PluginService.list_model_provider_bindings("tenant-1", client=client) + + assert result == [remote_binding] + client.fetch_model_provider_bindings.assert_called_once_with("tenant-1") + should_invalidate.assert_called_once_with("tenant-1", [remote_binding]) + invalidate_cache.assert_called_once_with("tenant-1") + store_marker.assert_called_once_with("tenant-1", remote_plugin_marker) + def test_get_debugging_key_does_not_invalidate_model_provider_cache(self) -> None: """Reading a debug key does not mean a debug runtime has registered a model provider.""" with ( @@ -925,7 +1023,13 @@ class TestPluginModelProviderCacheInvalidation: assert result is plugins installer_cls.return_value.list_plugins_by_category.assert_called_once_with( - "tenant-1", PluginCategory.Model, 1, 100 + "tenant-1", + PluginCategory.Model, + 1, + 100, + query="", + tags=(), + language="en_US", ) invalidate_cache.assert_called_once_with("tenant-1") store_marker.assert_called_once_with("tenant-1", remote_plugin_marker) @@ -992,14 +1096,42 @@ class TestPluginModelProviderCacheInvalidation: invalidate_cache.assert_not_called() store_marker.assert_called_once_with("tenant-1", remote_plugin_marker) - def test_list_model_category_invalidates_when_remote_model_plugin_disconnects(self) -> None: - """The current model category result clears provider cache when the previous debug model disappears.""" + @pytest.mark.parametrize(("page", "has_more"), [(1, True), (2, False)]) + def test_list_model_category_does_not_reconcile_partial_page(self, page: int, has_more: bool) -> None: + """Only an unfiltered, complete first page may write the remote model marker.""" installed_plugin = SimpleNamespace( plugin_id="langgenius/openai", plugin_unique_identifier="langgenius/openai:1.0.0", source=PluginInstallationSource.Marketplace, ) - plugins = SimpleNamespace(list=[installed_plugin], has_more=True) + plugins = SimpleNamespace(list=[installed_plugin], has_more=has_more) + + with ( + patch(f"{MODULE}.PluginInstaller") as installer_cls, + patch( + f"{MODULE}.PluginService._load_cached_remote_model_plugin_marker", + return_value="langgenius/debug-model:langgenius/debug-model:1.0.0", + ), + patch(f"{MODULE}.PluginService.invalidate_plugin_model_providers_cache") as invalidate_cache, + patch(f"{MODULE}.PluginService._store_cached_remote_model_plugin_marker") as store_marker, + ): + installer_cls.return_value.list_plugins_by_category.return_value = plugins + + from core.plugin.plugin_service import PluginService + + result = PluginService.list_by_category("tenant-1", PluginCategory.Model, page, 100) + + assert result is plugins + invalidate_cache.assert_not_called() + store_marker.assert_not_called() + + def test_list_model_category_complete_first_page_reconciles_remote_plugin_disconnect(self) -> None: + installed_plugin = SimpleNamespace( + plugin_id="langgenius/openai", + plugin_unique_identifier="langgenius/openai:1.0.0", + source=PluginInstallationSource.Marketplace, + ) + plugins = SimpleNamespace(list=[installed_plugin], has_more=False) with ( patch(f"{MODULE}.PluginInstaller") as installer_cls, diff --git a/api/tests/unit_tests/services/test_model_provider_service.py b/api/tests/unit_tests/services/test_model_provider_service.py index 12d9404ff79..3b57e46a7f9 100644 --- a/api/tests/unit_tests/services/test_model_provider_service.py +++ b/api/tests/unit_tests/services/test_model_provider_service.py @@ -5,12 +5,16 @@ from unittest.mock import MagicMock import pytest from core.entities.model_entities import ModelStatus +from core.entities.provider_entities import CredentialConfiguration +from core.plugin.entities.plugin import PluginInstallationSource +from core.plugin.entities.plugin_daemon import PluginModelProviderBinding from graphon.model_runtime.entities.common_entities import I18nObject from graphon.model_runtime.entities.model_entities import FetchFrom, ModelType, ParameterRule, ParameterType +from graphon.model_runtime.entities.provider_entities import ConfigurateMethod from models.provider import ProviderType from services import model_provider_service as service_module from services.errors.app_model_config import ProviderNotFoundError -from services.model_provider_service import ModelProviderService +from services.model_provider_service import ModelProviderService, _ProviderSummaryState def _create_service_with_mocked_manager() -> tuple[ModelProviderService, MagicMock]: @@ -59,6 +63,26 @@ def _build_provider_configuration( ) +def _build_model_provider_binding( + source: PluginInstallationSource, + *, + installation_id: str = "installation-1", + plugin_id: str = "langgenius/openai", + plugin_unique_identifier: str = "langgenius/openai:1.0.0@checksum", + verified: bool = True, +) -> PluginModelProviderBinding: + return PluginModelProviderBinding( + provider="openai", + installation_id=installation_id, + plugin_id=plugin_id, + plugin_unique_identifier=plugin_unique_identifier, + runtime_type="remote" if source == PluginInstallationSource.Remote else "local", + source=source, + version="1.0.0", + verified=verified, + ) + + class TestModelProviderServiceConfiguration: def test__get_provider_configuration_should_return_configuration_when_provider_exists(self) -> None: service, manager = _create_service_with_mocked_manager() @@ -96,6 +120,352 @@ class TestModelProviderServiceConfiguration: assert result[0].provider == "openai" assert result[0].custom_configuration.status.value == "no-configure" + def test_get_provider_summary_list_uses_lightweight_state_and_plugin_bindings( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + service = ModelProviderService() + provider = SimpleNamespace( + provider="langgenius/openai/openai", + label=I18nObject(en_US="OpenAI"), + description=I18nObject(en_US="OpenAI models"), + icon_small=I18nObject(en_US="icon.svg"), + icon_small_dark=I18nObject(en_US="icon-dark.svg"), + supported_model_types=[ModelType.LLM], + configurate_methods=[ConfigurateMethod.PREDEFINED_MODEL], + ) + binding = SimpleNamespace( + provider="openai", + plugin_id="langgenius/openai", + installation_id="installation-1", + plugin_unique_identifier="langgenius/openai:1.2.3@checksum", + runtime_type="local", + source=PluginInstallationSource.Marketplace, + version="1.2.3", + verified=True, + ) + state = _ProviderSummaryState( + has_custom_provider=True, + available_credentials=[ + CredentialConfiguration( + credential_id="credential-1", + credential_name="Production", + ), + CredentialConfiguration( + credential_id="credential-2", + credential_name="Backup", + ), + ], + has_custom_models=True, + current_credential_id="credential-1", + current_credential_name="Production", + current_credential_usable=True, + preferred_provider_type=ProviderType.CUSTOM, + ) + call_order: list[str] = [] + manager_constructor = MagicMock(side_effect=AssertionError("summary must not construct ProviderManager")) + monkeypatch.setattr(service, "_get_provider_manager", manager_constructor) + monkeypatch.setattr( + service_module.PluginService, + "list_model_provider_bindings", + MagicMock(side_effect=lambda *_args, **_kwargs: call_order.append("bindings") or [binding]), + ) + monkeypatch.setattr( + service_module.PluginService, + "fetch_plugin_model_providers", + MagicMock(side_effect=lambda *_args, **_kwargs: call_order.append("providers") or [provider]), + ) + monkeypatch.setattr( + service, + "_load_provider_summary_states", + MagicMock(return_value={provider.provider: state}), + ) + monkeypatch.setattr( + service_module.ext_hosting_provider.hosting_configuration, + "provider_map", + {provider.provider: SimpleNamespace(enabled=True, quotas=[SimpleNamespace()])}, + ) + + providers, plugins = service.get_provider_summary_list(tenant_id="tenant-1") + + assert len(providers) == 1 + assert providers[0].provider == provider.provider + assert providers[0].plugin_id == "langgenius/openai" + assert providers[0].is_configured is True + assert providers[0].custom_configuration.available_credentials == [ + CredentialConfiguration( + credential_id="credential-1", + credential_name="Production", + ), + CredentialConfiguration( + credential_id="credential-2", + credential_name="Backup", + ), + ] + assert providers[0].custom_configuration.has_custom_models is True + assert providers[0].custom_configuration.current_credential_name == "Production" + assert providers[0].custom_configuration.current_credential_usable is True + assert providers[0].system_configuration.enabled is True + assert plugins["langgenius/openai"].installation_id == "installation-1" + assert plugins["langgenius/openai"].version == "1.2.3" + assert call_order == ["bindings", "providers"] + manager_constructor.assert_not_called() + + def test_get_provider_summary_list_enables_system_only_for_verified_hosted_non_package_bindings( + self, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + marketplace_binding = _build_model_provider_binding(PluginInstallationSource.Marketplace) + package_binding = _build_model_provider_binding( + PluginInstallationSource.Package, + installation_id="installation-package", + plugin_id="langgenius/package", + plugin_unique_identifier="langgenius/package:1.0.0@checksum", + ) + unhosted_binding = _build_model_provider_binding( + PluginInstallationSource.Marketplace, + installation_id="installation-unhosted", + plugin_id="langgenius/unhosted", + plugin_unique_identifier="langgenius/unhosted:1.0.0@checksum", + ) + remote_binding = _build_model_provider_binding( + PluginInstallationSource.Remote, + installation_id="installation-remote", + plugin_id="langgenius/remote", + plugin_unique_identifier="langgenius/remote:1.0.0@checksum", + ) + unverified_binding = _build_model_provider_binding( + PluginInstallationSource.Marketplace, + installation_id="installation-unverified", + plugin_id="langgenius/unverified", + plugin_unique_identifier="langgenius/unverified:1.0.0@checksum", + verified=False, + ) + bindings = [ + marketplace_binding, + package_binding, + unhosted_binding, + remote_binding, + unverified_binding, + ] + provider_entities = [ + SimpleNamespace( + provider=f"{binding.plugin_id}/openai", + label=I18nObject(en_US=binding.plugin_id), + description=None, + icon_small=None, + icon_small_dark=None, + supported_model_types=[ModelType.LLM], + configurate_methods=[], + ) + for binding in bindings + ] + monkeypatch.setattr( + service_module.ext_hosting_provider.hosting_configuration, + "provider_map", + { + "langgenius/openai/openai": SimpleNamespace(enabled=True, quotas=[SimpleNamespace()]), + "langgenius/package/openai": SimpleNamespace(enabled=True, quotas=[SimpleNamespace()]), + "langgenius/remote/openai": SimpleNamespace(enabled=True, quotas=[SimpleNamespace()]), + "langgenius/unverified/openai": SimpleNamespace(enabled=True, quotas=[SimpleNamespace()]), + }, + ) + monkeypatch.setattr( + service_module.PluginService, + "list_model_provider_bindings", + MagicMock(return_value=bindings), + ) + monkeypatch.setattr( + service_module.PluginService, + "fetch_plugin_model_providers", + MagicMock(return_value=provider_entities), + ) + monkeypatch.setattr(ModelProviderService, "_load_provider_summary_states", MagicMock(return_value={})) + monkeypatch.setattr(service_module, "is_filtered", MagicMock(return_value=False)) + + providers, _ = ModelProviderService().get_provider_summary_list("tenant-1") + + assert {provider.provider: provider.system_configuration.enabled for provider in providers} == { + "langgenius/openai/openai": True, + "langgenius/package/openai": False, + "langgenius/unhosted/openai": False, + "langgenius/remote/openai": True, + "langgenius/unverified/openai": False, + } + + def test_model_provider_binding_without_verified_field_fails_closed(self) -> None: + binding = PluginModelProviderBinding.model_validate( + { + "provider": "openai", + "installation_id": "installation-1", + "plugin_id": "langgenius/openai", + "plugin_unique_identifier": "langgenius/openai:1.0.0@checksum", + "runtime_type": "local", + "source": "marketplace", + "version": "1.0.0", + } + ) + + assert binding.verified is False + + def test_get_provider_summary_list_returns_all_unique_provider_metadata( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + service = ModelProviderService() + llm_provider = SimpleNamespace( + provider="langgenius/openai/openai", + label=I18nObject(en_US="OpenAI"), + description=None, + icon_small=None, + icon_small_dark=None, + supported_model_types=[ModelType.LLM], + configurate_methods=[], + ) + embedding_provider = SimpleNamespace( + provider="langgenius/embedding/embedding", + label=I18nObject(en_US="Embedding"), + description=None, + icon_small=None, + icon_small_dark=None, + supported_model_types=[ModelType.TEXT_EMBEDDING], + configurate_methods=[], + ) + llm_binding = SimpleNamespace( + provider="openai", + plugin_id="langgenius/openai", + installation_id="installation-openai", + plugin_unique_identifier="langgenius/openai:1.0.0@checksum", + runtime_type="local", + source=service_module.PluginInstallationSource.Marketplace, + version="1.0.0", + verified=False, + ) + embedding_binding = SimpleNamespace( + provider="embedding", + plugin_id="langgenius/embedding", + installation_id="installation-embedding", + plugin_unique_identifier="langgenius/embedding:1.0.0@checksum", + runtime_type="local", + source=service_module.PluginInstallationSource.Marketplace, + version="1.0.0", + verified=False, + ) + monkeypatch.setattr( + service_module.PluginService, + "list_model_provider_bindings", + MagicMock(return_value=[llm_binding, embedding_binding]), + ) + monkeypatch.setattr( + service_module.PluginService, + "fetch_plugin_model_providers", + MagicMock(return_value=[llm_provider, llm_provider, embedding_provider]), + ) + monkeypatch.setattr(service, "_load_provider_summary_states", MagicMock(return_value={})) + monkeypatch.setattr(service_module, "is_filtered", MagicMock(return_value=False)) + + providers, plugins = service.get_provider_summary_list(tenant_id="tenant-1") + + assert [provider.provider for provider in providers] == [ + "langgenius/openai/openai", + "langgenius/embedding/embedding", + ] + assert providers[0].is_configured is False + assert providers[0].custom_configuration.status.value == "no-configure" + assert providers[0].custom_configuration.has_custom_models is False + assert providers[0].custom_configuration.available_credentials == [] + assert set(plugins) == {"langgenius/openai", "langgenius/embedding"} + + def test_preferred_provider_fallback_uses_custom_presence_not_configuration_status( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + monkeypatch.setattr(service_module.dify_config, "EDITION", "SELF_HOSTED") + state = _ProviderSummaryState(has_custom_provider=True) + + preferred_provider_type = ModelProviderService._get_preferred_provider_type( + state, + custom_present=True, + system_enabled=True, + ) + + assert preferred_provider_type == ProviderType.CUSTOM + + def test_load_provider_summary_states_reads_only_lightweight_columns(self, monkeypatch: pytest.MonkeyPatch) -> None: + canonical_provider = "langgenius/openai/openai" + session = MagicMock() + session.execute.side_effect = [ + SimpleNamespace( + all=lambda: [ + SimpleNamespace( + provider_name="openai", + credential_id="credential-legacy", + credential_provider_name="openai", + credential_name="Legacy", + ), + SimpleNamespace( + provider_name=canonical_provider, + credential_id="credential-current", + credential_provider_name=canonical_provider, + credential_name="Production", + ), + ] + ), + SimpleNamespace( + all=lambda: [ + SimpleNamespace( + id="credential-legacy", + provider_name="openai", + credential_name="Legacy", + ), + SimpleNamespace( + id="credential-current", + provider_name=canonical_provider, + credential_name="Production", + ), + ] + ), + SimpleNamespace(all=lambda: [SimpleNamespace(provider_name="openai")]), + SimpleNamespace( + all=lambda: [ + SimpleNamespace( + provider_name=canonical_provider, + preferred_provider_type=ProviderType.SYSTEM, + ) + ] + ), + ] + session_context = MagicMock() + session_context.__enter__.return_value = session + create_session = MagicMock(return_value=session_context) + monkeypatch.setattr(service_module.session_factory, "create_session", create_session) + + states = ModelProviderService._load_provider_summary_states("tenant-1") + + state = states[canonical_provider] + assert state.has_custom_provider is True + assert state.available_credentials == [ + CredentialConfiguration( + credential_id="credential-legacy", + credential_name="Legacy", + ), + CredentialConfiguration( + credential_id="credential-current", + credential_name="Production", + ), + ] + assert state.has_custom_models is True + assert state.current_credential_id == "credential-current" + assert state.current_credential_name == "Production" + assert state.current_credential_usable is True + assert state.preferred_provider_type == ProviderType.SYSTEM + + statements = [str(execute_call.args[0]) for execute_call in session.execute.call_args_list] + assert len(statements) == 4 + assert all("encrypted_config" not in statement for statement in statements) + assert "count(" not in statements[1].lower() + assert "provider_credentials.id" in statements[1] + assert "provider_credentials.credential_name" in statements[1] + assert "ORDER BY provider_credentials.created_at DESC, provider_credentials.id DESC" in statements[1] + assert "provider_model_credentials" in statements[2] + def test_get_models_by_provider_should_wrap_model_entities_with_tenant_context(self) -> None: service, manager = _create_service_with_mocked_manager() diff --git a/api/tests/unit_tests/services/test_workspace_credit_pool.py b/api/tests/unit_tests/services/test_workspace_credit_pool.py new file mode 100644 index 00000000000..dc3497a1a05 --- /dev/null +++ b/api/tests/unit_tests/services/test_workspace_credit_pool.py @@ -0,0 +1,76 @@ +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import pytest + +from enums.cloud_plan import CloudPlan +from services.credit_pool_service import CreditPoolBalance +from services.workspace_service import WorkspaceService + + +@pytest.mark.parametrize( + ("quota_limit", "quota_used", "remaining_credits", "is_unlimited"), + [(500, 120, 380, False), (-1, 999, -1, True)], +) +def test_get_effective_credit_pool_prefers_available_paid_pool( + quota_limit: int, quota_used: int, remaining_credits: int, is_unlimited: bool +) -> None: + session = MagicMock() + paid_pool = CreditPoolBalance( + tenant_id="tenant-1", + pool_type="paid", + quota_limit=quota_limit, + quota_used=quota_used, + ) + billing_info = { + "enabled": True, + "subscription": {"plan": CloudPlan.TEAM}, + "next_credit_reset_date": 1775001600, + } + config = SimpleNamespace(BILLING_ENABLED=True) + + with ( + patch("services.workspace_service.dify_config", config), + patch("services.workspace_service.BillingService.get_info", return_value=billing_info), + patch("services.credit_pool_service.CreditPoolService.get_pool", return_value=paid_pool) as get_pool, + ): + result = WorkspaceService.get_effective_credit_pool("tenant-1", session=session) + + get_pool.assert_called_once_with(tenant_id="tenant-1", pool_type="paid", session=session) + assert result.pool_type == "paid" + assert result.quota_limit == quota_limit + assert result.quota_used == quota_used + assert result.remaining_credits == remaining_credits + assert result.is_unlimited is is_unlimited + assert result.is_exhausted is False + assert result.next_credit_reset_date == 1775001600 + + +def test_get_effective_credit_pool_exposes_exhausted_trial_pool() -> None: + session = MagicMock() + trial_pool = CreditPoolBalance( + tenant_id="tenant-1", + pool_type="trial", + quota_limit=200, + quota_used=200, + exhausted_at=1772323200, + ) + billing_info = { + "enabled": True, + "subscription": {"plan": CloudPlan.SANDBOX}, + } + config = SimpleNamespace(BILLING_ENABLED=True) + + with ( + patch("services.workspace_service.dify_config", config), + patch("services.workspace_service.BillingService.get_info", return_value=billing_info), + patch("services.credit_pool_service.CreditPoolService.get_pool", return_value=trial_pool) as get_pool, + ): + result = WorkspaceService.get_effective_credit_pool("tenant-1", session=session) + + get_pool.assert_called_once_with(tenant_id="tenant-1", pool_type="trial", session=session) + assert result.pool_type == "trial" + assert result.remaining_credits == 0 + assert result.is_unlimited is False + assert result.is_exhausted is True + assert result.exhausted_at == 1772323200 diff --git a/oxlint-suppressions.json b/oxlint-suppressions.json index 9006044d536..c740319be0e 100644 --- a/oxlint-suppressions.json +++ b/oxlint-suppressions.json @@ -2891,14 +2891,6 @@ "count": 4 } }, - "web/app/components/header/account-setting/model-provider-page/hooks.ts": { - "eslint-react/no-unnecessary-use-prefix": { - "count": 1 - }, - "typescript/no-explicit-any": { - "count": 2 - } - }, "web/app/components/header/account-setting/model-provider-page/model-auth/add-custom-model.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 2 @@ -2923,14 +2915,6 @@ "count": 1 } }, - "web/app/components/header/account-setting/model-provider-page/model-auth/config-model.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/header/account-setting/model-provider-page/model-auth/credential-selector.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 1 @@ -2974,11 +2958,6 @@ "count": 5 } }, - "web/app/components/header/account-setting/model-provider-page/model-parameter-modal/configuration-button.tsx": { - "typescript/no-explicit-any": { - "count": 2 - } - }, "web/app/components/header/account-setting/model-provider-page/model-parameter-modal/index.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 1 @@ -3368,14 +3347,6 @@ "count": 25 } }, - "web/app/components/plugins/update-plugin/plugin-version-picker.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/rag-pipeline/components/__tests__/publish-as-knowledge-pipeline-modal.spec.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 1 diff --git a/packages/contracts/generated/api/console/workspaces/orpc.gen.ts b/packages/contracts/generated/api/console/workspaces/orpc.gen.ts index 730d89367a7..2d80fb7a4f1 100644 --- a/packages/contracts/generated/api/console/workspaces/orpc.gen.ts +++ b/packages/contracts/generated/api/console/workspaces/orpc.gen.ts @@ -69,8 +69,10 @@ import { zGetWorkspacesCurrentModelProvidersByProviderModelsParameterRulesResponse, zGetWorkspacesCurrentModelProvidersByProviderModelsPath, zGetWorkspacesCurrentModelProvidersByProviderModelsResponse, + zGetWorkspacesCurrentModelProvidersCreditsResponse, zGetWorkspacesCurrentModelProvidersQuery, zGetWorkspacesCurrentModelProvidersResponse, + zGetWorkspacesCurrentModelProvidersSummaryResponse, zGetWorkspacesCurrentModelsModelTypesByModelTypePath, zGetWorkspacesCurrentModelsModelTypesByModelTypeResponse, zGetWorkspacesCurrentPermissionResponse, @@ -86,6 +88,8 @@ import { zGetWorkspacesCurrentPluginFetchManifestResponse, zGetWorkspacesCurrentPluginIconQuery, zGetWorkspacesCurrentPluginIconResponse, + zGetWorkspacesCurrentPluginInstalledIdsQuery, + zGetWorkspacesCurrentPluginInstalledIdsResponse, zGetWorkspacesCurrentPluginListQuery, zGetWorkspacesCurrentPluginListResponse, zGetWorkspacesCurrentPluginMarketplacePkgQuery, @@ -1066,6 +1070,34 @@ export const members = { } export const get12 = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'getWorkspacesCurrentModelProvidersCredits', + path: '/workspaces/current/model-providers/credits', + tags: ['console'], + }) + .output(zGetWorkspacesCurrentModelProvidersCreditsResponse) + +export const credits = { + get: get12, +} + +export const get13 = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'getWorkspacesCurrentModelProvidersSummary', + path: '/workspaces/current/model-providers/summary', + tags: ['console'], + }) + .output(zGetWorkspacesCurrentModelProvidersSummaryResponse) + +export const summary = { + get: get13, +} + +export const get14 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1077,7 +1109,7 @@ export const get12 = oc .output(zGetWorkspacesCurrentModelProvidersByProviderCheckoutUrlResponse) export const checkoutUrl = { - get: get12, + get: get14, } export const post16 = oc @@ -1137,7 +1169,7 @@ export const delete5 = oc ) .output(zDeleteWorkspacesCurrentModelProvidersByProviderCredentialsResponse) -export const get13 = oc +export const get15 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1188,7 +1220,7 @@ export const put2 = oc export const credentials = { delete: delete5, - get: get13, + get: get15, post: post18, put: put2, switch: switch_, @@ -1252,7 +1284,7 @@ export const delete6 = oc ) .output(zDeleteWorkspacesCurrentModelProvidersByProviderModelsCredentialsResponse) -export const get14 = oc +export const get16 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1303,7 +1335,7 @@ export const put3 = oc export const credentials2 = { delete: delete6, - get: get14, + get: get16, post: post21, put: put3, switch: switch2, @@ -1407,7 +1439,7 @@ export const loadBalancingConfigs = { byConfigId, } -export const get15 = oc +export const get17 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1424,7 +1456,7 @@ export const get15 = oc .output(zGetWorkspacesCurrentModelProvidersByProviderModelsParameterRulesResponse) export const parameterRules = { - get: get15, + get: get17, } export const delete7 = oc @@ -1444,7 +1476,7 @@ export const delete7 = oc ) .output(zDeleteWorkspacesCurrentModelProvidersByProviderModelsResponse) -export const get16 = oc +export const get18 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1473,7 +1505,7 @@ export const post24 = oc export const models = { delete: delete7, - get: get16, + get: get18, post: post24, credentials: credentials2, disable: disable2, @@ -1509,7 +1541,7 @@ export const byProvider = { preferredProviderType, } -export const get17 = oc +export const get19 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1521,11 +1553,13 @@ export const get17 = oc .output(zGetWorkspacesCurrentModelProvidersResponse) export const modelProviders = { - get: get17, + get: get19, + credits, + summary, byProvider, } -export const get18 = oc +export const get20 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1537,7 +1571,7 @@ export const get18 = oc .output(zGetWorkspacesCurrentModelsModelTypesByModelTypeResponse) export const byModelType = { - get: get18, + get: get20, } export const modelTypes = { @@ -1553,7 +1587,7 @@ export const models2 = { * * Returns permission flags that control workspace features like member invitations and owner transfer. */ -export const get19 = oc +export const get21 = oc .route({ description: 'Returns permission flags that control workspace features like member invitations and owner transfer.', @@ -1567,10 +1601,10 @@ export const get19 = oc .output(zGetWorkspacesCurrentPermissionResponse) export const permission = { - get: get19, + get: get21, } -export const get20 = oc +export const get22 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1582,7 +1616,7 @@ export const get20 = oc .output(zGetWorkspacesCurrentPluginAssetResponse) export const asset = { - get: get20, + get: get22, } export const post26 = oc @@ -1615,7 +1649,7 @@ export const exclude = { post: post27, } -export const get21 = oc +export const get23 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1627,7 +1661,7 @@ export const get21 = oc .output(zGetWorkspacesCurrentPluginAutoUpgradeFetchResponse) export const fetch_ = { - get: get21, + get: get23, } export const autoUpgrade = { @@ -1636,7 +1670,7 @@ export const autoUpgrade = { fetch: fetch_, } -export const get22 = oc +export const get24 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1647,10 +1681,10 @@ export const get22 = oc .output(zGetWorkspacesCurrentPluginDebuggingKeyResponse) export const debuggingKey = { - get: get22, + get: get24, } -export const get23 = oc +export const get25 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1662,10 +1696,10 @@ export const get23 = oc .output(zGetWorkspacesCurrentPluginFetchManifestResponse) export const fetchManifest = { - get: get23, + get: get25, } -export const get24 = oc +export const get26 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1677,7 +1711,7 @@ export const get24 = oc .output(zGetWorkspacesCurrentPluginIconResponse) export const icon = { - get: get24, + get: get26, } export const post28 = oc @@ -1731,6 +1765,21 @@ export const install = { pkg, } +export const get27 = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'getWorkspacesCurrentPluginInstalledIds', + path: '/workspaces/current/plugin/installed-ids', + tags: ['console'], + }) + .input(z.object({ query: zGetWorkspacesCurrentPluginInstalledIdsQuery })) + .output(zGetWorkspacesCurrentPluginInstalledIdsResponse) + +export const installedIds = { + get: get27, +} + export const post31 = oc .route({ inputStructure: 'detailed', @@ -1765,7 +1814,7 @@ export const latestVersions = { post: post32, } -export const get25 = oc +export const get28 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1777,12 +1826,12 @@ export const get25 = oc .output(zGetWorkspacesCurrentPluginListResponse) export const list2 = { - get: get25, + get: get28, installations, latestVersions, } -export const get26 = oc +export const get29 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1794,14 +1843,14 @@ export const get26 = oc .output(zGetWorkspacesCurrentPluginMarketplacePkgResponse) export const pkg2 = { - get: get26, + get: get29, } export const marketplace2 = { pkg: pkg2, } -export const get27 = oc +export const get30 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1813,7 +1862,7 @@ export const get27 = oc .output(zGetWorkspacesCurrentPluginParametersDynamicOptionsResponse) export const dynamicOptions = { - get: get27, + get: get30, } /** @@ -1857,7 +1906,7 @@ export const change2 = { post: post34, } -export const get28 = oc +export const get31 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1868,7 +1917,7 @@ export const get28 = oc .output(zGetWorkspacesCurrentPluginPermissionFetchResponse) export const fetch2 = { - get: get28, + get: get31, } export const permission2 = { @@ -1876,7 +1925,7 @@ export const permission2 = { fetch: fetch2, } -export const get29 = oc +export const get32 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1888,7 +1937,7 @@ export const get29 = oc .output(zGetWorkspacesCurrentPluginReadmeResponse) export const readme = { - get: get29, + get: get32, } export const post35 = oc @@ -1936,7 +1985,7 @@ export const delete8 = { byIdentifier, } -export const get30 = oc +export const get33 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1948,11 +1997,11 @@ export const get30 = oc .output(zGetWorkspacesCurrentPluginTasksByTaskIdResponse) export const byTaskId = { - get: get30, + get: get33, delete: delete8, } -export const get31 = oc +export const get34 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1964,7 +2013,7 @@ export const get31 = oc .output(zGetWorkspacesCurrentPluginTasksResponse) export const tasks = { - get: get31, + get: get34, deleteAll, byTaskId, } @@ -2069,7 +2118,7 @@ export const upload = { pkg: pkg3, } -export const get32 = oc +export const get35 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2086,7 +2135,7 @@ export const get32 = oc .output(zGetWorkspacesCurrentPluginByCategoryListResponse) export const list3 = { - get: get32, + get: get35, } export const byCategory = { @@ -2100,6 +2149,7 @@ export const plugin2 = { fetchManifest, icon, install, + installedIds, list: list2, marketplace: marketplace2, parameters, @@ -2139,7 +2189,7 @@ export const delete9 = oc .input(z.object({ params: zDeleteWorkspacesCurrentRbacAccessPoliciesByPolicyIdPath })) .output(zDeleteWorkspacesCurrentRbacAccessPoliciesByPolicyIdResponse) -export const get33 = oc +export const get36 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2163,12 +2213,12 @@ export const put4 = oc export const byPolicyId = { delete: delete9, - get: get33, + get: get36, put: put4, copy, } -export const get34 = oc +export const get37 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2190,7 +2240,7 @@ export const post45 = oc .output(zPostWorkspacesCurrentRbacAccessPoliciesResponse) export const accessPolicies = { - get: get34, + get: get37, post: post45, byPolicyId, } @@ -2250,7 +2300,7 @@ export const delete10 = oc ) .output(zDeleteWorkspacesCurrentRbacAppsByAppIdAccessPoliciesByPolicyIdMemberBindingsResponse) -export const get35 = oc +export const get38 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2267,10 +2317,10 @@ export const get35 = oc export const memberBindings = { delete: delete10, - get: get35, + get: get38, } -export const get36 = oc +export const get39 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2286,7 +2336,7 @@ export const get36 = oc .output(zGetWorkspacesCurrentRbacAppsByAppIdAccessPoliciesByPolicyIdRoleBindingsResponse) export const roleBindings = { - get: get36, + get: get39, } export const byPolicyId2 = { @@ -2298,7 +2348,7 @@ export const accessPolicies2 = { byPolicyId: byPolicyId2, } -export const get37 = oc +export const get40 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2315,10 +2365,10 @@ export const get37 = oc .output(zGetWorkspacesCurrentRbacAppsByAppIdAccessPolicyResponse) export const accessPolicy = { - get: get37, + get: get40, } -export const get38 = oc +export const get41 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2335,7 +2385,7 @@ export const get38 = oc .output(zGetWorkspacesCurrentRbacAppsByAppIdUserAccessPoliciesResponse) export const userAccessPolicies = { - get: get38, + get: get41, } export const put7 = oc @@ -2366,7 +2416,7 @@ export const users = { byTargetAccountId, } -export const get39 = oc +export const get42 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2394,7 +2444,7 @@ export const put8 = oc .output(zPutWorkspacesCurrentRbacAppsByAppIdWhitelistResponse) export const whitelist = { - get: get39, + get: get42, put: put8, } @@ -2430,7 +2480,7 @@ export const delete11 = oc zDeleteWorkspacesCurrentRbacDatasetsByDatasetIdAccessPoliciesByPolicyIdMemberBindingsResponse, ) -export const get40 = oc +export const get43 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2451,10 +2501,10 @@ export const get40 = oc export const memberBindings2 = { delete: delete11, - get: get40, + get: get43, } -export const get41 = oc +export const get44 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2470,7 +2520,7 @@ export const get41 = oc .output(zGetWorkspacesCurrentRbacDatasetsByDatasetIdAccessPoliciesByPolicyIdRoleBindingsResponse) export const roleBindings2 = { - get: get41, + get: get44, } export const byPolicyId3 = { @@ -2482,7 +2532,7 @@ export const accessPolicies4 = { byPolicyId: byPolicyId3, } -export const get42 = oc +export const get45 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2499,10 +2549,10 @@ export const get42 = oc .output(zGetWorkspacesCurrentRbacDatasetsByDatasetIdAccessPolicyResponse) export const accessPolicy2 = { - get: get42, + get: get45, } -export const get43 = oc +export const get46 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2519,7 +2569,7 @@ export const get43 = oc .output(zGetWorkspacesCurrentRbacDatasetsByDatasetIdUserAccessPoliciesResponse) export const userAccessPolicies2 = { - get: get43, + get: get46, } export const put9 = oc @@ -2550,7 +2600,7 @@ export const users2 = { byTargetAccountId: byTargetAccountId2, } -export const get44 = oc +export const get47 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2578,7 +2628,7 @@ export const put10 = oc .output(zPutWorkspacesCurrentRbacDatasetsByDatasetIdWhitelistResponse) export const whitelist2 = { - get: get44, + get: get47, put: put10, } @@ -2594,7 +2644,7 @@ export const datasets = { byDatasetId, } -export const get45 = oc +export const get48 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2622,7 +2672,7 @@ export const put11 = oc .output(zPutWorkspacesCurrentRbacMembersByMemberIdRbacRolesResponse) export const rbacRoles = { - get: get45, + get: get48, put: put11, } @@ -2634,7 +2684,7 @@ export const members2 = { byMemberId: byMemberId2, } -export const get46 = oc +export const get49 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2645,10 +2695,10 @@ export const get46 = oc .output(zGetWorkspacesCurrentRbacMyPermissionsResponse) export const myPermissions = { - get: get46, + get: get49, } -export const get47 = oc +export const get50 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2659,10 +2709,10 @@ export const get47 = oc .output(zGetWorkspacesCurrentRbacRolePermissionsCatalogAppResponse) export const app = { - get: get47, + get: get50, } -export const get48 = oc +export const get51 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2673,10 +2723,10 @@ export const get48 = oc .output(zGetWorkspacesCurrentRbacRolePermissionsCatalogDatasetResponse) export const dataset = { - get: get48, + get: get51, } -export const get49 = oc +export const get52 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2687,7 +2737,7 @@ export const get49 = oc .output(zGetWorkspacesCurrentRbacRolePermissionsCatalogResponse) export const catalog = { - get: get49, + get: get52, app, dataset, } @@ -2712,7 +2762,7 @@ export const copy2 = { post: post46, } -export const get50 = oc +export const get53 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2724,7 +2774,7 @@ export const get50 = oc .output(zGetWorkspacesCurrentRbacRolesByRoleIdMembersResponse) export const members3 = { - get: get50, + get: get53, } export const delete12 = oc @@ -2738,7 +2788,7 @@ export const delete12 = oc .input(z.object({ params: zDeleteWorkspacesCurrentRbacRolesByRoleIdPath })) .output(zDeleteWorkspacesCurrentRbacRolesByRoleIdResponse) -export const get51 = oc +export const get54 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2762,13 +2812,13 @@ export const put12 = oc export const byRoleId = { delete: delete12, - get: get51, + get: get54, put: put12, copy: copy2, members: members3, } -export const get52 = oc +export const get55 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2790,7 +2840,7 @@ export const post47 = oc .output(zPostWorkspacesCurrentRbacRolesResponse) export const roles = { - get: get52, + get: get55, post: post47, byRoleId, } @@ -2815,7 +2865,7 @@ export const bindings = { put: put13, } -export const get53 = oc +export const get56 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2831,10 +2881,10 @@ export const get53 = oc .output(zGetWorkspacesCurrentRbacWorkspaceAppsAccessPoliciesByPolicyIdMemberBindingsResponse) export const memberBindings3 = { - get: get53, + get: get56, } -export const get54 = oc +export const get57 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2850,7 +2900,7 @@ export const get54 = oc .output(zGetWorkspacesCurrentRbacWorkspaceAppsAccessPoliciesByPolicyIdRoleBindingsResponse) export const roleBindings3 = { - get: get54, + get: get57, } export const byPolicyId4 = { @@ -2863,7 +2913,7 @@ export const accessPolicies6 = { byPolicyId: byPolicyId4, } -export const get55 = oc +export const get58 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2874,7 +2924,7 @@ export const get55 = oc .output(zGetWorkspacesCurrentRbacWorkspaceAppsAccessPolicyResponse) export const accessPolicy3 = { - get: get55, + get: get58, } export const apps2 = { @@ -2902,7 +2952,7 @@ export const bindings2 = { put: put14, } -export const get56 = oc +export const get59 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2918,10 +2968,10 @@ export const get56 = oc .output(zGetWorkspacesCurrentRbacWorkspaceDatasetsAccessPoliciesByPolicyIdMemberBindingsResponse) export const memberBindings4 = { - get: get56, + get: get59, } -export const get57 = oc +export const get60 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2937,7 +2987,7 @@ export const get57 = oc .output(zGetWorkspacesCurrentRbacWorkspaceDatasetsAccessPoliciesByPolicyIdRoleBindingsResponse) export const roleBindings4 = { - get: get57, + get: get60, } export const byPolicyId5 = { @@ -2950,7 +3000,7 @@ export const accessPolicies7 = { byPolicyId: byPolicyId5, } -export const get58 = oc +export const get61 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2961,7 +3011,7 @@ export const get58 = oc .output(zGetWorkspacesCurrentRbacWorkspaceDatasetsAccessPolicyResponse) export const accessPolicy4 = { - get: get58, + get: get61, } export const datasets2 = { @@ -2986,7 +3036,7 @@ export const rbac = { workspace, } -export const get59 = oc +export const get62 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2997,7 +3047,7 @@ export const get59 = oc .output(zGetWorkspacesCurrentToolLabelsResponse) export const toolLabels = { - get: get59, + get: get62, } export const post48 = oc @@ -3030,7 +3080,7 @@ export const delete13 = { post: post49, } -export const get60 = oc +export const get63 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3041,11 +3091,11 @@ export const get60 = oc .input(z.object({ query: zGetWorkspacesCurrentToolProviderApiGetQuery })) .output(zGetWorkspacesCurrentToolProviderApiGetResponse) -export const get61 = { - get: get60, +export const get64 = { + get: get63, } -export const get62 = oc +export const get65 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3057,7 +3107,7 @@ export const get62 = oc .output(zGetWorkspacesCurrentToolProviderApiRemoteResponse) export const remote = { - get: get62, + get: get65, } export const post50 = oc @@ -3094,7 +3144,7 @@ export const test = { pre, } -export const get63 = oc +export const get66 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3106,7 +3156,7 @@ export const get63 = oc .output(zGetWorkspacesCurrentToolProviderApiToolsResponse) export const tools = { - get: get63, + get: get66, } export const post52 = oc @@ -3127,7 +3177,7 @@ export const update2 = { export const api = { add, delete: delete13, - get: get61, + get: get64, remote, schema, test, @@ -3155,7 +3205,7 @@ export const add2 = { post: post53, } -export const get64 = oc +export const get67 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3172,10 +3222,10 @@ export const get64 = oc .output(zGetWorkspacesCurrentToolProviderBuiltinByProviderCredentialInfoResponse) export const info = { - get: get64, + get: get67, } -export const get65 = oc +export const get68 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3195,7 +3245,7 @@ export const get65 = oc ) export const byCredentialType = { - get: get65, + get: get68, } export const schema2 = { @@ -3207,7 +3257,7 @@ export const credential = { schema: schema2, } -export const get66 = oc +export const get69 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3224,7 +3274,7 @@ export const get66 = oc .output(zGetWorkspacesCurrentToolProviderBuiltinByProviderCredentialsResponse) export const credentials3 = { - get: get66, + get: get69, } export const post54 = oc @@ -3267,7 +3317,7 @@ export const delete14 = { post: post55, } -export const get67 = oc +export const get70 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3279,10 +3329,10 @@ export const get67 = oc .output(zGetWorkspacesCurrentToolProviderBuiltinByProviderIconResponse) export const icon2 = { - get: get67, + get: get70, } -export const get68 = oc +export const get71 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3294,10 +3344,10 @@ export const get68 = oc .output(zGetWorkspacesCurrentToolProviderBuiltinByProviderInfoResponse) export const info2 = { - get: get68, + get: get71, } -export const get69 = oc +export const get72 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3311,7 +3361,7 @@ export const get69 = oc .output(zGetWorkspacesCurrentToolProviderBuiltinByProviderOauthClientSchemaResponse) export const clientSchema = { - get: get69, + get: get72, } export const delete15 = oc @@ -3329,7 +3379,7 @@ export const delete15 = oc ) .output(zDeleteWorkspacesCurrentToolProviderBuiltinByProviderOauthCustomClientResponse) -export const get70 = oc +export const get73 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3360,7 +3410,7 @@ export const post56 = oc export const customClient = { delete: delete15, - get: get70, + get: get73, post: post56, } @@ -3369,7 +3419,7 @@ export const oauth = { customClient, } -export const get71 = oc +export const get74 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3381,7 +3431,7 @@ export const get71 = oc .output(zGetWorkspacesCurrentToolProviderBuiltinByProviderToolsResponse) export const tools2 = { - get: get71, + get: get74, } export const post57 = oc @@ -3436,7 +3486,7 @@ export const auth = { post: post58, } -export const get72 = oc +export const get75 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3448,14 +3498,14 @@ export const get72 = oc .output(zGetWorkspacesCurrentToolProviderMcpToolsByProviderIdResponse) export const byProviderId = { - get: get72, + get: get75, } export const tools3 = { byProviderId, } -export const get73 = oc +export const get76 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3467,7 +3517,7 @@ export const get73 = oc .output(zGetWorkspacesCurrentToolProviderMcpUpdateByProviderIdResponse) export const byProviderId2 = { - get: get73, + get: get76, } export const update4 = { @@ -3546,7 +3596,7 @@ export const delete17 = { post: post61, } -export const get74 = oc +export const get77 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3557,11 +3607,11 @@ export const get74 = oc .input(z.object({ query: zGetWorkspacesCurrentToolProviderWorkflowGetQuery.optional() })) .output(zGetWorkspacesCurrentToolProviderWorkflowGetResponse) -export const get75 = { - get: get74, +export const get78 = { + get: get77, } -export const get76 = oc +export const get79 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3573,7 +3623,7 @@ export const get76 = oc .output(zGetWorkspacesCurrentToolProviderWorkflowToolsResponse) export const tools4 = { - get: get76, + get: get79, } export const post62 = oc @@ -3594,7 +3644,7 @@ export const update5 = { export const workflow = { create: create2, delete: delete17, - get: get75, + get: get78, tools: tools4, update: update5, } @@ -3606,7 +3656,7 @@ export const toolProvider = { workflow, } -export const get77 = oc +export const get80 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3618,10 +3668,10 @@ export const get77 = oc .output(zGetWorkspacesCurrentToolProvidersResponse) export const toolProviders = { - get: get77, + get: get80, } -export const get78 = oc +export const get81 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3632,10 +3682,10 @@ export const get78 = oc .output(zGetWorkspacesCurrentToolsApiResponse) export const api2 = { - get: get78, + get: get81, } -export const get79 = oc +export const get82 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3646,10 +3696,10 @@ export const get79 = oc .output(zGetWorkspacesCurrentToolsBuiltinResponse) export const builtin2 = { - get: get79, + get: get82, } -export const get80 = oc +export const get83 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3660,10 +3710,10 @@ export const get80 = oc .output(zGetWorkspacesCurrentToolsMcpResponse) export const mcp2 = { - get: get80, + get: get83, } -export const get81 = oc +export const get84 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3674,7 +3724,7 @@ export const get81 = oc .output(zGetWorkspacesCurrentToolsWorkflowResponse) export const workflow2 = { - get: get81, + get: get84, } export const tools5 = { @@ -3684,7 +3734,7 @@ export const tools5 = { workflow: workflow2, } -export const get82 = oc +export const get85 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3696,13 +3746,13 @@ export const get82 = oc .output(zGetWorkspacesCurrentTriggerProviderByProviderIconResponse) export const icon3 = { - get: get82, + get: get85, } /** * Get info for a trigger provider */ -export const get83 = oc +export const get86 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3715,7 +3765,7 @@ export const get83 = oc .output(zGetWorkspacesCurrentTriggerProviderByProviderInfoResponse) export const info3 = { - get: get83, + get: get86, } /** @@ -3736,7 +3786,7 @@ export const delete18 = oc /** * Get OAuth client configuration for a provider */ -export const get84 = oc +export const get87 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3770,7 +3820,7 @@ export const post63 = oc export const client = { delete: delete18, - get: get84, + get: get87, post: post63, } @@ -3837,7 +3887,7 @@ export const create3 = { /** * Get the request logs for a subscription instance for a trigger provider */ -export const get85 = oc +export const get88 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3858,7 +3908,7 @@ export const get85 = oc ) export const bySubscriptionBuilderId2 = { - get: get85, + get: get88, } export const logs = { @@ -3932,7 +3982,7 @@ export const verifyAndUpdate = { /** * Get a subscription instance for a trigger provider */ -export const get86 = oc +export const get89 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3953,7 +4003,7 @@ export const get86 = oc ) export const bySubscriptionBuilderId5 = { - get: get86, + get: get89, } export const builder = { @@ -3968,7 +4018,7 @@ export const builder = { /** * List all trigger subscriptions for the current tenant's provider */ -export const get87 = oc +export const get90 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3981,13 +4031,13 @@ export const get87 = oc .output(zGetWorkspacesCurrentTriggerProviderByProviderSubscriptionsListResponse) export const list4 = { - get: get87, + get: get90, } /** * Initiate OAuth authorization flow for a trigger provider */ -export const get88 = oc +export const get91 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -4004,7 +4054,7 @@ export const get88 = oc .output(zGetWorkspacesCurrentTriggerProviderByProviderSubscriptionsOauthAuthorizeResponse) export const authorize = { - get: get88, + get: get91, } export const oauth3 = { @@ -4121,7 +4171,7 @@ export const triggerProvider = { /** * List all trigger providers for the current tenant */ -export const get89 = oc +export const get92 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -4133,7 +4183,7 @@ export const get89 = oc .output(zGetWorkspacesCurrentTriggersResponse) export const triggers = { - get: get89, + get: get92, } export const post71 = oc @@ -4234,7 +4284,7 @@ export const switch3 = { post: post75, } -export const get90 = oc +export const get93 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -4246,7 +4296,7 @@ export const get90 = oc .output(zGetWorkspacesByTenantIdModelProvidersByProviderByIconTypeByLangResponse) export const byLang = { - get: get90, + get: get93, } export const byIconType = { @@ -4265,7 +4315,7 @@ export const byTenantId = { modelProviders: modelProviders2, } -export const get91 = oc +export const get94 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -4276,7 +4326,7 @@ export const get91 = oc .output(zGetWorkspacesResponse) export const workspaces = { - get: get91, + get: get94, current, customConfig, info: info4, diff --git a/packages/contracts/generated/api/console/workspaces/types.gen.ts b/packages/contracts/generated/api/console/workspaces/types.gen.ts index 59b973b1ecd..70b1a6874f9 100644 --- a/packages/contracts/generated/api/console/workspaces/types.gen.ts +++ b/packages/contracts/generated/api/console/workspaces/types.gen.ts @@ -227,6 +227,24 @@ export type ModelProviderListResponse = { data: Array } +export type ModelProviderCreditsResponse = { + exhausted_at: number | null + is_exhausted: boolean + is_unlimited: boolean + next_credit_reset_date: number | null + pool_type: 'paid' | 'trial' | null + quota_limit: number | null + quota_used: number | null + remaining_credits: number | null +} + +export type ModelProviderSummaryListResponse = { + data: Array + plugins: { + [key: string]: ModelProviderPluginSummaryResponse + } +} + export type ModelProviderPaymentCheckoutUrlResponse = { payment_link: string } @@ -417,6 +435,10 @@ export type ParserPluginIdentifiers = { plugin_unique_identifiers: Array } +export type PluginInstalledIdsResponse = { + plugin_ids: Array +} + export type PluginListResponse = { plugins: Array total: number @@ -1192,6 +1214,30 @@ export type ProviderResponse = { tenant_id: string } +export type ModelProviderSummaryResponse = { + configurate_methods: Array + custom_configuration: ModelProviderCustomConfigurationSummaryResponse + description?: I18nObject | null + icon_small?: I18nObject | null + icon_small_dark?: I18nObject | null + is_configured: boolean + label: I18nObject + plugin_id: string + preferred_provider_type: ProviderType + provider: string + supported_model_types: Array + system_configuration: ModelProviderSystemConfigurationSummaryResponse +} + +export type ModelProviderPluginSummaryResponse = { + installation_id: string + plugin_id: string + plugin_unique_identifier: string + runtime_type: string + source: PluginInstallationSource + version: string +} + export type ModelType = 'llm' | 'moderation' | 'rerank' | 'speech2text' | 'text-embedding' | 'tts' export type ModelWithProviderEntityResponse = { @@ -1730,6 +1776,21 @@ export type SystemConfigurationResponse = { quota_configurations?: Array } +export type ModelProviderCustomConfigurationSummaryResponse = { + available_credentials: Array + current_credential_id?: string | null + current_credential_name?: string | null + current_credential_usable: boolean + has_custom_models: boolean + status: CustomConfigurationStatus +} + +export type ModelProviderSystemConfigurationSummaryResponse = { + enabled: boolean +} + +export type PluginInstallationSource = 'github' | 'marketplace' | 'package' | 'remote' + export type ModelFeature = | 'agent-thought' | 'audio' @@ -1900,8 +1961,6 @@ export type PluginInstallTaskPluginStatus = { export type PluginInstallTaskStatus = 'failed' | 'pending' | 'running' | 'success' -export type PluginInstallationSource = 'github' | 'marketplace' | 'package' | 'remote' - export type PluginDeclarationResponse = { agent_strategy?: { [key: string]: unknown @@ -3029,6 +3088,34 @@ export type GetWorkspacesCurrentModelProvidersResponses = { export type GetWorkspacesCurrentModelProvidersResponse = GetWorkspacesCurrentModelProvidersResponses[keyof GetWorkspacesCurrentModelProvidersResponses] +export type GetWorkspacesCurrentModelProvidersCreditsData = { + body?: never + path?: never + query?: never + url: '/workspaces/current/model-providers/credits' +} + +export type GetWorkspacesCurrentModelProvidersCreditsResponses = { + 200: ModelProviderCreditsResponse +} + +export type GetWorkspacesCurrentModelProvidersCreditsResponse = + GetWorkspacesCurrentModelProvidersCreditsResponses[keyof GetWorkspacesCurrentModelProvidersCreditsResponses] + +export type GetWorkspacesCurrentModelProvidersSummaryData = { + body?: never + path?: never + query?: never + url: '/workspaces/current/model-providers/summary' +} + +export type GetWorkspacesCurrentModelProvidersSummaryResponses = { + 200: ModelProviderSummaryListResponse +} + +export type GetWorkspacesCurrentModelProvidersSummaryResponse = + GetWorkspacesCurrentModelProvidersSummaryResponses[keyof GetWorkspacesCurrentModelProvidersSummaryResponses] + export type GetWorkspacesCurrentModelProvidersByProviderCheckoutUrlData = { body?: never path: { @@ -3575,6 +3662,22 @@ export type PostWorkspacesCurrentPluginInstallPkgResponses = { export type PostWorkspacesCurrentPluginInstallPkgResponse = PostWorkspacesCurrentPluginInstallPkgResponses[keyof PostWorkspacesCurrentPluginInstallPkgResponses] +export type GetWorkspacesCurrentPluginInstalledIdsData = { + body?: never + path?: never + query: { + category: 'agent-strategy' | 'datasource' | 'extension' | 'model' | 'tool' | 'trigger' + } + url: '/workspaces/current/plugin/installed-ids' +} + +export type GetWorkspacesCurrentPluginInstalledIdsResponses = { + 200: PluginInstalledIdsResponse +} + +export type GetWorkspacesCurrentPluginInstalledIdsResponse = + GetWorkspacesCurrentPluginInstalledIdsResponses[keyof GetWorkspacesCurrentPluginInstalledIdsResponses] + export type GetWorkspacesCurrentPluginListData = { body?: never path?: never @@ -3888,8 +3991,11 @@ export type GetWorkspacesCurrentPluginByCategoryListData = { category: string } query?: { + language?: 'en_US' | 'ja_JP' | 'pt_BR' | 'zh_Hans' page?: number page_size?: number + query?: string + tags?: Array } url: '/workspaces/current/plugin/{category}/list' } diff --git a/packages/contracts/generated/api/console/workspaces/zod.gen.ts b/packages/contracts/generated/api/console/workspaces/zod.gen.ts index 6f17f5c6d17..5f8dd9ba6c5 100644 --- a/packages/contracts/generated/api/console/workspaces/zod.gen.ts +++ b/packages/contracts/generated/api/console/workspaces/zod.gen.ts @@ -158,6 +158,20 @@ export const zMemberRoleUpdatePayload = z.object({ role: z.string(), }) +/** + * ModelProviderCreditsResponse + */ +export const zModelProviderCreditsResponse = z.object({ + exhausted_at: z.int().nullable(), + is_exhausted: z.boolean(), + is_unlimited: z.boolean(), + next_credit_reset_date: z.int().nullable(), + pool_type: z.enum(['paid', 'trial']).nullable(), + quota_limit: z.int().nullable(), + quota_used: z.int().nullable(), + remaining_credits: z.int().nullable(), +}) + /** * ModelProviderPaymentCheckoutUrlResponse */ @@ -283,6 +297,13 @@ export const zParserPluginIdentifiers = z.object({ plugin_unique_identifiers: z.array(z.string()), }) +/** + * PluginInstalledIdsResponse + */ +export const zPluginInstalledIdsResponse = z.object({ + plugin_ids: z.array(z.string()), +}) + /** * ParserLatest */ @@ -1613,6 +1634,30 @@ export const zConfigurateMethod = z.enum(['customizable-model', 'predefined-mode */ export const zProviderType = z.enum(['custom', 'system']) +/** + * ModelProviderSystemConfigurationSummaryResponse + */ +export const zModelProviderSystemConfigurationSummaryResponse = z.object({ + enabled: z.boolean(), +}) + +/** + * PluginInstallationSource + */ +export const zPluginInstallationSource = z.enum(['github', 'marketplace', 'package', 'remote']) + +/** + * ModelProviderPluginSummaryResponse + */ +export const zModelProviderPluginSummaryResponse = z.object({ + installation_id: z.string(), + plugin_id: z.string(), + plugin_unique_identifier: z.string(), + runtime_type: z.string(), + source: zPluginInstallationSource, + version: z.string(), +}) + /** * ModelFeature * @@ -1803,6 +1848,46 @@ export const zAvailableModelListResponse = z.object({ data: z.array(zProviderWithModelsResponse), }) +/** + * ModelProviderCustomConfigurationSummaryResponse + */ +export const zModelProviderCustomConfigurationSummaryResponse = z.object({ + available_credentials: z.array(zCredentialConfiguration), + current_credential_id: z.string().nullish(), + current_credential_name: z.string().nullish(), + current_credential_usable: z.boolean(), + has_custom_models: z.boolean(), + status: zCustomConfigurationStatus, +}) + +/** + * ModelProviderSummaryResponse + * + * Fields required to render the collapsed model-provider list. + */ +export const zModelProviderSummaryResponse = z.object({ + configurate_methods: z.array(zConfigurateMethod), + custom_configuration: zModelProviderCustomConfigurationSummaryResponse, + description: zI18nObject.nullish(), + icon_small: zI18nObject.nullish(), + icon_small_dark: zI18nObject.nullish(), + is_configured: z.boolean(), + label: zI18nObject, + plugin_id: z.string(), + preferred_provider_type: zProviderType, + provider: z.string(), + supported_model_types: z.array(zModelType), + system_configuration: zModelProviderSystemConfigurationSummaryResponse, +}) + +/** + * ModelProviderSummaryListResponse + */ +export const zModelProviderSummaryListResponse = z.object({ + data: z.array(zModelProviderSummaryResponse), + plugins: z.record(z.string(), zModelProviderPluginSummaryResponse), +}) + /** * TenantPluginAutoUpgradeStrategySetting */ @@ -1948,11 +2033,6 @@ export const zPluginTaskResponse = z.object({ task: zPluginInstallTask, }) -/** - * PluginInstallationSource - */ -export const zPluginInstallationSource = z.enum(['github', 'marketplace', 'package', 'remote']) - /** * PluginBundleDependencyType */ @@ -3696,6 +3776,16 @@ export const zGetWorkspacesCurrentModelProvidersQuery = z.object({ */ export const zGetWorkspacesCurrentModelProvidersResponse = zModelProviderListResponse +/** + * Model provider credits retrieved successfully + */ +export const zGetWorkspacesCurrentModelProvidersCreditsResponse = zModelProviderCreditsResponse + +/** + * Model provider summaries retrieved successfully + */ +export const zGetWorkspacesCurrentModelProvidersSummaryResponse = zModelProviderSummaryListResponse + export const zGetWorkspacesCurrentModelProvidersByProviderCheckoutUrlPath = z.object({ provider: z.string(), }) @@ -4071,6 +4161,15 @@ export const zPostWorkspacesCurrentPluginInstallPkgBody = zParserPluginIdentifie */ export const zPostWorkspacesCurrentPluginInstallPkgResponse = zPluginInstallTaskStartResponse +export const zGetWorkspacesCurrentPluginInstalledIdsQuery = z.object({ + category: z.enum(['agent-strategy', 'datasource', 'extension', 'model', 'tool', 'trigger']), +}) + +/** + * Success + */ +export const zGetWorkspacesCurrentPluginInstalledIdsResponse = zPluginInstalledIdsResponse + export const zGetWorkspacesCurrentPluginListQuery = z.object({ page: z.int().gte(1).optional().default(1), page_size: z.int().gte(1).lte(256).optional().default(256), @@ -4241,8 +4340,11 @@ export const zGetWorkspacesCurrentPluginByCategoryListPath = z.object({ }) export const zGetWorkspacesCurrentPluginByCategoryListQuery = z.object({ + language: z.enum(['en_US', 'ja_JP', 'pt_BR', 'zh_Hans']).optional().default('en_US'), page: z.int().gte(1).optional().default(1), page_size: z.int().gte(1).lte(256).optional().default(256), + query: z.string().max(256).optional().default(''), + tags: z.array(z.string()).max(128).optional(), }) /** diff --git a/web/__mocks__/provider-context.ts b/web/__mocks__/provider-context.ts index 6265352feee..0b2b2b90bf1 100644 --- a/web/__mocks__/provider-context.ts +++ b/web/__mocks__/provider-context.ts @@ -7,9 +7,12 @@ import { defaultPlan } from '@/app/components/billing/config' // Avoid being mocked in tests export const baseProviderContextValue: ProviderContextState = { modelProviders: [], + modelProviderPlugins: {}, refreshModelProviders: async () => {}, isLoadingModelProviders: false, + isSuccessModelProviders: false, textGenerationModelList: [], + supportRetrievalMethods: [], isAPIKeySet: true, plan: defaultPlan, isFetchedPlan: false, @@ -18,6 +21,7 @@ export const baseProviderContextValue: ProviderContextState = { onPlanInfoChanged: noop, enableReplaceWebAppLogo: false, modelLoadBalancingEnabled: false, + datasetOperatorEnabled: false, enableEducationPlan: false, isEducationWorkspace: false, isEducationAccount: false, @@ -45,7 +49,7 @@ export const createMockProviderContextValue = ( return { ...merged, - refreshModelProviders: merged.refreshModelProviders ?? (async () => {}), + refreshModelProviders: merged.refreshModelProviders ?? noop, onPlanInfoChanged: merged.onPlanInfoChanged ?? noop, refreshLicenseLimit: merged.refreshLicenseLimit ?? noop, } diff --git a/web/app/components/app/configuration/debug/debug-with-multiple-model/__tests__/model-parameter-trigger.spec.tsx b/web/app/components/app/configuration/debug/debug-with-multiple-model/__tests__/model-parameter-trigger.spec.tsx index ae993248a1a..f7ee57353a7 100644 --- a/web/app/components/app/configuration/debug/debug-with-multiple-model/__tests__/model-parameter-trigger.spec.tsx +++ b/web/app/components/app/configuration/debug/debug-with-multiple-model/__tests__/model-parameter-trigger.spec.tsx @@ -1,3 +1,4 @@ +import type { ModelProviderSummaryResponse } from '@dify/contracts/api/console/workspaces/types.gen' import type { ReactNode } from 'react' import type { ModelAndParameter } from '../../types' import type { @@ -158,7 +159,13 @@ describe('ModelParameterTrigger', () => { }) mockUseProviderContext.mockReturnValue( createMockProviderContextValue({ - modelProviders: [createModelProvider()], + modelProviders: [ + { + ...createModelProvider(), + is_configured: true, + plugin_id: 'langgenius/openai', + } as unknown as ModelProviderSummaryResponse, + ], }), ) mockUseCredentialPanelState.mockReturnValue({ diff --git a/web/app/components/app/overview/apikey-info-panel/__tests__/test-utils.tsx b/web/app/components/app/overview/apikey-info-panel/__tests__/test-utils.tsx index fe87c4071cd..90c6fcc6876 100644 --- a/web/app/components/app/overview/apikey-info-panel/__tests__/test-utils.tsx +++ b/web/app/components/app/overview/apikey-info-panel/__tests__/test-utils.tsx @@ -39,9 +39,12 @@ const mockUseProviderContext = actualUseProviderContext as MockedFunction< // Default mock data const defaultProviderContext = { modelProviders: [], + modelProviderPlugins: {}, refreshModelProviders: async () => {}, isLoadingModelProviders: false, + isSuccessModelProviders: false, textGenerationModelList: [], + supportRetrievalMethods: [], isAPIKeySet: false, plan: defaultPlan, isFetchedPlan: false, @@ -50,6 +53,7 @@ const defaultProviderContext = { onPlanInfoChanged: noop, enableReplaceWebAppLogo: false, modelLoadBalancingEnabled: false, + datasetOperatorEnabled: false, enableEducationPlan: false, isEducationWorkspace: false, isEducationAccount: false, diff --git a/web/app/components/base/features/new-feature-panel/moderation/__tests__/moderation-setting-modal.spec.tsx b/web/app/components/base/features/new-feature-panel/moderation/__tests__/moderation-setting-modal.spec.tsx index 0598c01ced1..392c79e689f 100644 --- a/web/app/components/base/features/new-feature-panel/moderation/__tests__/moderation-setting-modal.spec.tsx +++ b/web/app/components/base/features/new-feature-panel/moderation/__tests__/moderation-setting-modal.spec.tsx @@ -44,7 +44,7 @@ let mockModelProvidersData: { vi.mock('@/service/use-common', () => ({ useCodeBasedExtensions: () => mockCodeBasedExtensions, - useModelProviders: () => mockModelProvidersData, + useModelProviderDetails: () => mockModelProvidersData, })) vi.mock('@/app/components/header/account-setting/model-provider-page/declarations', () => ({ diff --git a/web/app/components/base/features/new-feature-panel/moderation/moderation-setting-modal.tsx b/web/app/components/base/features/new-feature-panel/moderation/moderation-setting-modal.tsx index 9406bcb6cdb..f7e968e4395 100644 --- a/web/app/components/base/features/new-feature-panel/moderation/moderation-setting-modal.tsx +++ b/web/app/components/base/features/new-feature-panel/moderation/moderation-setting-modal.tsx @@ -18,7 +18,7 @@ import { } from '@/app/components/header/account-setting/query-params' import { useDocLink, useLocale } from '@/context/i18n' import { LanguagesSupported } from '@/i18n-config/language' -import { useCodeBasedExtensions, useModelProviders } from '@/service/use-common' +import { useCodeBasedExtensions, useModelProviderDetails } from '@/service/use-common' import FormGeneration from './form-generation' import ModerationContent from './moderation-content' @@ -59,7 +59,7 @@ const ModerationSettingModal: FC = ({ data, onCance const { t } = useTranslation() const docLink = useDocLink() const locale = useLocale() - const { data: modelProviders, isPending: isLoading } = useModelProviders() + const { data: modelProviders, isPending: isLoading } = useModelProviderDetails() const localeDataRef = useRef(data) const [localeData, setLocaleData] = useState(data) const [, setSettingsDestination] = useQueryState(settingsQueryParamName, settingsQueryParser) diff --git a/web/app/components/header/account-setting/data-source-page-new/__tests__/card.spec.tsx b/web/app/components/header/account-setting/data-source-page-new/__tests__/card.spec.tsx index 35f629db634..fe56ab9dfdd 100644 --- a/web/app/components/header/account-setting/data-source-page-new/__tests__/card.spec.tsx +++ b/web/app/components/header/account-setting/data-source-page-new/__tests__/card.spec.tsx @@ -199,16 +199,15 @@ describe('Card Component', () => { // Act render() - // Assert // Assert expect(screen.getByText('Test Label'))!.toBeInTheDocument() expect(screen.queryByText(/Test Author/))!.not.toBeInTheDocument() expect(screen.queryByText(/test-name/))!.not.toBeInTheDocument() expect(screen.getByText('1.2.0'))!.toBeInTheDocument() - expect(screen.getByRole('img', { name: 'Test Label' }))!.toHaveAttribute( - 'src', - 'test-icon-url', - ) + const icon = screen.getByRole('img', { name: 'Test Label' }) + expect(icon).toHaveAttribute('src', 'test-icon-url') + expect(icon).toHaveAttribute('loading', 'lazy') + expect(icon).toHaveAttribute('decoding', 'async') expect(screen.getByText('Credential 1'))!.toBeInTheDocument() expect(screen.getByText(/plugin.auth.default/))!.toBeInTheDocument() diff --git a/web/app/components/header/account-setting/data-source-page-new/__tests__/index.spec.tsx b/web/app/components/header/account-setting/data-source-page-new/__tests__/index.spec.tsx index 6638a0b7a93..1fdb022ce29 100644 --- a/web/app/components/header/account-setting/data-source-page-new/__tests__/index.spec.tsx +++ b/web/app/components/header/account-setting/data-source-page-new/__tests__/index.spec.tsx @@ -5,6 +5,7 @@ import { fireEvent, screen } from '@testing-library/react' import { useTheme } from 'next-themes' import { usePluginsWithLatestVersion } from '@/app/components/plugins/hooks' import { usePluginAuthAction } from '@/app/components/plugins/plugin-auth' +import { PluginCategoryEnum } from '@/app/components/plugins/types' import { useRenderI18nObject } from '@/hooks/use-i18n' import { useGetDataSourceListAuth, @@ -301,6 +302,9 @@ describe('DataSourcePage Component', () => { // Assert expect(screen.getByTestId('plugin-actions-plugin-1')).toBeInTheDocument() + expect(useInstalledPluginList).toHaveBeenLastCalledWith({ + category: PluginCategoryEnum.datasource, + }) }) it('should filter installed data sources and pass search text to marketplace', () => { diff --git a/web/app/components/header/account-setting/data-source-page-new/card.tsx b/web/app/components/header/account-setting/data-source-page-new/card.tsx index 5225e1a69d4..be8f519f48d 100644 --- a/web/app/components/header/account-setting/data-source-page-new/card.tsx +++ b/web/app/components/header/account-setting/data-source-page-new/card.tsx @@ -108,6 +108,8 @@ const Card = ({ item, disabled, pluginDetail, onPluginUpdate }: CardProps) => { alt={providerLabel} width={20} height={20} + loading="lazy" + decoding="async" className="h-5 w-5 object-contain" /> diff --git a/web/app/components/header/account-setting/data-source-page-new/index.tsx b/web/app/components/header/account-setting/data-source-page-new/index.tsx index 01a73a19ee0..7347034d159 100644 --- a/web/app/components/header/account-setting/data-source-page-new/index.tsx +++ b/web/app/components/header/account-setting/data-source-page-new/index.tsx @@ -63,7 +63,9 @@ const DataSourcePage = ({ layout, onOpenMarketplace, stickyToolbar }: DataSource select: (s) => s.enable_marketplace, }) const { data, isLoading: isDataSourceListLoading } = useGetDataSourceListAuth() - const { data: installedPluginList } = useInstalledPluginList() + const { data: installedPluginList } = useInstalledPluginList({ + category: PluginCategoryEnum.datasource, + }) const pluginListWithLatestVersion = usePluginsWithLatestVersion(installedPluginList?.plugins) const invalidateInstalledPluginList = useInvalidateInstalledPluginList() const invalidateDataSourceListAuth = useInvalidDataSourceListAuth() diff --git a/web/app/components/header/account-setting/model-provider-page/__tests__/hooks.spec.ts b/web/app/components/header/account-setting/model-provider-page/__tests__/hooks.spec.ts index 53142bd552d..918f4f75f6f 100644 --- a/web/app/components/header/account-setting/model-provider-page/__tests__/hooks.spec.ts +++ b/web/app/components/header/account-setting/model-provider-page/__tests__/hooks.spec.ts @@ -21,7 +21,7 @@ import { PreferredProviderTypeEnum, } from '../declarations' import { - useCurrentProviderAndModel, + getCurrentProviderAndModel, useDefaultModel, useInvalidateDefaultModel, useLanguage, @@ -58,6 +58,7 @@ vi.mock('@/service/use-common', () => ({ commonQueryKeys: { modelList: (type: string) => ['model-list', type], modelProviders: ['model-providers'], + modelProviderDetails: ['model-provider-details'], defaultModel: (type: string) => ['default-model', type], }, })) @@ -293,6 +294,23 @@ describe('hooks', () => { expect(result.current.data).toEqual([]) }) + it('should keep the query disabled when requested', () => { + ;(useQuery as Mock).mockReturnValue({ + data: undefined, + isPending: true, + refetch: vi.fn(), + }) + + renderHook(() => useModelList(ModelTypeEnum.textEmbedding, { enabled: false })) + + expect(useQuery).toHaveBeenCalledWith( + expect.objectContaining({ + enabled: false, + queryKey: ['model-list', ModelTypeEnum.textEmbedding], + }), + ) + }) + it('should handle loading state', () => { ;(useQuery as Mock).mockReturnValue({ data: undefined, @@ -411,7 +429,7 @@ describe('hooks', () => { }) }) - describe('useCurrentProviderAndModel', () => { + describe('getCurrentProviderAndModel', () => { const createModelList = (): Model[] => [ { provider: 'openai', @@ -445,7 +463,7 @@ describe('hooks', () => { const modelList = createModelList() const defaultModel = { provider: 'openai', model: 'gpt-4' } - const { result } = renderHook(() => useCurrentProviderAndModel(modelList, defaultModel)) + const { result } = renderHook(() => getCurrentProviderAndModel(modelList, defaultModel)) expect(result.current.currentProvider?.provider).toBe('openai') expect(result.current.currentModel?.model).toBe('gpt-4') @@ -455,7 +473,7 @@ describe('hooks', () => { const modelList = createModelList() const defaultModel = { provider: 'anthropic', model: 'claude-3' } - const { result } = renderHook(() => useCurrentProviderAndModel(modelList, defaultModel)) + const { result } = renderHook(() => getCurrentProviderAndModel(modelList, defaultModel)) expect(result.current.currentProvider).toBeUndefined() expect(result.current.currentModel).toBeUndefined() @@ -465,7 +483,7 @@ describe('hooks', () => { const modelList = createModelList() const defaultModel = { provider: 'openai', model: 'gpt-5' } - const { result } = renderHook(() => useCurrentProviderAndModel(modelList, defaultModel)) + const { result } = renderHook(() => getCurrentProviderAndModel(modelList, defaultModel)) expect(result.current.currentProvider?.provider).toBe('openai') expect(result.current.currentModel).toBeUndefined() @@ -474,7 +492,7 @@ describe('hooks', () => { it('should handle undefined default model', () => { const modelList = createModelList() - const { result } = renderHook(() => useCurrentProviderAndModel(modelList, undefined)) + const { result } = renderHook(() => getCurrentProviderAndModel(modelList, undefined)) expect(result.current.currentProvider).toBeUndefined() expect(result.current.currentModel).toBeUndefined() @@ -483,7 +501,7 @@ describe('hooks', () => { it('should handle empty model list', () => { const defaultModel = { provider: 'openai', model: 'gpt-4' } - const { result } = renderHook(() => useCurrentProviderAndModel([], defaultModel)) + const { result } = renderHook(() => getCurrentProviderAndModel([], defaultModel)) expect(result.current.currentProvider).toBeUndefined() expect(result.current.currentModel).toBeUndefined() @@ -754,7 +772,10 @@ describe('hooks', () => { }) expect(invalidateQueries).toHaveBeenCalledWith({ - queryKey: ['model-providers'], + queryKey: consoleQuery.workspaces.current.modelProviders.summary.get.key(), + }) + expect(invalidateQueries).toHaveBeenCalledWith({ + queryKey: ['model-provider-details'], }) }) @@ -770,55 +791,17 @@ describe('hooks', () => { result.current() }) - expect(invalidateQueries).toHaveBeenCalledTimes(3) + expect(invalidateQueries).toHaveBeenCalledTimes(6) }) }) describe('useMarketplaceAllPlugins', () => { - const createMockProviders = (): ModelProvider[] => [ - { - provider: 'openai', - label: { en_US: 'OpenAI', zh_Hans: 'OpenAI' }, - icon_small: { en_US: 'icon', zh_Hans: 'icon' }, - supported_model_types: [ModelTypeEnum.textGeneration], - configurate_methods: [ConfigurationMethodEnum.predefinedModel], - provider_credential_schema: { credential_form_schemas: [] }, - model_credential_schema: { - model: { - label: { en_US: 'Model', zh_Hans: '模型' }, - placeholder: { en_US: 'Select model', zh_Hans: '选择模型' }, - }, - credential_form_schemas: [], - }, - preferred_provider_type: PreferredProviderTypeEnum.system, - custom_configuration: { - status: CustomConfigurationStatusEnum.noConfigure, - }, - system_configuration: { - enabled: true, - current_quota_type: CurrentSystemQuotaTypeEnum.trial, - quota_configurations: [], - }, - help: { - title: { - en_US: '', - zh_Hans: '', - }, - url: { - en_US: '', - zh_Hans: '', - }, - }, - }, - ] - const createMockPlugins = () => [ { plugin_id: 'plugin1', type: 'plugin' }, { plugin_id: 'plugin2', type: 'plugin' }, ] it('should combine collection and regular plugins', () => { - const providers = createMockProviders() const collectionPlugins = [{ plugin_id: 'collection1', type: 'plugin' }] const regularPlugins = createMockPlugins() ;(useMarketplacePluginsByCollectionId as Mock).mockReturnValue({ @@ -832,14 +815,13 @@ describe('hooks', () => { isLoading: false, }) - const { result } = renderHook(() => useMarketplaceAllPlugins(providers, '')) + const { result } = renderHook(() => useMarketplaceAllPlugins('', [])) expect(result.current.plugins).toHaveLength(3) expect(result.current.isLoading).toBe(false) }) it('should exclude installed providers', () => { - const providers = createMockProviders() const collectionPlugins = [ { plugin_id: 'openai', type: 'plugin' }, { plugin_id: 'other', type: 'plugin' }, @@ -849,16 +831,22 @@ describe('hooks', () => { isLoading: false, }) ;(useMarketplacePlugins as Mock).mockReturnValue({ - plugins: [], + plugins: [ + { plugin_id: 'openai', type: 'plugin' }, + { plugin_id: 'regular-only', type: 'plugin' }, + ], queryPlugins: vi.fn(), queryPluginsWithDebounced: vi.fn(), isLoading: false, }) - const { result } = renderHook(() => useMarketplaceAllPlugins(providers, '')) + const { result } = renderHook(() => useMarketplaceAllPlugins('', ['openai'])) - expect(result.current.plugins!).toHaveLength(1) - expect(result.current.plugins![0]!.plugin_id).toBe('other') + expect(result.current.plugins!).toHaveLength(2) + expect(result.current.plugins!.map((plugin) => plugin.plugin_id)).toEqual([ + 'other', + 'regular-only', + ]) }) it('should use search when searchText is provided', () => { @@ -874,7 +862,7 @@ describe('hooks', () => { isLoading: false, }) - renderHook(() => useMarketplaceAllPlugins([], 'test search')) + renderHook(() => useMarketplaceAllPlugins('test search', [])) expect(queryPluginsWithDebounced).toHaveBeenCalled() }) @@ -895,7 +883,7 @@ describe('hooks', () => { isLoading: false, }) - const { result } = renderHook(() => useMarketplaceAllPlugins([], '')) + const { result } = renderHook(() => useMarketplaceAllPlugins('', [])) expect(result.current.plugins!).toHaveLength(1) expect(result.current.plugins![0]!.plugin_id).toBe('plugin1') @@ -914,7 +902,7 @@ describe('hooks', () => { isLoading: false, }) - const { result } = renderHook(() => useMarketplaceAllPlugins([], '')) + const { result } = renderHook(() => useMarketplaceAllPlugins('', [])) expect(result.current.plugins).toHaveLength(2) expect(result.current.plugins!.filter((p) => p.plugin_id === 'shared-plugin')).toHaveLength(1) @@ -932,7 +920,7 @@ describe('hooks', () => { isLoading: true, }) - const { result } = renderHook(() => useMarketplaceAllPlugins([], '')) + const { result } = renderHook(() => useMarketplaceAllPlugins('', [])) expect(result.current.isLoading).toBe(true) }) @@ -949,7 +937,7 @@ describe('hooks', () => { isLoading: false, }) - const { result } = renderHook(() => useMarketplaceAllPlugins([], '')) + const { result } = renderHook(() => useMarketplaceAllPlugins('', [])) expect(result.current.plugins).toBeDefined() expect(result.current.isLoading).toBe(false) @@ -969,13 +957,33 @@ describe('hooks', () => { isLoading: false, }) - const { result } = renderHook(() => useMarketplaceAllPlugins([], 'openai')) + const { result } = renderHook(() => useMarketplaceAllPlugins('openai', [])) expect(result.current.plugins).toEqual(searchPlugins) expect(result.current.plugins?.some((p) => p.plugin_id === 'collection-only')).toBe(false) }) - it('should skip marketplace queries when disabled', () => { + it('should hide installed plugins when a search response is stale', () => { + ;(useMarketplacePluginsByCollectionId as Mock).mockReturnValue({ + plugins: [], + isLoading: false, + }) + ;(useMarketplacePlugins as Mock).mockReturnValue({ + plugins: [ + { plugin_id: 'langgenius/openai', type: 'plugin' }, + { plugin_id: 'langgenius/other', type: 'plugin' }, + ], + queryPlugins: vi.fn(), + queryPluginsWithDebounced: vi.fn(), + isLoading: false, + }) + + const { result } = renderHook(() => useMarketplaceAllPlugins('openai', ['langgenius/openai'])) + + expect(result.current.plugins).toEqual([{ plugin_id: 'langgenius/other', type: 'plugin' }]) + }) + + it('should preserve marketplace cache when disabled', () => { const queryPlugins = vi.fn() const queryPluginsWithDebounced = vi.fn() const cancelQueryPluginsWithDebounced = vi.fn() @@ -993,13 +1001,14 @@ describe('hooks', () => { isLoading: true, }) - const { result } = renderHook(() => useMarketplaceAllPlugins([], '', false)) + const { result } = renderHook(() => useMarketplaceAllPlugins('', [], false)) expect(useMarketplacePluginsByCollectionId).toHaveBeenCalledWith(undefined) expect(queryPlugins).not.toHaveBeenCalled() expect(queryPluginsWithDebounced).not.toHaveBeenCalled() expect(cancelQueryPluginsWithDebounced).toHaveBeenCalled() - expect(resetPlugins).toHaveBeenCalled() + expect(resetPlugins).not.toHaveBeenCalled() + expect(useMarketplacePlugins).toHaveBeenCalledWith(false) expect(result.current.plugins).toEqual([]) expect(result.current.isLoading).toBe(false) }) @@ -1065,7 +1074,12 @@ describe('hooks', () => { exact: true, refetchType: 'none', }) - expect(invalidateQueries).toHaveBeenCalledWith({ queryKey: ['model-providers'] }) + expect(invalidateQueries).toHaveBeenCalledWith({ + queryKey: consoleQuery.workspaces.current.modelProviders.summary.get.key(), + }) + expect(invalidateQueries).toHaveBeenCalledWith({ + queryKey: ['model-provider-details'], + }) expect(invalidateQueries).toHaveBeenCalledWith({ queryKey: ['model-list', ModelTypeEnum.textGeneration], }) @@ -1208,7 +1222,12 @@ describe('hooks', () => { result.current.handleRefreshModel(provider) }) - expect(invalidateQueries).toHaveBeenCalledWith({ queryKey: ['model-providers'] }) + expect(invalidateQueries).toHaveBeenCalledWith({ + queryKey: consoleQuery.workspaces.current.modelProviders.summary.get.key(), + }) + expect(invalidateQueries).toHaveBeenCalledWith({ + queryKey: ['model-provider-details'], + }) expect(invalidateQueries).toHaveBeenCalledWith({ queryKey: ['model-list', ModelTypeEnum.textGeneration], }) diff --git a/web/app/components/header/account-setting/model-provider-page/__tests__/index.spec.tsx b/web/app/components/header/account-setting/model-provider-page/__tests__/index.spec.tsx index 14e94d69d18..3cef8326aa3 100644 --- a/web/app/components/header/account-setting/model-provider-page/__tests__/index.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/__tests__/index.spec.tsx @@ -49,6 +49,18 @@ const { mockReferenceSetting, mockAutoUpgradeError } = vi.hoisted(() => ({ const { mockProviderContextState, mockRefreshModelProviders } = vi.hoisted(() => ({ mockProviderContextState: { isLoadingModelProviders: false, + isSuccessModelProviders: true, + modelProviderPlugins: {} as Record< + string, + { + installation_id: string + plugin_id: string + plugin_unique_identifier: string + runtime_type: string + source: 'github' | 'marketplace' | 'package' | 'remote' + version: string + } + >, }, mockRefreshModelProviders: vi.fn(), })) @@ -159,9 +171,22 @@ const createPluginDetail = (overrides: Partial = {}): PluginDetail } } -const mockProviders = [ +type MockProvider = { + provider: string + plugin_id?: string + label: { en_US: string } + custom_configuration: { status: CustomConfigurationStatusEnum } + system_configuration: { + enabled: boolean + current_quota_type: CurrentSystemQuotaTypeEnum + quota_configurations: (typeof mockQuotaConfig)[] + } +} + +const mockProviders: MockProvider[] = [ { provider: 'openai', + plugin_id: 'langgenius/openai', label: { en_US: 'OpenAI' }, custom_configuration: { status: CustomConfigurationStatusEnum.active }, system_configuration: { @@ -172,6 +197,7 @@ const mockProviders = [ }, { provider: 'anthropic', + plugin_id: 'langgenius/anthropic', label: { en_US: 'Anthropic' }, custom_configuration: { status: CustomConfigurationStatusEnum.noConfigure }, system_configuration: { @@ -184,8 +210,15 @@ const mockProviders = [ vi.mock('@/context/provider-context', () => ({ useProviderContext: () => ({ - modelProviders: mockProviders, + modelProviders: mockProviders.map((provider) => ({ + ...provider, + is_configured: + provider.custom_configuration.status === CustomConfigurationStatusEnum.active || + provider.system_configuration.enabled, + })), + modelProviderPlugins: mockProviderContextState.modelProviderPlugins, isLoadingModelProviders: mockProviderContextState.isLoadingModelProviders, + isSuccessModelProviders: mockProviderContextState.isSuccessModelProviders, refreshModelProviders: mockRefreshModelProviders, }), })) @@ -211,17 +244,17 @@ vi.mock('../provider-added-card', () => ({ default: ({ notConfigured, provider, - pluginDetail, + pluginSummary, }: { notConfigured?: boolean provider: { provider: string } - pluginDetail?: { plugin_id: string; source?: string } + pluginSummary?: { plugin_id: string; source?: string } }) => (
{provider.provider}
@@ -355,6 +388,8 @@ describe('ModelProviderPage', () => { mockRefreshModelProviders.mockClear() mockInstalledModelPlugins.value = [] mockProviderContextState.isLoadingModelProviders = false + mockProviderContextState.isSuccessModelProviders = true + mockProviderContextState.modelProviderPlugins = {} mockAutoUpgradeError.value = undefined mockReferenceSetting.auto_upgrade = { strategy_setting: 'latest', @@ -371,6 +406,7 @@ describe('ModelProviderPage', () => { mockProviders.length, { provider: 'openai', + plugin_id: 'langgenius/openai', label: { en_US: 'OpenAI' }, custom_configuration: { status: CustomConfigurationStatusEnum.active }, system_configuration: { @@ -381,6 +417,7 @@ describe('ModelProviderPage', () => { }, { provider: 'anthropic', + plugin_id: 'langgenius/anthropic', label: { en_US: 'Anthropic' }, custom_configuration: { status: CustomConfigurationStatusEnum.noConfigure }, system_configuration: { @@ -531,9 +568,10 @@ describe('ModelProviderPage', () => { expect(target).toContainElement(screen.getByText('common.modelProvider.emptyProviderTitle')) }) - it('should use the model plugin installation list to attach plugin detail to provider cards', () => { + it('should use the summary plugin map to attach plugin metadata to provider cards', () => { mockProviders.splice(0, mockProviders.length, { provider: 'langgenius/openai/openai', + plugin_id: 'langgenius/openai-marketplace', label: { en_US: 'OpenAI' }, custom_configuration: { status: CustomConfigurationStatusEnum.active }, system_configuration: { @@ -542,25 +580,23 @@ describe('ModelProviderPage', () => { quota_configurations: [mockQuotaConfig], }, }) - mockInstalledModelPlugins.value = [ - createPluginDetail({ - plugin_id: 'langgenius/openai', - declaration: createPluginDeclaration({ - plugin_unique_identifier: 'langgenius/openai:1.0.0', - name: 'openai', - label: { en_US: 'OpenAI Plugin' } as unknown as PluginDeclaration['label'], - }), - }), - ] + mockProviderContextState.modelProviderPlugins = { + 'langgenius/openai-marketplace': { + installation_id: 'openai-installation', + plugin_id: 'langgenius/openai-marketplace', + plugin_unique_identifier: 'langgenius/openai:1.0.0', + runtime_type: 'local', + source: 'marketplace', + version: '1.0.0', + }, + } renderModelProviderPage() - expect(mockUseInstalledPluginList).toHaveBeenCalledWith(false, 100, { - category: PluginCategoryEnum.model, - }) + expect(mockUseInstalledPluginList).not.toHaveBeenCalled() expect(screen.getByTestId('provider-card')).toHaveAttribute( 'data-plugin-id', - 'langgenius/openai', + 'langgenius/openai-marketplace', ) expect(screen.queryByText('OpenAI Plugin')).not.toBeInTheDocument() }) @@ -587,24 +623,27 @@ describe('ModelProviderPage', () => { ).not.toBeInTheDocument() }) - it('should refresh model providers once when a debugging model plugin is missing from providers', () => { - mockInstalledModelPlugins.value = [ - createPluginDetail({ + it('should not refresh providers when remote plugin metadata already comes from summary', () => { + mockProviderContextState.modelProviderPlugins = { + 'langgenius/debug-model': { + installation_id: 'debug-installation', plugin_id: 'langgenius/debug-model', - declaration: createPluginDeclaration({ - label: { en_US: 'Debug Model' } as unknown as PluginDeclaration['label'], - }), - }), - ] + plugin_unique_identifier: 'langgenius/debug-model:1.0.0', + runtime_type: 'remote', + source: 'remote', + version: '1.0.0', + }, + } renderModelProviderPage() - expect(mockRefreshModelProviders).toHaveBeenCalledTimes(1) + expect(mockRefreshModelProviders).not.toHaveBeenCalled() }) - it('should prefer debugging plugin detail when an installed model plugin shares the same plugin id', () => { + it('should render remote source from the authoritative summary plugin entry', () => { mockProviders.splice(0, mockProviders.length, { provider: 'langgenius/openai/openai', + plugin_id: 'langgenius/openai', label: { en_US: 'OpenAI' }, custom_configuration: { status: CustomConfigurationStatusEnum.active }, system_configuration: { @@ -613,26 +652,16 @@ describe('ModelProviderPage', () => { quota_configurations: [mockQuotaConfig], }, }) - mockInstalledModelPlugins.value = [ - createPluginDetail({ + mockProviderContextState.modelProviderPlugins = { + 'langgenius/openai': { + installation_id: 'openai-debug-installation', plugin_id: 'langgenius/openai', - declaration: createPluginDeclaration({ - plugin_unique_identifier: 'langgenius/openai:debug', - name: 'openai', - label: { en_US: 'OpenAI Debug Plugin' } as unknown as PluginDeclaration['label'], - }), - source: PluginSource.debugging, - }), - createPluginDetail({ - plugin_id: 'langgenius/openai', - declaration: createPluginDeclaration({ - plugin_unique_identifier: 'langgenius/openai:1.0.0', - name: 'openai', - label: { en_US: 'OpenAI Installed Plugin' } as unknown as PluginDeclaration['label'], - }), - source: PluginSource.marketplace, - }), - ] + plugin_unique_identifier: 'langgenius/openai:debug', + runtime_type: 'remote', + source: 'remote', + version: '1.0.0', + }, + } renderModelProviderPage() @@ -640,18 +669,17 @@ describe('ModelProviderPage', () => { 'data-plugin-id', 'langgenius/openai', ) - expect(screen.getByTestId('provider-card')).toHaveAttribute( - 'data-plugin-source', - PluginSource.debugging, - ) - expect(mockRefreshModelProviders).toHaveBeenCalledTimes(1) + expect(screen.getByTestId('provider-card')).toHaveAttribute('data-plugin-source', 'remote') + expect(mockRefreshModelProviders).not.toHaveBeenCalled() }) it('should show provider placeholders while model providers are loading', () => { mockProviderContextState.isLoadingModelProviders = true + mockProviderContextState.isSuccessModelProviders = false renderModelProviderPage() + expect(mockUseInstalledPluginList).not.toHaveBeenCalled() expect(screen.getByRole('status', { name: 'common.loading' })).toBeInTheDocument() expect(screen.queryByTestId('provider-card')).not.toBeInTheDocument() expect(screen.queryByTestId('install-from-marketplace')).not.toBeInTheDocument() @@ -857,6 +885,7 @@ describe('ModelProviderPage', () => { }, { provider: 'langgenius/debug-model/debug-model', + plugin_id: 'langgenius/debug-model', label: { en_US: 'Debug Model' }, custom_configuration: { status: CustomConfigurationStatusEnum.noConfigure }, system_configuration: { @@ -866,16 +895,16 @@ describe('ModelProviderPage', () => { }, }, ) - mockInstalledModelPlugins.value = [ - createPluginDetail({ + mockProviderContextState.modelProviderPlugins = { + 'langgenius/debug-model': { + installation_id: 'debug-installation', plugin_id: 'langgenius/debug-model', - declaration: createPluginDeclaration({ - plugin_unique_identifier: 'langgenius/debug-model:1.0.0', - name: 'debug-model', - label: { en_US: 'Debug Model' } as unknown as PluginDeclaration['label'], - }), - }), - ] + plugin_unique_identifier: 'langgenius/debug-model:1.0.0', + runtime_type: 'remote', + source: 'remote', + version: '1.0.0', + }, + } renderModelProviderPage() diff --git a/web/app/components/header/account-setting/model-provider-page/__tests__/install-from-marketplace.spec.tsx b/web/app/components/header/account-setting/model-provider-page/__tests__/install-from-marketplace.spec.tsx index 1cf527fc9da..f5a79f8ec7d 100644 --- a/web/app/components/header/account-setting/model-provider-page/__tests__/install-from-marketplace.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/__tests__/install-from-marketplace.spec.tsx @@ -1,14 +1,26 @@ import type { Mock } from 'vitest' -import type { ModelProvider } from '../declarations' -import { fireEvent, render, screen } from '@testing-library/react' -import { describe, expect, it, vi } from 'vitest' +import { act, fireEvent, screen } from '@testing-library/react' +import { afterEach, describe, expect, it, vi } from 'vitest' import { getStepByStepTourTargetSelector, STEP_BY_STEP_TOUR_TARGETS, } from '@/app/components/step-by-step-tour/target-registry' +import { consoleQuery } from '@/service/client' +import { createConsoleQueryClient, renderWithConsoleQuery } from '@/test/console/query-data' import { useMarketplaceAllPlugins } from '../hooks' import InstallFromMarketplace from '../install-from-marketplace' +const render = (ui: React.ReactElement) => { + const queryClient = createConsoleQueryClient() + queryClient.setQueryData( + consoleQuery.workspaces.current.plugin.installedIds.get.queryKey({ + input: { query: { category: 'model' } }, + }), + { plugin_ids: [] }, + ) + return renderWithConsoleQuery(ui, { queryClient }) +} + // Mock dependencies vi.mock('@/next/link', () => ({ default: ({ children, href }: { children: React.ReactNode; href: string }) => ( @@ -76,20 +88,53 @@ vi.mock('@/app/components/plugins/plugin-page/use-reference-setting', () => ({ })) describe('InstallFromMarketplace', () => { - const mockProviders = [] as ModelProvider[] - beforeEach(() => { vi.clearAllMocks() }) + afterEach(() => { + vi.unstubAllGlobals() + }) + + it('should wait until the section approaches the viewport before loading marketplace data', () => { + let intersectionCallback: IntersectionObserverCallback | undefined + const disconnect = vi.fn() + function MockIntersectionObserver(callback: IntersectionObserverCallback) { + intersectionCallback = callback + return { + disconnect, + observe: vi.fn(), + takeRecords: vi.fn(), + unobserve: vi.fn(), + root: null, + thresholds: [0], + } + } + vi.stubGlobal('IntersectionObserver', vi.fn(MockIntersectionObserver)) + + render() + + expect(useMarketplaceAllPlugins).toHaveBeenLastCalledWith('', [], false) + + act(() => { + intersectionCallback?.( + [{ isIntersecting: true } as IntersectionObserverEntry], + {} as IntersectionObserver, + ) + }) + + expect(useMarketplaceAllPlugins).toHaveBeenLastCalledWith('', [], true) + expect(disconnect).toHaveBeenCalled() + }) + it('should render expanded by default', () => { - render() + render() expect(screen.getByText('common.modelProvider.installProvider')).toBeInTheDocument() expect(screen.getByTestId('plugin-list')).toBeInTheDocument() }) it('should collapse when clicked', () => { - render() + render() const toggle = screen.getByRole('button', { name: /common\.modelProvider\.installProvider/ }) fireEvent.click(toggle) @@ -107,7 +152,7 @@ describe('InstallFromMarketplace', () => { isLoading: true, }) - render() + render() // It's expanded by default, so loading should show immediately expect(screen.getByTestId('loading')).toBeInTheDocument() }) @@ -118,7 +163,7 @@ describe('InstallFromMarketplace', () => { isLoading: false, }) - render() + render() // Expanded by default expect(screen.getByText('Plugin 1')).toBeInTheDocument() }) @@ -136,7 +181,6 @@ describe('InstallFromMarketplace', () => { render( , @@ -164,14 +208,14 @@ describe('InstallFromMarketplace', () => { isLoading: false, }) - render() + render() expect(screen.getByText('Plugin 1')).toBeInTheDocument() expect(screen.queryByText('Bundle 1')).not.toBeInTheDocument() }) it('should render discovery link', () => { - render() + render() expect(screen.getByText('plugin.marketplace.difyMarketplace')).toHaveAttribute( 'href', 'https://marketplace.test/plugins/model?theme=light', @@ -181,13 +225,7 @@ describe('InstallFromMarketplace', () => { it('should use the marketplace callback action when provided', () => { const onOpenMarketplace = vi.fn() - render( - , - ) + render() fireEvent.click(screen.getByRole('button', { name: 'plugin.marketplace.difyMarketplace' })) diff --git a/web/app/components/header/account-setting/model-provider-page/__tests__/utils.spec.ts b/web/app/components/header/account-setting/model-provider-page/__tests__/utils.spec.ts index 33f556cc9c2..698e14d170b 100644 --- a/web/app/components/header/account-setting/model-provider-page/__tests__/utils.spec.ts +++ b/web/app/components/header/account-setting/model-provider-page/__tests__/utils.spec.ts @@ -4,7 +4,6 @@ import { genModelNameFormSchema, genModelTypeFormSchema, modelTypeFormat, - providerToPluginId, sizeFormat, } from '../utils' @@ -23,16 +22,6 @@ describe('utils', () => { }) }) - describe('providerToPluginId', () => { - it('should return the plugin id prefix when the provider key contains a provider segment', () => { - expect(providerToPluginId('langgenius/openai/openai')).toBe('langgenius/openai') - }) - - it('should return an empty string when the provider key has no plugin prefix', () => { - expect(providerToPluginId('openai')).toBe('') - }) - }) - describe('modelTypeFormat', () => { it('should format text embedding type', () => { expect(modelTypeFormat(ModelTypeEnum.textEmbedding)).toBe('TEXT EMBEDDING') diff --git a/web/app/components/header/account-setting/model-provider-page/derive-model-status.ts b/web/app/components/header/account-setting/model-provider-page/derive-model-status.ts index 9079dba81fd..874eaefb8fa 100644 --- a/web/app/components/header/account-setting/model-provider-page/derive-model-status.ts +++ b/web/app/components/header/account-setting/model-provider-page/derive-model-status.ts @@ -1,3 +1,4 @@ +import type { ModelProviderSummaryResponse } from '@dify/contracts/api/console/workspaces/types.gen' import type { Model, ModelItem, ModelProvider } from './declarations' import type { CredentialPanelState } from './provider-added-card/use-credential-panel-state' import { ModelStatusEnum } from './declarations' @@ -28,7 +29,7 @@ export const DERIVED_MODEL_STATUS_TOOLTIP_I18N = { export const deriveModelStatus = ( modelId: string | undefined, providerName: string | undefined, - currentModelProvider: ModelProvider | Model | undefined, + currentModelProvider: ModelProvider | ModelProviderSummaryResponse | Model | undefined, currentModel: ModelItem | undefined, credentialState: CredentialPanelState, ): DerivedModelStatus => { diff --git a/web/app/components/header/account-setting/model-provider-page/hooks.ts b/web/app/components/header/account-setting/model-provider-page/hooks.ts index b4a0ee81370..5edc03781ec 100644 --- a/web/app/components/header/account-setting/model-provider-page/hooks.ts +++ b/web/app/components/header/account-setting/model-provider-page/hooks.ts @@ -1,3 +1,4 @@ +import type { ModelType } from '@dify/contracts/api/console/workspaces/types.gen' import type { ConfigurationMethodEnum, Credential, @@ -10,6 +11,7 @@ import type { ModelProvider, ModelTypeEnum, } from './declarations' +import type { ModelModalType } from '@/context/modal-context' import { useQuery, useQueryClient } from '@tanstack/react-query' import { useCallback, useEffect, useMemo, useState } from 'react' import { @@ -22,7 +24,7 @@ import { useModalContextSelector } from '@/context/modal-context' import { useProviderContext } from '@/context/provider-context' import { consoleQuery } from '@/service/client' import { fetchDefaultModal, fetchModelList } from '@/service/common' -import { commonQueryKeys } from '@/service/use-common' +import { commonQueryKeys, modelProviderDetailsQueryOptions } from '@/service/use-common' import { useExpandModelProviderList } from './atoms' import { CustomConfigurationStatusEnum, ModelStatusEnum } from './declarations' @@ -74,10 +76,16 @@ export const useLanguage = () => { const locale = useLocale() return locale.replace('-', '_') } -export const useModelList = (type: ModelTypeEnum) => { + +type UseModelListOptions = { + enabled?: boolean +} + +export const useModelList = (type: ModelTypeEnum, { enabled = true }: UseModelListOptions = {}) => { const { data, refetch, isPending } = useQuery({ queryKey: commonQueryKeys.modelList(type), queryFn: () => fetchModelList(`/workspaces/current/models/model-types/${type}`), + enabled, }) return { @@ -100,7 +108,7 @@ export const useDefaultModel = (type: ModelTypeEnum) => { } } -export const useCurrentProviderAndModel = (modelList: Model[], defaultModel?: DefaultModel) => { +export const getCurrentProviderAndModel = (modelList: Model[], defaultModel?: DefaultModel) => { const currentProvider = modelList.find((provider) => provider.provider === defaultModel?.provider) const currentModel = currentProvider?.models.find((model) => model.model === defaultModel?.model) @@ -110,6 +118,8 @@ export const useCurrentProviderAndModel = (modelList: Model[], defaultModel?: De } } +export { getCurrentProviderAndModel as useCurrentProviderAndModel } + export const useTextGenerationCurrentProviderAndModelAndModelList = ( defaultModel?: DefaultModel, ) => { @@ -117,7 +127,7 @@ export const useTextGenerationCurrentProviderAndModelAndModelList = ( const activeTextGenerationModelList = textGenerationModelList.filter( (model) => model.status === ModelStatusEnum.active, ) - const { currentProvider, currentModel } = useCurrentProviderAndModel( + const { currentProvider, currentModel } = getCurrentProviderAndModel( textGenerationModelList, defaultModel, ) @@ -142,7 +152,7 @@ export const useModelListAndDefaultModel = (type: ModelTypeEnum) => { export const useModelListAndDefaultModelAndCurrentProviderAndModel = (type: ModelTypeEnum) => { const { modelList, defaultModel } = useModelListAndDefaultModel(type) - const { currentProvider, currentModel } = useCurrentProviderAndModel(modelList, { + const { currentProvider, currentModel } = getCurrentProviderAndModel(modelList, { provider: defaultModel?.provider.provider || '', model: defaultModel?.model || '', }) @@ -159,7 +169,7 @@ export const useUpdateModelList = () => { const queryClient = useQueryClient() const updateModelList = useCallback( - (type: ModelTypeEnum) => { + (type: ModelTypeEnum | ModelType) => { queryClient.invalidateQueries({ queryKey: commonQueryKeys.modelList(type) }) }, [queryClient], @@ -182,20 +192,48 @@ export const useUpdateModelProviders = () => { const queryClient = useQueryClient() const updateModelProviders = useCallback(() => { - queryClient.invalidateQueries({ queryKey: commonQueryKeys.modelProviders }) + queryClient.invalidateQueries({ + queryKey: consoleQuery.workspaces.current.modelProviders.summary.get.key(), + }) + queryClient.invalidateQueries({ queryKey: commonQueryKeys.modelProviderDetails }) }, [queryClient]) return updateModelProviders } +export const useLazyModelProviderDetail = (providerName: string) => { + const [enabled, setEnabled] = useState(false) + const queryClient = useQueryClient() + const { data, isFetching } = useQuery({ + ...modelProviderDetailsQueryOptions(), + enabled, + }) + const providerDetail = data?.data.find((provider) => provider.provider === providerName) + + const loadProviderDetail = useCallback(async () => { + setEnabled(true) + try { + const response = await queryClient.fetchQuery(modelProviderDetailsQueryOptions()) + return response.data.find((provider) => provider.provider === providerName) + } catch { + return undefined + } + }, [providerName, queryClient]) + + return { + providerDetail, + loadProviderDetail, + isProviderDetailEnabled: enabled, + isLoadingProviderDetail: enabled && isFetching, + } +} + export const useMarketplaceAllPlugins = ( - providers: ModelProvider[], searchText: string, + installedPluginIds: string[], enabled = true, ) => { - const exclude = useMemo(() => { - return providers.map((provider) => provider.provider.replace(/(.+)\/([^/]+)$/, '$1')) - }, [providers]) + const exclude = installedPluginIds const { plugins: collectionPlugins = [], isLoading: isCollectionLoading } = useMarketplacePluginsByCollectionId(enabled ? '__model-settings-pinned-models' : undefined) const { @@ -203,14 +241,12 @@ export const useMarketplaceAllPlugins = ( queryPlugins, queryPluginsWithDebounced, cancelQueryPluginsWithDebounced = () => {}, - resetPlugins = () => {}, isLoading: isPluginsLoading, - } = useMarketplacePlugins() + } = useMarketplacePlugins(enabled) useEffect(() => { if (!enabled) { cancelQueryPluginsWithDebounced() - resetPlugins() return } @@ -239,7 +275,6 @@ export const useMarketplaceAllPlugins = ( enabled, queryPlugins, queryPluginsWithDebounced, - resetPlugins, searchText, exclude, ]) @@ -253,7 +288,11 @@ export const useMarketplaceAllPlugins = ( for (let i = 0; i < plugins.length; i++) { const plugin = plugins[i] - if (plugin!.type !== 'bundle' && !allPlugins.find((p) => p.plugin_id === plugin!.plugin_id)) + if ( + !exclude.includes(plugin!.plugin_id) && + plugin!.type !== 'bundle' && + !allPlugins.find((p) => p.plugin_id === plugin!.plugin_id) + ) allPlugins.push(plugin!) } } @@ -262,7 +301,10 @@ export const useMarketplaceAllPlugins = ( }, [enabled, plugins, collectionPlugins, exclude]) return { - plugins: enabled && searchText ? plugins : allPlugins, + plugins: + enabled && searchText + ? plugins?.filter((plugin) => !exclude.includes(plugin.plugin_id)) + : allPlugins, isLoading: enabled && (isCollectionLoading || isPluginsLoading), } } @@ -332,7 +374,7 @@ export const useModelModalHandler = () => { isModelCredential?: boolean credential?: Credential model?: CustomModel - onUpdate?: (newPayload: any, formValues?: Record) => void + onUpdate?: (newPayload?: ModelModalType, formValues?: Record) => void mode?: ModelModalModeEnum } = {}, ) => { diff --git a/web/app/components/header/account-setting/model-provider-page/index.tsx b/web/app/components/header/account-setting/model-provider-page/index.tsx index e0c0717d806..07cbccd0152 100644 --- a/web/app/components/header/account-setting/model-provider-page/index.tsx +++ b/web/app/components/header/account-setting/model-provider-page/index.tsx @@ -1,20 +1,21 @@ +import type { + ModelProviderPluginSummaryResponse, + ModelProviderSummaryResponse, +} from '@dify/contracts/api/console/workspaces/types.gen' import type { ReactNode } from 'react' -import type { ModelProvider } from './declarations' -import type { PluginDetail } from '@/app/components/plugins/types' -import { useSuspenseQuery } from '@tanstack/react-query' +import { useQuery, useSuspenseQuery } from '@tanstack/react-query' import { useDebounce } from 'ahooks' import { noop } from 'es-toolkit/function' -import { useEffect, useMemo, useRef } from 'react' +import { useMemo } from 'react' import { useTranslation } from 'react-i18next' import { SearchInput } from '@/app/components/base/search-input' -import { usePluginsWithLatestVersion } from '@/app/components/plugins/hooks' import { usePluginSettingsAccess } from '@/app/components/plugins/plugin-page/use-reference-setting' -import { PluginCategoryEnum, PluginSource } from '@/app/components/plugins/types' +import { PluginCategoryEnum } from '@/app/components/plugins/types' import { useProviderContext } from '@/context/provider-context' import { systemFeaturesQueryOptions } from '@/features/system-features/client' -import { useInstalledPluginList } from '@/service/use-plugins' +import { consoleQuery } from '@/service/client' import UpdateSettingDialog from '../update-setting-dialog' -import { CustomConfigurationStatusEnum, ModelTypeEnum } from './declarations' +import { ModelTypeEnum } from './declarations' import { useDefaultModel } from './hooks' import ModelProviderPageBody from './model-provider-page-body' import SystemModelSelector from './system-model-selector' @@ -36,6 +37,11 @@ type Props = Readonly<{ const FixedModelProvider = ['langgenius/openai/openai', 'langgenius/anthropic/anthropic'] +export type ModelProviderPluginSummary = ModelProviderPluginSummaryResponse & { + latestVersion?: string + latestUniqueIdentifier?: string +} + const ModelProviderPage = ({ layout, onOpenMarketplace, @@ -61,44 +67,36 @@ const ModelProviderPage = ({ ) const { modelProviders: providers, + modelProviderPlugins = {}, isLoadingModelProviders, - refreshModelProviders, } = useProviderContext() const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) - const { data: installedModelPlugins } = useInstalledPluginList(false, 100, { - category: PluginCategoryEnum.model, - }) - const enrichedPlugins = usePluginsWithLatestVersion(installedModelPlugins?.plugins) - const pluginDetailMap = useMemo(() => { - const map = new Map() - for (const plugin of enrichedPlugins) { - const existingPlugin = map.get(plugin.plugin_id) - if (!existingPlugin || plugin.source === PluginSource.debugging) - map.set(plugin.plugin_id, plugin) + const marketplacePluginIds = useMemo( + () => + Object.values(modelProviderPlugins) + .filter((plugin) => plugin.source === 'marketplace') + .map((plugin) => plugin.plugin_id), + [modelProviderPlugins], + ) + const { data: latestVersionData } = useQuery( + consoleQuery.workspaces.current.plugin.list.latestVersions.post.queryOptions({ + input: { body: { plugin_ids: marketplacePluginIds } }, + enabled: !!marketplacePluginIds.length, + }), + ) + const pluginSummaryMap = useMemo(() => { + const map = new Map() + for (const plugin of Object.values(modelProviderPlugins)) { + const latestVersion = latestVersionData?.versions[plugin.plugin_id] + map.set(plugin.plugin_id, { + ...plugin, + latestVersion: latestVersion?.version, + latestUniqueIdentifier: latestVersion?.unique_identifier, + }) } return map - }, [enrichedPlugins]) - const debuggingModelPluginKey = useMemo(() => { - const debuggingModelPluginIds = enrichedPlugins - .filter((plugin) => plugin.source === PluginSource.debugging) - .map((plugin) => `${plugin.plugin_id}:${plugin.plugin_unique_identifier}`) - .sort() - - return debuggingModelPluginIds.join(',') - }, [enrichedPlugins]) - const refreshedDebuggingModelPluginKeyRef = useRef('') - useEffect(() => { - if (!debuggingModelPluginKey) { - refreshedDebuggingModelPluginKeyRef.current = '' - return - } - - if (refreshedDebuggingModelPluginKeyRef.current === debuggingModelPluginKey) return - - refreshedDebuggingModelPluginKeyRef.current = debuggingModelPluginKey - refreshModelProviders?.() - }, [debuggingModelPluginKey, refreshModelProviders]) + }, [latestVersionData, modelProviderPlugins]) const enableMarketplace = systemFeatures.enable_marketplace const isDefaultModelLoading = isTextGenerationDefaultModelLoading || @@ -107,17 +105,11 @@ const ModelProviderPage = ({ isSpeech2textDefaultModelLoading || isTTSDefaultModelLoading const [configuredProviders, notConfiguredProviders] = useMemo(() => { - const configuredProviders: ModelProvider[] = [] - const notConfiguredProviders: ModelProvider[] = [] + const configuredProviders: ModelProviderSummaryResponse[] = [] + const notConfiguredProviders: ModelProviderSummaryResponse[] = [] providers.forEach((provider) => { - if ( - provider.custom_configuration.status === CustomConfigurationStatusEnum.active || - (provider.system_configuration.enabled === true && - provider.system_configuration.quota_configurations.some( - (item) => item.quota_type === provider.system_configuration.current_quota_type, - )) - ) { + if (provider.is_configured) { configuredProviders.push(provider) } else { notConfiguredProviders.push(provider) @@ -183,14 +175,14 @@ const ModelProviderPage = ({ (provider) => provider.provider.toLowerCase().includes(debouncedSearchText.toLowerCase()) || Object.values(provider.label).some((text) => - text.toLowerCase().includes(debouncedSearchText.toLowerCase()), + text?.toLowerCase().includes(debouncedSearchText.toLowerCase()), ), ) const filteredNotConfiguredProviders = notConfiguredProviders.filter( (provider) => provider.provider.toLowerCase().includes(debouncedSearchText.toLowerCase()) || Object.values(provider.label).some((text) => - text.toLowerCase().includes(debouncedSearchText.toLowerCase()), + text?.toLowerCase().includes(debouncedSearchText.toLowerCase()), ), ) @@ -257,7 +249,7 @@ const ModelProviderPage = ({ showMarketplace={showMarketplace} enableMarketplace={enableMarketplace} searchText={searchText} - pluginDetailMap={pluginDetailMap} + pluginSummaryMap={pluginSummaryMap} onOpenMarketplace={onOpenMarketplace} /> ) diff --git a/web/app/components/header/account-setting/model-provider-page/install-from-marketplace.tsx b/web/app/components/header/account-setting/model-provider-page/install-from-marketplace.tsx index 82f78121b5b..21d7b5e9fbb 100644 --- a/web/app/components/header/account-setting/model-provider-page/install-from-marketplace.tsx +++ b/web/app/components/header/account-setting/model-provider-page/install-from-marketplace.tsx @@ -1,8 +1,8 @@ -import type { ModelProvider } from './declarations' import type { Plugin } from '@/app/components/plugins/types' import { cn } from '@langgenius/dify-ui/cn' +import { useQuery } from '@tanstack/react-query' import { useTheme } from 'next-themes' -import { useCallback, useState } from 'react' +import { useCallback, useEffect, useRef, useState } from 'react' import { useTranslation } from 'react-i18next' import Divider from '@/app/components/base/divider' import Loading from '@/app/components/base/loading' @@ -12,17 +12,16 @@ import { usePluginSettingsAccess } from '@/app/components/plugins/plugin-page/us import ProviderCard from '@/app/components/plugins/provider-card' import { PluginCategoryEnum } from '@/app/components/plugins/types' import Link from '@/next/link' +import { consoleQuery } from '@/service/client' import { useMarketplaceAllPlugins } from './hooks' type InstallFromMarketplaceProps = { onOpenMarketplace?: () => void - providers: ModelProvider[] searchText: string stepByStepTourTarget?: string } const InstallFromMarketplace = ({ onOpenMarketplace, - providers, searchText, stepByStepTourTarget, }: InstallFromMarketplaceProps) => { @@ -30,10 +29,44 @@ const InstallFromMarketplace = ({ const { theme } = useTheme() const { canInstallPlugin } = usePluginSettingsAccess() const [collapse, setCollapse] = useState(false) - const { plugins: allPlugins, isLoading: isAllPluginsLoading } = useMarketplaceAllPlugins( - providers, - searchText, + const [hasEnteredViewport, setHasEnteredViewport] = useState( + () => !globalThis.IntersectionObserver, ) + const [hasBeenReopened, setHasBeenReopened] = useState(false) + const sectionRef = useRef(null) + const shouldLoadMarketplace = !collapse && (hasEnteredViewport || !!searchText || hasBeenReopened) + const { data: installedPluginIds, isSuccess: hasLoadedInstalledPluginIds } = useQuery({ + ...consoleQuery.workspaces.current.plugin.installedIds.get.queryOptions({ + input: { query: { category: 'model' } }, + enabled: shouldLoadMarketplace, + }), + select: (data) => data.plugin_ids, + }) + const { plugins: allPlugins, isLoading: isAllPluginsLoading } = useMarketplaceAllPlugins( + searchText, + installedPluginIds ?? [], + shouldLoadMarketplace && hasLoadedInstalledPluginIds, + ) + + useEffect(() => { + const section = sectionRef.current + if (!section || hasEnteredViewport) return + + const observer = new IntersectionObserver(([entry]) => { + if (!entry?.isIntersecting) return + setHasEnteredViewport(true) + observer.disconnect() + }) + observer.observe(section) + return () => observer.disconnect() + }, [hasEnteredViewport]) + + const handleToggle = () => { + setCollapse((previous) => { + if (previous) setHasBeenReopened(true) + return !previous + }) + } const cardRender = useCallback((plugin: Plugin) => { if (plugin.type === 'bundle') return null @@ -42,7 +75,11 @@ const InstallFromMarketplace = ({ }, []) return ( -
+
setCollapse((prev) => !prev)} + onClick={handleToggle} aria-expanded={!collapse} > @@ -86,8 +123,11 @@ const InstallFromMarketplace = ({ )}
- {!collapse && isAllPluginsLoading && } - {!isAllPluginsLoading && !collapse && ( + {!collapse && shouldLoadMarketplace && !hasLoadedInstalledPluginIds && ( + + )} + {!collapse && hasLoadedInstalledPluginIds && isAllPluginsLoading && } + {!isAllPluginsLoading && !collapse && hasLoadedInstalledPluginIds && ( { + it.each([ + { + name: 'common.operation.config', + props: {}, + }, + { + name: 'common.modelProvider.auth.authorizationError', + props: { loadBalancingInvalid: true }, + }, + ])('announces loading for the $name action', ({ name, props }) => { + render() + + const action = screen.getByRole('button', { name }) + expect(action).toHaveAttribute('aria-busy', 'true') + expect(action).toHaveAttribute('aria-disabled', 'true') + }) +}) diff --git a/web/app/components/header/account-setting/model-provider-page/model-auth/__tests__/manage-custom-model-credentials.spec.tsx b/web/app/components/header/account-setting/model-provider-page/model-auth/__tests__/manage-custom-model-credentials.spec.tsx index 4583efcc17b..f96164ebbae 100644 --- a/web/app/components/header/account-setting/model-provider-page/model-auth/__tests__/manage-custom-model-credentials.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/model-auth/__tests__/manage-custom-model-credentials.spec.tsx @@ -1,5 +1,6 @@ import type { ModelProvider } from '@/app/components/header/account-setting/model-provider-page/declarations' import { render, screen } from '@testing-library/react' +import userEvent from '@testing-library/user-event' import ManageCustomModelCredentials from '../manage-custom-model-credentials' // Mock hooks @@ -17,6 +18,8 @@ vi.mock('../authorized', () => ({ renderTrigger, items, popupTitle, + isOpen, + onOpenChange, }: { renderTrigger: (o?: boolean) => React.ReactNode items: Array<{ @@ -24,12 +27,17 @@ vi.mock('../authorized', () => ({ selectedCredential?: { credential_id?: string } }> popupTitle: string + isOpen?: boolean + onOpenChange?: (open: boolean) => void }) => (
{renderTrigger()}
{renderTrigger(true)}
{popupTitle}
{items.length}
+
{items.map((item, index) => ( { expect(screen.getByTestId('trigger-open')).toBeInTheDocument() }) + it('should forward controlled open state', async () => { + const user = userEvent.setup() + const onOpenChange = vi.fn() + mockUseCustomModels.mockReturnValue([{ model: 'gpt-4' }]) + + render( + , + ) + + const closeButton = screen.getByRole('button', { name: 'close' }) + expect(closeButton).toHaveAttribute('data-open', 'true') + await user.click(closeButton) + expect(onOpenChange).toHaveBeenCalledWith(false) + }) + it('should pass undefined selectedCredential when model has no current_credential_id', () => { const mockModels = [ { diff --git a/web/app/components/header/account-setting/model-provider-page/model-auth/add-custom-model.tsx b/web/app/components/header/account-setting/model-provider-page/model-auth/add-custom-model.tsx index 8eb3adbf9c7..4ffcc550992 100644 --- a/web/app/components/header/account-setting/model-provider-page/model-auth/add-custom-model.tsx +++ b/web/app/components/header/account-setting/model-provider-page/model-auth/add-custom-model.tsx @@ -26,9 +26,16 @@ const AddCustomModel = ({ provider, configurationMethod, currentCustomConfigurationModelFixedFields, + open: controlledOpen, + onOpenChange, }: AddCustomModelProps) => { const { t } = useTranslation() - const [open, setOpen] = useState(false) + const [localOpen, setLocalOpen] = useState(false) + const open = controlledOpen ?? localOpen + const setOpen = (nextOpen: boolean) => { + setLocalOpen(nextOpen) + onOpenChange?.(nextOpen) + } const canAddedModels = useCanAddedModels(provider) const noModels = !canAddedModels.length const { canUseCredential, canCreateCredential } = useCredentialPermissions() diff --git a/web/app/components/header/account-setting/model-provider-page/model-auth/config-model.tsx b/web/app/components/header/account-setting/model-provider-page/model-auth/config-model.tsx index a78ad53d0d8..14c32514537 100644 --- a/web/app/components/header/account-setting/model-provider-page/model-auth/config-model.tsx +++ b/web/app/components/header/account-setting/model-provider-page/model-auth/config-model.tsx @@ -7,12 +7,16 @@ import { useTranslation } from 'react-i18next' type ConfigModelProps = { onClick?: () => void + loading?: boolean + disabled?: boolean loadBalancingEnabled?: boolean loadBalancingInvalid?: boolean credentialRemoved?: boolean } const ConfigModel = ({ onClick, + loading, + disabled, loadBalancingEnabled, loadBalancingInvalid, credentialRemoved, @@ -21,14 +25,19 @@ const ConfigModel = ({ if (loadBalancingInvalid) { return ( -
{t(($) => $['modelProvider.auth.authorizationError'], { ns: 'common' })} -
+ ) } @@ -36,6 +45,9 @@ const ConfigModel = ({
) } diff --git a/web/app/components/header/account-setting/model-provider-page/model-parameter-modal/__tests__/agent-model-trigger.spec.tsx b/web/app/components/header/account-setting/model-provider-page/model-parameter-modal/__tests__/agent-model-trigger.spec.tsx index 0c3c773451f..556739641d3 100644 --- a/web/app/components/header/account-setting/model-provider-page/model-parameter-modal/__tests__/agent-model-trigger.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/model-parameter-modal/__tests__/agent-model-trigger.spec.tsx @@ -1,6 +1,6 @@ import type { MouseEvent } from 'react' import type { ModelProvider } from '../../declarations' -import { fireEvent, render, screen } from '@testing-library/react' +import { fireEvent, render, screen, waitFor } from '@testing-library/react' import { vi } from 'vitest' import { PluginCategoryEnum } from '@/app/components/plugins/types' import { @@ -19,6 +19,8 @@ const invalidateInstalledPluginList = vi.fn() const handleOpenModal = vi.fn() const updateModelProviders = vi.fn() const updateModelList = vi.fn() +const loadProviderDetail = vi.fn() +let isLoadingProviderDetail = false vi.mock('@/context/provider-context', () => ({ useProviderContext: () => ({ @@ -33,6 +35,10 @@ vi.mock('@/service/use-plugins', () => ({ })) vi.mock('../../hooks', () => ({ + useLazyModelProviderDetail: () => ({ + loadProviderDetail, + isLoadingProviderDetail, + }), useModelModalHandler: () => handleOpenModal, useUpdateModelList: () => updateModelList, useUpdateModelProviders: () => updateModelProviders, @@ -76,6 +82,7 @@ describe('AgentModelTrigger', () => { pluginInfo = null pluginLoading = false inModelList = true + isLoadingProviderDetail = false }) it('should render loading state when plugin info is still fetching', () => { @@ -127,6 +134,14 @@ describe('AgentModelTrigger', () => { expect(invalidateInstalledPluginList).toHaveBeenCalledWith(PluginCategoryEnum.model) }) + it('should not render the install action when the plugin has no package identifier', () => { + pluginInfo = { latest_package_identifier: '' } + + render() + + expect(screen.queryByText('Install Plugin')).not.toBeInTheDocument() + }) + it('should show configuration action when provider requires setup', () => { modelProviders = [ { @@ -145,6 +160,19 @@ describe('AgentModelTrigger', () => { expect(screen.getByText('workflow.nodes.agent.notAuthorized')).toBeInTheDocument() }) + it('should load complete provider detail before opening configuration', async () => { + const providerDetail = { provider: 'openai' } as ModelProvider + modelProviders = [{ provider: 'openai', is_configured: false }] as unknown as ModelProvider[] + loadProviderDetail.mockResolvedValue(providerDetail) + + render() + fireEvent.click(screen.getByRole('button', { name: /notAuthorized/ })) + + await waitFor(() => { + expect(handleOpenModal).toHaveBeenCalledWith(providerDetail, 'predefined-model', undefined) + }) + }) + it('should render unconfigured state when model is not selected', () => { render() expect(screen.getByText('workflow.nodes.agent.configureModel')).toBeInTheDocument() diff --git a/web/app/components/header/account-setting/model-provider-page/model-parameter-modal/__tests__/status-indicators.spec.tsx b/web/app/components/header/account-setting/model-provider-page/model-parameter-modal/__tests__/status-indicators.spec.tsx index 30e3e07da7f..0ad3eae95a7 100644 --- a/web/app/components/header/account-setting/model-provider-page/model-parameter-modal/__tests__/status-indicators.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/model-parameter-modal/__tests__/status-indicators.spec.tsx @@ -6,15 +6,12 @@ import { withSelectorKey } from '@/test/i18n-mock' import StatusIndicators from '../status-indicators' let installedPlugins = [{ name: 'demo-plugin', plugin_unique_identifier: 'demo@1.0.0' }] -const { mockUseInstalledPluginList } = vi.hoisted(() => ({ - mockUseInstalledPluginList: vi.fn(), +const mockUseInstalledPluginList = vi.fn((_options: unknown) => ({ + data: { plugins: installedPlugins }, })) vi.mock('@/service/use-plugins', () => ({ - useInstalledPluginList: (...args: unknown[]) => { - mockUseInstalledPluginList(...args) - return { data: { plugins: installedPlugins } } - }, + useInstalledPluginList: (options: unknown) => mockUseInstalledPluginList(options), })) vi.mock('@/app/components/workflow/nodes/_base/components/switch-plugin-version', () => ({ @@ -29,23 +26,7 @@ describe('StatusIndicators', () => { beforeEach(() => { vi.clearAllMocks() installedPlugins = [{ name: 'demo-plugin', plugin_unique_identifier: 'demo@1.0.0' }] - }) - - it('reads the model-specific installed plugin list', () => { - render( - , - ) - - expect(mockUseInstalledPluginList).toHaveBeenCalledWith(false, 100, { - category: PluginCategoryEnum.model, - }) + mockUseInstalledPluginList.mockReturnValue({ data: { plugins: installedPlugins } }) }) const getPopoverTrigger = (name: string) => { @@ -66,6 +47,10 @@ describe('StatusIndicators', () => { />, ) expect(container).toBeEmptyDOMElement() + expect(mockUseInstalledPluginList).toHaveBeenLastCalledWith({ + category: PluginCategoryEnum.model, + enabled: false, + }) }) it('should render deprecated tooltip when provider model is disabled and in model list', async () => { @@ -80,6 +65,10 @@ describe('StatusIndicators', () => { t={t} />, ) + expect(mockUseInstalledPluginList).toHaveBeenLastCalledWith({ + category: PluginCategoryEnum.model, + enabled: false, + }) await user.hover(getPopoverTrigger('nodes.agent.modelSelectorTooltips.deprecated')) @@ -100,6 +89,10 @@ describe('StatusIndicators', () => { t={t} />, ) + expect(mockUseInstalledPluginList).toHaveBeenLastCalledWith({ + category: PluginCategoryEnum.model, + enabled: false, + }) await user.hover(getPopoverTrigger('nodes.agent.modelNotSupport.title')) @@ -119,6 +112,10 @@ describe('StatusIndicators', () => { ) expect(screen.getByText('SwitchVersion:demo@1.0.0')).toBeInTheDocument() + expect(mockUseInstalledPluginList).toHaveBeenLastCalledWith({ + category: PluginCategoryEnum.model, + enabled: true, + }) }) it('should render nothing when needsConfiguration is true even with disabled and modelProvider', () => { @@ -133,6 +130,10 @@ describe('StatusIndicators', () => { />, ) expect(container).toBeEmptyDOMElement() + expect(mockUseInstalledPluginList).toHaveBeenLastCalledWith({ + category: PluginCategoryEnum.model, + enabled: false, + }) }) it('should render SwitchVersion with empty identifier when plugin is not in installed list', () => { diff --git a/web/app/components/header/account-setting/model-provider-page/model-parameter-modal/agent-model-trigger.tsx b/web/app/components/header/account-setting/model-provider-page/model-parameter-modal/agent-model-trigger.tsx index afae0d1e623..44efcbc92c3 100644 --- a/web/app/components/header/account-setting/model-provider-page/model-parameter-modal/agent-model-trigger.tsx +++ b/web/app/components/header/account-setting/model-provider-page/model-parameter-modal/agent-model-trigger.tsx @@ -13,8 +13,13 @@ import { useModelInList, usePluginInfo, } from '@/service/use-plugins' -import { CustomConfigurationStatusEnum, ModelTypeEnum } from '../declarations' -import { useModelModalHandler, useUpdateModelList, useUpdateModelProviders } from '../hooks' +import { ConfigurationMethodEnum, ModelTypeEnum } from '../declarations' +import { + useLazyModelProviderDetail, + useModelModalHandler, + useUpdateModelList, + useUpdateModelProviders, +} from '../hooks' import ModelIcon from '../model-icon' import ConfigurationButton from './configuration-button' import ModelDisplay from './model-display' @@ -47,14 +52,7 @@ const AgentModelTrigger: FC = ({ const updateModelList = useUpdateModelList() const { modelProvider, needsConfiguration } = useMemo(() => { const modelProvider = modelProviders.find((item) => item.provider === providerName) - const needsConfiguration = - modelProvider?.custom_configuration.status === CustomConfigurationStatusEnum.noConfigure && - !( - modelProvider.system_configuration.enabled === true && - modelProvider.system_configuration.quota_configurations.find( - (item) => item.quota_type === modelProvider.system_configuration.current_quota_type, - ) - ) + const needsConfiguration = modelProvider ? !modelProvider.is_configured : false return { modelProvider, needsConfiguration, @@ -63,10 +61,22 @@ const AgentModelTrigger: FC = ({ const [installed, setInstalled] = useState(false) const invalidateInstalledPluginList = useInvalidateInstalledPluginList() const handleOpenModal = useModelModalHandler() + const { loadProviderDetail, isLoadingProviderDetail } = useLazyModelProviderDetail( + providerName ?? '', + ) const { data: inModelList = false } = useModelInList(currentProvider, modelId) const { data: pluginInfo, isLoading: isPluginLoading } = usePluginInfo(providerName) + const handleConfigure = async () => { + if (!providerName) return + + const providerDetail = await loadProviderDetail() + if (!providerDetail) return + + handleOpenModal(providerDetail, ConfigurationMethodEnum.predefinedModel, undefined) + } + if (modelId && isPluginLoading) return return ( @@ -85,7 +95,7 @@ const AgentModelTrigger: FC = ({ /> {needsConfiguration && ( - + )} = ({ pluginInfo={pluginInfo} t={translateWorkflow} /> - {!installed && !modelProvider && pluginInfo && ( + {!installed && !modelProvider && pluginInfo?.latest_package_identifier && ( e.stopPropagation()} size="small" diff --git a/web/app/components/header/account-setting/model-provider-page/model-parameter-modal/configuration-button.tsx b/web/app/components/header/account-setting/model-provider-page/model-parameter-modal/configuration-button.tsx index 7d66ada99b1..efc2d519aac 100644 --- a/web/app/components/header/account-setting/model-provider-page/model-parameter-modal/configuration-button.tsx +++ b/web/app/components/header/account-setting/model-provider-page/model-parameter-modal/configuration-button.tsx @@ -1,21 +1,20 @@ import { Button } from '@langgenius/dify-ui/button' import { StatusDot } from '@langgenius/dify-ui/status-dot' import { useTranslation } from 'react-i18next' -import { ConfigurationMethodEnum } from '../declarations' - type ConfigurationButtonProps = { - modelProvider: any - handleOpenModal: any + loading: boolean + onConfigure: () => void } -const ConfigurationButton = ({ modelProvider, handleOpenModal }: ConfigurationButtonProps) => { +const ConfigurationButton = ({ loading, onConfigure }: ConfigurationButtonProps) => { const { t } = useTranslation() return ( - + {isUsingCredits ? ( hasCredits ? ( <> - + {t(($) => $['modelProvider.selector.aiCredits'], { ns: 'common' })} @@ -160,13 +178,15 @@ function PopupItem({ } /> - + {providerDetail && ( + + )}
diff --git a/web/app/components/header/account-setting/model-provider-page/model-selector/popup.tsx b/web/app/components/header/account-setting/model-provider-page/model-selector/popup.tsx index bd20c530e26..e2e6a02e96b 100644 --- a/web/app/components/header/account-setting/model-provider-page/model-selector/popup.tsx +++ b/web/app/components/header/account-setting/model-provider-page/model-selector/popup.tsx @@ -1,6 +1,6 @@ import type { DefaultModel, Model } from '../declarations' import type { ModelSelectorPreviewPayload } from './popup-item' -import type { ModelSelectorModelPredicate } from './types' +import type { ModelSelectorModelPredicate, ModelSelectorProvider } from './types' import type { ModelProviderQuotaGetPaid } from '@/types/model-provider' import { ComboboxList } from '@langgenius/dify-ui/combobox' import { @@ -20,17 +20,14 @@ import { import checkTaskStatus from '@/app/components/plugins/install-plugin/base/check-task-status' import useRefreshPluginList from '@/app/components/plugins/install-plugin/hooks/use-refresh-plugin-list' import useWorkspacePluginInstallPermission from '@/app/components/plugins/install-plugin/hooks/use-workspace-plugin-install-permission' +import { PluginCategoryEnum } from '@/app/components/plugins/types' import { useProviderContext } from '@/context/provider-context' import { systemFeaturesQueryOptions } from '@/features/system-features/client' import { consoleQuery } from '@/service/client' +import { fetchPluginInfoFromMarketPlace } from '@/service/plugins' import { useInstallPackageFromMarketPlace } from '@/service/use-plugins' -import { - CustomConfigurationStatusEnum, - ModelFeatureEnum, - ModelStatusEnum, - ModelTypeEnum, -} from '../declarations' -import { useLanguage, useMarketplaceAllPlugins } from '../hooks' +import { CustomConfigurationStatusEnum, ModelFeatureEnum, ModelTypeEnum } from '../declarations' +import { useLanguage } from '../hooks' import ModelBadge from '../model-badge' import ModelIcon from '../model-icon' import CreditsExhaustedAlert from '../provider-added-card/model-auth-dropdown/credits-exhausted-alert' @@ -89,10 +86,13 @@ function Popup({ ) const { theme } = useTheme() const language = useLanguage() - const [previewCardHandle] = useState(() => createPreviewCardHandle()) + const previewCardHandle = useMemo( + () => createPreviewCardHandle(), + [], + ) const [marketplaceCollapsed, setMarketplaceCollapsed] = useState(false) const [showIncompatibleModels, setShowIncompatibleModels] = useState(false) - const { modelProviders } = useProviderContext() + const { modelProviders, modelProviderPlugins = {} } = useProviderContext() const { data: enableMarketplace } = useSuspenseQuery({ ...systemFeaturesQueryOptions(), select: (systemFeatures) => systemFeatures.enable_marketplace, @@ -101,11 +101,6 @@ function Popup({ ...systemFeaturesQueryOptions(), select: ({ deployment_edition }) => deployment_edition, }) - const { plugins: allPlugins, isLoading: isMarketplacePluginsLoading } = useMarketplaceAllPlugins( - modelProviders, - '', - enableMarketplace, - ) const { mutateAsync: installPackageFromMarketPlace } = useInstallPackageFromMarketPlace() const { refreshPluginList } = useRefreshPluginList() const { canInstallPlugin } = useWorkspacePluginInstallPermission() @@ -141,71 +136,68 @@ function Popup({ const hasApiKeyFallback = modelProviders.some((provider) => { const isApiKeyActive = provider.custom_configuration?.status === CustomConfigurationStatusEnum.active - return isApiKeyActive && providerSupportsCredits(provider, trialModels, deploymentEdition) + return ( + isApiKeyActive && + provider.custom_configuration.current_credential_usable && + providerSupportsCredits(provider, trialModels, deploymentEdition) + ) }) const handleInstallPlugin = useCallback( async (key: ModelProviderQuotaGetPaid) => { - if ( - !enableMarketplace || - !canInstallPlugin || - !allPlugins || - isMarketplacePluginsLoading || - installingProvider - ) - return + if (!enableMarketplace || !canInstallPlugin || installingProvider) return const pluginId = providerKeyToPluginId[key] - const plugin = allPlugins.find((p) => p.plugin_id === pluginId) - if (!plugin) return - - const uniqueIdentifier = plugin.latest_package_identifier + const [org, name] = pluginId.split('/') + if (!org || !name) return setInstallingProvider(key) try { + const pluginInfo = await fetchPluginInfoFromMarketPlace({ org, name }) + const uniqueIdentifier = pluginInfo.data.plugin.latest_package_identifier + if (!uniqueIdentifier) return const { all_installed, task_id } = await installPackageFromMarketPlace(uniqueIdentifier) if (!all_installed) { const { check } = checkTaskStatus() await check({ taskId: task_id, pluginUniqueIdentifier: uniqueIdentifier }) } - refreshPluginList(plugin) + refreshPluginList({ category: PluginCategoryEnum.model }) } catch { } finally { setInstallingProvider(null) } }, [ - allPlugins, enableMarketplace, canInstallPlugin, installPackageFromMarketPlace, installingProvider, - isMarketplacePluginsLoading, refreshPluginList, ], ) const installedModelList = useMemo(() => { const modelMap = new Map(modelList.map((model) => [model.provider, model])) - const installedMarketplaceModels = MODEL_PROVIDER_QUOTA_GET_PAID.flatMap((providerKey) => { - const installedProvider = installedProviderMap.get(providerKey) + const installedMarketplaceModels = MODEL_PROVIDER_QUOTA_GET_PAID.flatMap( + (providerKey) => { + const installedProvider = installedProviderMap.get(providerKey) - if (!installedProvider) return [] + if (!installedProvider) return [] - const matchedModel = modelMap.get(providerKey) - if (matchedModel) return [matchedModel] + const matchedModel = modelMap.get(providerKey) + if (matchedModel) return [matchedModel] - if (!aiCreditVisibleProviders.has(providerKey)) return [] + if (!aiCreditVisibleProviders.has(providerKey)) return [] - return [ - { - provider: installedProvider.provider, - icon_small: installedProvider.icon_small, - icon_small_dark: installedProvider.icon_small_dark, - label: installedProvider.label, - models: [], - status: ModelStatusEnum.active, - }, - ] - }) + return [ + { + provider: installedProvider.provider, + icon_small: installedProvider.icon_small, + icon_small_dark: installedProvider.icon_small_dark, + label: installedProvider.label, + models: [], + }, + ] + }, + ) const otherModels = modelList.filter( (model) => !MODEL_PROVIDER_QUOTA_GET_PAID.includes(model.provider as ModelProviderQuotaGetPaid), @@ -245,9 +237,13 @@ function Popup({ const marketplaceProviders = useMemo(() => { if (!enableMarketplace) return [] - const installedProviders = new Set(modelProviders.map((provider) => provider.provider)) - return MODEL_PROVIDER_QUOTA_GET_PAID.filter((key) => !installedProviders.has(key)) - }, [enableMarketplace, modelProviders]) + const installedPluginIds = new Set( + Object.values(modelProviderPlugins).map((plugin) => plugin.plugin_id), + ) + return MODEL_PROVIDER_QUOTA_GET_PAID.filter( + (key) => !installedPluginIds.has(providerKeyToPluginId[key]), + ) + }, [enableMarketplace, modelProviderPlugins]) const handleOpenSettings = useCallback(() => { onHide() @@ -307,7 +303,6 @@ function Popup({ marketplaceProviders={marketplaceProviders} marketplaceCollapsed={marketplaceCollapsed} installingProvider={installingProvider} - isMarketplacePluginsLoading={isMarketplacePluginsLoading} canInstallPlugin={canInstallPlugin} theme={theme} onMarketplaceCollapsedChange={setMarketplaceCollapsed} @@ -322,7 +317,7 @@ function Popup({ $['model.capabilities'], { ns: 'common' })} language={language} - payload={payload} + payload={payload as ModelSelectorPreviewPayload | undefined} /> )} diff --git a/web/app/components/header/account-setting/model-provider-page/model-selector/types.ts b/web/app/components/header/account-setting/model-provider-page/model-selector/types.ts index ce2d6135a91..7b115157d0a 100644 --- a/web/app/components/header/account-setting/model-provider-page/model-selector/types.ts +++ b/web/app/components/header/account-setting/model-provider-page/model-selector/types.ts @@ -1,11 +1,23 @@ -import type { Model, ModelItem } from '../declarations' +import type { I18nObject } from '@dify/contracts/api/console/workspaces/types.gen' +import type { ModelItem } from '../declarations' export type ModelSelectorValue = { provider: string model: string } -export type ModelSelectorModelPredicate = (provider: Model, modelItem: ModelItem) => boolean +export type ModelSelectorProvider = { + provider: string + icon_small?: I18nObject | null + icon_small_dark?: I18nObject | null + label: I18nObject + models: ModelItem[] +} + +export type ModelSelectorModelPredicate = ( + provider: ModelSelectorProvider, + modelItem: ModelItem, +) => boolean export const isSameModelSelectorValue = ( itemValue: ModelSelectorValue, diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/index.spec.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/index.spec.tsx index f7a17b9611c..4a08bb1105b 100644 --- a/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/index.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/index.spec.tsx @@ -1,6 +1,6 @@ import type { ReactElement } from 'react' import type { ModelProvider } from '../../declarations' -import type { PluginDetail } from '@/app/components/plugins/types' +import type { ModelProviderPluginSummary } from '../../index' import { QueryClient } from '@tanstack/react-query' import { act, fireEvent, screen, waitFor } from '@testing-library/react' import { PluginCategoryEnum } from '@/app/components/plugins/types' @@ -229,15 +229,11 @@ describe('ProviderAddedCard', () => { renderWithQueryClient( , ) @@ -352,29 +348,44 @@ describe('ProviderAddedCard', () => { const customConfigProvider = { ...mockProvider, configurate_methods: [ConfigurationMethodEnum.customizableModel], + custom_configuration: { has_custom_models: true }, } as unknown as ModelProvider const { unmount } = renderWithQueryClient() - expect(screen.getByTestId('manage-custom-model')).toBeInTheDocument() - expect(screen.getByTestId('add-custom-model')).toBeInTheDocument() + expect( + screen.getByRole('button', { name: 'common.modelProvider.auth.manageCredentials' }), + ).toBeInTheDocument() + expect( + screen.getByRole('button', { name: 'common.modelProvider.addModel' }), + ).toBeInTheDocument() unmount() mockIsCurrentWorkspaceManager = false mockWorkspacePermissionKeys = ['credential.use', 'credential.create', 'credential.manage'] renderWithQueryClient() - expect(screen.queryByTestId('manage-custom-model')).not.toBeInTheDocument() + expect( + screen.queryByRole('button', { name: 'common.modelProvider.auth.manageCredentials' }), + ).not.toBeInTheDocument() + expect( + screen.queryByRole('button', { name: 'common.modelProvider.addModel' }), + ).not.toBeInTheDocument() }) it('should render custom model actions when user can configure models without credential permissions', () => { const customConfigProvider = { ...mockProvider, configurate_methods: [ConfigurationMethodEnum.customizableModel], + custom_configuration: { has_custom_models: false }, } as unknown as ModelProvider mockWorkspacePermissionKeys = ['plugin.model_config'] renderWithQueryClient() - expect(screen.getByTestId('manage-custom-model')).toBeInTheDocument() - expect(screen.getByTestId('add-custom-model')).toBeInTheDocument() + expect( + screen.queryByRole('button', { name: 'common.modelProvider.auth.manageCredentials' }), + ).not.toBeInTheDocument() + expect( + screen.getByRole('button', { name: 'common.modelProvider.addModel' }), + ).toBeInTheDocument() }) }) diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/lazy-custom-model-actions.spec.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/lazy-custom-model-actions.spec.tsx new file mode 100644 index 00000000000..60b6018c771 --- /dev/null +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/lazy-custom-model-actions.spec.tsx @@ -0,0 +1,90 @@ +import type { ModelProviderSummaryResponse } from '@dify/contracts/api/console/workspaces/types.gen' +import { screen } from '@testing-library/react' +import userEvent from '@testing-library/user-event' +import { render } from '@/test/console/render' +import LazyCustomModelActions from '../lazy-custom-model-actions' + +const { mockLoadProviderDetail, mockOpenModelModal, providerDetail } = vi.hoisted(() => { + const providerDetail = { + provider: 'langgenius/openai/openai', + custom_configuration: { + custom_models: [{ model: 'custom-gpt' }], + can_added_models: [{ model: 'custom-gpt-2', model_type: 'llm' }], + }, + } + + return { + mockLoadProviderDetail: vi.fn().mockResolvedValue(providerDetail), + mockOpenModelModal: vi.fn(), + providerDetail, + } +}) + +vi.mock('../../hooks', async () => { + const { useState } = await import('react') + + return { + useLazyModelProviderDetail: () => { + const [detail, setDetail] = useState() + + return { + providerDetail: detail, + loadProviderDetail: async () => { + const loadedDetail = await mockLoadProviderDetail() + setDetail(loadedDetail) + return loadedDetail + }, + isLoadingProviderDetail: false, + } + }, + useModelModalHandler: () => mockOpenModelModal, + } +}) + +vi.mock('@/app/components/header/account-setting/model-provider-page/model-auth', () => ({ + AddCustomModel: ({ open }: { open?: boolean }) => ( +
+ ), + ManageCustomModelCredentials: ({ isOpen }: { isOpen?: boolean }) => ( +
+ ), +})) + +const createProvider = (hasCustomModels: boolean) => + ({ + provider: 'langgenius/openai/openai', + custom_configuration: { + has_custom_models: hasCustomModels, + }, + }) as ModelProviderSummaryResponse + +describe('LazyCustomModelActions', () => { + beforeEach(() => { + vi.clearAllMocks() + }) + + it('loads provider detail and opens credential management on demand', async () => { + const user = userEvent.setup() + render() + + expect(mockLoadProviderDetail).not.toHaveBeenCalled() + + await user.click( + screen.getByRole('button', { name: 'common.modelProvider.auth.manageCredentials' }), + ) + + expect(await screen.findByTestId('manage-custom-model')).toHaveAttribute('data-open', 'true') + expect(mockLoadProviderDetail).toHaveBeenCalledTimes(1) + }) + + it('does not render credential management without custom models', () => { + render() + + expect( + screen.queryByRole('button', { name: 'common.modelProvider.auth.manageCredentials' }), + ).not.toBeInTheDocument() + expect( + screen.getByRole('button', { name: 'common.modelProvider.addModel' }), + ).toBeInTheDocument() + }) +}) diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/model-list.dynamic-import.spec.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/model-list.dynamic-import.spec.tsx new file mode 100644 index 00000000000..32ee443c2f2 --- /dev/null +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/model-list.dynamic-import.spec.tsx @@ -0,0 +1,75 @@ +import type { ModelItem, ModelProvider } from '../../declarations' +import { screen, waitFor } from '@testing-library/react' +import userEvent from '@testing-library/user-event' +import { renderWithConsoleQuery as render } from '@/test/console/query-data' +import ModelList from '../model-list' + +const { mockToastError } = vi.hoisted(() => ({ + mockToastError: vi.fn(), +})) + +vi.mock('@/context/permission-state', async () => { + const { createPermissionStateModuleMock } = await import('@/test/console/state-fixture') + return createPermissionStateModuleMock(() => ({ + workspacePermissionKeys: ['plugin.model_config'], + })) +}) + +vi.mock('../../hooks', () => ({ + useLazyModelProviderDetail: () => ({ + loadProviderDetail: vi.fn(), + }), +})) + +vi.mock('@langgenius/dify-ui/toast', () => ({ + toast: { + error: mockToastError, + }, +})) + +vi.mock('../model-load-balancing-modal', () => { + throw new Error('Failed to load model load balancing modal') +}) + +vi.mock('../model-list-item', () => ({ + default: ({ + model, + onModifyLoadBalancing, + }: { + model: ModelItem + onModifyLoadBalancing: (model: ModelItem) => void + }) => ( + + ), +})) + +describe('ModelList dynamic import failure', () => { + const provider = { + provider: 'test-provider', + configurate_methods: [], + } as unknown as ModelProvider + const model = { + model: 'gpt-4', + model_type: 'llm', + fetch_from: 'system', + } as unknown as ModelItem + + beforeEach(() => { + vi.clearAllMocks() + }) + + it('clears the loading dialog and reports an error when the modal module cannot load', async () => { + render() + + const user = userEvent.setup() + await user.click(screen.getByRole('button', { name: 'gpt-4' })) + + await waitFor(() => { + expect(mockToastError).toHaveBeenCalledWith('common.api.actionFailed') + }) + expect(screen.queryByRole('status')).not.toBeInTheDocument() + expect(screen.queryByRole('dialog')).not.toBeInTheDocument() + }) +}) diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/model-list.spec.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/model-list.spec.tsx index 31b71e4845c..cd34dcf3190 100644 --- a/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/model-list.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/model-list.spec.tsx @@ -1,12 +1,18 @@ +import type { ModelProviderSummaryResponse } from '@dify/contracts/api/console/workspaces/types.gen' import type { ModelItem, ModelProvider } from '../../declarations' -import type { ModelLoadBalancingModalProps } from '../model-load-balancing-modal' -import { fireEvent, screen, waitFor, within } from '@testing-library/react' +import { fireEvent, screen, waitFor } from '@testing-library/react' import userEvent from '@testing-library/user-event' -import * as React from 'react' -import { render } from '@/test/console/render' +import { renderWithConsoleQuery as render } from '@/test/console/query-data' import { ConfigurationMethodEnum } from '../../declarations' import ModelList from '../model-list' +const mockLoadProviderDetail = vi.fn() +const { mockLoadModelLoadBalancingModal } = vi.hoisted(() => ({ + mockLoadModelLoadBalancingModal: vi.fn(), +})) +const { mockToastError } = vi.hoisted(() => ({ + mockToastError: vi.fn(), +})) let mockWorkspacePermissionKeys: string[] = [ 'plugin.model_config', 'credential.manage', @@ -20,51 +26,65 @@ vi.mock('@/context/permission-state', async () => { })) }) -vi.mock('@/next/dynamic', () => ({ - default: (loader: () => Promise<{ default: React.ComponentType }>) => { - const LazyComponent = React.lazy(loader) - return function DynamicComponent(props: Record) { - return React.createElement( - React.Suspense, - { fallback: null }, - React.createElement(LazyComponent, props), - ) - } +vi.mock('../../hooks', () => ({ + useLazyModelProviderDetail: () => ({ + loadProviderDetail: mockLoadProviderDetail, + }), +})) + +vi.mock('@langgenius/dify-ui/toast', () => ({ + toast: { + error: mockToastError, }, })) -vi.mock('../model-load-balancing-modal', () => ({ - default: ({ model, onClose, onSave, open, provider }: ModelLoadBalancingModalProps) => - open ? ( -
- {provider.provider} - - -
- ) : null, +const MockModelLoadBalancingModal = ({ + onClose, + onSave, +}: { + onClose?: () => void + onSave?: (provider: string) => void +}) => ( +
+ + +
+) + +vi.mock('../model-load-balancing-modal', () => mockLoadModelLoadBalancingModal()) + +vi.mock('../lazy-custom-model-actions', () => ({ + default: ({ provider }: { provider: ModelProvider }) => ( + <> + {(provider.custom_configuration.custom_models?.length ?? 0) > 0 && ( +
+ )} + + + ), })) vi.mock('../model-list-item', () => ({ default: ({ model, + isLoadingLoadBalancing, + isLoadBalancingDisabled, onModifyLoadBalancing, }: { model: ModelItem + isLoadingLoadBalancing?: boolean + isLoadBalancingDisabled?: boolean onModifyLoadBalancing: (model: ModelItem) => void }) => (
- {pluginDetail && ( - + {pluginSummary && ( + )}
@@ -182,7 +179,7 @@ const ProviderAddedCard: FC = ({ {(showModelProvider || !notConfigured) && (
@@ -253,8 +242,12 @@ const ProviderAddedCard: FC = ({
- {pluginDetail && ( - + {pluginSummary && ( + )}
@@ -270,7 +263,7 @@ const ProviderAddedCard: FC = ({ {(showModelProvider || !notConfigured) && (
diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/lazy-custom-model-actions.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/lazy-custom-model-actions.tsx new file mode 100644 index 00000000000..85412e87b1a --- /dev/null +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/lazy-custom-model-actions.tsx @@ -0,0 +1,96 @@ +import type { ModelProviderSummaryResponse } from '@dify/contracts/api/console/workspaces/types.gen' +import type { ModelProvider } from '../declarations' +import { Button } from '@langgenius/dify-ui/button' +import { useState } from 'react' +import { useTranslation } from 'react-i18next' +import { + AddCustomModel, + ManageCustomModelCredentials, +} from '@/app/components/header/account-setting/model-provider-page/model-auth' +import { ConfigurationMethodEnum, ModelModalModeEnum } from '../declarations' +import { useLazyModelProviderDetail, useModelModalHandler } from '../hooks' + +type ProviderSummary = ModelProviderSummaryResponse | ModelProvider + +export default function LazyCustomModelActions({ provider }: { provider: ProviderSummary }) { + const { t } = useTranslation() + const handleOpenModelModal = useModelModalHandler() + const [isAddOpen, setIsAddOpen] = useState(false) + const [isManageOpen, setIsManageOpen] = useState(false) + const { providerDetail, loadProviderDetail, isLoadingProviderDetail } = + useLazyModelProviderDetail(provider.provider) + const hasCustomModels = + 'has_custom_models' in provider.custom_configuration + ? provider.custom_configuration.has_custom_models + : !!provider.custom_configuration.custom_models?.length + + const handleAddClick = async () => { + const detail = await loadProviderDetail() + if (!detail) return + + if (detail.custom_configuration.can_added_models?.length) { + setIsAddOpen(true) + return + } + + handleOpenModelModal(detail, ConfigurationMethodEnum.customizableModel, undefined, { + isModelCredential: true, + mode: ModelModalModeEnum.configCustomModel, + }) + } + + const handleManageClick = async () => { + const detail = await loadProviderDetail() + if (!detail) return + + setIsManageOpen(true) + } + + if (providerDetail) { + return ( + <> + {hasCustomModels && ( + + )} + + + ) + } + + return ( + <> + {hasCustomModels && ( + + )} + + + ) +} diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-auth-dropdown/__tests__/index.spec.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-auth-dropdown/__tests__/index.spec.tsx index e43b501a099..3708ccaa169 100644 --- a/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-auth-dropdown/__tests__/index.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-auth-dropdown/__tests__/index.spec.tsx @@ -1,9 +1,14 @@ +import type { ModelProviderSummaryResponse } from '@dify/contracts/api/console/workspaces/types.gen' import type { ModelProvider } from '../../../declarations' import type { CredentialPanelState } from '../../use-credential-panel-state' -import { fireEvent, render, screen, waitFor } from '@testing-library/react' +import { act, fireEvent, screen, waitFor } from '@testing-library/react' +import { commonQueryKeys } from '@/service/use-common' +import { renderWithConsoleQuery } from '@/test/console/query-data' import { CustomConfigurationStatusEnum, PreferredProviderTypeEnum } from '../../../declarations' import ModelAuthDropdown from '../index' +const render = (ui: React.ReactElement) => renderWithConsoleQuery(ui) + vi.mock('../../../model-auth/hooks', () => ({ useAuth: () => ({ openConfirmDelete: vi.fn(), @@ -190,6 +195,64 @@ describe('ModelAuthDropdown', () => { }) describe('Popover behavior', () => { + it('should keep the popover open and allow retrying when loading a summary detail fails', async () => { + const providerSummary = { + provider: 'test', + plugin_id: 'test-plugin', + label: { en_US: 'Test', zh_Hans: 'Test' }, + supported_model_types: ['llm'], + configurate_methods: [], + preferred_provider_type: 'system', + is_configured: true, + custom_configuration: { + status: 'active', + has_custom_models: false, + available_credentials: [], + current_credential_usable: false, + }, + system_configuration: { enabled: true }, + } satisfies ModelProviderSummaryResponse + const fullProvider = createProvider() + render( + , + ) + let resolveFirstRequest: ((response: Response) => void) | undefined + const fetchMock = vi.mocked(globalThis.fetch) + fetchMock.mockImplementationOnce( + () => + new Promise((resolve) => { + resolveFirstRequest = resolve + }), + ) + + fireEvent.click(screen.getByRole('button', { name: /addApiKey/i })) + + expect(await screen.findByRole('status')).toBeInTheDocument() + await act(async () => { + resolveFirstRequest?.(new Response(null, { status: 500 })) + }) + expect(await screen.findByRole('alert')).toHaveTextContent('common.api.actionFailed') + + fetchMock.mockResolvedValueOnce( + new Response(JSON.stringify({ data: [fullProvider] }), { + status: 200, + headers: { 'Content-Type': 'application/json' }, + }), + ) + + fireEvent.click(screen.getByRole('button', { name: 'common.operation.retry' })) + + await waitFor(() => { + expect(screen.getByText('common.modelProvider.card.noApiKeysTitle')).toBeInTheDocument() + }) + expect(fetchMock).toHaveBeenCalledTimes(2) + }) + it('should open popover on button click and show dropdown content', async () => { render( { expect(screen.getByText('Key 1')).toBeInTheDocument() }) }) + + it('should load provider detail on first click and reuse it on reopen', async () => { + const fullProvider = createProvider({ + custom_configuration: { + status: CustomConfigurationStatusEnum.active, + available_credentials: [{ credential_id: 'c1', credential_name: 'Key 1' }], + current_credential_id: 'c1', + current_credential_name: 'Key 1', + }, + }) + const providerSummary = { + provider: 'test', + is_configured: true, + custom_configuration: { + available_credentials: [{ credential_id: 'c1', credential_name: 'Key 1' }], + current_credential_id: 'c1', + current_credential_name: 'Key 1', + current_credential_usable: true, + }, + } as unknown as ModelProviderSummaryResponse + const rendered = render( + , + ) + const fetchQuery = vi + .spyOn(rendered.queryClient, 'fetchQuery') + .mockImplementation(async () => { + const response = { data: [fullProvider] } + rendered.queryClient.setQueryData(commonQueryKeys.modelProviderDetails, response) + return response + }) + const trigger = screen.getByRole('button', { name: /config/i }) + + fireEvent.click(trigger) + + await waitFor(() => { + expect(screen.getByText('Key 1')).toBeInTheDocument() + }) + expect(fetchQuery).toHaveBeenCalledTimes(1) + + fireEvent.click(trigger) + await waitFor(() => { + expect(screen.queryByText('Key 1')).not.toBeInTheDocument() + }) + fireEvent.click(trigger) + + await waitFor(() => { + expect(screen.getByText('Key 1')).toBeInTheDocument() + }) + expect(fetchQuery).toHaveBeenCalledTimes(1) + }) + + it('should render updated provider detail from the query cache', async () => { + const providerSummary = { + provider: 'test', + is_configured: true, + custom_configuration: { + available_credentials: [], + current_credential_usable: true, + }, + } as unknown as ModelProviderSummaryResponse + const firstProvider = createProvider({ + custom_configuration: { + status: CustomConfigurationStatusEnum.active, + available_credentials: [{ credential_id: 'c1', credential_name: 'Key 1' }], + current_credential_id: 'c1', + current_credential_name: 'Key 1', + }, + }) + const nextProvider = createProvider({ + custom_configuration: { + status: CustomConfigurationStatusEnum.active, + available_credentials: [{ credential_id: 'c2', credential_name: 'Key 2' }], + current_credential_id: 'c2', + current_credential_name: 'Key 2', + }, + }) + const rendered = render( + , + ) + vi.spyOn(rendered.queryClient, 'fetchQuery').mockImplementation(async () => { + const response = { data: [firstProvider] } + rendered.queryClient.setQueryData(commonQueryKeys.modelProviderDetails, response) + return response + }) + + fireEvent.click(screen.getByRole('button', { name: /config/i })) + await waitFor(() => expect(screen.getByText('Key 1')).toBeInTheDocument()) + + rendered.queryClient.setQueryData(commonQueryKeys.modelProviderDetails, { + data: [nextProvider], + }) + + await waitFor(() => { + expect(screen.getByText('Key 2')).toBeInTheDocument() + }) + expect(screen.queryByText('Key 1')).not.toBeInTheDocument() + }) }) }) diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-auth-dropdown/credits-exhausted-alert.stories.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-auth-dropdown/credits-exhausted-alert.stories.tsx index 9228f6cfd55..0d9f7c771e8 100644 --- a/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-auth-dropdown/credits-exhausted-alert.stories.tsx +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-auth-dropdown/credits-exhausted-alert.stories.tsx @@ -1,29 +1,26 @@ +import type { ModelProviderCreditsResponse } from '@dify/contracts/api/console/workspaces/types.gen' import type { Meta, StoryObj } from '@storybook/nextjs-vite' -import type { ICurrentWorkspace } from '@/models/common' import { QueryClient, QueryClientProvider } from '@tanstack/react-query' import { consoleQuery } from '@/service/client' import CreditsExhaustedAlert from './credits-exhausted-alert' -const baseWorkspace: ICurrentWorkspace = { - id: 'ws-1', - name: 'Test Workspace', - plan: 'sandbox', - status: 'normal', - created_at: Date.now(), - role: 'owner', - providers: [], - trial_credits: 200, - trial_credits_used: 200, - trial_credits_exhausted_at: 0, +const baseCredits: ModelProviderCreditsResponse = { + pool_type: 'trial', + quota_limit: 200, + quota_used: 200, + remaining_credits: 0, + is_unlimited: false, + is_exhausted: true, + exhausted_at: 0, next_credit_reset_date: Date.now() + 86400000, } -function createSeededQueryClient(overrides?: Partial) { +function createSeededQueryClient(overrides?: Partial) { const qc = new QueryClient({ defaultOptions: { queries: { refetchOnWindowFocus: false, retry: false } }, }) - qc.setQueryData(consoleQuery.workspaces.current.post.queryKey(), { - ...baseWorkspace, + qc.setQueryData(consoleQuery.workspaces.current.modelProviders.credits.get.queryKey(), { + ...baseCredits, ...overrides, }) return qc @@ -74,7 +71,12 @@ export const PartialUsage: Story = { (Story) => { return (
diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-auth-dropdown/index.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-auth-dropdown/index.tsx index 36e9ef19515..90c67bfaf2a 100644 --- a/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-auth-dropdown/index.tsx +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-auth-dropdown/index.tsx @@ -1,14 +1,16 @@ +import type { ModelProviderSummaryResponse } from '@dify/contracts/api/console/workspaces/types.gen' import type { ModelProvider, PreferredProviderTypeEnum } from '../../declarations' import type { CredentialPanelState } from '../use-credential-panel-state' import { Button } from '@langgenius/dify-ui/button' import { Popover, PopoverContent, PopoverTrigger } from '@langgenius/dify-ui/popover' import { memo, useCallback, useState } from 'react' import { useTranslation } from 'react-i18next' +import { useLazyModelProviderDetail } from '../../hooks' import { getButtonConfig } from './button-config' import DropdownContent from './dropdown-content' type ModelAuthDropdownProps = { - provider: ModelProvider + provider: ModelProviderSummaryResponse | ModelProvider state: CredentialPanelState isChangingPriority: boolean onChangePriority: (key: PreferredProviderTypeEnum) => void @@ -22,13 +24,41 @@ function ModelAuthDropdown({ }: ModelAuthDropdownProps) { const { t } = useTranslation() const [open, setOpen] = useState(false) + const [isProviderDetailError, setIsProviderDetailError] = useState(false) + const isFullProvider = !('is_configured' in provider) || 'provider_credential_schema' in provider + const { providerDetail, loadProviderDetail, isProviderDetailEnabled, isLoadingProviderDetail } = + useLazyModelProviderDetail(provider.provider) + const currentProvider = isFullProvider ? provider : providerDetail - const handleClose = useCallback(() => setOpen(false), []) + const handleClose = useCallback(() => { + setOpen(false) + setIsProviderDetailError(false) + }, []) const buttonConfig = getButtonConfig(state.variant, state.hasCredentials, t) + const loadDetail = useCallback(async () => { + setIsProviderDetailError(false) + const detail = await loadProviderDetail() + if (!detail) setIsProviderDetailError(true) + }, [loadProviderDetail]) + + const handleOpenChange = (nextOpen: boolean) => { + if (!nextOpen) { + handleClose() + return + } + + setOpen(true) + + if ( + !isFullProvider && + !(currentProvider && isProviderDetailEnabled && !isLoadingProviderDetail) + ) + void loadDetail() + } return ( - + {buttonConfig.text} @@ -43,13 +74,38 @@ function ModelAuthDropdown({ } /> - + {currentProvider ? ( + + ) : isProviderDetailError ? ( +
+ + {t(($) => $['api.actionFailed'], { ns: 'common' })} + + +
+ ) : ( +
+ + + {t(($) => $.loading, { ns: 'common' })} + +
+ )}
) diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-list-item.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-list-item.tsx index 76d84e772da..5df9fb0b02a 100644 --- a/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-list-item.tsx +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-list-item.tsx @@ -1,3 +1,4 @@ +import type { ModelProviderSummaryResponse } from '@dify/contracts/api/console/workspaces/types.gen' import type { ModelItem, ModelProvider } from '../declarations' import { cn } from '@langgenius/dify-ui/cn' import { Popover, PopoverContent, PopoverTrigger } from '@langgenius/dify-ui/popover' @@ -23,8 +24,10 @@ import ModelName from '../model-name' type ModelListItemProps = { model: ModelItem - provider: ModelProvider + provider: ModelProvider | ModelProviderSummaryResponse isConfigurable: boolean + isLoadingLoadBalancing?: boolean + isLoadBalancingDisabled?: boolean onChange?: (provider: string) => void onModifyLoadBalancing?: (model: ModelItem) => void } @@ -33,6 +36,8 @@ const ModelListItem = ({ model, provider, isConfigurable, + isLoadingLoadBalancing, + isLoadBalancingDisabled, onChange, onModifyLoadBalancing, }: ModelListItemProps) => { @@ -132,6 +137,8 @@ const ModelListItem = ({ [ModelStatusEnum.active, ModelStatusEnum.disabled].includes(model.status) && ( onModifyLoadBalancing?.(model)} + loading={isLoadingLoadBalancing} + disabled={isLoadBalancingDisabled} loadBalancingEnabled={model.load_balancing_enabled} loadBalancingInvalid={model.has_invalid_load_balancing_configs} credentialRemoved={model.status === ModelStatusEnum.credentialRemoved} diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-list.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-list.tsx index 3823d366781..327c77aebc7 100644 --- a/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-list.tsx +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/model-list.tsx @@ -1,30 +1,69 @@ -import type { FC } from 'react' +import type { ModelProviderSummaryResponse } from '@dify/contracts/api/console/workspaces/types.gen' +import type { ComponentType, FC } from 'react' import type { Credential, ModelItem, ModelProvider } from '../declarations' import type { ModelLoadBalancingModalProps } from './model-load-balancing-modal' +import { Dialog, DialogContent, DialogTitle } from '@langgenius/dify-ui/dialog' +import { toast } from '@langgenius/dify-ui/toast' import { useAtomValue } from 'jotai' import { useCallback, useState } from 'react' import { useTranslation } from 'react-i18next' -import { - AddCustomModel, - ManageCustomModelCredentials, -} from '@/app/components/header/account-setting/model-provider-page/model-auth' import { workspacePermissionKeysAtom } from '@/context/permission-state' -import dynamic from '@/next/dynamic' import { hasPermission } from '@/utils/permission' import { ConfigurationMethodEnum } from '../declarations' +import { useLazyModelProviderDetail } from '../hooks' +import LazyCustomModelActions from './lazy-custom-model-actions' // import Tab from './tab' import ModelListItem from './model-list-item' -const ModelLoadBalancingModal = dynamic(() => import('./model-load-balancing-modal'), { - ssr: false, -}) +const ModelLoadBalancingLoadingDialog = ({ + onClose, +}: Pick) => { + const { t } = useTranslation() + + return ( + !open && onClose?.()}> + + + {t(($) => $['modelProvider.auth.configModel'], { ns: 'common' })} + +
+ + + {t(($) => $.loading, { ns: 'common' })} + +
+
+
+ ) +} type ModelListProps = { - provider: ModelProvider + provider: ModelProvider | ModelProviderSummaryResponse models: ModelItem[] onCollapse: () => void onChange?: (provider: string) => void } + +const getModelKey = (model: ModelItem) => `${model.model}-${model.model_type}-${model.fetch_from}` + +let ModelLoadBalancingModal: ComponentType | undefined +let modelLoadBalancingModalPromise: Promise | undefined + +const loadModelLoadBalancingModal = () => { + if (ModelLoadBalancingModal) return Promise.resolve() + + modelLoadBalancingModalPromise ??= import('./model-load-balancing-modal').then( + ({ default: Modal }) => { + ModelLoadBalancingModal = Modal + }, + ) + + return modelLoadBalancingModalPromise +} + const ModelList: FC = ({ provider, models, onCollapse, onChange }) => { const { t } = useTranslation() const configurativeMethods = provider.configurate_methods.filter( @@ -35,17 +74,52 @@ const ModelList: FC = ({ provider, models, onCollapse, onChange const isConfigurable = configurativeMethods.includes(ConfigurationMethodEnum.customizableModel) const [modelLoadBalancingModalProps, setModelLoadBalancingModalProps] = useState(null) + const [loadingModelKey, setLoadingModelKey] = useState(null) + const [isModelLoadBalancingModalLoading, setIsModelLoadBalancingModalLoading] = useState(false) + const { loadProviderDetail } = useLazyModelProviderDetail(provider.provider) const onModifyLoadBalancing = useCallback( - (model: ModelItem, credential?: Credential) => { + async (model: ModelItem, credential?: Credential) => { + if (loadingModelKey) return + + let providerDetail: ModelProvider | undefined + if ('is_configured' in provider) { + setLoadingModelKey(getModelKey(model)) + try { + providerDetail = await loadProviderDetail() + } finally { + setLoadingModelKey(null) + } + } else { + providerDetail = provider + } + + if (!providerDetail) { + toast.error(t(($) => $['api.actionFailed'], { ns: 'common' })) + return + } + setModelLoadBalancingModalProps({ - provider, + provider: providerDetail, credential, configurateMethod: model.fetch_from, model, open: true, + onSave: onChange, }) + + if (ModelLoadBalancingModal) return + + setIsModelLoadBalancingModalLoading(true) + try { + await loadModelLoadBalancingModal() + } catch { + setModelLoadBalancingModalProps(null) + toast.error(t(($) => $['api.actionFailed'], { ns: 'common' })) + } finally { + setIsModelLoadBalancingModalLoading(false) + } }, - [provider], + [loadingModelKey, loadProviderDetail, onChange, provider, t], ) return ( @@ -60,7 +134,7 @@ const ModelList: FC = ({ provider, models, onCollapse, onChange
{models.map((model) => ( = ({ provider, models, onCollapse, onChange ))}
- {modelLoadBalancingModalProps && ( + {isModelLoadBalancingModalLoading && modelLoadBalancingModalProps && ( + setModelLoadBalancingModalProps(null)} /> + )} + {ModelLoadBalancingModal && modelLoadBalancingModalProps && ( setModelLoadBalancingModalProps(null)} - onSave={onChange} /> )} diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/provider-card-actions.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/provider-card-actions.tsx index 60aea8819f1..dcad34b8884 100644 --- a/web/app/components/header/account-setting/model-provider-page/provider-added-card/provider-card-actions.tsx +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/provider-card-actions.tsx @@ -1,9 +1,21 @@ import type { FC } from 'react' +import type { ModelProviderPluginSummary } from '../index' import type { PluginDetail } from '@/app/components/plugins/types' +import { + AlertDialog, + AlertDialogActions, + AlertDialogCancelButton, + AlertDialogConfirmButton, + AlertDialogContent, + AlertDialogDescription, + AlertDialogTitle, +} from '@langgenius/dify-ui/alert-dialog' import { Button } from '@langgenius/dify-ui/button' import { cn } from '@langgenius/dify-ui/cn' import { Tooltip, TooltipContent, TooltipTrigger } from '@langgenius/dify-ui/tooltip' -import { useMemo } from 'react' +import { useQueryClient } from '@tanstack/react-query' +import { useBoolean } from 'ahooks' +import { useCallback, useEffect, useMemo, useRef, useState } from 'react' import { useTranslation } from 'react-i18next' import Badge from '@/app/components/base/badge' import { HeaderModals } from '@/app/components/plugins/plugin-detail-panel/detail-header/components' @@ -17,24 +29,272 @@ import { PluginSource } from '@/app/components/plugins/types' import PluginVersionPicker from '@/app/components/plugins/update-plugin/plugin-version-picker' import { useLocale } from '@/context/i18n' import useTheme from '@/hooks/use-theme' +import { consoleQuery } from '@/service/client' +import { uninstallPlugin } from '@/service/plugins' +import { commonQueryKeys } from '@/service/use-common' +import { normalizeInstalledPluginDetail } from '@/service/use-plugins' import { getMarketplaceUrl } from '@/utils/var' -type Props = Readonly<{ +type DetailAction = 'version' | 'latest' | 'info' | 'check' + +type Props = + | Readonly<{ + summary: ModelProviderPluginSummary + providerLabel: string + onUpdate?: () => void + detail?: never + }> + | Readonly<{ + detail: PluginDetail + onUpdate?: () => void + summary?: never + providerLabel?: never + }> + +const pluginSourceMap: Record = { + github: PluginSource.github, + marketplace: PluginSource.marketplace, + package: PluginSource.local, + remote: PluginSource.debugging, +} + +const ProviderCardActions: FC = (props) => { + const queryClient = useQueryClient() + const { onUpdate } = props + const handlePluginChanged = useCallback(async () => { + await Promise.all([ + queryClient.invalidateQueries({ + queryKey: consoleQuery.workspaces.current.modelProviders.summary.get.key(), + }), + queryClient.invalidateQueries({ + queryKey: consoleQuery.workspaces.current.plugin.installedIds.get.key(), + }), + queryClient.invalidateQueries({ + queryKey: consoleQuery.workspaces.current.plugin.list.installations.ids.post.key(), + }), + queryClient.invalidateQueries({ + queryKey: consoleQuery.workspaces.current.plugin.list.latestVersions.post.key(), + }), + queryClient.invalidateQueries({ queryKey: commonQueryKeys.modelProviderDetails }), + queryClient.invalidateQueries({ queryKey: ['marketplacePlugins'] }), + queryClient.invalidateQueries({ queryKey: ['marketplaceCollectionPlugins'] }), + ]) + onUpdate?.() + }, [onUpdate, queryClient]) + + if (props.detail) { + return ( + {}} + onUpdate={handlePluginChanged} + /> + ) + } + + return +} + +type SummaryProps = Extract + +function SummaryProviderCardActions({ summary, providerLabel, onUpdate }: SummaryProps) { + const { t } = useTranslation() + const queryClient = useQueryClient() + const { canDeletePlugin, canUpdatePlugin } = usePluginSettingsAccess() + const [detail, setDetail] = useState() + const [detailAction, setDetailAction] = useState() + const [loadingAction, setLoadingAction] = useState() + const [showDeleteConfirm, { setTrue: openDeleteConfirm, setFalse: closeDeleteConfirm }] = + useBoolean(false) + const [deleting, { setTrue: startDeleting, setFalse: finishDeleting }] = useBoolean(false) + const source = pluginSourceMap[summary.source] + const isFromMarketplace = source === PluginSource.marketplace + const isFromGitHub = source === PluginSource.github + const canChangeVersion = canUpdatePlugin && isFromMarketplace + const hasNewVersion = + isFromMarketplace && !!summary.latestVersion && summary.latestVersion !== summary.version + const [author, name] = summary.plugin_id.split('/') + const detailUrl = + isFromMarketplace && author && name ? getMarketplaceUrl(`/plugins/${author}/${name}`) : '' + + const loadDetail = async (action: DetailAction) => { + setLoadingAction(action) + try { + const response = await queryClient.fetchQuery( + consoleQuery.workspaces.current.plugin.list.installations.ids.post.queryOptions({ + input: { body: { plugin_ids: [summary.plugin_id] } }, + }), + ) + const nextDetail = response.plugins[0] + if (!nextDetail) return + const normalizedDetail = normalizeInstalledPluginDetail(nextDetail) + const detailWithLatestVersion = + isFromMarketplace && summary.latestVersion && summary.latestUniqueIdentifier + ? { + ...normalizedDetail, + latest_version: summary.latestVersion, + latest_unique_identifier: summary.latestUniqueIdentifier, + } + : normalizedDetail + setDetail(detailWithLatestVersion) + setDetailAction(action) + } catch { + } finally { + setLoadingAction(undefined) + } + } + + const handleDelete = async () => { + startDeleting() + try { + const response = await uninstallPlugin(summary.installation_id) + if (!response.success) return + closeDeleteConfirm() + await onUpdate?.() + } finally { + finishDeleting() + } + } + + if (detail) { + return ( + setDetailAction(undefined)} + onUpdate={onUpdate} + /> + ) + } + + return ( + <> + {!!summary.version && ( + <> + {canChangeVersion ? ( + + ) : ( + {summary.version}} + hasRedCornerMark={hasNewVersion} + /> + )} + {source === PluginSource.debugging && ( + $['operation.debugConfig'], { ns: 'appDebug' })} + /> + )} + + )} + + {canUpdatePlugin && (hasNewVersion || isFromGitHub) && ( + + loadDetail('latest')} + > + {t(($) => $['detailPanel.operation.update'], { ns: 'plugin' })} + + } + /> + + {t(($) => $['detailPanel.operation.updateTooltip'], { ns: 'plugin' })} + + + )} + + loadDetail('info')} + onCheckVersion={() => loadDetail('check')} + onRemove={openDeleteConfirm} + detailUrl={detailUrl} + placement="bottom-start" + destructiveRemove + showCheckVersion={canUpdatePlugin} + showRemove={canDeletePlugin} + /> + + { + if (!open) closeDeleteConfirm() + }} + > + +
+ + {t(($) => $['action.delete'], { ns: 'plugin' })} + + + {t(($) => $['action.deleteContentLeft'], { ns: 'plugin' })} + {providerLabel} + {t(($) => $['action.deleteContentRight'], { ns: 'plugin' })} + +
+ + + {t(($) => $['operation.cancel'], { ns: 'common' })} + + + {t(($) => $['operation.confirm'], { ns: 'common' })} + + +
+
+ + ) +} + +type LoadedProps = Readonly<{ detail: PluginDetail - onUpdate?: () => void | Promise + initialAction?: DetailAction + onInitialActionHandled: () => void + onUpdate?: () => void }> -const ProviderCardActions: FC = ({ detail, onUpdate }) => { +function LoadedProviderCardActions({ + detail, + initialAction, + onInitialActionHandled, + onUpdate, +}: LoadedProps) { const { t } = useTranslation() const { theme } = useTheme() const locale = useLocale() const { canDeletePlugin, canUpdatePlugin } = usePluginSettingsAccess() - const { source, version, latest_version, latest_unique_identifier, meta } = detail const author = detail.declaration?.author ?? '' const name = detail.declaration?.name ?? detail.name const isDebuggingPlugin = source === PluginSource.debugging - const { modalStates, versionPicker, @@ -43,7 +303,6 @@ const ProviderCardActions: FC = ({ detail, onUpdate }) => { isFromMarketplace, isFromGitHub, } = useDetailHeaderState(detail) - const { handleUpdate, handleUpdatedFromMarketplace, handleDelete } = usePluginOperations({ detail, modalStates, @@ -63,15 +322,28 @@ const ProviderCardActions: FC = ({ detail, onUpdate }) => { handleUpdate(state.isDowngrade) } - const handleTriggerLatestUpdate = () => { + const handleTriggerLatestUpdate = useCallback(() => { if (isFromMarketplace) { + if (!latest_unique_identifier) return versionPicker.setTargetVersion({ version: latest_version, unique_identifier: latest_unique_identifier, }) } handleUpdate() - } + }, [handleUpdate, isFromMarketplace, latest_unique_identifier, latest_version, versionPicker]) + + const pendingInitialActionRef = useRef(initialAction) + useEffect(() => { + const pendingAction = pendingInitialActionRef.current + if (!pendingAction) return + pendingInitialActionRef.current = undefined + if (pendingAction === 'version') versionPicker.setIsShow(true) + if (pendingAction === 'latest') handleTriggerLatestUpdate() + if (pendingAction === 'info') modalStates.showPluginInfo() + if (pendingAction === 'check') handleUpdate() + onInitialActionHandled() + }, [handleTriggerLatestUpdate, handleUpdate, modalStates, onInitialActionHandled, versionPicker]) const detailUrl = useMemo(() => { if (source === PluginSource.github) return meta?.repo ? `https://github.com/${meta.repo}` : '' diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/quota-panel.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/quota-panel.tsx index 11433e2f64e..258e7051440 100644 --- a/web/app/components/header/account-setting/model-provider-page/provider-added-card/quota-panel.tsx +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/quota-panel.tsx @@ -1,6 +1,7 @@ +import type { ModelProviderSummaryResponse } from '@dify/contracts/api/console/workspaces/types.gen' import type { FC, MouseEvent } from 'react' import type { ModelProvider } from '../declarations' -import type { Plugin } from '@/app/components/plugins/types' +import type { PluginManifestInMarket } from '@/app/components/plugins/types' import type { ModelProviderQuotaGetPaid } from '@/types/model-provider' import { cn } from '@langgenius/dify-ui/cn' import { Popover, PopoverContent, PopoverTrigger } from '@langgenius/dify-ui/popover' @@ -17,9 +18,9 @@ import InstallFromMarketplace from '@/app/components/plugins/install-plugin/inst import { systemFeaturesQueryOptions } from '@/features/system-features/client' import useTimestamp from '@/hooks/use-timestamp' import { consoleQuery } from '@/service/client' +import { fetchManifestFromMarketPlace, fetchPluginInfoFromMarketPlace } from '@/service/plugins' import { formatNumber } from '@/utils/format' import { PreferredProviderTypeEnum } from '../declarations' -import { useMarketplaceAllPlugins } from '../hooks' import { MODEL_PROVIDER_QUOTA_GET_PAID, modelNameMap, @@ -69,16 +70,29 @@ const QuotaInfotip: FC = ({ tipText }) => { } type QuotaPanelProps = { - providers: ModelProvider[] + providers: Array } + +type MarketplacePluginToInstall = { + manifest: PluginManifestInMarket + uniqueIdentifier: string +} + const QuotaPanel: FC = ({ providers }) => { const { t } = useTranslation() const { data: deploymentEdition } = useSuspenseQuery({ ...systemFeaturesQueryOptions(), select: ({ deployment_edition }) => deployment_edition, }) - const { usedCredits, totalCredits, isExhausted, isLoading, exhaustedAt, nextCreditResetDate } = - useTrialCredits() + const { + usedCredits, + totalCredits, + isUnlimited, + isExhausted, + isLoading, + exhaustedAt, + nextCreditResetDate, + } = useTrialCredits() const { data: trialModels = [] } = useQuery( consoleQuery.trialModels.get.queryOptions({ enabled: deploymentEdition === 'CLOUD', @@ -94,8 +108,8 @@ const QuotaPanel: FC = ({ providers }) => { [providers], ) const { formatMonthDay } = useTimestamp() - const { plugins: allPlugins } = useMarketplaceAllPlugins(providers, '') - const [selectedPlugin, setSelectedPlugin] = useState(null) + const [selectedPlugin, setSelectedPlugin] = useState(null) + const [loadingPluginId, setLoadingPluginId] = useState(null) const [ isShowInstallModal, { setTrue: showInstallFromMarketplace, setFalse: hideInstallFromMarketplace }, @@ -105,19 +119,41 @@ const QuotaPanel: FC = ({ providers }) => { const selectedPluginIdRef = useRef(null) const handleIconClick = useCallback( - (key: ModelProviderQuotaGetPaid) => { + async (key: ModelProviderQuotaGetPaid) => { + if (loadingPluginId) return + const isInstalled = providerMap.get(key) - if (!isInstalled && allPlugins && canInstallPlugin) { + if (!isInstalled && canInstallPlugin) { const pluginId = providerKeyToPluginId[key] - const plugin = allPlugins.find((p) => p.plugin_id === pluginId) - if (plugin) { - setSelectedPlugin(plugin) + const [org, name] = pluginId.split('/') + if (!org || !name) return + + setLoadingPluginId(pluginId) + try { + const pluginInfo = await fetchPluginInfoFromMarketPlace({ org, name }) + const uniqueIdentifier = pluginInfo.data.plugin.latest_package_identifier + if (!uniqueIdentifier) return + const manifest = await fetchManifestFromMarketPlace(uniqueIdentifier) + setSelectedPlugin({ + manifest: { + ...manifest.data.plugin, + org, + name, + from: 'marketplace', + icon: manifest.data.plugin.icon || 'marketplace', + }, + uniqueIdentifier, + }) selectedPluginIdRef.current = pluginId showInstallFromMarketplace() + } catch { + // Keep the provider actionable so the user can retry the marketplace request. + } finally { + setLoadingPluginId(null) } } }, - [allPlugins, canInstallPlugin, providerMap, showInstallFromMarketplace], + [canInstallPlugin, loadingPluginId, providerMap, showInstallFromMarketplace], ) useEffect(() => { @@ -166,16 +202,24 @@ const QuotaPanel: FC = ({ providers }) => {
- - {formatNumber(usedCredits)} - - / - - {formatNumber(totalCredits)} - - - {t(($) => $['modelProvider.used'], { ns: 'common' })} - + {isUnlimited ? ( + + {t(($) => $['license.unlimited'], { ns: 'common' })} + + ) : ( + <> + + {formatNumber(usedCredits)} + + / + + {formatNumber(totalCredits)} + + + {t(($) => $['modelProvider.used'], { ns: 'common' })} + + + )}
{isExhausted && exhaustedAt ? ( <> @@ -211,6 +255,7 @@ const QuotaPanel: FC = ({ providers }) => { .filter(({ key }) => trialModels.includes(key)) .map(({ key, Icon }) => { const providerType = providerMap.get(key) + const isLoadingPlugin = loadingPluginId === providerKeyToPluginId[key] const isConfigured = (installedProvidersMap.get(key)?.length ?? 0) > 0 const getTooltipKey = () => { if (!providerType) return 'modelProvider.card.modelNotSupported' @@ -229,13 +274,21 @@ const QuotaPanel: FC = ({ providers }) => { render={ diff --git a/web/app/components/header/account-setting/model-provider-page/utils.ts b/web/app/components/header/account-setting/model-provider-page/utils.ts index 5598385449c..af2f2405849 100644 --- a/web/app/components/header/account-setting/model-provider-page/utils.ts +++ b/web/app/components/header/account-setting/model-provider-page/utils.ts @@ -32,11 +32,6 @@ import { export { ModelProviderQuotaGetPaid } from '@/types/model-provider' -export const providerToPluginId = (providerKey: string): string => { - const lastSlash = providerKey.lastIndexOf('/') - return lastSlash > 0 ? providerKey.slice(0, lastSlash) : '' -} - export const MODEL_PROVIDER_QUOTA_GET_PAID = [ ModelProviderQuotaGetPaid.OPENAI, ModelProviderQuotaGetPaid.ANTHROPIC, @@ -86,7 +81,7 @@ export const sizeFormat = (size: number) => { else return `${remainder}K` } -export const modelTypeFormat = (modelType: ModelTypeEnum) => { +export const modelTypeFormat = (modelType: ModelTypeEnum | ModelType) => { if (modelType === ModelTypeEnum.textEmbedding) return 'TEXT EMBEDDING' return modelType.toLocaleUpperCase() diff --git a/web/app/components/integrations/__tests__/tool-provider-list.spec.tsx b/web/app/components/integrations/__tests__/tool-provider-list.spec.tsx index 936f889b071..47d4c6931ee 100644 --- a/web/app/components/integrations/__tests__/tool-provider-list.spec.tsx +++ b/web/app/components/integrations/__tests__/tool-provider-list.spec.tsx @@ -94,8 +94,14 @@ let mockCollectionData: ReturnType = [] let mockIsLoadingToolProviders = false const mockRefetch = vi.fn() const mockUseAllToolProviders = vi.hoisted(() => vi.fn()) +const mockUseAllCustomTools = vi.hoisted(() => vi.fn()) +const mockUseAllMCPTools = vi.hoisted(() => vi.fn()) +const mockUseAllWorkflowTools = vi.hoisted(() => vi.fn()) vi.mock('@/service/use-tools', () => ({ useAllToolProviders: (enabled?: boolean) => mockUseAllToolProviders(enabled), + useAllCustomTools: (enabled?: boolean) => mockUseAllCustomTools(enabled), + useAllMCPTools: (enabled?: boolean) => mockUseAllMCPTools(enabled), + useAllWorkflowTools: (enabled?: boolean) => mockUseAllWorkflowTools(enabled), })) const mockConsoleState = vi.hoisted(() => ({ @@ -386,6 +392,23 @@ describe('ProviderList', () => { isLoading: enabled ? mockIsLoadingToolProviders : false, refetch: mockRefetch, })) + mockUseAllCustomTools.mockImplementation((enabled = true) => ({ + data: enabled ? mockCollectionData.filter((collection) => collection.type === 'api') : [], + isLoading: enabled ? mockIsLoadingToolProviders : false, + refetch: mockRefetch, + })) + mockUseAllMCPTools.mockImplementation((enabled = true) => ({ + data: enabled ? mockCollectionData.filter((collection) => collection.type === 'mcp') : [], + isLoading: enabled ? mockIsLoadingToolProviders : false, + refetch: mockRefetch, + })) + mockUseAllWorkflowTools.mockImplementation((enabled = true) => ({ + data: enabled + ? mockCollectionData.filter((collection) => collection.type === 'workflow') + : [], + isLoading: enabled ? mockIsLoadingToolProviders : false, + refetch: mockRefetch, + })) mockCheckedInstalledData = null mockCanSetPermissions.mockReturnValue(true) mockReferenceSetting.mockReturnValue({ @@ -447,7 +470,10 @@ describe('ProviderList', () => { renderProviderList({ category }) - expect(mockUseAllToolProviders).toHaveBeenCalledWith(undefined) + expect(mockUseAllToolProviders).toHaveBeenCalledWith(false) + expect(mockUseAllCustomTools).toHaveBeenCalledWith(category === 'api') + expect(mockUseAllWorkflowTools).toHaveBeenCalledWith(category === 'workflow') + expect(mockUseAllMCPTools).toHaveBeenCalledWith(false) expect(screen.getByTestId(cardTestId)).toBeInTheDocument() expect(screen.queryByTestId('custom-create-card')).not.toBeInTheDocument() expect(screen.queryByTestId('toolbar-add-custom-tool')).not.toBeInTheDocument() @@ -940,7 +966,10 @@ describe('ProviderList', () => { renderProviderList({ category: 'mcp' }) - expect(mockUseAllToolProviders).toHaveBeenCalledWith(undefined) + expect(mockUseAllToolProviders).toHaveBeenCalledWith(false) + expect(mockUseAllCustomTools).toHaveBeenCalledWith(false) + expect(mockUseAllWorkflowTools).toHaveBeenCalledWith(false) + expect(mockUseAllMCPTools).toHaveBeenCalledWith(true) expect(screen.getByTestId('mcp-list')).toBeInTheDocument() expect(screen.getByTestId('mcp-list')).toHaveAttribute('data-show-create-card', 'false') expect(screen.queryByTestId('toolbar-add-mcp')).not.toBeInTheDocument() diff --git a/web/app/components/integrations/tool-provider-list.tsx b/web/app/components/integrations/tool-provider-list.tsx index 9e6d7886835..b53fb6e797e 100644 --- a/web/app/components/integrations/tool-provider-list.tsx +++ b/web/app/components/integrations/tool-provider-list.tsx @@ -1,5 +1,5 @@ 'use client' -import type { ReactNode, RefObject } from 'react' +import type { ReactNode } from 'react' import type { ToolCategory } from '@/app/components/integrations/routes' import type { ToolsContentInset } from '@/app/components/tools/content-inset' import type { Collection } from '@/app/components/tools/types' @@ -28,14 +28,18 @@ import { useCanManageMCP, useCanManageTools, } from '@/app/components/tools/hooks/use-tool-permissions' -import Marketplace from '@/app/components/tools/marketplace' +import { BuiltinMarketplacePanel } from '@/app/components/tools/marketplace/builtin-marketplace-panel' import MCPList from '@/app/components/tools/mcp' import ProviderDetail from '@/app/components/tools/provider/detail' import { ToolProviderGrid } from '@/app/components/tools/tool-provider-grid' import { systemFeaturesQueryOptions } from '@/features/system-features/client' import { useCheckInstalled, useInvalidateInstalledPluginList } from '@/service/use-plugins' -import { useAllToolProviders } from '@/service/use-tools' -import { useToolMarketplacePanel } from './hooks/use-tool-marketplace-panel' +import { + useAllCustomTools, + useAllMCPTools, + useAllToolProviders, + useAllWorkflowTools, +} from '@/service/use-tools' import { useToolProviderCategory } from './hooks/use-tool-provider-category' import ToolProviderCreateAction from './tool-provider-create-action' import { ToolProviderToolbar } from './tool-provider-toolbar' @@ -46,40 +50,7 @@ type ProviderListProps = { layout?: (parts: { body: ReactNode; toolbar: ReactNode }) => ReactNode } -type BuiltinMarketplacePanelProps = { - containerRef: RefObject - contentInset: ToolsContentInset - keywords: string - tagFilterValue: string[] -} - -const BuiltinMarketplacePanel = ({ - containerRef, - contentInset, - keywords, - tagFilterValue, -}: BuiltinMarketplacePanelProps) => { - const { isMarketplaceArrowVisible, marketplaceContext, showMarketplacePanel, toolListTailRef } = - useToolMarketplacePanel({ - containerRef, - keywords, - tagFilterValue, - }) - - return ( - <> -
- - - ) -} +const EMPTY_COLLECTIONS: Collection[] = [] const ProviderList = ({ category, contentInset = 'default', layout }: ProviderListProps) => { // const searchParams = useSearchParams() @@ -119,11 +90,25 @@ const ProviderList = ({ category, contentInset = 'default', layout }: ProviderLi const handleCreatedMCPProviderHandled = useCallback(() => { setCreatedMCPProviderId(undefined) }, []) - const { - data: collectionList = [], - isLoading: isCollectionListLoading, - refetch, - } = useAllToolProviders() + const allToolProvidersQuery = useAllToolProviders(activeTab === 'builtin') + const customToolsQuery = useAllCustomTools(activeTab === 'api') + const workflowToolsQuery = useAllWorkflowTools(activeTab === 'workflow') + const mcpToolsQuery = useAllMCPTools(activeTab === 'mcp') + const { refetch: refetchMcpTools } = mcpToolsQuery + const activeToolsQuery = + activeTab === 'api' + ? customToolsQuery + : activeTab === 'workflow' + ? workflowToolsQuery + : activeTab === 'mcp' + ? mcpToolsQuery + : allToolProvidersQuery + const collectionList = activeToolsQuery.data ?? EMPTY_COLLECTIONS + const isCollectionListLoading = activeToolsQuery.isLoading + const refetch = activeToolsQuery.refetch + const refreshMcpTools = useCallback(async () => { + await refetchMcpTools() + }, [refetchMcpTools]) const activeTabCollectionList = useMemo(() => { return collectionList.filter((collection) => collection.type === activeTab) }, [activeTab, collectionList]) @@ -250,10 +235,13 @@ const ProviderList = ({ category, contentInset = 'default', layout }: ProviderLi )} {activeTab === 'mcp' && ( )} diff --git a/web/app/components/plugins/card/base/__tests__/card-icon.spec.tsx b/web/app/components/plugins/card/base/__tests__/card-icon.spec.tsx new file mode 100644 index 00000000000..10b88f36a10 --- /dev/null +++ b/web/app/components/plugins/card/base/__tests__/card-icon.spec.tsx @@ -0,0 +1,21 @@ +import { fireEvent, render, screen } from '@testing-library/react' +import { describe, expect, it } from 'vitest' +import Icon from '../card-icon' + +describe('Plugin card icon', () => { + it('lazy-loads URL icons and hides a failed image', () => { + render() + + const image = screen.getByAltText('') + + expect(image).toHaveAttribute('src', 'https://example.com/plugin-icon.png') + expect(image).toHaveAttribute('loading', 'lazy') + expect(image).toHaveAttribute('decoding', 'async') + expect(image).toHaveAttribute('width', '40') + expect(image).toHaveAttribute('height', '40') + + fireEvent.error(image) + + expect(image).toHaveStyle({ display: 'none' }) + }) +}) diff --git a/web/app/components/plugins/card/base/card-icon.tsx b/web/app/components/plugins/card/base/card-icon.tsx index cf3e16b5726..79804f25a28 100644 --- a/web/app/components/plugins/card/base/card-icon.tsx +++ b/web/app/components/plugins/card/base/card-icon.tsx @@ -11,6 +11,17 @@ const iconSizeMap = { medium: 'w-9 h-9', large: 'w-10 h-10', } + +const iconPixelSizeMap = { + xs: 16, + tiny: 24, + small: 32, + medium: 36, + large: 40, +} + +type IconSize = keyof typeof iconSizeMap + const Icon = ({ className, src, @@ -27,7 +38,7 @@ const Icon = ({ } installed?: boolean installFailed?: boolean - size?: 'xs' | 'tiny' | 'small' | 'medium' | 'large' + size?: IconSize }) => { const iconClassName = 'flex justify-center items-center gap-2 absolute bottom-[-4px] right-[-4px] w-[18px] h-[18px] rounded-full border-2 border-components-panel-bg' @@ -51,16 +62,19 @@ const Icon = ({ } return ( -
+
+ { + currentTarget.style.display = 'none' + }} + /> {installed && (
diff --git a/web/app/components/plugins/marketplace/__tests__/hooks.spec.tsx b/web/app/components/plugins/marketplace/__tests__/hooks.spec.tsx new file mode 100644 index 00000000000..46c770694b2 --- /dev/null +++ b/web/app/components/plugins/marketplace/__tests__/hooks.spec.tsx @@ -0,0 +1,151 @@ +import type { ReactNode } from 'react' +import { QueryClient, QueryClientProvider } from '@tanstack/react-query' +import { act, renderHook, waitFor } from '@testing-library/react' + +const getMarketplacePluginsByCollectionId = vi.hoisted(() => vi.fn()) +const getMarketplaceCollectionsAndPlugins = vi.hoisted(() => vi.fn()) + +vi.mock('@/service/base', () => ({ + postMarketplace: vi.fn(), +})) + +vi.mock('../utils', () => ({ + getFormattedPlugin: (plugin: unknown) => plugin, + getMarketplaceCollectionsAndPlugins, + getMarketplacePluginsByCollectionId, +})) + +const createWrapper = () => { + const queryClient = new QueryClient({ + defaultOptions: { + queries: { retry: false }, + }, + }) + const Wrapper = ({ children }: { children: ReactNode }) => ( + {children} + ) + + return { Wrapper, queryClient } +} + +describe('useMarketplacePluginsByCollectionId', () => { + beforeEach(() => { + vi.clearAllMocks() + }) + + it('should show loading while the first collection request is pending', async () => { + getMarketplacePluginsByCollectionId.mockImplementation(() => new Promise(() => {})) + const { useMarketplacePluginsByCollectionId } = await import('../hooks') + const { Wrapper } = createWrapper() + const { result } = renderHook( + () => useMarketplacePluginsByCollectionId('__model-settings-pinned-models'), + { wrapper: Wrapper }, + ) + + await waitFor(() => { + expect(getMarketplacePluginsByCollectionId).toHaveBeenCalledTimes(1) + }) + + expect(result.current.isLoading).toBe(true) + }) + + it('should retain collection results while refreshing cached data', async () => { + let resolveRefresh: (() => void) | undefined + getMarketplacePluginsByCollectionId + .mockResolvedValueOnce([{ plugin_id: 'cached-plugin', type: 'plugin' }]) + .mockImplementationOnce( + () => + new Promise<{ plugin_id: string; type: string }[]>((resolve) => { + resolveRefresh = () => resolve([{ plugin_id: 'refreshed-plugin', type: 'plugin' }]) + }), + ) + + const { useMarketplacePluginsByCollectionId } = await import('../hooks') + const { Wrapper, queryClient } = createWrapper() + const { result } = renderHook( + () => useMarketplacePluginsByCollectionId('__model-settings-pinned-models'), + { wrapper: Wrapper }, + ) + + await waitFor(() => { + expect(result.current.plugins).toEqual([{ plugin_id: 'cached-plugin', type: 'plugin' }]) + }) + + act(() => { + void queryClient.invalidateQueries({ queryKey: ['marketplaceCollectionPlugins'] }) + }) + + await waitFor(() => { + expect(getMarketplacePluginsByCollectionId).toHaveBeenCalledTimes(2) + }) + + expect(result.current.plugins).toEqual([{ plugin_id: 'cached-plugin', type: 'plugin' }]) + expect(result.current.isLoading).toBe(false) + + await act(async () => { + resolveRefresh?.() + }) + + await waitFor(() => { + expect(result.current.plugins).toEqual([{ plugin_id: 'refreshed-plugin', type: 'plugin' }]) + }) + }) +}) + +describe('useMarketplaceCollectionsAndPlugins', () => { + beforeEach(() => { + vi.clearAllMocks() + }) + + it('should retain collection results while refreshing cached data', async () => { + let resolveRefresh: (() => void) | undefined + getMarketplaceCollectionsAndPlugins + .mockResolvedValueOnce({ + marketplaceCollections: [{ id: 'cached-collection' }], + marketplaceCollectionPluginsMap: {}, + }) + .mockImplementationOnce( + () => + new Promise((resolve) => { + resolveRefresh = () => + resolve({ + marketplaceCollections: [{ id: 'refreshed-collection' }], + marketplaceCollectionPluginsMap: {}, + }) + }), + ) + + const { useMarketplaceCollectionsAndPlugins } = await import('../hooks') + const { Wrapper, queryClient } = createWrapper() + const { result } = renderHook(() => useMarketplaceCollectionsAndPlugins(), { + wrapper: Wrapper, + }) + + act(() => { + result.current.queryMarketplaceCollectionsAndPlugins() + }) + + await waitFor(() => { + expect(result.current.marketplaceCollections).toEqual([{ id: 'cached-collection' }]) + }) + + act(() => { + void queryClient.invalidateQueries({ queryKey: ['marketplaceCollectionsAndPlugins'] }) + }) + + await waitFor(() => { + expect(getMarketplaceCollectionsAndPlugins).toHaveBeenCalledTimes(2) + }) + + expect(result.current.marketplaceCollections).toEqual([{ id: 'cached-collection' }]) + expect(result.current.isLoading).toBe(false) + + await act(async () => { + resolveRefresh?.() + }) + + await waitFor(() => { + expect(result.current.marketplaceCollections).toEqual([{ id: 'refreshed-collection' }]) + }) + }) +}) diff --git a/web/app/components/plugins/marketplace/hooks.ts b/web/app/components/plugins/marketplace/hooks.ts index 20f04a9a08c..455ae83dd92 100644 --- a/web/app/components/plugins/marketplace/hooks.ts +++ b/web/app/components/plugins/marketplace/hooks.ts @@ -41,7 +41,7 @@ export const useMarketplaceCollectionsAndPlugins = () => { }, [], ) - const isLoading = !!queryParams && (isFetching || isPending) + const isLoading = !!queryParams && (isPending || (isFetching && !data)) return { marketplaceCollections: marketplaceCollectionsOverride ?? data?.marketplaceCollections, @@ -73,14 +73,14 @@ export const useMarketplacePluginsByCollectionId = ( return { plugins: data || [], - isLoading: !!collectionId && (isFetching || isPending), + isLoading: !!collectionId && (isPending || (isFetching && !data)), isSuccess, } } /** * @deprecated Use useMarketplacePlugins from query.ts instead */ -export const useMarketplacePlugins = () => { +export const useMarketplacePlugins = (enabled = true) => { const queryClient = useQueryClient() const [queryParams, setQueryParams] = useState() @@ -150,7 +150,7 @@ export const useMarketplacePlugins = () => { return loaded < (lastPage.total || 0) ? nextPage : undefined }, initialPageParam: 1, - enabled: !!queryParams, + enabled: enabled && !!queryParams, staleTime: 1000 * 60 * 5, gcTime: 1000 * 60 * 10, retry: false, @@ -187,6 +187,7 @@ export const useMarketplacePlugins = () => { : undefined const total = hasQuery && hasData ? marketplacePluginsQuery.data.pages?.[0]?.total : undefined const isPluginsLoading = + enabled && hasQuery && (marketplacePluginsQuery.isPending || (marketplacePluginsQuery.isFetching && !marketplacePluginsQuery.data)) diff --git a/web/app/components/plugins/plugin-item/__tests__/index.spec.tsx b/web/app/components/plugins/plugin-item/__tests__/index.spec.tsx index 2320cdce7b6..3abf9a20a65 100644 --- a/web/app/components/plugins/plugin-item/__tests__/index.spec.tsx +++ b/web/app/components/plugins/plugin-item/__tests__/index.spec.tsx @@ -213,6 +213,10 @@ describe('PluginItem', () => { // Assert const img = screen.getByRole('img') expect(img).toHaveAttribute('alt', `plugin-${plugin.plugin_unique_identifier}-logo`) + expect(img).toHaveAttribute('loading', 'lazy') + expect(img).toHaveAttribute('decoding', 'async') + expect(img).toHaveAttribute('width', '40') + expect(img).toHaveAttribute('height', '40') }) it('should not render category label in corner mark', () => { diff --git a/web/app/components/plugins/plugin-item/index.tsx b/web/app/components/plugins/plugin-item/index.tsx index 40bd028bf8b..e6590f4116d 100644 --- a/web/app/components/plugins/plugin-item/index.tsx +++ b/web/app/components/plugins/plugin-item/index.tsx @@ -137,8 +137,12 @@ const PluginItem: FC = ({
{`plugin-${plugin_unique_identifier}-logo`}
diff --git a/web/app/components/plugins/plugin-page/__tests__/plugins-panel.spec.tsx b/web/app/components/plugins/plugin-page/__tests__/plugins-panel.spec.tsx index cb9f5945ea0..eba15948cef 100644 --- a/web/app/components/plugins/plugin-page/__tests__/plugins-panel.spec.tsx +++ b/web/app/components/plugins/plugin-page/__tests__/plugins-panel.spec.tsx @@ -1,6 +1,6 @@ import type { PluginDetail } from '../../types' import type { Collection } from '@/app/components/tools/types' -import { fireEvent, render, screen } from '@testing-library/react' +import { act, fireEvent, render, screen } from '@testing-library/react' import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' import { getStepByStepTourTargetSelector, @@ -22,8 +22,12 @@ const mockSetFilters = vi.fn() const mockSetCurrentPluginID = vi.fn() const mockLoadNextPage = vi.fn() const mockInvalidateInstalledPluginList = vi.fn() +const mockRemoveFilteredInstalledPluginPageOnUnmount = vi.fn() const mockUseInstalledPluginList = vi.fn() const mockPluginListWithLatestVersion = vi.fn<() => PluginDetail[]>(() => []) +const intersectionObserverCallbacks: IntersectionObserverCallback[] = [] +const mockObserve = vi.fn() +const mockDisconnect = vi.fn() vi.mock('@tanstack/react-query', () => ({ queryOptions: (options: unknown) => options, @@ -34,8 +38,11 @@ vi.mock('@/i18n-config', () => ({ })) vi.mock('@/service/use-plugins', () => ({ + normalizePluginCategoryListLanguage: (locale: string) => locale, useInstalledPluginList: (...args: unknown[]) => mockUseInstalledPluginList(...args), useInvalidateInstalledPluginList: () => mockInvalidateInstalledPluginList, + useRemoveFilteredInstalledPluginPageOnUnmount: (...args: unknown[]) => + mockRemoveFilteredInstalledPluginPageOnUnmount(...args), })) vi.mock('../../hooks', () => ({ @@ -160,7 +167,7 @@ vi.mock('@/app/components/integrations/tool-provider-card', () => ({ ), })) -vi.mock('@/app/components/integrations/hooks/use-tool-marketplace-panel', () => ({ +vi.mock('@/app/components/tools/marketplace/use-tool-marketplace-panel', () => ({ useToolMarketplacePanel: () => ({ isMarketplaceArrowVisible: true, marketplaceContext: {}, @@ -265,6 +272,18 @@ describe('PluginsPanel', () => { beforeEach(() => { vi.clearAllMocks() vi.useFakeTimers() + intersectionObserverCallbacks.length = 0 + vi.stubGlobal( + 'IntersectionObserver', + class { + constructor(callback: IntersectionObserverCallback) { + intersectionObserverCallbacks.push(callback) + } + + observe = mockObserve + disconnect = mockDisconnect + }, + ) mockState.filters = { categories: [], tags: [], searchQuery: '' } mockState.currentPluginID = undefined mockUseInstalledPluginList.mockReturnValue({ @@ -279,6 +298,7 @@ describe('PluginsPanel', () => { afterEach(() => { vi.useRealTimers() + vi.unstubAllGlobals() }) it('renders the loading state while the plugin list is pending', () => { @@ -351,23 +371,49 @@ describe('PluginsPanel', () => { expect(screen.getByTestId('plugin-list')).not.toHaveTextContent('tool-plugin') }) - it('loads the scoped plugin category list whenever an integrations category panel mounts', () => { - render() + it.each([ + PluginCategoryEnum.tool, + PluginCategoryEnum.trigger, + PluginCategoryEnum.agent, + PluginCategoryEnum.extension, + ])('loads %s Integration Plugins in Studio-sized pages', (category) => { + render() - expect(mockUseInstalledPluginList).toHaveBeenCalledWith(false, 100, { - category: PluginCategoryEnum.trigger, - refetchOnMount: 'always', - }) + expect(mockUseInstalledPluginList).toHaveBeenCalledWith( + expect.objectContaining({ + category, + gcTime: 10 * 60 * 1000, + pageSize: 30, + staleTime: 5 * 60 * 1000, + }), + ) + expect(mockUseInstalledPluginList.mock.calls.at(-1)?.[0]).not.toHaveProperty('refetchOnMount') + }) + + it('configures filtered cache cleanup for an Integration category panel', () => { + render() + + expect(mockRemoveFilteredInstalledPluginPageOnUnmount).toHaveBeenCalledWith( + PluginCategoryEnum.tool, + 30, + expect.any(Object), + ) + }) + + it('does not configure filtered cache cleanup for the standalone Plugin page', () => { + render() + + expect(mockRemoveFilteredInstalledPluginPageOnUnmount).toHaveBeenCalledWith( + undefined, + 30, + undefined, + ) }) it('loads the scoped tool plugin category list when fixed to tool plugins', () => { render() expect(screen.getByTestId('filter-management')).toHaveAttribute('data-hide-tag-filter', 'false') - expect(mockUseInstalledPluginList).toHaveBeenCalledWith(false, 100, { - category: PluginCategoryEnum.tool, - refetchOnMount: 'always', - }) }) it('filters tool plugins, builtin tools, and marketplace suggestions by selected tags', () => { @@ -711,6 +757,56 @@ describe('PluginsPanel', () => { }) }) + it.each([ + PluginCategoryEnum.tool, + PluginCategoryEnum.trigger, + PluginCategoryEnum.agent, + PluginCategoryEnum.extension, + ])('automatically loads more %s Plugins near the list end once per intersection', (category) => { + mockPluginListWithLatestVersion.mockReturnValue([ + createPlugin('category-plugin', 'Category Plugin', [], category), + ]) + mockUseInstalledPluginList.mockReturnValue({ + data: { plugins: [] }, + isLoading: false, + isFetching: false, + isLastPage: false, + loadNextPage: mockLoadNextPage, + }) + + render() + + expect(intersectionObserverCallbacks).toHaveLength(1) + + act(() => { + intersectionObserverCallbacks[0]?.( + [{ isIntersecting: true } as IntersectionObserverEntry], + {} as IntersectionObserver, + ) + intersectionObserverCallbacks[0]?.( + [{ isIntersecting: true } as IntersectionObserverEntry], + {} as IntersectionObserver, + ) + }) + + expect(mockLoadNextPage).toHaveBeenCalledTimes(1) + }) + + it('does not observe the Tool Plugin list while the next page is loading', () => { + mockPluginListWithLatestVersion.mockReturnValue([createPlugin('tool-plugin', 'Tool Plugin')]) + mockUseInstalledPluginList.mockReturnValue({ + data: { plugins: [] }, + isLoading: false, + isFetching: true, + isLastPage: false, + loadNextPage: mockLoadNextPage, + }) + + render() + + expect(intersectionObserverCallbacks).toHaveLength(0) + }) + it('renders the empty state and keeps the current plugin detail in sync', () => { mockState.currentPluginID = 'beta-tool' mockState.filters.searchQuery = 'missing' diff --git a/web/app/components/plugins/plugin-page/plugins-panel-results.tsx b/web/app/components/plugins/plugin-page/plugins-panel-results.tsx index d37491e645e..e3693b4079d 100644 --- a/web/app/components/plugins/plugin-page/plugins-panel-results.tsx +++ b/web/app/components/plugins/plugin-page/plugins-panel-results.tsx @@ -11,49 +11,15 @@ import { ScrollAreaThumb, ScrollAreaViewport, } from '@langgenius/dify-ui/scroll-area' +import { useEffect, useRef } from 'react' import { useTranslation } from 'react-i18next' import Loading from '@/app/components/base/loading' -import { useToolMarketplacePanel } from '@/app/components/integrations/hooks/use-tool-marketplace-panel' import IntegrationsToolProviderCard from '@/app/components/integrations/tool-provider-card' -import Marketplace from '@/app/components/tools/marketplace' +import { BuiltinMarketplacePanel } from '@/app/components/tools/marketplace/builtin-marketplace-panel' import List from './list' -type BuiltinMarketplacePanelProps = { - containerRef: RefObject - contentInset: PluginPageContentInset - keywords: string - tagFilterValue: string[] -} - -const BuiltinMarketplacePanel = ({ - containerRef, - contentInset, - keywords, - tagFilterValue, -}: BuiltinMarketplacePanelProps) => { - const { isMarketplaceArrowVisible, marketplaceContext, showMarketplacePanel, toolListTailRef } = - useToolMarketplacePanel({ - containerRef, - keywords, - tagFilterValue, - }) - - return ( - <> -
- - - ) -} - type PluginsPanelResultsProps = { + autoLoadNextPage: boolean canDeletePlugin: boolean canUpdatePlugin: boolean containerRef: RefObject @@ -78,6 +44,7 @@ type PluginsPanelResultsProps = { } const PluginsPanelResults = ({ + autoLoadNextPage, canDeletePlugin, canUpdatePlugin, containerRef, @@ -101,6 +68,42 @@ const PluginsPanelResults = ({ tagFilterValue, }: PluginsPanelResultsProps) => { const { t } = useTranslation() + const loadMoreAnchorRef = useRef(null) + const loadNextPageRequestedRef = useRef(false) + + useEffect(() => { + const anchor = loadMoreAnchorRef.current + const root = containerRef.current + + if (!isFetching) loadNextPageRequestedRef.current = false + + if ( + !autoLoadNextPage || + !anchor || + !root || + isFetching || + isLastPage || + !globalThis.IntersectionObserver + ) + return + + const observer = new IntersectionObserver( + (entries) => { + if (!entries[0]?.isIntersecting || loadNextPageRequestedRef.current) return + + loadNextPageRequestedRef.current = true + loadNextPage() + }, + { + root, + rootMargin: '200px', + threshold: 0.1, + }, + ) + + observer.observe(anchor) + return () => observer.disconnect() + }, [autoLoadNextPage, containerRef, isFetching, isLastPage, loadNextPage]) return ( {isFetching ? ( - ) : ( + ) : autoLoadNextPage ? null : ( )} + {autoLoadNextPage && ( +
+ )}
)} {hasToolMarketplacePanel && ( diff --git a/web/app/components/plugins/plugin-page/plugins-panel.tsx b/web/app/components/plugins/plugin-page/plugins-panel.tsx index da4ddcfaae3..548e783a914 100644 --- a/web/app/components/plugins/plugin-page/plugins-panel.tsx +++ b/web/app/components/plugins/plugin-page/plugins-panel.tsx @@ -15,7 +15,12 @@ import ProviderDetail from '@/app/components/tools/provider/detail' import { useGetLanguage } from '@/context/i18n' import { systemFeaturesQueryOptions } from '@/features/system-features/client' import { renderI18nObject } from '@/i18n-config' -import { useInstalledPluginList, useInvalidateInstalledPluginList } from '@/service/use-plugins' +import { + normalizePluginCategoryListLanguage, + useInstalledPluginList, + useInvalidateInstalledPluginList, + useRemoveFilteredInstalledPluginPageOnUnmount, +} from '@/service/use-plugins' import { usePluginsWithLatestVersion } from '../hooks' import { PluginCategoryEnum } from '../types' import { pluginPageContentFrameClassNames, pluginPageContentInsetClassNames } from './content-inset' @@ -26,6 +31,10 @@ import PluginListSkeleton from './plugin-list-skeleton' import PluginsPanelResults from './plugins-panel-results' import { EMPTY_BUILTIN_TOOLS, filterBuiltinTools } from './plugins-panel-utils' +const INTEGRATION_PLUGIN_PAGE_SIZE = 30 +const INTEGRATION_PLUGIN_STALE_TIME = 5 * 60 * 1000 +const INTEGRATION_PLUGIN_GC_TIME = 10 * 60 * 1000 + const matchesSearchQuery = ( plugin: PluginDetail & { latest_version: string }, query: string, @@ -84,6 +93,17 @@ const PluginsPanel = ({ isAgentStrategyIntegrationPage || isExtensionIntegrationPage const supportsTagFilter = !fixedCategory || isToolIntegrationPage || isTriggerIntegrationPage + const installedPluginFilters = useMemo( + () => + isIntegrationCategoryPage + ? { + language: normalizePluginCategoryListLanguage(locale), + query: filters.searchQuery, + tags: supportsTagFilter ? filters.tags : [], + } + : undefined, + [filters.searchQuery, filters.tags, isIntegrationCategoryPage, locale, supportsTagFilter], + ) const { data: enableMarketplace } = useSuspenseQuery({ ...systemFeaturesQueryOptions(), select: (s) => s.enable_marketplace, @@ -94,18 +114,20 @@ const PluginsPanel = ({ isFetching, isLastPage, loadNextPage, - } = useInstalledPluginList( - false, - 100, - fixedCategory - ? { - category: fixedCategory, - refetchOnMount: isIntegrationCategoryPage ? 'always' : undefined, - } - : undefined, - ) + } = useInstalledPluginList({ + category: fixedCategory, + filters: installedPluginFilters, + gcTime: isIntegrationCategoryPage ? INTEGRATION_PLUGIN_GC_TIME : undefined, + pageSize: isIntegrationCategoryPage ? INTEGRATION_PLUGIN_PAGE_SIZE : 100, + staleTime: isIntegrationCategoryPage ? INTEGRATION_PLUGIN_STALE_TIME : undefined, + }) const pluginListWithLatestVersion = usePluginsWithLatestVersion(pluginList?.plugins) const invalidateInstalledPluginList = useInvalidateInstalledPluginList() + useRemoveFilteredInstalledPluginPageOnUnmount( + isIntegrationCategoryPage ? fixedCategory : undefined, + INTEGRATION_PLUGIN_PAGE_SIZE, + installedPluginFilters, + ) const currentPluginID = usePluginPageContext((v) => v.currentPluginID) const setCurrentPluginID = usePluginPageContext((v) => v.setCurrentPluginID) const [currentBuiltinToolID, setCurrentBuiltinToolID] = useState() @@ -236,6 +258,7 @@ const PluginsPanel = ({ <> {hasVisiblePlugins || hasVisibleBuiltinTools || hasToolMarketplacePanel ? ( { expect(onShowChange).not.toHaveBeenCalled() }) - it('should call onSelect with correct params when a version is selected', () => { + it('should call onSelect with correct params when a version is selected', async () => { // Arrange const onSelect = vi.fn() const onShowChange = vi.fn() + const user = userEvent.setup() // Act render( @@ -995,12 +997,7 @@ describe('update-plugin', () => { onShowChange={onShowChange} />, ) - // Click on version 2.0.0 - const versionElements = screen.getAllByText(/^\d+\.\d+\.\d+$/) - const version2Element = versionElements.find((el) => el.textContent === '2.0.0') - if (version2Element) { - fireEvent.click(version2Element.closest('div[class*="cursor-pointer"]')!) - } + await user.click(screen.getByRole('button', { name: /2\.0\.0/ })) // Assert expect(onSelect).toHaveBeenCalledWith({ @@ -1011,9 +1008,10 @@ describe('update-plugin', () => { expect(onShowChange).toHaveBeenCalledWith(false) }) - it('should not call onSelect when clicking on current version', () => { + it('should not call onSelect when clicking on current version', async () => { // Arrange const onSelect = vi.fn() + const user = userEvent.setup() // Act render( @@ -1024,20 +1022,16 @@ describe('update-plugin', () => { onSelect={onSelect} />, ) - // Click on current version 1.0.0 - const versionElements = screen.getAllByText(/^\d+\.\d+\.\d+$/) - const version1Element = versionElements.find((el) => el.textContent === '1.0.0') - if (version1Element) { - fireEvent.click(version1Element.closest('div[class*="cursor"]')!) - } + await user.click(screen.getByRole('button', { name: /1\.0\.0/ })) // Assert expect(onSelect).not.toHaveBeenCalled() }) - it('should indicate downgrade when selecting a lower version', () => { + it('should indicate downgrade when selecting a lower version', async () => { // Arrange const onSelect = vi.fn() + const user = userEvent.setup() // Act render( @@ -1048,12 +1042,7 @@ describe('update-plugin', () => { onSelect={onSelect} />, ) - // Click on version 1.0.0 (downgrade) - const versionElements = screen.getAllByText(/^\d+\.\d+\.\d+$/) - const version1Element = versionElements.find((el) => el.textContent === '1.0.0') - if (version1Element) { - fireEvent.click(version1Element.closest('div[class*="cursor-pointer"]')!) - } + await user.click(screen.getByRole('button', { name: /1\.0\.0/ })) // Assert expect(onSelect).toHaveBeenCalledWith({ diff --git a/web/app/components/plugins/update-plugin/__tests__/plugin-version-picker.spec.tsx b/web/app/components/plugins/update-plugin/__tests__/plugin-version-picker.spec.tsx index eac9fbadb88..6e180059814 100644 --- a/web/app/components/plugins/update-plugin/__tests__/plugin-version-picker.spec.tsx +++ b/web/app/components/plugins/update-plugin/__tests__/plugin-version-picker.spec.tsx @@ -1,4 +1,5 @@ -import { fireEvent, render, screen } from '@testing-library/react' +import { render, screen } from '@testing-library/react' +import userEvent from '@testing-library/user-event' import { beforeEach, describe, expect, it, vi } from 'vitest' import PluginVersionPicker from '../plugin-version-picker' @@ -14,6 +15,8 @@ const mockVersionList = vi.hoisted(() => ({ }, })) +const mockUseVersionListOfPlugin = vi.hoisted(() => vi.fn()) + vi.mock('@/hooks/use-timestamp', () => ({ default: () => ({ formatDate: (value: string, format: string) => `${value}:${format}`, @@ -21,14 +24,19 @@ vi.mock('@/hooks/use-timestamp', () => ({ })) vi.mock('@/service/use-plugins', () => ({ - useVersionListOfPlugin: () => ({ + useVersionListOfPlugin: mockUseVersionListOfPlugin.mockImplementation(() => ({ data: mockVersionList, - }), + isLoading: false, + })), })) describe('PluginVersionPicker', () => { beforeEach(() => { vi.clearAllMocks() + mockUseVersionListOfPlugin.mockReturnValue({ + data: mockVersionList, + isLoading: false, + }) mockVersionList.data.versions = [ { version: '2.0.0', @@ -43,6 +51,55 @@ describe('PluginVersionPicker', () => { ] }) + it('loads versions only while the popover is open', () => { + const { rerender } = render( + trigger} + onSelect={vi.fn()} + />, + ) + + expect(mockUseVersionListOfPlugin).toHaveBeenLastCalledWith('plugin-1', false) + + rerender( + trigger} + onSelect={vi.fn()} + />, + ) + + expect(mockUseVersionListOfPlugin).toHaveBeenLastCalledWith('plugin-1', true) + }) + + it('shows a loading state while versions are loading', () => { + mockUseVersionListOfPlugin.mockReturnValue({ + data: undefined, + isLoading: true, + }) + + render( + trigger} + onSelect={vi.fn()} + />, + ) + + expect(screen.getByRole('status', { name: 'common.loading' })).toBeInTheDocument() + expect(screen.queryByText('2.0.0')).not.toBeInTheDocument() + }) + it('renders version options and highlights the current version', () => { render( { expect(currentBadge).toHaveClass('bg-components-badge-bg-dimm') }) - it('calls onSelect with downgrade metadata and closes the picker', () => { + it('calls onSelect with downgrade metadata and closes the picker', async () => { const onSelect = vi.fn() const onShowChange = vi.fn() + const user = userEvent.setup() render( { />, ) - fireEvent.click(screen.getByText('1.0.0')) + await user.click(screen.getByText('1.0.0')) expect(onSelect).toHaveBeenCalledWith({ version: '1.0.0', @@ -130,8 +188,9 @@ describe('PluginVersionPicker', () => { expect(onShowChange).toHaveBeenCalledWith(false) }) - it('does not call onSelect when the current version is clicked', () => { + it('does not call onSelect when the current version is clicked', async () => { const onSelect = vi.fn() + const user = userEvent.setup() render( { />, ) - fireEvent.click(screen.getByText('2.0.0')) + await user.click(screen.getByText('2.0.0')) expect(onSelect).not.toHaveBeenCalled() }) diff --git a/web/app/components/plugins/update-plugin/plugin-version-picker.tsx b/web/app/components/plugins/update-plugin/plugin-version-picker.tsx index c2c286f8273..6d8aabf5d51 100644 --- a/web/app/components/plugins/update-plugin/plugin-version-picker.tsx +++ b/web/app/components/plugins/update-plugin/plugin-version-picker.tsx @@ -48,7 +48,7 @@ const PluginVersionPicker: FC = ({ const format = t(($) => $.dateTimeFormat, { ns: 'appLog' }).split(' ')[0] const { formatDate } = useTimestamp() - const { data: res } = useVersionListOfPlugin(pluginID) + const { data: res, isLoading } = useVersionListOfPlugin(pluginID, isShow && !disabled) const handleSelect = useCallback( ({ @@ -99,36 +99,50 @@ const PluginVersionPicker: FC = ({
{t(($) => $['detailPanel.switchVersion'], { ns: 'plugin' })}
-
- {res?.data.versions.map((version) => ( +
+ {isLoading ? (
- handleSelect({ - version: version.version, - unique_identifier: version.unique_identifier, - isDowngrade: isEarlierThanVersion(version.version, currentVersion), - }) - } + role="status" + aria-label={t(($) => $.loading, { ns: 'common' })} + className="flex h-12 items-center justify-center" > -
-
- {version.version} -
- {currentVersion === version.version && ( - - )} -
- {formatDate(version.created_at, format!)} -
-
+
- ))} + ) : ( + res?.data.versions.map((version) => ( + + )) + )}
diff --git a/web/app/components/tools/marketplace/__tests__/builtin-marketplace-panel.spec.tsx b/web/app/components/tools/marketplace/__tests__/builtin-marketplace-panel.spec.tsx new file mode 100644 index 00000000000..fb1e712fefc --- /dev/null +++ b/web/app/components/tools/marketplace/__tests__/builtin-marketplace-panel.spec.tsx @@ -0,0 +1,95 @@ +import { act, fireEvent, render, screen } from '@testing-library/react' +import { createRef } from 'react' +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' +import { BuiltinMarketplacePanel } from '../builtin-marketplace-panel' + +const mockUseMarketplace = vi.fn() +const intersectionObserverCallbacks: IntersectionObserverCallback[] = [] + +vi.mock('@/app/components/tools/marketplace/hooks', () => ({ + useMarketplace: (...args: unknown[]) => mockUseMarketplace(...args), +})) + +vi.mock('@/app/components/tools/marketplace', () => ({ + default: ({ showMarketplacePanel }: { showMarketplacePanel: () => void }) => ( + + ), +})) + +describe('BuiltinMarketplacePanel', () => { + beforeEach(() => { + vi.clearAllMocks() + intersectionObserverCallbacks.length = 0 + mockUseMarketplace.mockReturnValue({ handleScroll: vi.fn() }) + vi.stubGlobal( + 'IntersectionObserver', + class { + constructor(callback: IntersectionObserverCallback) { + intersectionObserverCallbacks.push(callback) + } + + observe = vi.fn() + disconnect = vi.fn() + }, + ) + }) + + afterEach(() => { + vi.unstubAllGlobals() + }) + + const renderPanel = (overrides?: { keywords?: string; tagFilterValue?: string[] }) => { + const containerRef = createRef() + const { container } = render( +
+ +
, + ) + + return { container, containerRef } + } + + it('defers Marketplace queries until the section enters the viewport', () => { + renderPanel() + + expect(mockUseMarketplace).toHaveBeenLastCalledWith('', [], false) + + act(() => { + intersectionObserverCallbacks[0]?.( + [{ isIntersecting: true }] as IntersectionObserverEntry[], + {} as IntersectionObserver, + ) + }) + + expect(mockUseMarketplace).toHaveBeenLastCalledWith('', [], true) + }) + + it('loads Marketplace immediately when the user searches', () => { + renderPanel({ keywords: 'weather' }) + + expect(mockUseMarketplace).toHaveBeenLastCalledWith('weather', [], true) + }) + + it('loads Marketplace immediately when the user filters by tag', () => { + renderPanel({ tagFilterValue: ['search'] }) + + expect(mockUseMarketplace).toHaveBeenLastCalledWith('', ['search'], true) + }) + + it('activates Marketplace before scrolling to it from the arrow action', () => { + const { containerRef } = renderPanel() + containerRef.current!.scrollTo = vi.fn() + + fireEvent.click(screen.getByRole('button', { name: 'Marketplace' })) + + expect(mockUseMarketplace).toHaveBeenLastCalledWith('', [], true) + expect(containerRef.current!.scrollTo).toHaveBeenCalledWith({ top: -80, behavior: 'smooth' }) + }) +}) diff --git a/web/app/components/tools/marketplace/__tests__/hooks.spec.ts b/web/app/components/tools/marketplace/__tests__/hooks.spec.ts index 45fc64d87b9..b78f59faa85 100644 --- a/web/app/components/tools/marketplace/__tests__/hooks.spec.ts +++ b/web/app/components/tools/marketplace/__tests__/hooks.spec.ts @@ -1,11 +1,12 @@ +import type { ReactNode } from 'react' import type { Plugin } from '@/app/components/plugins/types' -import type { Collection } from '@/app/components/tools/types' +import { QueryClient, QueryClientProvider } from '@tanstack/react-query' import { act, renderHook, waitFor } from '@testing-library/react' +import { createElement } from 'react' import { beforeEach, describe, expect, it, vi } from 'vitest' import { SCROLL_BOTTOM_THRESHOLD } from '@/app/components/plugins/marketplace/constants' import { getMarketplaceListCondition } from '@/app/components/plugins/marketplace/utils' import { PluginCategoryEnum } from '@/app/components/plugins/types' -import { CollectionType } from '@/app/components/tools/types' import { useMarketplace } from '../hooks' // ==================== Mock Setup ==================== @@ -18,15 +19,27 @@ const mockFetchNextPage = vi.fn() const mockUseMarketplaceCollectionsAndPlugins = vi.fn() const mockUseMarketplacePlugins = vi.fn() +const mockInstalledIdsQueryOptions = vi.fn() vi.mock('@/app/components/plugins/marketplace/hooks', () => ({ useMarketplaceCollectionsAndPlugins: (...args: unknown[]) => mockUseMarketplaceCollectionsAndPlugins(...args), useMarketplacePlugins: (...args: unknown[]) => mockUseMarketplacePlugins(...args), })) -const mockUseAllToolProviders = vi.fn() -vi.mock('@/service/use-tools', () => ({ - useAllToolProviders: (...args: unknown[]) => mockUseAllToolProviders(...args), +vi.mock('@/service/client', () => ({ + consoleQuery: { + workspaces: { + current: { + plugin: { + installedIds: { + get: { + queryOptions: (...args: unknown[]) => mockInstalledIdsQueryOptions(...args), + }, + }, + }, + }, + }, + }, })) vi.mock('@/utils/var', () => ({ @@ -35,20 +48,12 @@ vi.mock('@/utils/var', () => ({ // ==================== Test Utilities ==================== -const createToolProvider = (overrides: Partial = {}): Collection => ({ - id: 'provider-1', - name: 'Provider 1', - author: 'Author', - description: { en_US: 'desc', zh_Hans: '描述' }, - icon: 'icon', - label: { en_US: 'label', zh_Hans: '标签' }, - type: CollectionType.custom, - team_credentials: {}, - is_team_authorization: false, - allow_delete: false, - labels: [], - ...overrides, -}) +const createWrapper = () => { + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }) + return function Wrapper({ children }: { children: ReactNode }) { + return createElement(QueryClientProvider, { client: queryClient }, children) + } +} const setupHookMocks = (overrides?: { isLoading?: boolean @@ -80,27 +85,37 @@ const setupHookMocks = (overrides?: { describe('useMarketplace', () => { beforeEach(() => { vi.clearAllMocks() - mockUseAllToolProviders.mockReturnValue({ - data: [], - isSuccess: true, + mockInstalledIdsQueryOptions.mockReturnValue({ + queryKey: ['installed-plugin-ids', 'tool'], + queryFn: () => Promise.resolve({ plugin_ids: [] }), }) setupHookMocks() }) describe('Queries', () => { - it('should query plugins with debounce when search text is provided', async () => { - mockUseAllToolProviders.mockReturnValue({ - data: [ - createToolProvider({ plugin_id: 'plugin-a' }), - createToolProvider({ plugin_id: undefined }), - ], - isSuccess: true, - }) - - renderHook(() => useMarketplace('alpha', [])) + it('does not query installed IDs or Marketplace data before activation', async () => { + renderHook(() => useMarketplace('', [], false), { wrapper: createWrapper() }) await waitFor(() => { - expect(mockQueryPluginsWithDebounced).toHaveBeenCalledWith({ + expect(mockInstalledIdsQueryOptions).toHaveBeenCalledWith({ + input: { query: { category: 'tool' } }, + enabled: false, + }) + }) + expect(mockQueryMarketplaceCollectionsAndPlugins).not.toHaveBeenCalled() + expect(mockQueryPlugins).not.toHaveBeenCalled() + }) + + it('should query plugins when the debounced page filter provides search text', async () => { + mockInstalledIdsQueryOptions.mockReturnValue({ + queryKey: ['installed-plugin-ids', 'tool'], + queryFn: () => Promise.resolve({ plugin_ids: ['plugin-a'] }), + }) + + renderHook(() => useMarketplace('alpha', []), { wrapper: createWrapper() }) + + await waitFor(() => { + expect(mockQueryPlugins).toHaveBeenCalledWith({ category: PluginCategoryEnum.tool, query: 'alpha', tags: [], @@ -108,17 +123,18 @@ describe('useMarketplace', () => { type: 'plugin', }) }) + expect(mockQueryPluginsWithDebounced).not.toHaveBeenCalled() expect(mockQueryMarketplaceCollectionsAndPlugins).not.toHaveBeenCalled() expect(mockResetPlugins).not.toHaveBeenCalled() }) it('should query plugins immediately when only tags are provided', async () => { - mockUseAllToolProviders.mockReturnValue({ - data: [createToolProvider({ plugin_id: 'plugin-b' })], - isSuccess: true, + mockInstalledIdsQueryOptions.mockReturnValue({ + queryKey: ['installed-plugin-ids', 'tool'], + queryFn: () => Promise.resolve({ plugin_ids: ['plugin-b'] }), }) - renderHook(() => useMarketplace('', ['tag-1'])) + renderHook(() => useMarketplace('', ['tag-1']), { wrapper: createWrapper() }) await waitFor(() => { expect(mockQueryPlugins).toHaveBeenCalledWith({ @@ -132,12 +148,12 @@ describe('useMarketplace', () => { }) it('should query collections and reset plugins when no filters are provided', async () => { - mockUseAllToolProviders.mockReturnValue({ - data: [createToolProvider({ plugin_id: 'plugin-c' })], - isSuccess: true, + mockInstalledIdsQueryOptions.mockReturnValue({ + queryKey: ['installed-plugin-ids', 'tool'], + queryFn: () => Promise.resolve({ plugin_ids: ['plugin-c'] }), }) - renderHook(() => useMarketplace('', [])) + renderHook(() => useMarketplace('', []), { wrapper: createWrapper() }) await waitFor(() => { expect(mockQueryMarketplaceCollectionsAndPlugins).toHaveBeenCalledWith({ @@ -155,7 +171,7 @@ describe('useMarketplace', () => { it('should expose combined loading state and fallback page value', () => { setupHookMocks({ isLoading: true, isPluginsLoading: false, pluginsPage: undefined }) - const { result } = renderHook(() => useMarketplace('', [])) + const { result } = renderHook(() => useMarketplace('', []), { wrapper: createWrapper() }) expect(result.current.isLoading).toBe(true) expect(result.current.page).toBe(1) @@ -165,7 +181,9 @@ describe('useMarketplace', () => { describe('Scroll', () => { it('should fetch next page when scrolling near bottom with filters', () => { setupHookMocks({ hasNextPage: true }) - const { result } = renderHook(() => useMarketplace('search', [])) + const { result } = renderHook(() => useMarketplace('search', []), { + wrapper: createWrapper(), + }) const event = { target: { scrollTop: 100, @@ -183,7 +201,7 @@ describe('useMarketplace', () => { it('should not fetch next page when no filters are applied', () => { setupHookMocks({ hasNextPage: true }) - const { result } = renderHook(() => useMarketplace('', [])) + const { result } = renderHook(() => useMarketplace('', []), { wrapper: createWrapper() }) const event = { target: { scrollTop: 100, diff --git a/web/app/components/tools/marketplace/__tests__/index.spec.tsx b/web/app/components/tools/marketplace/__tests__/index.spec.tsx index eed9f38862b..5c46a9d78ba 100644 --- a/web/app/components/tools/marketplace/__tests__/index.spec.tsx +++ b/web/app/components/tools/marketplace/__tests__/index.spec.tsx @@ -200,7 +200,7 @@ describe('Marketplace', () => { // Arrange const marketplaceContext = createMarketplaceContext() const showMarketplacePanel = vi.fn() - const { container } = render( + render( { ) // Act - const arrowIcon = container.querySelector('svg.cursor-pointer') - expect(arrowIcon).toBeTruthy() - await user.click(arrowIcon as SVGElement) + await user.click(screen.getByRole('button', { name: /plugin.marketplace.moreFrom/i })) // Assert expect(showMarketplacePanel).toHaveBeenCalledTimes(1) diff --git a/web/app/components/tools/marketplace/builtin-marketplace-panel.tsx b/web/app/components/tools/marketplace/builtin-marketplace-panel.tsx new file mode 100644 index 00000000000..f012618caea --- /dev/null +++ b/web/app/components/tools/marketplace/builtin-marketplace-panel.tsx @@ -0,0 +1,39 @@ +import type { RefObject } from 'react' +import type { ToolsContentInset } from '../content-inset' +import Marketplace from '.' +import { useToolMarketplacePanel } from './use-tool-marketplace-panel' + +type BuiltinMarketplacePanelProps = { + containerRef: RefObject + contentInset: ToolsContentInset + keywords: string + tagFilterValue: string[] +} + +export function BuiltinMarketplacePanel({ + containerRef, + contentInset, + keywords, + tagFilterValue, +}: BuiltinMarketplacePanelProps) { + const { isMarketplaceArrowVisible, marketplaceContext, showMarketplacePanel, toolListTailRef } = + useToolMarketplacePanel({ + containerRef, + keywords, + tagFilterValue, + }) + + return ( + <> +
+ + + ) +} diff --git a/web/app/components/tools/marketplace/hooks.ts b/web/app/components/tools/marketplace/hooks.ts index b893a63fbdd..1b692200c0b 100644 --- a/web/app/components/tools/marketplace/hooks.ts +++ b/web/app/components/tools/marketplace/hooks.ts @@ -1,4 +1,5 @@ -import { useCallback, useEffect, useMemo, useRef } from 'react' +import { useQuery } from '@tanstack/react-query' +import { useCallback, useEffect, useRef } from 'react' import { SCROLL_BOTTOM_THRESHOLD } from '@/app/components/plugins/marketplace/constants' import { useMarketplaceCollectionsAndPlugins, @@ -6,16 +7,20 @@ import { } from '@/app/components/plugins/marketplace/hooks' import { getMarketplaceListCondition } from '@/app/components/plugins/marketplace/utils' import { PluginCategoryEnum } from '@/app/components/plugins/types' -import { useAllToolProviders } from '@/service/use-tools' +import { consoleQuery } from '@/service/client' -export const useMarketplace = (searchPluginText: string, filterPluginTags: string[]) => { - const { data: toolProvidersData, isSuccess } = useAllToolProviders() - const exclude = useMemo(() => { - if (isSuccess) - return toolProvidersData - ?.filter((toolProvider) => !!toolProvider.plugin_id) - .map((toolProvider) => toolProvider.plugin_id!) - }, [isSuccess, toolProvidersData]) +export const useMarketplace = ( + searchPluginText: string, + filterPluginTags: string[], + enabled = true, +) => { + const { data: installedPluginIds, isSuccess } = useQuery( + consoleQuery.workspaces.current.plugin.installedIds.get.queryOptions({ + input: { query: { category: 'tool' } }, + enabled, + }), + ) + const exclude = installedPluginIds?.plugin_ids const { isLoading, marketplaceCollections, @@ -26,7 +31,6 @@ export const useMarketplace = (searchPluginText: string, filterPluginTags: strin plugins, resetPlugins, queryPlugins, - queryPluginsWithDebounced, isLoading: isPluginsLoading, fetchNextPage, hasNextPage, @@ -40,9 +44,11 @@ export const useMarketplace = (searchPluginText: string, filterPluginTags: strin filterPluginTagsRef.current = filterPluginTags }, [searchPluginText, filterPluginTags]) useEffect(() => { + if (!enabled) return + if ((searchPluginText || filterPluginTags.length) && isSuccess) { if (searchPluginText) { - queryPluginsWithDebounced({ + queryPlugins({ category: PluginCategoryEnum.tool, query: searchPluginText, tags: filterPluginTags, @@ -74,9 +80,9 @@ export const useMarketplace = (searchPluginText: string, filterPluginTags: strin filterPluginTags, queryPlugins, queryMarketplaceCollectionsAndPlugins, - queryPluginsWithDebounced, resetPlugins, exclude, + enabled, isSuccess, ]) @@ -87,14 +93,15 @@ export const useMarketplace = (searchPluginText: string, filterPluginTags: strin if (scrollTop + clientHeight >= scrollHeight - SCROLL_BOTTOM_THRESHOLD && scrollTop > 0) { const searchPluginText = searchPluginTextRef.current const filterPluginTags = filterPluginTagsRef.current - if (hasNextPage && (!!searchPluginText || !!filterPluginTags.length)) fetchNextPage() + if (enabled && hasNextPage && (!!searchPluginText || !!filterPluginTags.length)) + fetchNextPage() } }, - [exclude, fetchNextPage, hasNextPage, plugins, queryPlugins], + [enabled, fetchNextPage, hasNextPage], ) return { - isLoading: isLoading || isPluginsLoading, + isLoading: enabled && (isLoading || isPluginsLoading), marketplaceCollections, marketplaceCollectionPluginsMap, plugins, diff --git a/web/app/components/tools/marketplace/index.tsx b/web/app/components/tools/marketplace/index.tsx index 19629089e66..dd83ab9e8a9 100644 --- a/web/app/components/tools/marketplace/index.tsx +++ b/web/app/components/tools/marketplace/index.tsx @@ -1,8 +1,8 @@ import type { SearchParamsFromCollection } from '@dify/contracts/marketplace' import type { ToolsContentInset } from '../content-inset' import type { useMarketplace } from './hooks' +import { Button } from '@langgenius/dify-ui/button' import { cn } from '@langgenius/dify-ui/cn' -import { RiArrowRightUpLine, RiArrowUpDoubleLine } from '@remixicon/react' import { useTheme } from 'next-themes' import { useTranslation } from 'react-i18next' import { useLocale } from '#i18n' @@ -53,10 +53,15 @@ const Marketplace = ({ <>
{isMarketplaceArrowVisible && ( - $['marketplace.moreFrom'], { ns: 'plugin' })} + className="absolute top-2 left-1/2 z-10 size-6 -translate-x-1/2 p-0 text-text-quaternary" onClick={showMarketplacePanel} - /> + size="small" + variant="ghost" + > +