mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
refactor(api): migrate workspace model endpoints to BaseModel (#37963)
Co-authored-by: Asuka Minato <i@asukaminato.eu.org> Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> Co-authored-by: Byron Wang <byron@dify.ai>
This commit is contained in:
co-authored by
Asuka Minato
autofix-ci[bot]
Byron Wang
parent
c5cef80ea4
commit
3cd8d850fa
@@ -1,10 +1,14 @@
|
||||
from unittest.mock import MagicMock, patch
|
||||
from inspect import unwrap
|
||||
from types import SimpleNamespace
|
||||
from typing import cast
|
||||
from unittest.mock import ANY, patch
|
||||
|
||||
import pytest
|
||||
from flask import Flask
|
||||
from pydantic_core import ValidationError
|
||||
from werkzeug.exceptions import Forbidden
|
||||
|
||||
from configs import dify_config
|
||||
from controllers.console.workspace.model_providers import (
|
||||
ModelProviderCredentialApi,
|
||||
ModelProviderCredentialSwitchApi,
|
||||
@@ -14,30 +18,126 @@ from controllers.console.workspace.model_providers import (
|
||||
ModelProviderValidateApi,
|
||||
PreferredProviderTypeUpdateApi,
|
||||
)
|
||||
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
|
||||
from graphon.model_runtime.errors.validate import CredentialsValidateFailedError
|
||||
from models import Account
|
||||
from models.provider import ProviderType
|
||||
from services.entities.model_provider_entities import (
|
||||
CustomConfigurationResponse,
|
||||
CustomConfigurationStatus,
|
||||
ProviderResponse,
|
||||
SystemConfigurationResponse,
|
||||
)
|
||||
|
||||
VALID_UUID = "123e4567-e89b-12d3-a456-426614174000"
|
||||
INVALID_UUID = "123"
|
||||
|
||||
|
||||
from inspect import unwrap
|
||||
def make_account() -> Account:
|
||||
return cast(Account, SimpleNamespace(id="account-1", email="owner@example.com"))
|
||||
|
||||
|
||||
def make_provider_response() -> ProviderResponse:
|
||||
return ProviderResponse(
|
||||
tenant_id="tenant1",
|
||||
provider="openai",
|
||||
label=I18nObject(en_US="OpenAI", zh_Hans="OpenAI"),
|
||||
description=I18nObject(en_US="OpenAI models", zh_Hans="OpenAI models zh"),
|
||||
icon_small=I18nObject(en_US="icon.svg", zh_Hans="icon.svg"),
|
||||
icon_small_dark=I18nObject(en_US="icon-dark.svg", zh_Hans="icon-dark.svg"),
|
||||
background="#ffffff",
|
||||
supported_model_types=[ModelType.LLM, ModelType.TEXT_EMBEDDING],
|
||||
configurate_methods=[ConfigurateMethod.PREDEFINED_MODEL, ConfigurateMethod.CUSTOMIZABLE_MODEL],
|
||||
preferred_provider_type=ProviderType.CUSTOM,
|
||||
custom_configuration=CustomConfigurationResponse(
|
||||
status=CustomConfigurationStatus.ACTIVE,
|
||||
current_credential_id=VALID_UUID,
|
||||
current_credential_name="production",
|
||||
available_credentials=[],
|
||||
custom_models=[],
|
||||
can_added_models=[],
|
||||
),
|
||||
system_configuration=SystemConfigurationResponse(
|
||||
enabled=True,
|
||||
current_quota_type=None,
|
||||
quota_configurations=[],
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def expected_provider_payload() -> dict[str, object]:
|
||||
icon_url_prefix = f"{dify_config.CONSOLE_API_URL}/console/api/workspaces/tenant1/model-providers/openai"
|
||||
return {
|
||||
"tenant_id": "tenant1",
|
||||
"provider": "openai",
|
||||
"label": {"zh_Hans": "OpenAI", "en_US": "OpenAI"},
|
||||
"description": {"zh_Hans": "OpenAI models zh", "en_US": "OpenAI models"},
|
||||
"icon_small": {
|
||||
"zh_Hans": f"{icon_url_prefix}/icon_small/zh_Hans",
|
||||
"en_US": f"{icon_url_prefix}/icon_small/en_US",
|
||||
},
|
||||
"icon_small_dark": {
|
||||
"zh_Hans": f"{icon_url_prefix}/icon_small_dark/zh_Hans",
|
||||
"en_US": f"{icon_url_prefix}/icon_small_dark/en_US",
|
||||
},
|
||||
"background": "#ffffff",
|
||||
"help": None,
|
||||
"supported_model_types": ["llm", "text-embedding"],
|
||||
"configurate_methods": ["predefined-model", "customizable-model"],
|
||||
"provider_credential_schema": None,
|
||||
"model_credential_schema": None,
|
||||
"preferred_provider_type": "custom",
|
||||
"custom_configuration": {
|
||||
"status": "active",
|
||||
"current_credential_id": VALID_UUID,
|
||||
"current_credential_name": "production",
|
||||
"available_credentials": [],
|
||||
"custom_models": [],
|
||||
"can_added_models": [],
|
||||
},
|
||||
"system_configuration": {
|
||||
"enabled": True,
|
||||
"current_quota_type": None,
|
||||
"quota_configurations": [],
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
class TestModelProviderListApi:
|
||||
def test_get_success(self, app: Flask):
|
||||
api = ModelProviderListApi()
|
||||
method = unwrap(api.get)
|
||||
provider = make_provider_response()
|
||||
|
||||
with (
|
||||
app.test_request_context("/?model_type=llm"),
|
||||
patch(
|
||||
"controllers.console.workspace.model_providers.ModelProviderService.get_provider_list",
|
||||
return_value=[{"name": "openai"}],
|
||||
),
|
||||
return_value=[provider],
|
||||
) as get_provider_list,
|
||||
):
|
||||
result = method(api, "tenant1")
|
||||
|
||||
assert "data" in result
|
||||
get_provider_list.assert_called_once_with(tenant_id="tenant1", model_type=ModelType.LLM)
|
||||
assert result == {"data": [expected_provider_payload()]}
|
||||
|
||||
def test_get_without_model_type_passes_none(self, app: Flask):
|
||||
api = ModelProviderListApi()
|
||||
method = unwrap(api.get)
|
||||
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch(
|
||||
"controllers.console.workspace.model_providers.ModelProviderService.get_provider_list",
|
||||
return_value=[],
|
||||
) as get_provider_list,
|
||||
):
|
||||
result = method(api, "tenant1")
|
||||
|
||||
get_provider_list.assert_called_once_with(tenant_id="tenant1", model_type=None)
|
||||
assert result == {"data": []}
|
||||
|
||||
|
||||
class TestModelProviderCredentialApi:
|
||||
@@ -49,12 +149,41 @@ class TestModelProviderCredentialApi:
|
||||
app.test_request_context(f"/?credential_id={VALID_UUID}"),
|
||||
patch(
|
||||
"controllers.console.workspace.model_providers.ModelProviderService.get_provider_credential",
|
||||
return_value={"key": "value"},
|
||||
),
|
||||
return_value={
|
||||
"api_key": "sk-test",
|
||||
"endpoint": "https://api.example.com",
|
||||
"nested": {"region": "us-east-1"},
|
||||
},
|
||||
) as get_provider_credential,
|
||||
):
|
||||
result = method(api, "tenant1", provider="openai")
|
||||
|
||||
assert "credentials" in result
|
||||
get_provider_credential.assert_called_once_with(
|
||||
tenant_id="tenant1", provider="openai", credential_id=VALID_UUID
|
||||
)
|
||||
assert result == {
|
||||
"credentials": {
|
||||
"api_key": "sk-test",
|
||||
"endpoint": "https://api.example.com",
|
||||
"nested": {"region": "us-east-1"},
|
||||
}
|
||||
}
|
||||
|
||||
def test_get_current_credential_without_id(self, app: Flask):
|
||||
api = ModelProviderCredentialApi()
|
||||
method = unwrap(api.get)
|
||||
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch(
|
||||
"controllers.console.workspace.model_providers.ModelProviderService.get_provider_credential",
|
||||
return_value=None,
|
||||
) as get_provider_credential,
|
||||
):
|
||||
result = method(api, "tenant1", provider="openai")
|
||||
|
||||
get_provider_credential.assert_called_once_with(tenant_id="tenant1", provider="openai", credential_id=None)
|
||||
assert result == {"credentials": None}
|
||||
|
||||
def test_get_invalid_uuid(self, app: Flask):
|
||||
api = ModelProviderCredentialApi()
|
||||
@@ -75,11 +204,17 @@ class TestModelProviderCredentialApi:
|
||||
patch(
|
||||
"controllers.console.workspace.model_providers.ModelProviderService.create_provider_credential",
|
||||
return_value=None,
|
||||
),
|
||||
) as create_provider_credential,
|
||||
):
|
||||
result, status = method(api, "tenant1", provider="openai")
|
||||
|
||||
assert result["result"] == "success"
|
||||
create_provider_credential.assert_called_once_with(
|
||||
tenant_id="tenant1",
|
||||
provider="openai",
|
||||
credentials={"a": "b"},
|
||||
credential_name="test",
|
||||
)
|
||||
assert result == {"result": "success"}
|
||||
assert status == 201
|
||||
|
||||
def test_post_create_validation_error(self, app: Flask):
|
||||
@@ -109,11 +244,18 @@ class TestModelProviderCredentialApi:
|
||||
patch(
|
||||
"controllers.console.workspace.model_providers.ModelProviderService.update_provider_credential",
|
||||
return_value=None,
|
||||
),
|
||||
) as update_provider_credential,
|
||||
):
|
||||
result = method(api, "tenant1", provider="openai")
|
||||
|
||||
assert result["result"] == "success"
|
||||
update_provider_credential.assert_called_once_with(
|
||||
tenant_id="tenant1",
|
||||
provider="openai",
|
||||
credentials={"a": "b"},
|
||||
credential_id=VALID_UUID,
|
||||
credential_name=None,
|
||||
)
|
||||
assert result == {"result": "success"}
|
||||
|
||||
def test_put_invalid_uuid(self, app: Flask):
|
||||
api = ModelProviderCredentialApi()
|
||||
@@ -136,10 +278,13 @@ class TestModelProviderCredentialApi:
|
||||
patch(
|
||||
"controllers.console.workspace.model_providers.ModelProviderService.remove_provider_credential",
|
||||
return_value=None,
|
||||
),
|
||||
) as remove_provider_credential,
|
||||
):
|
||||
result, status = method(api, "tenant1", provider="openai")
|
||||
|
||||
remove_provider_credential.assert_called_once_with(
|
||||
tenant_id="tenant1", provider="openai", credential_id=VALID_UUID
|
||||
)
|
||||
assert status == 204
|
||||
assert result == ""
|
||||
|
||||
@@ -156,11 +301,16 @@ class TestModelProviderCredentialSwitchApi:
|
||||
patch(
|
||||
"controllers.console.workspace.model_providers.ModelProviderService.switch_active_provider_credential",
|
||||
return_value=None,
|
||||
),
|
||||
) as switch_active_provider_credential,
|
||||
):
|
||||
result = method(api, "tenant1", provider="openai")
|
||||
|
||||
assert result["result"] == "success"
|
||||
switch_active_provider_credential.assert_called_once_with(
|
||||
tenant_id="tenant1",
|
||||
provider="openai",
|
||||
credential_id=VALID_UUID,
|
||||
)
|
||||
assert result == {"result": "success"}
|
||||
|
||||
def test_switch_invalid_uuid(self, app: Flask):
|
||||
api = ModelProviderCredentialSwitchApi()
|
||||
@@ -185,11 +335,14 @@ class TestModelProviderValidateApi:
|
||||
patch(
|
||||
"controllers.console.workspace.model_providers.ModelProviderService.validate_provider_credentials",
|
||||
return_value=None,
|
||||
),
|
||||
) as validate_provider_credentials,
|
||||
):
|
||||
result = method(api, "tenant1", provider="openai")
|
||||
|
||||
assert result["result"] == "success"
|
||||
validate_provider_credentials.assert_called_once_with(
|
||||
tenant_id="tenant1", provider="openai", credentials={"a": "b"}
|
||||
)
|
||||
assert result == {"result": "success", "error": None}
|
||||
|
||||
def test_validate_failure(self, app: Flask):
|
||||
api = ModelProviderValidateApi()
|
||||
@@ -206,7 +359,7 @@ class TestModelProviderValidateApi:
|
||||
):
|
||||
result = method(api, "tenant1", provider="openai")
|
||||
|
||||
assert result["result"] == "error"
|
||||
assert result == {"result": "error", "error": "bad"}
|
||||
|
||||
|
||||
class TestModelProviderIconApi:
|
||||
@@ -218,11 +371,14 @@ class TestModelProviderIconApi:
|
||||
patch(
|
||||
"controllers.console.workspace.model_providers.ModelProviderService.get_model_provider_icon",
|
||||
return_value=(b"123", "image/png"),
|
||||
),
|
||||
) as get_model_provider_icon,
|
||||
):
|
||||
response = api.get("t1", "openai", "logo", "en")
|
||||
|
||||
get_model_provider_icon.assert_called_once_with(tenant_id="t1", provider="openai", icon_type="logo", lang="en")
|
||||
assert response.mimetype == "image/png"
|
||||
response.direct_passthrough = False
|
||||
assert response.get_data() == b"123"
|
||||
|
||||
def test_icon_not_found(self, app: Flask):
|
||||
api = ModelProviderIconApi()
|
||||
@@ -250,11 +406,14 @@ class TestPreferredProviderTypeUpdateApi:
|
||||
patch(
|
||||
"controllers.console.workspace.model_providers.ModelProviderService.switch_preferred_provider",
|
||||
return_value=None,
|
||||
),
|
||||
) as switch_preferred_provider,
|
||||
):
|
||||
result = method(api, "tenant1", provider="openai")
|
||||
|
||||
assert result["result"] == "success"
|
||||
switch_preferred_provider.assert_called_once_with(
|
||||
tenant_id="tenant1", provider="openai", preferred_provider_type="custom"
|
||||
)
|
||||
assert result == {"result": "success"}
|
||||
|
||||
def test_invalid_enum(self, app: Flask):
|
||||
api = PreferredProviderTypeUpdateApi()
|
||||
@@ -272,22 +431,29 @@ class TestModelProviderPaymentCheckoutUrlApi:
|
||||
api = ModelProviderPaymentCheckoutUrlApi()
|
||||
method = unwrap(api.get)
|
||||
|
||||
user = MagicMock(id="u1", email="x@test.com")
|
||||
user = make_account()
|
||||
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch(
|
||||
"controllers.console.workspace.model_providers.BillingService.is_tenant_owner_or_admin",
|
||||
return_value=None,
|
||||
),
|
||||
) as is_tenant_owner_or_admin,
|
||||
patch(
|
||||
"controllers.console.workspace.model_providers.BillingService.get_model_provider_payment_link",
|
||||
return_value={"url": "x"},
|
||||
),
|
||||
return_value={"payment_link": "https://payment.example.com/provider"},
|
||||
) as get_model_provider_payment_link,
|
||||
):
|
||||
result = method(api, "tenant1", user, provider="anthropic")
|
||||
|
||||
assert "url" in result
|
||||
is_tenant_owner_or_admin.assert_called_once_with(user, session=ANY)
|
||||
get_model_provider_payment_link.assert_called_once_with(
|
||||
provider_name="anthropic",
|
||||
tenant_id="tenant1",
|
||||
account_id="account-1",
|
||||
prefilled_email="owner@example.com",
|
||||
)
|
||||
assert result == {"payment_link": "https://payment.example.com/provider"}
|
||||
|
||||
def test_invalid_provider(self, app: Flask):
|
||||
api = ModelProviderPaymentCheckoutUrlApi()
|
||||
@@ -295,13 +461,13 @@ class TestModelProviderPaymentCheckoutUrlApi:
|
||||
|
||||
with app.test_request_context("/"):
|
||||
with pytest.raises(ValueError):
|
||||
method(api, "tenant1", MagicMock(), provider="openai")
|
||||
method(api, "tenant1", make_account(), provider="openai")
|
||||
|
||||
def test_permission_denied(self, app: Flask):
|
||||
api = ModelProviderPaymentCheckoutUrlApi()
|
||||
method = unwrap(api.get)
|
||||
|
||||
user = MagicMock(id="u1", email="x@test.com")
|
||||
user = make_account()
|
||||
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
|
||||
@@ -32,7 +32,16 @@ class TestDefaultModelApi:
|
||||
),
|
||||
patch("controllers.console.workspace.models.ModelProviderService") as service_mock,
|
||||
):
|
||||
service_mock.return_value.get_default_model_of_model_type.return_value = {"model": "gpt-4"}
|
||||
service_mock.return_value.get_default_model_of_model_type.return_value = {
|
||||
"model": "gpt-4",
|
||||
"model_type": ModelType.LLM,
|
||||
"provider": {
|
||||
"tenant_id": "tenant1",
|
||||
"provider": "openai",
|
||||
"label": {"en_US": "OpenAI", "zh_Hans": "OpenAI"},
|
||||
"supported_model_types": [ModelType.LLM],
|
||||
},
|
||||
}
|
||||
|
||||
result = method(api, "tenant1")
|
||||
|
||||
|
||||
@@ -42,8 +42,11 @@ from controllers.console.workspace.plugin import (
|
||||
PluginUploadFromGithubApi,
|
||||
PluginUploadFromPkgApi,
|
||||
)
|
||||
from core.plugin.entities.plugin import PluginInstallation
|
||||
from core.plugin.entities.parameters import PluginParameterOption
|
||||
from core.plugin.entities.plugin import 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 models.account import (
|
||||
Account,
|
||||
TenantAccountRole,
|
||||
@@ -118,6 +121,215 @@ def _builtin_tool_provider_item() -> dict[str, Any]:
|
||||
}
|
||||
|
||||
|
||||
def _plugin_declaration_payload() -> dict[str, Any]:
|
||||
return {
|
||||
"version": "1.2.3",
|
||||
"author": "langgenius",
|
||||
"name": "demo_plugin",
|
||||
"description": {"en_US": "Demo plugin"},
|
||||
"icon": "icon.svg",
|
||||
"icon_dark": None,
|
||||
"label": {"en_US": "Demo Plugin"},
|
||||
"created_at": "2024-01-02T03:04:05",
|
||||
"resource": {"memory": 268435456, "permission": None},
|
||||
"plugins": {"tools": ["provider/demo.yaml"]},
|
||||
"tags": ["search", "demo"],
|
||||
"repo": "https://github.com/langgenius/demo",
|
||||
"verified": True,
|
||||
"meta": {"minimum_dify_version": "0.15.0", "version": "1.2.3"},
|
||||
}
|
||||
|
||||
|
||||
def _expected_i18n(en_us: str) -> dict[str, str]:
|
||||
return {"en_US": en_us, "zh_Hans": en_us, "pt_BR": en_us, "ja_JP": en_us}
|
||||
|
||||
|
||||
def _expected_plugin_declaration_dump() -> dict[str, Any]:
|
||||
return {
|
||||
"version": "1.2.3",
|
||||
"author": "langgenius",
|
||||
"name": "demo_plugin",
|
||||
"description": _expected_i18n("Demo plugin"),
|
||||
"icon": "icon.svg",
|
||||
"icon_dark": None,
|
||||
"label": _expected_i18n("Demo Plugin"),
|
||||
"category": "extension",
|
||||
"created_at": "2024-01-02T03:04:05",
|
||||
"resource": {"memory": 268435456, "permission": None},
|
||||
"plugins": {
|
||||
"tools": ["provider/demo.yaml"],
|
||||
"models": [],
|
||||
"endpoints": [],
|
||||
"datasources": [],
|
||||
"triggers": [],
|
||||
},
|
||||
"tags": ["search", "demo"],
|
||||
"repo": "https://github.com/langgenius/demo",
|
||||
"verified": True,
|
||||
"tool": None,
|
||||
"model": None,
|
||||
"endpoint": None,
|
||||
"agent_strategy": None,
|
||||
"datasource": None,
|
||||
"trigger": None,
|
||||
"meta": {"minimum_dify_version": "0.15.0", "version": "1.2.3"},
|
||||
}
|
||||
|
||||
|
||||
def _plugin_installation_payload() -> dict[str, Any]:
|
||||
return {
|
||||
"id": "installation-row-1",
|
||||
"created_at": "2024-01-02T03:04:05",
|
||||
"updated_at": "2024-01-03T04:05:06",
|
||||
"tenant_id": "tenant-1",
|
||||
"endpoints_setups": 2,
|
||||
"endpoints_active": 1,
|
||||
"runtime_type": "remote",
|
||||
"source": "marketplace",
|
||||
"meta": {"from": "marketplace"},
|
||||
"plugin_id": "langgenius/demo_plugin",
|
||||
"plugin_unique_identifier": "langgenius/demo_plugin:1.2.3@sha256:abc",
|
||||
"version": "1.2.3",
|
||||
"checksum": "sha256:abc",
|
||||
"declaration": _plugin_declaration_payload(),
|
||||
}
|
||||
|
||||
|
||||
def _plugin_entity_payload() -> dict[str, Any]:
|
||||
return {
|
||||
**_plugin_installation_payload(),
|
||||
"name": "demo_plugin",
|
||||
"installation_id": "installation-row-1",
|
||||
}
|
||||
|
||||
|
||||
def _plugin_declaration() -> PluginDeclaration:
|
||||
return PluginDeclaration.model_validate(_plugin_declaration_payload())
|
||||
|
||||
|
||||
def _plugin_installation() -> PluginInstallation:
|
||||
return PluginInstallation.model_validate(_plugin_installation_payload())
|
||||
|
||||
|
||||
def _plugin_entity() -> PluginEntity:
|
||||
return PluginEntity.model_validate(_plugin_entity_payload())
|
||||
|
||||
|
||||
def _expected_plugin_installation_dump() -> dict[str, Any]:
|
||||
return {
|
||||
"id": "installation-row-1",
|
||||
"created_at": "2024-01-02T03:04:05",
|
||||
"updated_at": "2024-01-03T04:05:06",
|
||||
"tenant_id": "tenant-1",
|
||||
"endpoints_setups": 2,
|
||||
"endpoints_active": 1,
|
||||
"runtime_type": "remote",
|
||||
"source": "marketplace",
|
||||
"meta": {"from": "marketplace"},
|
||||
"plugin_id": "langgenius/demo_plugin",
|
||||
"plugin_unique_identifier": "langgenius/demo_plugin:1.2.3@sha256:abc",
|
||||
"version": "1.2.3",
|
||||
"checksum": "sha256:abc",
|
||||
"declaration": _expected_plugin_declaration_dump(),
|
||||
}
|
||||
|
||||
|
||||
def _expected_plugin_entity_dump() -> dict[str, Any]:
|
||||
return {
|
||||
**_expected_plugin_installation_dump(),
|
||||
"name": "demo_plugin",
|
||||
"installation_id": "installation-row-1",
|
||||
}
|
||||
|
||||
|
||||
def _plugin_task_payload() -> dict[str, Any]:
|
||||
return {
|
||||
"id": "task-1",
|
||||
"created_at": "2024-02-03T04:05:06",
|
||||
"updated_at": "2024-02-03T04:06:07",
|
||||
"status": "running",
|
||||
"total_plugins": 2,
|
||||
"completed_plugins": 1,
|
||||
"plugins": [
|
||||
{
|
||||
"plugin_unique_identifier": "langgenius/demo_plugin:1.2.3@sha256:abc",
|
||||
"plugin_id": "langgenius/demo_plugin",
|
||||
"status": "success",
|
||||
"message": "installed",
|
||||
"icon": "icon.svg",
|
||||
"labels": {"en_US": "Demo Plugin"},
|
||||
"source": "marketplace",
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def _plugin_task() -> PluginInstallTask:
|
||||
return PluginInstallTask.model_validate(_plugin_task_payload())
|
||||
|
||||
|
||||
def _expected_plugin_task_dump() -> dict[str, Any]:
|
||||
return {
|
||||
"id": "task-1",
|
||||
"created_at": "2024-02-03T04:05:06",
|
||||
"updated_at": "2024-02-03T04:06:07",
|
||||
"status": "running",
|
||||
"total_plugins": 2,
|
||||
"completed_plugins": 1,
|
||||
"plugins": [
|
||||
{
|
||||
"plugin_unique_identifier": "langgenius/demo_plugin:1.2.3@sha256:abc",
|
||||
"plugin_id": "langgenius/demo_plugin",
|
||||
"status": "success",
|
||||
"message": "installed",
|
||||
"icon": "icon.svg",
|
||||
"labels": _expected_i18n("Demo Plugin"),
|
||||
"source": "marketplace",
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def _latest_plugin_cache() -> PluginService.LatestPluginCache:
|
||||
return PluginService.LatestPluginCache(
|
||||
plugin_id="langgenius/demo_plugin",
|
||||
version="1.3.0",
|
||||
unique_identifier="langgenius/demo_plugin:1.3.0@sha256:def",
|
||||
status="active",
|
||||
deprecated_reason="",
|
||||
alternative_plugin_id="",
|
||||
)
|
||||
|
||||
|
||||
def _expected_latest_plugin_cache_dump() -> dict[str, str]:
|
||||
return {
|
||||
"plugin_id": "langgenius/demo_plugin",
|
||||
"version": "1.3.0",
|
||||
"unique_identifier": "langgenius/demo_plugin:1.3.0@sha256:def",
|
||||
"status": "active",
|
||||
"deprecated_reason": "",
|
||||
"alternative_plugin_id": "",
|
||||
}
|
||||
|
||||
|
||||
def _dynamic_option() -> PluginParameterOption:
|
||||
return PluginParameterOption.model_validate(
|
||||
{
|
||||
"value": 101,
|
||||
"label": {"en_US": "Dataset 101"},
|
||||
"icon": None,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _expected_dynamic_option_dump() -> dict[str, Any]:
|
||||
return {
|
||||
"value": "101",
|
||||
"label": _expected_i18n("Dataset 101"),
|
||||
"icon": None,
|
||||
}
|
||||
|
||||
|
||||
def _account(role: TenantAccountRole = TenantAccountRole.OWNER) -> Account:
|
||||
account = Account(name="Test User", email="u1@example.com")
|
||||
account.id = "u1"
|
||||
@@ -140,17 +352,24 @@ class TestPluginListLatestVersionsApi:
|
||||
api = PluginListLatestVersionsApi()
|
||||
method = unwrap(api.post)
|
||||
|
||||
payload = {"plugin_ids": ["p1"]}
|
||||
payload = {"plugin_ids": ["langgenius/demo_plugin", "langgenius/missing_plugin"]}
|
||||
versions = {
|
||||
"langgenius/demo_plugin": _latest_plugin_cache(),
|
||||
"langgenius/missing_plugin": None,
|
||||
}
|
||||
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
patch(
|
||||
"controllers.console.workspace.plugin.PluginService.list_latest_versions", return_value={"p1": "1.0"}
|
||||
),
|
||||
patch("controllers.console.workspace.plugin.PluginService.list_latest_versions", return_value=versions),
|
||||
):
|
||||
result = method(api)
|
||||
|
||||
assert "versions" in result
|
||||
assert result == {
|
||||
"versions": {
|
||||
"langgenius/demo_plugin": _expected_latest_plugin_cache_dump(),
|
||||
"langgenius/missing_plugin": None,
|
||||
}
|
||||
}
|
||||
|
||||
def test_daemon_error(self, app: Flask):
|
||||
api = PluginListLatestVersionsApi()
|
||||
@@ -202,18 +421,18 @@ class TestPluginListApi:
|
||||
api = PluginListApi()
|
||||
method = unwrap(api.get)
|
||||
|
||||
mock_list = MagicMock(list=[{"id": 1}], total=1)
|
||||
plugins_with_total = MagicMock(list=[_plugin_entity()], total=1)
|
||||
|
||||
with (
|
||||
app.test_request_context("/?page=1&page_size=10"),
|
||||
patch(
|
||||
"controllers.console.workspace.plugin.PluginService.list_with_total",
|
||||
return_value=mock_list,
|
||||
return_value=plugins_with_total,
|
||||
) as mock_list_with_total,
|
||||
):
|
||||
result = method(api, "t1", "u1")
|
||||
|
||||
assert result["total"] == 1
|
||||
assert result == {"plugins": [_expected_plugin_entity_dump()], "total": 1}
|
||||
mock_list_with_total.assert_called_once_with("t1", "u1", 1, 10)
|
||||
|
||||
|
||||
@@ -454,12 +673,12 @@ class TestPluginFetchDynamicSelectOptionsApi:
|
||||
app.test_request_context("/?plugin_id=p&provider=x&action=y¶meter=z&provider_type=tool"),
|
||||
patch(
|
||||
"controllers.console.workspace.plugin.PluginParameterService.get_dynamic_select_options",
|
||||
return_value=[1, 2],
|
||||
return_value=[_dynamic_option()],
|
||||
),
|
||||
):
|
||||
result = method(api, "t1", user)
|
||||
|
||||
assert result["options"] == [1, 2]
|
||||
assert result == {"options": [_expected_dynamic_option_dump()]}
|
||||
|
||||
|
||||
class TestPluginReadmeApi:
|
||||
@@ -481,29 +700,18 @@ class TestPluginListInstallationsFromIdsApi:
|
||||
api = PluginListInstallationsFromIdsApi()
|
||||
method = unwrap(api.post)
|
||||
|
||||
payload = {"plugin_ids": ["p1", "p2"]}
|
||||
payload = {"plugin_ids": ["langgenius/demo_plugin"]}
|
||||
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
patch(
|
||||
"controllers.console.workspace.plugin.PluginService.list_installations_from_ids",
|
||||
return_value=[PluginInstallation.model_validate(_plugin_category_list_item())],
|
||||
return_value=[_plugin_installation()],
|
||||
),
|
||||
):
|
||||
result = method(api, "t1")
|
||||
|
||||
assert result["plugins"][0]["id"] == "entity-1"
|
||||
assert result["plugins"][0]["plugin_id"] == "test-author/test-plugin"
|
||||
assert result["plugins"][0]["plugin_unique_identifier"] == "test-author/test-plugin:1.0.0@checksum"
|
||||
assert result["plugins"][0]["version"] == "1.0.0"
|
||||
assert result["plugins"][0]["declaration"]["name"] == "test-plugin"
|
||||
assert "name" not in result["plugins"][0]
|
||||
assert "installation_id" not in result["plugins"][0]
|
||||
assert "latest_version" not in result["plugins"][0]
|
||||
assert "latest_unique_identifier" not in result["plugins"][0]
|
||||
assert "status" not in result["plugins"][0]
|
||||
assert "deprecated_reason" not in result["plugins"][0]
|
||||
assert "alternative_plugin_id" not in result["plugins"][0]
|
||||
assert result == {"plugins": [_expected_plugin_installation_dump()]}
|
||||
|
||||
def test_daemon_error(self, app: Flask):
|
||||
api = PluginListInstallationsFromIdsApi()
|
||||
@@ -689,11 +897,14 @@ class TestPluginFetchMarketplacePkgApi:
|
||||
|
||||
with (
|
||||
app.test_request_context("/?plugin_unique_identifier=p"),
|
||||
patch("controllers.console.workspace.plugin.PluginService.fetch_marketplace_pkg", return_value={"m": 1}),
|
||||
patch(
|
||||
"controllers.console.workspace.plugin.PluginService.fetch_marketplace_pkg",
|
||||
return_value=_plugin_declaration(),
|
||||
),
|
||||
):
|
||||
result = method(api, "t1")
|
||||
|
||||
assert "manifest" in result
|
||||
assert result == {"manifest": _expected_plugin_declaration_dump()}
|
||||
|
||||
def test_daemon_error(self, app: Flask):
|
||||
api = PluginFetchMarketplacePkgApi()
|
||||
@@ -715,8 +926,7 @@ class TestPluginFetchManifestApi:
|
||||
api = PluginFetchManifestApi()
|
||||
method = unwrap(api.get)
|
||||
|
||||
manifest = MagicMock()
|
||||
manifest.model_dump.return_value = {"x": 1}
|
||||
manifest = _plugin_declaration()
|
||||
|
||||
with (
|
||||
app.test_request_context("/?plugin_unique_identifier=p"),
|
||||
@@ -724,7 +934,7 @@ class TestPluginFetchManifestApi:
|
||||
):
|
||||
result = method(api, "t1")
|
||||
|
||||
assert "manifest" in result
|
||||
assert result == {"manifest": _expected_plugin_declaration_dump()}
|
||||
|
||||
def test_daemon_error(self, app: Flask):
|
||||
api = PluginFetchManifestApi()
|
||||
@@ -748,11 +958,14 @@ class TestPluginFetchInstallTasksApi:
|
||||
|
||||
with (
|
||||
app.test_request_context("/?page=1&page_size=10"),
|
||||
patch("controllers.console.workspace.plugin.PluginService.fetch_install_tasks", return_value=[{"id": 1}]),
|
||||
patch(
|
||||
"controllers.console.workspace.plugin.PluginService.fetch_install_tasks",
|
||||
return_value=[_plugin_task()],
|
||||
),
|
||||
):
|
||||
result = method(api, "t1")
|
||||
|
||||
assert "tasks" in result
|
||||
assert result == {"tasks": [_expected_plugin_task_dump()]}
|
||||
|
||||
def test_daemon_error(self, app: Flask):
|
||||
api = PluginFetchInstallTasksApi()
|
||||
@@ -776,11 +989,11 @@ class TestPluginFetchInstallTaskApi:
|
||||
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch("controllers.console.workspace.plugin.PluginService.fetch_install_task", return_value={"id": "x"}),
|
||||
patch("controllers.console.workspace.plugin.PluginService.fetch_install_task", return_value=_plugin_task()),
|
||||
):
|
||||
result = method(api, "t1", "x")
|
||||
|
||||
assert "task" in result
|
||||
assert result == {"task": _expected_plugin_task_dump()}
|
||||
|
||||
def test_daemon_error(self, app: Flask):
|
||||
api = PluginFetchInstallTaskApi()
|
||||
@@ -989,12 +1202,12 @@ class TestPluginFetchDynamicSelectOptionsWithCredentialsApi:
|
||||
app.test_request_context("/", json=payload),
|
||||
patch(
|
||||
"controllers.console.workspace.plugin.PluginParameterService.get_dynamic_select_options_with_credentials",
|
||||
return_value=[1],
|
||||
return_value=[_dynamic_option()],
|
||||
),
|
||||
):
|
||||
result = method(api, "t1", user)
|
||||
|
||||
assert result["options"] == [1]
|
||||
assert result == {"options": [_expected_dynamic_option_dump()]}
|
||||
|
||||
def test_daemon_error(self, app: Flask, user):
|
||||
api = PluginFetchDynamicSelectOptionsWithCredentialsApi()
|
||||
|
||||
Reference in New Issue
Block a user