mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
feat: optimize integrations initial loading (#40118)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
autofix-ci[bot]
parent
6f2e4d72d3
commit
5dfd83e9ab
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user