diff --git a/api/controllers/console/feature.py b/api/controllers/console/feature.py index 41aa6513a63..edc9a9f19ab 100644 --- a/api/controllers/console/feature.py +++ b/api/controllers/console/feature.py @@ -75,18 +75,18 @@ class FeatureVectorSpaceApi(Resource): @console_ns.route("/trial-models") class TrialModelsApi(Resource): @console_ns.doc("get_trial_models") - @console_ns.doc(description="Get hosted trial model provider configuration") + @console_ns.doc(description="Get hosted credit model provider configuration for the current workspace") @console_ns.response( 200, "Success", console_ns.models[TrialModelsResponse.__name__], ) @console_account_admission() - def get(self, _request_context: RequestContext): - """Get hosted trial model provider configuration for model-provider pages.""" + def get(self, request_context: RequestContext): + """Get hosted credit provider configuration for the current workspace.""" return dump_response( TrialModelsResponse, - {"trial_models": application_services().feature_queries.get_trial_models()}, + {"trial_models": application_services().feature_queries.get_trial_models(request_context)}, ) diff --git a/api/enums/__init__.py b/api/enums/__init__.py index a56af59be25..fdcc5e94795 100644 --- a/api/enums/__init__.py +++ b/api/enums/__init__.py @@ -14,6 +14,10 @@ class CloudPlan(StrEnum): PROFESSIONAL = auto() TEAM = auto() + @property + def is_paid(self) -> bool: + return self in (CloudPlan.PROFESSIONAL, CloudPlan.TEAM) + class DeploymentEdition(StrEnum): """Enum representing the deployment edition of the platform.""" diff --git a/api/extensions/ext_application_services.py b/api/extensions/ext_application_services.py index dc4015767ca..5829cad861a 100644 --- a/api/extensions/ext_application_services.py +++ b/api/extensions/ext_application_services.py @@ -298,7 +298,6 @@ def build_application_services( ), feature_queries=FeatureQueryService( features=feature_gateway, - trial_models=FeatureService.get_trial_models(), app_dsl_version=CURRENT_APP_DSL_VERSION, ), init_validation=InitValidationService( diff --git a/api/openapi/markdown/console-openapi.md b/api/openapi/markdown/console-openapi.md index 848fe656dc8..6a4c9baee78 100644 --- a/api/openapi/markdown/console-openapi.md +++ b/api/openapi/markdown/console-openapi.md @@ -9566,9 +9566,9 @@ Returns the site configuration for the application including theme, icons, and t | 200 | Success | **application/json**: [SimpleResultResponse](#simpleresultresponse)
| ### [GET] /trial-models -**Get hosted trial model provider configuration for model-provider pages** +**Get hosted credit provider configuration for the current workspace** -Get hosted trial model provider configuration +Get hosted credit model provider configuration for the current workspace #### Responses diff --git a/api/services/feature_query_service.py b/api/services/feature_query_service.py index f24e387bb0f..8e0c82e6e25 100644 --- a/api/services/feature_query_service.py +++ b/api/services/feature_query_service.py @@ -1,6 +1,5 @@ """Application service for feature queries exposed by API adapters.""" -from collections.abc import Sequence from typing import Protocol from machinery.context import RequestContext @@ -17,6 +16,8 @@ class FeatureQueryGateway(Protocol): def get_workspace_features(self, workspace_id: str) -> FeatureModel: ... + def get_trial_models(self, workspace_id: str) -> list[str]: ... + def get_vector_space(self, workspace_id: str) -> VectorSpaceLimitationModel: ... def get_public_system_features(self) -> SystemFeatureModel: ... @@ -29,11 +30,9 @@ class FeatureQueryService: self, *, features: FeatureQueryGateway, - trial_models: Sequence[str], app_dsl_version: str, ) -> None: self._features = features - self._trial_models = tuple(trial_models) self._app_dsl_version = app_dsl_version def get_features(self, context: RequestContext) -> FeatureModel: @@ -42,8 +41,8 @@ class FeatureQueryService: def get_vector_space(self, context: RequestContext) -> VectorSpaceLimitationModel: return self._features.get_vector_space(self._require_active_workspace(context)) - def get_trial_models(self) -> list[str]: - return list(self._trial_models) + def get_trial_models(self, context: RequestContext) -> list[str]: + return self._features.get_trial_models(self._require_active_workspace(context)) def get_app_dsl_version(self) -> str: return self._app_dsl_version diff --git a/api/services/feature_service.py b/api/services/feature_service.py index 840aa4dda06..ad716a71f6e 100644 --- a/api/services/feature_service.py +++ b/api/services/feature_service.py @@ -20,6 +20,16 @@ class _EnterprisePluginInstallationPermission(BaseModel): class FeatureService: + @classmethod + def get_workspace_plan(cls, tenant_id: str) -> CloudPlan: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD: + return CloudPlan.SANDBOX + + billing_info = BillingService.get_info(tenant_id, exclude_vector_space=True) + if not billing_info["enabled"]: + return CloudPlan.SANDBOX + return CloudPlan(billing_info["subscription"]["plan"]) + @classmethod def get_features(cls, tenant_id: str, exclude_vector_space: bool = False) -> feature_entities.FeatureModel: features = feature_entities.FeatureModel() @@ -76,11 +86,8 @@ class FeatureService: if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD or not tenant_id: return default_limit - billing_info = BillingService.get_info(tenant_id, exclude_vector_space=True) - if billing_info["enabled"] and billing_info["subscription"]["plan"] in ( - CloudPlan.PROFESSIONAL, - CloudPlan.TEAM, - ): + subscription_plan = cls.get_workspace_plan(tenant_id) + if subscription_plan.is_paid: return max(default_limit, dify_config.KNOWLEDGE_UPLOAD_FILE_SIZE_LIMIT_FOR_PAID_PLAN) return default_limit @@ -93,10 +100,7 @@ class FeatureService: return True if not tenant_id: return False - return features.billing.enabled and features.billing.subscription.plan in ( - CloudPlan.PROFESSIONAL, - CloudPlan.TEAM, - ) + return features.billing.enabled and features.billing.subscription.plan.is_paid @classmethod def get_system_features(cls) -> feature_entities.SystemFeatureModel: @@ -183,20 +187,26 @@ class FeatureService: system_features.knowledge_fs_enabled = dify_config.KNOWLEDGE_FS_ENABLED @classmethod - def _fulfill_trial_models_from_env(cls) -> list[str]: + def _fulfill_trial_models_from_env(cls, quota_types: tuple[str, ...] | None = None) -> list[str]: + allowed_quota_types = quota_types or ("PAID", "TRIAL") return [ provider.value for provider in HostedTrialProvider - if ( - getattr(dify_config, f"HOSTED_{provider.config_key}_PAID_ENABLED", False) - and getattr(dify_config, f"HOSTED_{provider.config_key}_TRIAL_ENABLED", False) + if any( + getattr(dify_config, f"HOSTED_{provider.config_key}_{quota_type}_ENABLED", False) + for quota_type in allowed_quota_types ) ] @classmethod - def get_trial_models(cls) -> list[str]: - """Return hosted trial provider ids without requiring the full system-features payload.""" - return cls._fulfill_trial_models_from_env() + def get_trial_models(cls, tenant_id: str) -> list[str]: + """Return hosted credit providers filtered by the workspace subscription plan.""" + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD: + return cls._fulfill_trial_models_from_env() + + subscription_plan = cls.get_workspace_plan(tenant_id) + quota_types = ("PAID", "TRIAL") if subscription_plan.is_paid else ("TRIAL",) + return cls._fulfill_trial_models_from_env(quota_types) @classmethod def _fulfill_params_from_env(cls, features: feature_entities.FeatureModel): diff --git a/api/services/feature_service_gateway.py b/api/services/feature_service_gateway.py index 3afa5096442..1a93b380dd5 100644 --- a/api/services/feature_service_gateway.py +++ b/api/services/feature_service_gateway.py @@ -19,6 +19,10 @@ class FeatureServiceGateway(FeatureQueryGateway): def get_workspace_features(self, workspace_id: str) -> FeatureModel: return FeatureService.get_features(workspace_id, exclude_vector_space=True) + @override + def get_trial_models(self, workspace_id: str) -> list[str]: + return FeatureService.get_trial_models(workspace_id) + @override def get_vector_space(self, workspace_id: str) -> VectorSpaceLimitationModel: return FeatureService.get_vector_space(workspace_id) diff --git a/api/tests/unit_tests/controllers/console/test_feature.py b/api/tests/unit_tests/controllers/console/test_feature.py index 0f2f9970cb7..a97f366da5b 100644 --- a/api/tests/unit_tests/controllers/console/test_feature.py +++ b/api/tests/unit_tests/controllers/console/test_feature.py @@ -107,10 +107,11 @@ class TestTrialModelsApi: api = TrialModelsApi() raw_get = unwrap(TrialModelsApi.get) - result = raw_get(api, _request_context()) + request_context = _request_context() + result = raw_get(api, request_context) assert result == {"trial_models": ["langgenius/openai/openai"]} - get_trial_models.assert_called_once_with() + get_trial_models.assert_called_once_with(request_context) class TestAppDslVersionApi: diff --git a/api/tests/unit_tests/services/test_feature_query_service.py b/api/tests/unit_tests/services/test_feature_query_service.py index febbfbc6af3..9e99acd1c1f 100644 --- a/api/tests/unit_tests/services/test_feature_query_service.py +++ b/api/tests/unit_tests/services/test_feature_query_service.py @@ -28,13 +28,16 @@ def test_workspace_queries_use_workspace_from_request_context() -> None: vector_space = VectorSpaceLimitationModel(size=1, limit=5) gateway.get_workspace_features.return_value = features gateway.get_vector_space.return_value = vector_space - service = FeatureQueryService(features=gateway, trial_models=(), app_dsl_version="0.7.0") + gateway.get_trial_models.return_value = ["langgenius/openai/openai"] + service = FeatureQueryService(features=gateway, app_dsl_version="0.7.0") context = _request_context() assert service.get_features(context) is features assert service.get_vector_space(context) is vector_space + assert service.get_trial_models(context) == ["langgenius/openai/openai"] gateway.get_workspace_features.assert_called_once_with("workspace_123") gateway.get_vector_space.assert_called_once_with("workspace_123") + gateway.get_trial_models.assert_called_once_with("workspace_123") def test_deployment_queries_delegate_without_request_context() -> None: @@ -45,11 +48,9 @@ def test_deployment_queries_delegate_without_request_context() -> None: gateway.get_license.return_value = license_model service = FeatureQueryService( features=gateway, - trial_models=["langgenius/openai/openai"], app_dsl_version="0.6.0", ) - assert service.get_trial_models() == ["langgenius/openai/openai"] assert service.get_app_dsl_version() == "0.6.0" assert service.get_system_features() is system_features assert service.get_license() is license_model @@ -57,7 +58,7 @@ def test_deployment_queries_delegate_without_request_context() -> None: def test_workspace_queries_require_active_workspace() -> None: gateway = create_autospec(FeatureQueryGateway, instance=True, spec_set=True) - service = FeatureQueryService(features=gateway, trial_models=(), app_dsl_version="0.7.0") + service = FeatureQueryService(features=gateway, app_dsl_version="0.7.0") with pytest.raises(RuntimeError, match="did not resolve an active workspace"): service.get_features(_request_context(active_workspace_id=None)) diff --git a/api/tests/unit_tests/services/test_feature_service_gateway.py b/api/tests/unit_tests/services/test_feature_service_gateway.py index 253df785743..8dd981e11f8 100644 --- a/api/tests/unit_tests/services/test_feature_service_gateway.py +++ b/api/tests/unit_tests/services/test_feature_service_gateway.py @@ -24,3 +24,13 @@ def test_workspace_features_exclude_independently_queried_vector_space(mocker: M assert result is features get_features.assert_called_once_with("workspace_123", exclude_vector_space=True) + + +def test_trial_models_delegate_to_workspace_aware_feature_service(mocker: MockerFixture) -> None: + trial_models = ["langgenius/openai/openai"] + get_trial_models = mocker.patch.object(FeatureService, "get_trial_models", return_value=trial_models) + + result = FeatureServiceGateway().get_trial_models("workspace_123") + + assert result == trial_models + get_trial_models.assert_called_once_with("workspace_123") diff --git a/api/tests/unit_tests/services/test_feature_service_trial_models.py b/api/tests/unit_tests/services/test_feature_service_trial_models.py index 599623aaba0..6c5c2c556da 100644 --- a/api/tests/unit_tests/services/test_feature_service_trial_models.py +++ b/api/tests/unit_tests/services/test_feature_service_trial_models.py @@ -1,6 +1,8 @@ +from unittest.mock import Mock + import pytest -from enums import HostedTrialProvider +from enums import CloudPlan, DeploymentEdition, HostedTrialProvider from services import feature_service as feature_service_module from services.feature_service import FeatureService @@ -11,7 +13,7 @@ def test_get_system_features_excludes_trial_models(): assert "trial_models" not in result -def test_get_trial_models_returns_providers_enabled_for_paid_and_trial(monkeypatch: pytest.MonkeyPatch): +def test_get_trial_models_returns_providers_with_paid_or_trial_enabled(monkeypatch: pytest.MonkeyPatch): for provider in HostedTrialProvider: monkeypatch.setattr( feature_service_module.dify_config, @@ -28,11 +30,50 @@ def test_get_trial_models_returns_providers_enabled_for_paid_and_trial(monkeypat monkeypatch.setattr(feature_service_module.dify_config, "HOSTED_OPENAI_PAID_ENABLED", True, raising=False) monkeypatch.setattr(feature_service_module.dify_config, "HOSTED_OPENAI_TRIAL_ENABLED", True, raising=False) - monkeypatch.setattr(feature_service_module.dify_config, "HOSTED_ANTHROPIC_PAID_ENABLED", True, raising=False) - monkeypatch.setattr(feature_service_module.dify_config, "HOSTED_ANTHROPIC_TRIAL_ENABLED", False, raising=False) - monkeypatch.setattr(feature_service_module.dify_config, "HOSTED_GEMINI_PAID_ENABLED", False, raising=False) - monkeypatch.setattr(feature_service_module.dify_config, "HOSTED_GEMINI_TRIAL_ENABLED", True, raising=False) + monkeypatch.setattr(feature_service_module.dify_config, "HOSTED_XAI_PAID_ENABLED", True, raising=False) + monkeypatch.setattr(feature_service_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) - result = FeatureService.get_trial_models() + result = FeatureService.get_trial_models("tenant_1") - assert result == [HostedTrialProvider.OPENAI.value] + assert result == [ + HostedTrialProvider.OPENAI.value, + HostedTrialProvider.X.value, + ] + + +@pytest.mark.parametrize( + ("plan", "expected"), + [ + (CloudPlan.SANDBOX, [HostedTrialProvider.OPENAI.value]), + (CloudPlan.PROFESSIONAL, [HostedTrialProvider.OPENAI.value, HostedTrialProvider.X.value]), + ], +) +def test_get_trial_models_filters_providers_by_workspace_plan( + monkeypatch: pytest.MonkeyPatch, + plan: CloudPlan, + expected: list[str], +) -> None: + for provider in HostedTrialProvider: + monkeypatch.setattr( + feature_service_module.dify_config, + f"HOSTED_{provider.config_key}_PAID_ENABLED", + False, + raising=False, + ) + monkeypatch.setattr( + feature_service_module.dify_config, + f"HOSTED_{provider.config_key}_TRIAL_ENABLED", + False, + raising=False, + ) + + monkeypatch.setattr(feature_service_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) + monkeypatch.setattr(feature_service_module.dify_config, "HOSTED_OPENAI_TRIAL_ENABLED", True, raising=False) + monkeypatch.setattr(feature_service_module.dify_config, "HOSTED_XAI_PAID_ENABLED", True, raising=False) + get_workspace_plan = Mock(return_value=plan) + monkeypatch.setattr(feature_service_module.FeatureService, "get_workspace_plan", get_workspace_plan) + + result = FeatureService.get_trial_models("tenant_1") + + assert result == expected + get_workspace_plan.assert_called_once_with("tenant_1") diff --git a/packages/contracts/generated/api/console/trial-models/orpc.gen.ts b/packages/contracts/generated/api/console/trial-models/orpc.gen.ts index e0aa34bf478..ce3fb8d09a4 100644 --- a/packages/contracts/generated/api/console/trial-models/orpc.gen.ts +++ b/packages/contracts/generated/api/console/trial-models/orpc.gen.ts @@ -4,18 +4,18 @@ import { oc } from '@orpc/contract' import { zGetTrialModelsResponse } from './zod.gen' /** - * Get hosted trial model provider configuration for model-provider pages + * Get hosted credit provider configuration for the current workspace * - * Get hosted trial model provider configuration + * Get hosted credit model provider configuration for the current workspace */ export const get = oc .route({ - description: 'Get hosted trial model provider configuration', + description: 'Get hosted credit model provider configuration for the current workspace', inputStructure: 'detailed', method: 'GET', operationId: 'getTrialModels', path: '/trial-models', - summary: 'Get hosted trial model provider configuration for model-provider pages', + summary: 'Get hosted credit provider configuration for the current workspace', tags: ['console'], }) .output(zGetTrialModelsResponse)