diff --git a/api/core/plugin/plugin_service.py b/api/core/plugin/plugin_service.py index 6b306e2df86..89274b635ac 100644 --- a/api/core/plugin/plugin_service.py +++ b/api/core/plugin/plugin_service.py @@ -92,6 +92,7 @@ class PluginService: PLUGIN_MODEL_PROVIDERS_REDIS_KEY_PREFIX = "plugin_model_providers:tenant_id:" PLUGIN_MODEL_PROVIDERS_GENERATION_REDIS_KEY_PREFIX = "plugin_model_providers_generation:tenant_id:" PLUGIN_MODEL_PROVIDERS_LOCK_REDIS_KEY_PREFIX = "plugin_model_providers_refresh_lock:tenant_id:" + PLUGIN_MODEL_PROVIDERS_REMOTE_DEBUG_REDIS_KEY_PREFIX = "plugin_model_providers_remote_debug:tenant_id:" PLUGIN_MODEL_PROVIDERS_LOCK_TTL = 30 PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_TIMEOUT = 2.0 PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_INTERVAL = 0.05 @@ -117,6 +118,10 @@ class PluginService: def _get_plugin_model_providers_lock_key(cls, tenant_id: str, generation: int) -> str: return f"{cls.PLUGIN_MODEL_PROVIDERS_LOCK_REDIS_KEY_PREFIX}{tenant_id}:generation:{generation}" + @classmethod + def _get_plugin_model_providers_remote_debug_cache_key(cls, tenant_id: str) -> str: + return f"{cls.PLUGIN_MODEL_PROVIDERS_REMOTE_DEBUG_REDIS_KEY_PREFIX}{tenant_id}" + @staticmethod def _get_provider_short_name_alias(provider: PluginModelProviderEntity) -> str: """ @@ -259,6 +264,111 @@ class PluginService: except (RedisError, RuntimeError): 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: + remote_model_plugins = sorted( + f"{plugin.plugin_id}:{plugin.plugin_unique_identifier}" + for plugin in plugins + if plugin.source == PluginInstallationSource.Remote + ) + if not remote_model_plugins: + return None + + return "\n".join(remote_model_plugins) + + @classmethod + def _load_cached_remote_model_plugin_marker(cls, tenant_id: str) -> str | None: + cache_key = cls._get_plugin_model_providers_remote_debug_cache_key(tenant_id) + try: + cached_marker = redis_client.get(cache_key) + except (RedisError, RuntimeError): + logger.warning("Failed to read remote debug model plugin marker for tenant %s.", tenant_id, exc_info=True) + return None + + if cached_marker is None: + return None + if isinstance(cached_marker, bytes): + try: + return cached_marker.decode() + except UnicodeDecodeError: + logger.warning( + "Invalid remote debug model plugin marker for tenant %s; deleting cache marker.", + tenant_id, + exc_info=True, + ) + try: + redis_client.delete(cache_key) + except (RedisError, RuntimeError): + logger.warning( + "Failed to delete invalid remote debug model plugin marker for tenant %s.", + tenant_id, + exc_info=True, + ) + return None + if isinstance(cached_marker, str): + return cached_marker + + logger.warning("Invalid remote debug model plugin marker for tenant %s; deleting cache marker.", tenant_id) + try: + redis_client.delete(cache_key) + except (RedisError, RuntimeError): + logger.warning( + "Failed to delete invalid remote debug model plugin marker for tenant %s.", + tenant_id, + exc_info=True, + ) + return None + + @classmethod + def _store_cached_remote_model_plugin_marker(cls, tenant_id: str, marker: str | None) -> None: + cache_key = cls._get_plugin_model_providers_remote_debug_cache_key(tenant_id) + try: + if marker is None: + redis_client.delete(cache_key) + else: + redis_client.setex(cache_key, dify_config.PLUGIN_MODEL_PROVIDERS_CACHE_TTL, marker) + except (RedisError, RuntimeError): + logger.warning("Failed to cache remote debug model plugin marker for tenant %s.", tenant_id, exc_info=True) + + @classmethod + def _load_cached_plugin_model_provider_plugin_ids(cls, tenant_id: str) -> set[str] | None: + """Return plugin ids represented by the current provider cache, or None when no usable cache exists.""" + generation = cls._load_plugin_model_providers_generation(tenant_id) + cached_providers, _ = cls._load_cached_plugin_model_providers_for_generation(tenant_id, generation) + if cached_providers is None: + return None + + plugin_ids: set[str] = set() + for provider in cached_providers: + last_slash = provider.provider.rfind("/") + if last_slash > 0: + plugin_ids.add(provider.provider[:last_slash]) + + return plugin_ids + + @classmethod + def _should_invalidate_model_provider_cache_for_remote_model_plugins( + cls, + tenant_id: str, + plugins: Sequence[PluginEntity], + ) -> 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) + if remote_model_plugin_marker is None: + return cached_remote_model_plugin_marker is not None + + if remote_model_plugin_marker != cached_remote_model_plugin_marker: + return True + + remote_model_plugin_ids = { + plugin.plugin_id for plugin in plugins if plugin.source == PluginInstallationSource.Remote + } + cached_plugin_ids = cls._load_cached_plugin_model_provider_plugin_ids(tenant_id) + if cached_plugin_ids is None: + return False + + return not remote_model_plugin_ids.issubset(cached_plugin_ids) + @classmethod @contextmanager def _plugin_model_providers_refresh_lock( @@ -571,7 +681,21 @@ class PluginService: This keeps pagination usable before category is persisted on installation rows. """ manager = PluginInstaller() - return manager.list_plugins_by_category(tenant_id, category, page, page_size) + plugins = manager.list_plugins_by_category(tenant_id, category, page, page_size) + if category == PluginCategory.Model: + should_invalidate_model_provider_cache = ( + PluginService._should_invalidate_model_provider_cache_for_remote_model_plugins( + tenant_id, + plugins.list, + ) + ) + if should_invalidate_model_provider_cache: + PluginService.invalidate_plugin_model_providers_cache(tenant_id) + + remote_model_plugin_marker = PluginService._get_remote_model_plugin_cache_marker(plugins.list) + PluginService._store_cached_remote_model_plugin_marker(tenant_id, remote_model_plugin_marker) + + return plugins @staticmethod def _normalize_endpoint_count(value: object) -> int: 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 a8922154a95..278898926b9 100644 --- a/api/tests/unit_tests/services/plugin/test_plugin_service.py +++ b/api/tests/unit_tests/services/plugin/test_plugin_service.py @@ -8,7 +8,7 @@ import zstandard from pydantic import TypeAdapter from redis import RedisError -from core.plugin.entities.plugin import PluginInstallationSource +from core.plugin.entities.plugin import PluginCategory, PluginInstallationSource from core.plugin.entities.plugin_daemon import PluginInstallTask, PluginInstallTaskStatus, PluginModelProviderEntity from graphon.model_runtime.entities.common_entities import I18nObject from graphon.model_runtime.entities.provider_entities import ConfigurateMethod, ProviderEntity @@ -71,6 +71,16 @@ def _build_install_task(*, task_id: str = "task-1", status: PluginInstallTaskSta ) +def _build_remote_model_plugin( + *, plugin_id: str = "langgenius/debug-model", plugin_unique_identifier: str = "langgenius/debug-model:1.0.0" +) -> SimpleNamespace: + return SimpleNamespace( + plugin_id=plugin_id, + plugin_unique_identifier=plugin_unique_identifier, + source=PluginInstallationSource.Remote, + ) + + def _provider_cache_key(tenant_id: str, generation: int | None = None) -> str: if generation is None: return f"plugin_model_providers:tenant_id:{tenant_id}" @@ -797,6 +807,144 @@ class TestPluginListEndpointCounts: class TestPluginModelProviderCacheInvalidation: + 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 ( + patch(f"{MODULE}.PluginDebuggingClient") as debugging_client_cls, + patch(f"{MODULE}.PluginService.invalidate_plugin_model_providers_cache") as invalidate_cache, + ): + debugging_client_cls.return_value.get_debugging_key.return_value = "debug-key" + + from core.plugin.plugin_service import PluginService + + result = PluginService.get_debugging_key("tenant-1") + + assert result == "debug-key" + debugging_client_cls.return_value.get_debugging_key.assert_called_once_with("tenant-1") + invalidate_cache.assert_not_called() + + def test_list_model_category_invalidates_when_remote_model_plugin_is_missing_from_provider_cache(self) -> None: + """Remote model plugins are daemon-registered, so category reads repair a stale provider cache.""" + remote_plugin = _build_remote_model_plugin() + remote_plugin_marker = "langgenius/debug-model:langgenius/debug-model:1.0.0" + plugins = SimpleNamespace(list=[remote_plugin], has_more=False) + + with ( + patch(f"{MODULE}.PluginInstaller") as installer_cls, + patch( + f"{MODULE}.PluginService._load_cached_remote_model_plugin_marker", + return_value=remote_plugin_marker, + ), + patch( + f"{MODULE}.PluginService._load_cached_plugin_model_provider_plugin_ids", + return_value={"langgenius/openai"}, + ), + 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) + + assert result is plugins + installer_cls.return_value.list_plugins_by_category.assert_called_once_with( + "tenant-1", PluginCategory.Model, 1, 100 + ) + invalidate_cache.assert_called_once_with("tenant-1") + store_marker.assert_called_once_with("tenant-1", remote_plugin_marker) + + def test_list_model_category_invalidates_when_remote_model_plugin_identity_changes(self) -> None: + """A debug model plugin can share plugin_id with an installed plugin, so identity changes bust cache too.""" + remote_plugin = _build_remote_model_plugin( + plugin_id="langgenius/openai", + plugin_unique_identifier="langgenius/openai:debug", + ) + remote_plugin_marker = "langgenius/openai:langgenius/openai:debug" + plugins = SimpleNamespace(list=[remote_plugin], has_more=False) + + with ( + patch(f"{MODULE}.PluginInstaller") as installer_cls, + patch( + f"{MODULE}.PluginService._load_cached_remote_model_plugin_marker", + return_value="langgenius/openai:langgenius/openai:1.0.0", + ), + patch( + f"{MODULE}.PluginService._load_cached_plugin_model_provider_plugin_ids", + return_value={"langgenius/openai"}, + ) as load_cached_provider_plugin_ids, + 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) + + assert result is plugins + invalidate_cache.assert_called_once_with("tenant-1") + load_cached_provider_plugin_ids.assert_not_called() + store_marker.assert_called_once_with("tenant-1", remote_plugin_marker) + + def test_list_model_category_keeps_provider_cache_when_remote_model_plugin_is_already_cached(self) -> None: + """A connected remote model plugin should not force provider cache churn once represented.""" + remote_plugin = _build_remote_model_plugin() + remote_plugin_marker = "langgenius/debug-model:langgenius/debug-model:1.0.0" + plugins = SimpleNamespace(list=[remote_plugin], has_more=False) + + with ( + patch(f"{MODULE}.PluginInstaller") as installer_cls, + patch( + f"{MODULE}.PluginService._load_cached_remote_model_plugin_marker", + return_value=remote_plugin_marker, + ), + patch( + f"{MODULE}.PluginService._load_cached_plugin_model_provider_plugin_ids", + return_value={"langgenius/debug-model"}, + ), + 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) + + assert result is plugins + 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.""" + 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) + + 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, 1, 100) + + assert result is plugins + invalidate_cache.assert_called_once_with("tenant-1") + store_marker.assert_called_once_with("tenant-1", None) + def test_fetch_install_task_invalidates_model_provider_cache_when_finished(self) -> None: """Finished plugin install tasks invalidate tenant provider cache.""" task = _build_install_task(status=PluginInstallTaskStatus.Success) diff --git a/web/app/components/header/account-setting/model-provider-page/__tests__/index.non-cloud.spec.tsx b/web/app/components/header/account-setting/model-provider-page/__tests__/index.non-cloud.spec.tsx index f7bf0d1ce05..8896c4575c8 100644 --- a/web/app/components/header/account-setting/model-provider-page/__tests__/index.non-cloud.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/__tests__/index.non-cloud.spec.tsx @@ -41,6 +41,7 @@ vi.mock('@/context/provider-context', () => ({ vi.mock('../hooks', () => ({ useDefaultModel: () => ({ data: null, isLoading: false }), + useLanguage: () => 'en_US', })) vi.mock('../provider-added-card', () => ({ @@ -84,9 +85,11 @@ vi.mock('@/app/components/plugins/plugin-page/use-reference-setting', () => ({ })) vi.mock('@/service/use-plugins', () => ({ - useCheckInstalled: () => ({ + useInstalledPluginList: () => ({ data: { plugins: [] }, }), + useInvalidateInstalledPluginList: () => vi.fn(), + useInvalidateCheckInstalled: () => vi.fn(), usePluginAutoUpgradeSettings: () => ({ data: { category: 'model', 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 426276cfc5f..a648f4288af 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 @@ -1,7 +1,9 @@ import type { ReactNode } from 'react' +import type { PluginDeclaration, PluginDetail } from '@/app/components/plugins/types' import { act, fireEvent, screen } from '@testing-library/react' import { describe, expect, it, vi } from 'vitest' import { renderWithSystemFeatures } from '@/__tests__/utils/mock-system-features' +import { PluginCategoryEnum, PluginSource } from '@/app/components/plugins/types' import { CurrentSystemQuotaTypeEnum, CustomConfigurationStatusEnum, @@ -41,10 +43,18 @@ const { mockReferenceSetting, mockAutoUpgradeError } = vi.hoisted(() => ({ }, })) -const { mockProviderContextState } = vi.hoisted(() => ({ +const { mockProviderContextState, mockRefreshModelProviders } = vi.hoisted(() => ({ mockProviderContextState: { isLoadingModelProviders: false, }, + mockRefreshModelProviders: vi.fn(), +})) + +const { mockInstalledModelPlugins, mockUseInstalledPluginList } = vi.hoisted(() => ({ + mockInstalledModelPlugins: { + value: [] as PluginDetail[], + }, + mockUseInstalledPluginList: vi.fn(), })) const mockQuotaConfig = { @@ -79,6 +89,70 @@ const saveUpdateSettings = () => { fireEvent.click(screen.getByRole('button', { name: 'common.operation.save' })) } +const createPluginDeclaration = (overrides: Partial = {}): PluginDeclaration => ({ + plugin_unique_identifier: 'langgenius/debug-model:1.0.0', + version: '1.0.0', + author: 'langgenius', + icon: 'debug-model.png', + icon_dark: 'debug-model-dark.png', + name: 'debug-model', + category: PluginCategoryEnum.model, + label: { en_US: 'Debug Model' } as unknown as PluginDeclaration['label'], + description: { en_US: 'Debug model provider' } as unknown as PluginDeclaration['description'], + created_at: '2024-01-01', + resource: null, + plugins: null, + verified: false, + endpoint: null, + tool: undefined, + datasource: undefined, + model: {}, + tags: [], + agent_strategy: null, + meta: { + version: '1.0.0', + }, + trigger: {} as unknown as PluginDeclaration['trigger'], + ...overrides, +}) + +const createPluginDetail = (overrides: Partial = {}): PluginDetail => { + const { + declaration: overrideDeclaration, + plugin_id: overridePluginId, + ...restOverrides + } = overrides + const declaration = overrideDeclaration ?? createPluginDeclaration() + const pluginId = overridePluginId ?? 'langgenius/debug-model' + + return { + id: 'plugin-installation-id', + created_at: '2024-01-01', + updated_at: '2024-01-01', + name: declaration.name, + plugin_id: pluginId, + plugin_unique_identifier: declaration.plugin_unique_identifier, + declaration, + installation_id: 'plugin-installation-id', + tenant_id: 'tenant-id', + endpoints_setups: 0, + endpoints_active: 0, + version: '1.0.0', + latest_version: '1.0.0', + latest_unique_identifier: declaration.plugin_unique_identifier, + source: PluginSource.debugging, + meta: { + repo: '', + version: '1.0.0', + package: '', + }, + status: 'active', + deprecated_reason: '', + alternative_plugin_id: '', + ...restOverrides, + } +} + const mockProviders = [ { provider: 'openai', @@ -106,6 +180,7 @@ vi.mock('@/context/provider-context', () => ({ useProviderContext: () => ({ modelProviders: mockProviders, isLoadingModelProviders: mockProviderContextState.isLoadingModelProviders, + refreshModelProviders: mockRefreshModelProviders, }), })) @@ -119,6 +194,7 @@ const mockDefaultModels: Record = vi.mock('../hooks', () => ({ useDefaultModel: (type: string) => mockDefaultModels[type] ?? { data: null, isLoading: false }, + useLanguage: () => 'en_US', })) vi.mock('../install-from-marketplace', () => ({ @@ -126,7 +202,24 @@ vi.mock('../install-from-marketplace', () => ({ })) vi.mock('../provider-added-card', () => ({ - default: ({ provider }: { provider: { provider: string } }) =>
{provider.provider}
, + default: ({ + notConfigured, + provider, + pluginDetail, + }: { + notConfigured?: boolean + provider: { provider: string } + pluginDetail?: { plugin_id: string, source?: string } + }) => ( +
+ {provider.provider} +
+ ), })) vi.mock('../provider-added-card/quota-panel', () => ({ @@ -160,12 +253,12 @@ vi.mock('@/app/components/plugins/plugin-page/use-reference-setting', () => ({ })) vi.mock('@/service/use-plugins', () => ({ - useInstalledPluginList: () => ({ - data: { plugins: [] }, - }), - useCheckInstalled: () => ({ - data: { plugins: [] }, - }), + useInstalledPluginList: (...args: unknown[]) => { + mockUseInstalledPluginList(...args) + return { + data: { plugins: mockInstalledModelPlugins.value }, + } + }, usePluginAutoUpgradeSettings: () => ({ data: mockReferenceSetting.auto_upgrade ? { @@ -280,6 +373,9 @@ describe('ModelProviderPage', () => { beforeEach(() => { vi.useFakeTimers() vi.clearAllMocks() + mockUseInstalledPluginList.mockClear() + mockRefreshModelProviders.mockClear() + mockInstalledModelPlugins.value = [] mockProviderContextState.isLoadingModelProviders = false mockAutoUpgradeError.value = undefined mockReferenceSetting.auto_upgrade = { @@ -418,6 +514,107 @@ describe('ModelProviderPage', () => { expect(screen.getByText('anthropic')).toBeInTheDocument() }) + it('should use the model plugin installation list to attach plugin detail to provider cards', () => { + mockProviders.splice(0, mockProviders.length, { + provider: 'langgenius/openai/openai', + label: { en_US: 'OpenAI' }, + custom_configuration: { status: CustomConfigurationStatusEnum.active }, + system_configuration: { + enabled: false, + current_quota_type: CurrentSystemQuotaTypeEnum.free, + 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'], + }), + }), + ] + + renderModelProviderPage() + + expect(mockUseInstalledPluginList).toHaveBeenCalledWith(false, 100, { category: PluginCategoryEnum.model }) + expect(screen.getByTestId('provider-card')).toHaveAttribute('data-plugin-id', 'langgenius/openai') + expect(screen.queryByText('OpenAI Plugin')).not.toBeInTheDocument() + }) + + it('should not render installed model plugins that are not registered as model providers', () => { + mockInstalledModelPlugins.value = [ + createPluginDetail({ + plugin_id: 'langgenius/debug-model', + declaration: createPluginDeclaration({ + label: { en_US: 'Debug Model' } as unknown as PluginDeclaration['label'], + description: { en_US: 'Debug model provider' } as unknown as PluginDeclaration['description'], + }), + }), + ] + + renderModelProviderPage() + + expect(screen.queryByText('Debug Model')).not.toBeInTheDocument() + expect(screen.queryByText('langgenius/debug-model')).not.toBeInTheDocument() + expect(screen.queryByRole('button', { name: 'plugin actions langgenius/debug-model' })).not.toBeInTheDocument() + }) + + it('should refresh model providers once when a debugging model plugin is missing from providers', () => { + mockInstalledModelPlugins.value = [ + createPluginDetail({ + plugin_id: 'langgenius/debug-model', + declaration: createPluginDeclaration({ + label: { en_US: 'Debug Model' } as unknown as PluginDeclaration['label'], + }), + }), + ] + + renderModelProviderPage() + + expect(mockRefreshModelProviders).toHaveBeenCalledTimes(1) + }) + + it('should prefer debugging plugin detail when an installed model plugin shares the same plugin id', () => { + mockProviders.splice(0, mockProviders.length, { + provider: 'langgenius/openai/openai', + label: { en_US: 'OpenAI' }, + custom_configuration: { status: CustomConfigurationStatusEnum.active }, + system_configuration: { + enabled: false, + current_quota_type: CurrentSystemQuotaTypeEnum.free, + quota_configurations: [mockQuotaConfig], + }, + }) + mockInstalledModelPlugins.value = [ + createPluginDetail({ + 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, + }), + ] + + renderModelProviderPage() + + expect(screen.getByTestId('provider-card')).toHaveAttribute('data-plugin-id', 'langgenius/openai') + expect(screen.getByTestId('provider-card')).toHaveAttribute('data-plugin-source', PluginSource.debugging) + expect(mockRefreshModelProviders).toHaveBeenCalledTimes(1) + }) + it('should show provider placeholders while model providers are loading', () => { mockProviderContextState.isLoadingModelProviders = true @@ -572,4 +769,66 @@ describe('ModelProviderPage', () => { ]) expect(screen.queryByText('common.modelProvider.toBeConfigured')).not.toBeInTheDocument() }) + + it('should prioritize debugging model plugins within their provider section', () => { + mockProviders.splice(0, mockProviders.length, { + provider: 'langgenius/openai/openai', + label: { en_US: 'OpenAI Fixed' }, + custom_configuration: { status: CustomConfigurationStatusEnum.active }, + system_configuration: { + enabled: false, + current_quota_type: CurrentSystemQuotaTypeEnum.free, + quota_configurations: [mockQuotaConfig], + }, + }, { + provider: 'zeta-provider', + label: { en_US: 'Zeta Provider' }, + custom_configuration: { status: CustomConfigurationStatusEnum.active }, + system_configuration: { + enabled: false, + current_quota_type: CurrentSystemQuotaTypeEnum.free, + quota_configurations: [mockQuotaConfig], + }, + }, { + provider: 'langgenius/normal-model/normal-model', + label: { en_US: 'Normal Model' }, + custom_configuration: { status: CustomConfigurationStatusEnum.noConfigure }, + system_configuration: { + enabled: false, + current_quota_type: CurrentSystemQuotaTypeEnum.free, + quota_configurations: [mockQuotaConfig], + }, + }, { + provider: 'langgenius/debug-model/debug-model', + label: { en_US: 'Debug Model' }, + custom_configuration: { status: CustomConfigurationStatusEnum.noConfigure }, + system_configuration: { + enabled: false, + current_quota_type: CurrentSystemQuotaTypeEnum.free, + quota_configurations: [mockQuotaConfig], + }, + }) + mockInstalledModelPlugins.value = [ + createPluginDetail({ + 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'], + }), + }), + ] + + renderModelProviderPage() + + const renderedProviders = screen.getAllByTestId('provider-card').map(item => item.textContent) + expect(renderedProviders).toEqual([ + 'langgenius/openai/openai', + 'zeta-provider', + 'langgenius/debug-model/debug-model', + 'langgenius/normal-model/normal-model', + ]) + expect(screen.getAllByTestId('provider-card')[2]).toHaveAttribute('data-not-configured', 'true') + expect(screen.getByText('common.modelProvider.toBeConfigured')).toBeInTheDocument() + }) }) 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 5d8ddfd8f13..9bc43082349 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 @@ -6,15 +6,15 @@ import type { PluginDetail } from '@/app/components/plugins/types' import { useSuspenseQuery } from '@tanstack/react-query' import { useDebounce } from 'ahooks' import { noop } from 'es-toolkit/function' -import { useMemo } from 'react' +import { useEffect, useMemo, useRef } 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 } from '@/app/components/plugins/types' +import { PluginCategoryEnum, PluginSource } from '@/app/components/plugins/types' import { useProviderContext } from '@/context/provider-context' import { systemFeaturesQueryOptions } from '@/features/system-features/client' -import { useCheckInstalled } from '@/service/use-plugins' +import { useInstalledPluginList } from '@/service/use-plugins' import UpdateSettingDialog from '../update-setting-dialog' import { CustomConfigurationStatusEnum, @@ -25,7 +25,6 @@ import { } from './hooks' import ModelProviderPageBody from './model-provider-page-body' import SystemModelSelector from './system-model-selector' -import { providerToPluginId } from './utils' type SystemModelConfigStatus = 'no-provider' | 'none-configured' | 'partially-configured' | 'fully-configured' @@ -58,23 +57,43 @@ const ModelProviderPage = ({ const { data: rerankDefaultModel, isLoading: isRerankDefaultModelLoading } = useDefaultModel(ModelTypeEnum.rerank) const { data: speech2textDefaultModel, isLoading: isSpeech2textDefaultModelLoading } = useDefaultModel(ModelTypeEnum.speech2text) const { data: ttsDefaultModel, isLoading: isTTSDefaultModelLoading } = useDefaultModel(ModelTypeEnum.tts) - const { modelProviders: providers, isLoadingModelProviders } = useProviderContext() + const { modelProviders: providers, isLoadingModelProviders, refreshModelProviders } = useProviderContext() const { data: systemFeatures } = useSuspenseQuery(systemFeaturesQueryOptions()) - const allPluginIds = useMemo(() => { - return [...new Set(providers.map(p => providerToPluginId(p.provider)).filter(Boolean))] - }, [providers]) - const { data: installedPlugins } = useCheckInstalled({ - pluginIds: allPluginIds, - enabled: allPluginIds.length > 0, + const { data: installedModelPlugins } = useInstalledPluginList(false, 100, { + category: PluginCategoryEnum.model, }) - const enrichedPlugins = usePluginsWithLatestVersion(installedPlugins?.plugins) + const enrichedPlugins = usePluginsWithLatestVersion(installedModelPlugins?.plugins) const pluginDetailMap = useMemo(() => { const map = new Map() - for (const plugin of enrichedPlugins) - map.set(plugin.plugin_id, plugin) + for (const plugin of enrichedPlugins) { + const existingPlugin = map.get(plugin.plugin_id) + if (!existingPlugin || plugin.source === PluginSource.debugging) + map.set(plugin.plugin_id, plugin) + } 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]) const enableMarketplace = systemFeatures.enable_marketplace const isDefaultModelLoading = isTextGenerationDefaultModelLoading || isEmbeddingsDefaultModelLoading diff --git a/web/app/components/header/account-setting/model-provider-page/model-provider-page-body.tsx b/web/app/components/header/account-setting/model-provider-page/model-provider-page-body.tsx index 625731d29a3..e614b2c917c 100644 --- a/web/app/components/header/account-setting/model-provider-page/model-provider-page-body.tsx +++ b/web/app/components/header/account-setting/model-provider-page/model-provider-page-body.tsx @@ -3,6 +3,7 @@ import type { ModelProvider } from './declarations' import type { PluginDetail } from '@/app/components/plugins/types' import { Trans, useTranslation } from 'react-i18next' import { SkeletonContainer, SkeletonRectangle, SkeletonRow } from '@/app/components/base/skeleton' +import { PluginSource } from '@/app/components/plugins/types' import { IS_CLOUD_EDITION } from '@/config' import InstallFromMarketplace from './install-from-marketplace' import ProviderAddedCard from './provider-added-card' @@ -99,21 +100,40 @@ type ProviderCardListProps = { notConfigured?: boolean } +function isDebuggingProvider(provider: ModelProvider, pluginDetailMap: Map) { + return pluginDetailMap.get(providerToPluginId(provider.provider))?.source === PluginSource.debugging +} + function ProviderCardList({ providers, pluginDetailMap, notConfigured, }: ProviderCardListProps) { + const sortedProviders = [...providers] + .sort((a, b) => { + const aIsDebuggingPlugin = isDebuggingProvider(a, pluginDetailMap) + const bIsDebuggingPlugin = isDebuggingProvider(b, pluginDetailMap) + + if (aIsDebuggingPlugin === bIsDebuggingPlugin) + return 0 + + return aIsDebuggingPlugin ? -1 : 1 + }) + return (
- {providers.map(provider => ( - - ))} + {sortedProviders.map((provider) => { + const pluginDetail = pluginDetailMap.get(providerToPluginId(provider.provider)) + + return ( + + ) + })}
) } @@ -157,8 +177,8 @@ const ModelProviderPageBody: FC = ({
{t('modelProvider.toBeConfigured', { ns: 'common' })}
diff --git a/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/provider-card-actions.spec.tsx b/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/provider-card-actions.spec.tsx index b62b56a8a6f..935292a599a 100644 --- a/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/provider-card-actions.spec.tsx +++ b/web/app/components/header/account-setting/model-provider-page/provider-added-card/__tests__/provider-card-actions.spec.tsx @@ -158,6 +158,15 @@ describe('ProviderCardActions', () => { expect(mockHandleUpdate).toHaveBeenCalledWith(true) }) + it('should show a compact debug badge after the version for debugging plugins', () => { + render() + + const version = screen.getByText('1.0.0') + const debugBadge = screen.getByText('appDebug.operation.debugConfig') + + expect(version.compareDocumentPosition(debugBadge) & Node.DOCUMENT_POSITION_FOLLOWING).toBeTruthy() + }) + it('should trigger the latest marketplace update when clicking the update button', () => { render() 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 b2511b18670..4538cd9d16e 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 @@ -30,6 +30,7 @@ const ProviderCardActions: FC = ({ detail, onUpdate }) => { 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, @@ -80,31 +81,41 @@ const ProviderCardActions: FC = ({ detail, onUpdate }) => { return ( <> {!!version && ( - + + {version} + {canUpdatePlugin && isFromMarketplace && } + + )} + hasRedCornerMark={hasNewVersion} + /> + )} + /> + {isDebuggingPlugin && ( - {version} - {canUpdatePlugin && isFromMarketplace && } - - )} - hasRedCornerMark={hasNewVersion} + text={t('operation.debugConfig', { ns: 'appDebug' })} /> )} - /> + )} {canUpdatePlugin && (hasNewVersion || isFromGitHub) && (