mirror of
https://github.com/langgenius/dify.git
synced 2026-09-29 17:07:38 +08:00
fix(api): scope hosted credits by tenant plan (#41131)
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] <114827586+autofix-ci[bot]@users.noreply.github.com>
parent
bec4ef6ac7
commit
70a513fc62
@@ -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)},
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -9566,9 +9566,9 @@ Returns the site configuration for the application including theme, icons, and t
|
||||
| 200 | Success | **application/json**: [SimpleResultResponse](#simpleresultresponse)<br> |
|
||||
|
||||
### [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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user